dml.rsannotateddml.rssource161 lines · 7.1 KB · raw
1//! DML and row locks (contract *SQL surface*, the refusal list): the scan
2//! batches the calls of UPDATE, DELETE, INSERT … SELECT and SELECT … FOR
3//! UPDATE, and answers an EvalPlanQual recheck from the statement's
4//! judgments or the cache, never sending a second time.
5
6use std::time::Duration;
7
8use serde_json::{Value, json};
9use support::jev_instance;
10use support::mock_jev::{MockJev, Recorded, Reply};
11use tokio_postgres::error::SqlState;
12
13/// Answers every Noul with the row's id / 100.
14fn answer(req: &Recorded) -> Reply {
15    let body: Value = serde_json::from_slice(&req.body).unwrap();
16    let id = body["state"]["id"].as_i64().unwrap();
17    let answers: serde_json::Map<String, Value> = body["questions"]
18        .as_object()
19        .unwrap()
20        .keys()
21        .map(|key| (key.clone(), json!({ "type": "noul", "noul": id as f64 / 100.0 })))
22        .collect();
23    Reply::json(200, json!({ "model": "jev-1.13.0", "answers": answers, "usage": { "input_tokens": 1, "output_tokens": 1 } }))
24}
25
26async fn tickets(client: &tokio_postgres::Client) {
27    client
28        .batch_execute(
29            "ALTER TABLE tickets ADD p float8;
30             INSERT INTO tickets SELECT g, 'ticket ' || g FROM generate_series(2, 20) g; ANALYZE tickets;",
31        )
32        .await
33        .unwrap();
34}
35
36/// The per-row reference: the mock's answer for each id at or above 0.1.
37fn want(ids: impl Iterator<Item = i32>) -> Vec<(i32, Option<f64>)> {
38    ids.map(|id| (id, (id >= 10).then_some(id as f64 / 100.0))).collect()
39}
40
41async fn ps(client: &tokio_postgres::Client, table: &str) -> Vec<(i32, Option<f64>)> {
42    let rows = client.query(&format!("SELECT id, p FROM {table} ORDER BY id"), &[]).await.unwrap();
43    rows.iter().map(|r| (r.get(0), r.get(1))).collect()
44}
45
46#[tokio::test(flavor = "multi_thread")]
47async fn update_set_and_where_share_one_batched_judgment() {
48    // Held until all 20 arrive: a per-row UPDATE would hang.
49    let mock = MockJev::start(|req| answer(req).after_requests(20)).await;
50    let (_pg, client) = jev_instance(&mock).await;
51    tickets(&client).await;
52
53    let plan: Vec<String> = client
54        .query("EXPLAIN UPDATE tickets t SET p = jev_prob(t, 'q') WHERE jev(t, 'q', 0.1)", &[])
55        .await
56        .unwrap()
57        .iter()
58        .map(|r| r.get(0))
59        .collect();
60    assert!(plan.join("\n").contains("Custom Scan (JevScan)"), "{plan:?}");
61
62    let n = client.execute("UPDATE tickets t SET p = jev_prob(t, 'q') WHERE jev(t, 'q', 0.1)", &[]).await.unwrap();
63    assert_eq!(n, 11);
64    assert_eq!(mock.requests().len(), 20, "each row judged once, SET and WHERE sharing it");
65    assert_eq!(mock.peak_in_flight(), 20);
66    assert_eq!(ps(&client, "tickets").await, want(1..=20));
67}
68
69#[tokio::test(flavor = "multi_thread")]
70async fn delete_where_is_batched_and_returning_is_refused() {
71    let mock = MockJev::start(|req| answer(req).after_requests(20)).await;
72    let (_pg, client) = jev_instance(&mock).await;
73    tickets(&client).await;
74
75    // RETURNING is evaluated by ModifyTable, row by row, above the scan.
76    let err = client
77        .query("DELETE FROM tickets t WHERE id = 1 RETURNING jev_prob(t, 'q')", &[])
78        .await
79        .expect_err("refused");
80    assert_eq!(err.as_db_error().unwrap().code(), &SqlState::FEATURE_NOT_SUPPORTED, "{err:?}");
81    assert!(mock.requests().is_empty(), "nothing sent");
82
83    let n = client.execute("DELETE FROM tickets t WHERE jev(t, 'q', 0.1) RETURNING id", &[]).await.unwrap();
84    assert_eq!(n, 11);
85    assert_eq!(mock.requests().len(), 20);
86    let left: Vec<i32> = client.query("SELECT id FROM tickets ORDER BY id", &[]).await.unwrap().iter().map(|r| r.get(0)).collect();
87    assert_eq!(left, (1..10).collect::<Vec<_>>());
88}
89
90#[tokio::test(flavor = "multi_thread")]
91async fn insert_select_and_select_for_update_are_batched() {
92    let mock = MockJev::start(|req| answer(req).after_requests(20)).await;
93    let (_pg, client) = jev_instance(&mock).await;
94    tickets(&client).await;
95    client.batch_execute("CREATE TABLE judged (id int, p float8)").await.unwrap();
96
97    let n = client
98        .execute("INSERT INTO judged SELECT id, jev_prob(t, 'q') FROM tickets t WHERE jev(t, 'q', 0.1)", &[])
99        .await
100        .unwrap();
101    assert_eq!(n, 11);
102    assert_eq!(mock.requests().len(), 20);
103    assert_eq!(ps(&client, "judged").await, want(10..=20));
104
105    // Every answer is cached now, so the locking read sends nothing.
106    let locked = client.query("SELECT id FROM tickets t WHERE jev(t, 'q', 0.1) ORDER BY id FOR UPDATE", &[]).await.unwrap();
107    assert_eq!(locked.len(), 11);
108    assert_eq!(mock.requests().len(), 20);
109}
110
111/// Holds row 15 in a second session's open transaction, runs `update` in
112/// the test's session until it waits on that row, then runs `then` in the
113/// holder and commits, which makes the update recheck row 15.
114async fn recheck(
115    hold: &str,
116) -> (MockJev, Result<u64, tokio_postgres::Error>, tokio_postgres::Client, support::postgres::Instance) {
117    let mock = MockJev::start(answer).await;
118    let (pg, client) = jev_instance(&mock).await;
119    tickets(&client).await;
120    let (holder, conn) = tokio_postgres::connect(&pg.conn_str(), tokio_postgres::NoTls).await.unwrap();
121    tokio::spawn(conn);
122    holder.batch_execute(&format!("SET statement_timeout = '20s'; BEGIN; {hold}")).await.unwrap();
123
124    let client = std::sync::Arc::new(client);
125    let c = client.clone();
126    let update =
127        tokio::spawn(async move { c.execute("UPDATE tickets t SET p = jev_prob(t, 'q') WHERE jev(t, 'q', 0.1)", &[]).await });
128    for _ in 0..200 {
129        let waiting: i64 =
130            holder.query_one("SELECT count(*) FROM pg_locks WHERE NOT granted", &[]).await.unwrap().get(0);
131        if waiting > 0 {
132            break;
133        }
134        tokio::time::sleep(Duration::from_millis(50)).await;
135    }
136    assert_eq!(mock.requests().len(), 20, "every row judged before the update waits");
137    holder.batch_execute("COMMIT").await.unwrap();
138    let result = update.await.unwrap();
139    let client = std::sync::Arc::into_inner(client).expect("the update is done");
140    (mock, result, client, pg)
141}
142
143#[tokio::test(flavor = "multi_thread")]
144async fn a_recheck_of_an_unchanged_row_is_answered_without_sending() {
145    // The concurrent update leaves the row's values as they were, so the
146    // recheck asks the question the statement already asked.
147    let (mock, result, client, _pg) = recheck("UPDATE tickets SET body = body WHERE id = 15;").await;
148    assert_eq!(result.unwrap(), 11);
149    assert_eq!(mock.requests().len(), 20, "the recheck sent nothing");
150    assert_eq!(ps(&client, "tickets").await, want(1..=20));
151}
152
153#[tokio::test(flavor = "multi_thread")]
154async fn a_recheck_of_a_changed_row_is_a_serialization_failure() {
155    let (mock, result, client, _pg) = recheck("UPDATE tickets SET body = 'changed' WHERE id = 15;").await;
156    let err = result.expect_err("refused");
157    assert_eq!(err.as_db_error().unwrap().code(), &SqlState::T_R_SERIALIZATION_FAILURE, "{err:?}");
158    assert_eq!(mock.requests().len(), 20, "the new version was not sent");
159    let p: Option<f64> = client.query_one("SELECT p FROM tickets WHERE id = 16", &[]).await.unwrap().get(0);
160    assert_eq!(p, None, "the statement rolled back");
161}