Skip to main content

ironwork_rt/sql/
replay.rs

1//! Recordings: each call a run made and the answer it got, as text a person can read and edit.
2//! [`Replay`] answers from a recording; [`Recorder`] writes one while another backend answers.
3//!
4//! ```text
5//! # ironwork sql recording 1
6//! @ 1 PAYROLL:3:9f2a41c0 SELECT
7//! > char:"00123"
8//! < 0 00000 rows=1
9//! = dec:1234.50 | char:"SMITH" | null
10//! ```
11
12use super::{Abandoned, Answer, Call, Database, Outcome, Value};
13use std::io::Write;
14use super::fingerprint;
15
16const HEADER: &str = "# ironwork sql recording 1";
17
18#[derive(Debug)]
19struct Entry {
20    program: String,
21    ordinal: u32,
22    hash: u32,
23    verb: String,
24    cursor: Option<String>,
25    inputs: Vec<Value>,
26    outcome: Outcome,
27}
28
29impl Entry {
30    fn answers(&self, call: &Call) -> bool {
31        self.program == call.program
32            && self.ordinal == call.ordinal
33            && self.hash == fingerprint(call.text)
34            && self.verb == call.verb
35            && self.cursor.as_deref() == call.cursor
36            && self.inputs == call.inputs
37    }
38}
39
40fn describe(program: &str, ordinal: u32, hash: u32, verb: &str, cursor: Option<&str>, inputs: &[Value]) -> String {
41    let cursor = cursor.map(|c| format!(" {c}")).unwrap_or_default();
42    let inputs = if inputs.is_empty() { String::new() } else { format!(" with {}", values_text(inputs)) };
43    format!("{program}:{ordinal}:{hash:08x} {verb}{cursor}{inputs}")
44}
45
46/// Answers calls from a recording. Strict replay answers call n from entry n; keyed replay answers
47/// each call from the first unused entry with the same statement and inputs. A call the recording
48/// does not hold ends the run.
49pub struct Replay {
50    entries: Vec<Entry>,
51    used: Vec<bool>,
52    next: usize,
53    keyed: bool,
54}
55
56impl Replay {
57    pub fn parse(text: &str, keyed: bool) -> Result<Self, String> {
58        let mut entries: Vec<Entry> = Vec::new();
59        let mut header = false;
60        let mut awaiting_outcome = false;
61        for (n, line) in text.lines().enumerate().map(|(i, l)| (i + 1, l.trim_end())) {
62            let fail = |why: String| format!("line {n}: {why}");
63            if line.is_empty() {
64                continue;
65            }
66            if line.starts_with('#') {
67                header |= line == HEADER;
68                continue;
69            }
70            if !header {
71                return Err(fail(format!("a recording starts with \"{HEADER}\"")));
72            }
73            let (mark, rest) = line.split_at(1);
74            let rest = rest.trim_start();
75            match mark {
76                "@" => {
77                    if awaiting_outcome {
78                        return Err(fail("the call before this one has no < line".into()));
79                    }
80                    let words: Vec<&str> = rest.split_whitespace().collect();
81                    let [_, id, verb, cursor @ ..] = words.as_slice() else { return Err(fail("@ takes a number, PROGRAM:ORDINAL:HASH and a verb".into())) };
82                    let parts: Vec<&str> = id.split(':').collect();
83                    let [program, ordinal, hash] = parts.as_slice() else { return Err(fail(format!("{id} is not PROGRAM:ORDINAL:HASH"))) };
84                    let ordinal = ordinal.parse().map_err(|_| fail(format!("{ordinal} is not an ordinal")))?;
85                    let hash = u32::from_str_radix(hash, 16).map_err(|_| fail(format!("{hash} is not a hexadecimal hash")))?;
86                    let cursor = match cursor {
87                        [] => None,
88                        [c] => Some((*c).to_owned()),
89                        _ => return Err(fail("@ takes at most one cursor after the verb".into())),
90                    };
91                    entries.push(Entry { program: (*program).into(), ordinal, hash, verb: (*verb).into(), cursor, inputs: Vec::new(), outcome: Outcome::ok() });
92                    awaiting_outcome = true;
93                }
94                ">" => match entries.last_mut() {
95                    Some(e) if awaiting_outcome && e.inputs.is_empty() => e.inputs = parse_values(rest).map_err(fail)?,
96                    _ => return Err(fail("a > line belongs after an @ line, before its < line".into())),
97                },
98                "<" => match entries.last_mut() {
99                    Some(e) if awaiting_outcome => {
100                        e.outcome = parse_outcome(rest).map_err(fail)?;
101                        awaiting_outcome = false;
102                    }
103                    _ => return Err(fail("a < line belongs after an @ line".into())),
104                },
105                "=" => match entries.last_mut() {
106                    Some(e) if !awaiting_outcome => e.outcome.rows.push(parse_values(rest).map_err(fail)?),
107                    _ => return Err(fail("an = line belongs after a < line".into())),
108                },
109                _ => return Err(fail(format!("{mark} starts no kind of line"))),
110            }
111        }
112        if awaiting_outcome {
113            return Err("the last call has no < line".into());
114        }
115        if !header {
116            return Err(format!("a recording starts with \"{HEADER}\""));
117        }
118        let used = vec![false; entries.len()];
119        Ok(Self { entries, used, next: 0, keyed })
120    }
121
122    fn answer(&mut self, call: &Call) -> Answer {
123        let found = if self.keyed {
124            (0..self.entries.len()).find(|&i| !self.used[i] && self.entries[i].answers(call))
125        } else {
126            (self.next < self.entries.len() && self.entries[self.next].answers(call)).then_some(self.next)
127        };
128        let actual = describe(call.program, call.ordinal, fingerprint(call.text), call.verb, call.cursor, call.inputs);
129        let Some(i) = found else {
130            let expected = match self.entries.get(self.next).filter(|_| !self.keyed) {
131                Some(e) => format!("the recording's call {} is {}", self.next + 1, describe(&e.program, e.ordinal, e.hash, &e.verb, e.cursor.as_deref(), &e.inputs)),
132                None => "the recording holds no such call".into(),
133            };
134            return Err(Abandoned { code: "SQLR", message: format!("{expected}, and the run made {actual}") });
135        };
136        self.used[i] = true;
137        self.next = i + 1;
138        Ok(self.entries[i].outcome.clone())
139    }
140}
141
142impl Database for Replay {
143    fn execute(&mut self, call: &Call) -> Answer {
144        self.answer(call)
145    }
146    fn prepare(&mut self, call: &Call) -> Answer {
147        self.answer(call)
148    }
149    fn open(&mut self, call: &Call) -> Answer {
150        self.answer(call)
151    }
152    fn fetch(&mut self, call: &Call) -> Answer {
153        self.answer(call)
154    }
155    fn close(&mut self, call: &Call) -> Answer {
156        self.answer(call)
157    }
158    fn commit(&mut self, call: &Call) -> Answer {
159        self.answer(call)
160    }
161    fn rollback(&mut self, call: &Call) -> Answer {
162        self.answer(call)
163    }
164}
165
166/// Writes every call and its answer while `inner` answers.
167pub struct Recorder<'w> {
168    inner: Box<dyn Database + 'w>,
169    out: Box<dyn Write + 'w>,
170    seq: u64,
171}
172
173impl<'w> Recorder<'w> {
174    /// `source` names what answered, and goes in the recording's header.
175    pub fn new(inner: Box<dyn Database + 'w>, mut out: Box<dyn Write + 'w>, source: &str) -> std::io::Result<Self> {
176        writeln!(out, "{HEADER}\n# source: {source}")?;
177        Ok(Self { inner, out, seq: 0 })
178    }
179
180    fn record(&mut self, call: &Call, answer: Answer) -> Answer {
181        let outcome = answer?;
182        self.seq += 1;
183        let text = entry_text(self.seq, call, &outcome);
184        // Flushed per call: a served session ends when the server is interrupted, not by returning.
185        self.out.write_all(text.as_bytes()).and_then(|()| self.out.flush()).map_err(|e| Abandoned { code: "SQLR", message: format!("the recording could not be written: {e}") })?;
186        Ok(outcome)
187    }
188}
189
190impl Database for Recorder<'_> {
191    fn execute(&mut self, call: &Call) -> Answer {
192        let a = self.inner.execute(call);
193        self.record(call, a)
194    }
195    fn prepare(&mut self, call: &Call) -> Answer {
196        let a = self.inner.prepare(call);
197        self.record(call, a)
198    }
199    fn open(&mut self, call: &Call) -> Answer {
200        let a = self.inner.open(call);
201        self.record(call, a)
202    }
203    fn fetch(&mut self, call: &Call) -> Answer {
204        let a = self.inner.fetch(call);
205        self.record(call, a)
206    }
207    fn close(&mut self, call: &Call) -> Answer {
208        let a = self.inner.close(call);
209        self.record(call, a)
210    }
211    fn commit(&mut self, call: &Call) -> Answer {
212        let a = self.inner.commit(call);
213        self.record(call, a)
214    }
215    fn rollback(&mut self, call: &Call) -> Answer {
216        let a = self.inner.rollback(call);
217        self.record(call, a)
218    }
219    fn close_all(&mut self) -> Result<(), Abandoned> {
220        self.inner.close_all()
221    }
222}
223
224fn entry_text(seq: u64, call: &Call, outcome: &Outcome) -> String {
225    let cursor = call.cursor.map(|c| format!(" {c}")).unwrap_or_default();
226    let mut text = format!("@ {seq} {}:{}:{:08x} {}{cursor}\n", call.program, call.ordinal, fingerprint(call.text), call.verb);
227    if !call.inputs.is_empty() {
228        text += &format!("> {}\n", values_text(call.inputs));
229    }
230    text += &format!("< {} {} rows={}", outcome.sqlcode, outcome.sqlstate, outcome.affected);
231    if !outcome.tokens.is_empty() {
232        text += &format!(" tokens={}", value_text(&Value::Char(outcome.tokens.clone())));
233    }
234    text.push('\n');
235    for row in &outcome.rows {
236        text += &format!("= {}\n", values_text(row));
237    }
238    text
239}
240
241fn values_text(values: &[Value]) -> String {
242    values.iter().map(value_text).collect::<Vec<_>>().join(" | ")
243}
244
245fn value_text(v: &Value) -> String {
246    match v {
247        Value::Null => "null".into(),
248        Value::Int(i) => format!("int:{i}"),
249        Value::Decimal { value, scale } => format!("dec:{}", Value::decimal_text(*value, *scale)),
250        Value::Double(f) => format!("double:{f:?}"),
251        Value::Char(s) => {
252            let mut q = String::from("char:\"");
253            for c in s.chars() {
254                match c {
255                    '"' => q += "\\\"",
256                    '\\' => q += "\\\\",
257                    c if c.is_control() => q += &format!("\\x{:02X}", c as u32),
258                    c => q.push(c),
259                }
260            }
261            q + "\""
262        }
263        Value::Binary(b) => format!("hex:{}", b.iter().map(|x| format!("{x:02X}")).collect::<String>()),
264    }
265}
266
267fn parse_outcome(text: &str) -> Result<Outcome, String> {
268    let mut words = text.splitn(4, ' ');
269    let (Some(code), Some(state), Some(rows)) = (words.next(), words.next(), words.next()) else { return Err("< takes SQLCODE, SQLSTATE and rows=N".into()) };
270    let sqlcode = code.parse().map_err(|_| format!("{code} is not an SQLCODE"))?;
271    if state.len() != 5 {
272        return Err(format!("{state} is not a five-character SQLSTATE"));
273    }
274    let affected = rows.strip_prefix("rows=").and_then(|n| n.parse().ok()).ok_or_else(|| format!("{rows} is not rows=N"))?;
275    let tokens = match words.next().map(str::trim) {
276        None | Some("") => String::new(),
277        Some(t) => match t.strip_prefix("tokens=").map(parse_value) {
278            Some(Ok((Value::Char(s), rest))) if rest.trim().is_empty() => s,
279            _ => return Err("the rest of a < line is tokens=char:\"...\"".into()),
280        },
281    };
282    Ok(Outcome { sqlcode, sqlstate: state.into(), affected, rows: Vec::new(), tokens })
283}
284
285fn parse_values(text: &str) -> Result<Vec<Value>, String> {
286    let (mut out, mut rest) = (Vec::new(), text.trim_start());
287    while !rest.is_empty() {
288        let (v, after) = parse_value(rest)?;
289        out.push(v);
290        rest = after.trim_start();
291        if let Some(next) = rest.strip_prefix('|') {
292            rest = next.trim_start();
293            if rest.is_empty() {
294                return Err("a value is missing after |".into());
295            }
296        } else if !rest.is_empty() {
297            return Err(format!("values are separated by |, not \"{rest}\""));
298        }
299    }
300    Ok(out)
301}
302
303/// One value literal from the start of `text`, and what follows it.
304fn parse_value(text: &str) -> Result<(Value, &str), String> {
305    if let Some(quoted) = text.strip_prefix("char:\"") {
306        let mut s = String::new();
307        let mut chars = quoted.char_indices();
308        while let Some((i, c)) = chars.next() {
309            match c {
310                '"' => return Ok((Value::Char(s), &quoted[i + 1..])),
311                '\\' => match chars.next() {
312                    Some((_, '"')) => s.push('"'),
313                    Some((_, '\\')) => s.push('\\'),
314                    Some((j, 'x')) => {
315                        let hex = quoted.get(j + 1..j + 3).ok_or("\\x takes two hexadecimal digits")?;
316                        let code = u32::from_str_radix(hex, 16).map_err(|_| format!("\\x{hex} is not hexadecimal"))?;
317                        s.push(char::from_u32(code).ok_or("\\x names no character")?);
318                        chars.next();
319                        chars.next();
320                    }
321                    _ => return Err("a backslash in char:\"...\" escapes \", \\ or xNN".into()),
322                },
323                c => s.push(c),
324            }
325        }
326        return Err("char:\" is not closed".into());
327    }
328    let end = text.find(|c: char| c.is_whitespace() || c == '|').unwrap_or(text.len());
329    let (word, rest) = text.split_at(end);
330    let value = match word.split_once(':') {
331        None if word == "null" => Value::Null,
332        Some(("int", n)) => Value::Int(n.parse().map_err(|_| format!("{word} is not an integer"))?),
333        Some(("dec", n)) => Value::parse_decimal(n).ok_or_else(|| format!("{word} is not a decimal"))?,
334        Some(("double", n)) => Value::Double(n.parse().map_err(|_| format!("{word} is not a double"))?),
335        Some(("hex", h)) if h.len() % 2 == 0 => {
336            let bytes: Result<Vec<u8>, _> = (0..h.len()).step_by(2).map(|i| u8::from_str_radix(&h[i..i + 2], 16)).collect();
337            Value::Binary(bytes.map_err(|_| format!("{word} is not hexadecimal"))?)
338        }
339        _ => return Err(format!("{word} is not a value: null, int:, dec:, double:, char:\"...\" or hex:")),
340    };
341    Ok((value, rest))
342}
343
344#[cfg(test)]
345mod tests {
346    use super::*;
347
348    fn call<'a>(verb: &'a str, text: &'a str, inputs: &'a [Value]) -> Call<'a> {
349        Call { program: "P", ordinal: 2, verb, cursor: None, text, inputs }
350    }
351
352    #[test]
353    fn values_round_trip() {
354        let values = vec![
355            Value::Null,
356            Value::Int(-42),
357            Value::Decimal { value: -123_450, scale: 2 },
358            Value::Decimal { value: 5, scale: 2 },
359            Value::Decimal { value: 7, scale: 0 },
360            Value::Double(0.1),
361            Value::Double(6.02e23),
362            Value::Char("say \"hi\" | x \\ é\n".into()),
363            Value::Binary(vec![0xC1, 0x00]),
364        ];
365        let text = values_text(&values);
366        assert!(text.contains("dec:-1234.50") && text.contains("dec:0.05") && text.contains("dec:7"), "{text}");
367        assert_eq!(parse_values(&text), Ok(values));
368    }
369
370    #[test]
371    fn a_recording_answers_the_calls_it_holds_in_order() {
372        let inputs = [Value::Int(7)];
373        let first = entry_text(1, &call("SELECT", "SELECT A FROM T WHERE K = ?", &inputs), &Outcome::rows(vec![vec![Value::Char("X".into())]]));
374        let second = entry_text(2, &call("COMMIT", "COMMIT", &[]), &Outcome::ok());
375        let mut replay = Replay::parse(&format!("{HEADER}\n{first}{second}"), false).expect("parses");
376        assert_eq!(replay.execute(&call("SELECT", "SELECT A FROM T WHERE K = ?", &inputs)).unwrap().rows, [[Value::Char("X".into())]]);
377        assert_eq!(replay.commit(&call("COMMIT", "COMMIT", &[])), Ok(Outcome::ok()));
378        let beyond = replay.commit(&call("COMMIT", "COMMIT", &[])).unwrap_err();
379        assert_eq!((beyond.code, beyond.message.starts_with("the recording holds no such call")), ("SQLR", true));
380    }
381
382    #[test]
383    fn strict_replay_refuses_a_different_call_and_names_both() {
384        let text = format!("{HEADER}\n{}", entry_text(1, &call("SELECT", "SELECT A FROM T WHERE K = ?", &[Value::Int(7)]), &Outcome::ok()));
385        let mut replay = Replay::parse(&text, false).unwrap();
386        let err = replay.execute(&call("SELECT", "SELECT A FROM T WHERE K = ?", &[Value::Int(8)])).unwrap_err();
387        assert_eq!(err.code, "SQLR");
388        assert!(err.message.contains("with int:7") && err.message.contains("with int:8"), "{}", err.message);
389    }
390
391    #[test]
392    fn keyed_replay_takes_calls_in_any_order() {
393        let (a, b) = ([Value::Int(1)], [Value::Int(2)]);
394        let text = format!(
395            "{HEADER}\n{}{}",
396            entry_text(1, &call("SELECT", "Q", &a), &Outcome::rows(vec![vec![Value::Int(10)]])),
397            entry_text(2, &call("SELECT", "Q", &b), &Outcome::rows(vec![vec![Value::Int(20)]]))
398        );
399        let mut replay = Replay::parse(&text, true).unwrap();
400        assert_eq!(replay.execute(&call("SELECT", "Q", &b)).unwrap().rows, [[Value::Int(20)]]);
401        assert_eq!(replay.execute(&call("SELECT", "Q", &a)).unwrap().rows, [[Value::Int(10)]]);
402        assert!(replay.execute(&call("SELECT", "Q", &a)).is_err());
403    }
404
405    #[test]
406    fn a_recorder_writes_what_replay_reads() {
407        struct Fixed;
408        impl Database for Fixed {
409            fn execute(&mut self, _: &Call) -> Answer {
410                Ok(Outcome { tokens: "T1".into(), ..Outcome::rows(vec![vec![Value::Decimal { value: 150, scale: 2 }, Value::Null]]) })
411            }
412            fn prepare(&mut self, _: &Call) -> Answer {
413                Ok(Outcome::ok())
414            }
415            fn open(&mut self, _: &Call) -> Answer {
416                Ok(Outcome::ok())
417            }
418            fn fetch(&mut self, _: &Call) -> Answer {
419                Ok(Outcome::error(100, "02000"))
420            }
421            fn close(&mut self, _: &Call) -> Answer {
422                Ok(Outcome::ok())
423            }
424            fn commit(&mut self, _: &Call) -> Answer {
425                Ok(Outcome::ok())
426            }
427            fn rollback(&mut self, _: &Call) -> Answer {
428                Ok(Outcome::ok())
429            }
430        }
431        let written = std::rc::Rc::new(std::cell::RefCell::new(Vec::new()));
432        struct Sink(std::rc::Rc<std::cell::RefCell<Vec<u8>>>);
433        impl Write for Sink {
434            fn write(&mut self, b: &[u8]) -> std::io::Result<usize> {
435                self.0.borrow_mut().extend_from_slice(b);
436                Ok(b.len())
437            }
438            fn flush(&mut self) -> std::io::Result<()> {
439                Ok(())
440            }
441        }
442        let inputs = [Value::Char("A|B".into())];
443        let mut recorder = Recorder::new(Box::new(Fixed), Box::new(Sink(written.clone())), "a test double").unwrap();
444        let live = recorder.execute(&call("SELECT", "SELECT X, Y FROM T WHERE Z = ?", &inputs)).unwrap();
445        let text = String::from_utf8(written.borrow().clone()).unwrap();
446        assert!(text.starts_with(&format!("{HEADER}\n# source: a test double\n")), "{text}");
447        let mut replay = Replay::parse(&text, false).unwrap();
448        assert_eq!(replay.execute(&call("SELECT", "SELECT X, Y FROM T WHERE Z = ?", &inputs)), Ok(live));
449    }
450
451    #[test]
452    fn malformed_recordings_name_the_line() {
453        assert_eq!(Replay::parse("@ 1 P:1:0 SELECT\n< 0 00000 rows=0\n", false).err().unwrap(), format!("line 1: a recording starts with \"{HEADER}\""));
454        let err = Replay::parse(&format!("{HEADER}\n@ 1 P:1:0 SELECT\n= int:1\n"), false).err().unwrap();
455        assert_eq!(err, "line 3: an = line belongs after a < line");
456        assert!(Replay::parse(&format!("{HEADER}\n@ 1 P:1:0 SELECT\n< 0 00000 rows=0\n= int:x\n"), false).err().unwrap().starts_with("line 4: "));
457        assert_eq!(Replay::parse(&format!("{HEADER}\n@ 1 P:1:0 SELECT\n"), false).err().unwrap(), "the last call has no < line");
458    }
459}