postjevsql.git / tests / batching.rs
1//! The batch scan (contract *Execution*): every `jev*` call in a query is
2//! evaluated by one plan node, which judges rows concurrently as streams
3//! on the backend's one HTTP/2 connection, and each row alone.
4
5use serde_json::{Value, json};
6use support::jev_instance;
7use support::mock_jev::{MockJev, Recorded, Reply};
8use tokio_postgres::error::SqlState;
9
10/// Answers every question in a request with `p(row id, question)`.
11fn answer(req: &Recorded, p: impl Fn(i64, &str) -> f64) -> Reply {
12    let body: Value = serde_json::from_slice(&req.body).unwrap();
13    let id = body["state"]["id"].as_i64().unwrap();
14    let answers: serde_json::Map<String, Value> = body["questions"]
15        .as_object()
16        .unwrap()
17        .iter()
18        .map(|(key, q)| (key.clone(), json!({ "type": "noul", "noul": p(id, q["instructions"].as_str().unwrap()) })))
19        .collect();
20    Reply::json(200, json!({ "model": "jev-1.13.0", "answers": answers, "usage": { "input_tokens": 1, "output_tokens": 1 } }))
21}
22
23fn state_id(req: &Recorded) -> i64 {
24    serde_json::from_slice::<Value>(&req.body).unwrap()["state"]["id"].as_i64().unwrap()
25}
26
27async fn tickets(client: &tokio_postgres::Client, n: i32) {
28    client
29        .batch_execute(&format!("INSERT INTO tickets SELECT g, 'ticket ' || g FROM generate_series(2, {n}) g; ANALYZE tickets;"))
30        .await
31        .unwrap();
32}
33
34#[tokio::test(flavor = "multi_thread")]
35async fn rows_are_judged_concurrently_on_one_connection() {
36    // No answer until all 20 have arrived: a scan that waited for one
37    // answer before sending the next would hang until statement_timeout.
38    // The peak is counted at the mock, so a slow machine cannot fail it.
39    let mock = MockJev::start(|req| answer(req, |id, _| id as f64 / 100.0).after_requests(20)).await;
40    let (_pg, client) = jev_instance(&mock).await;
41    tickets(&client, 20).await;
42
43    let rows = client.query("SELECT id, jev_prob(t, 'q') FROM tickets t ORDER BY id", &[]).await.unwrap();
44
45    let got: Vec<(i32, f64)> = rows.iter().map(|r| (r.get(0), r.get(1))).collect();
46    let want: Vec<(i32, f64)> = (1..=20).map(|id| (id, id as f64 / 100.0)).collect();
47    assert_eq!(got, want, "each row gets its own answer");
48    assert_eq!(mock.requests().len(), 20, "one request per row");
49    assert_eq!(mock.connections(), 1);
50    assert_eq!(mock.peak_in_flight(), 20, "every row in flight at once");
51}
52
53#[tokio::test(flavor = "multi_thread")]
54async fn rows_the_sql_filters_out_are_never_judged() {
55    let mock = MockJev::start(|req| answer(req, |id, _| if id % 2 == 0 { 0.9 } else { 0.1 })).await;
56    let (_pg, client) = jev_instance(&mock).await;
57    tickets(&client, 20).await;
58
59    let rows = client
60        .query("SELECT id FROM tickets t WHERE id <= 10 AND jev_prob(t, 'q') > 0.5 ORDER BY id", &[])
61        .await
62        .unwrap();
63    let ids: Vec<i32> = rows.iter().map(|r| r.get(0)).collect();
64    assert_eq!(ids, [2, 4, 6, 8, 10]);
65    let mut judged: Vec<i64> = mock.requests().iter().map(state_id).collect();
66    judged.sort();
67    assert_eq!(judged, (1..=10).collect::<Vec<_>>());
68}
69
70#[tokio::test(flavor = "multi_thread")]
71async fn a_limit_stops_judging_after_the_window() {
72    // No answer until the whole window has arrived, so the count cannot
73    // depend on how fast the requests reach the mock. A scan that sent
74    // fewer would hang here until statement_timeout.
75    const WINDOW: usize = 4;
76    let mock = MockJev::start(|req| answer(req, |_, _| 0.9).after_requests(WINDOW)).await;
77    let (_pg, client) = jev_instance(&mock).await;
78    tickets(&client, 50).await;
79
80    client.batch_execute(&format!("SET jev.concurrency = {WINDOW}")).await.unwrap();
81    let rows = client.query("SELECT id FROM tickets t WHERE jev_prob(t, 'q') > 0.5 LIMIT 1", &[]).await.unwrap();
82    assert_eq!(rows.len(), 1);
83    assert_eq!(mock.requests().len(), WINDOW, "only the in-flight window was sent");
84}
85
86#[tokio::test(flavor = "multi_thread")]
87async fn questions_about_one_row_share_its_request() {
88    let mock = MockJev::start(|req| answer(req, |_, q| if q == "a" { 0.25 } else { 0.75 })).await;
89    let (_pg, client) = jev_instance(&mock).await;
90
91    let row = client
92        .query_one("SELECT jev_prob(t, 'a'), jev_prob(t, 'b') FROM tickets t WHERE jev_prob(t, 'a') < 0.5", &[])
93        .await
94        .unwrap();
95    assert_eq!((row.get::<_, f64>(0), row.get::<_, f64>(1)), (0.25, 0.75));
96    let requests = mock.requests();
97    assert_eq!(requests.len(), 1, "the row is the state; its questions are branches");
98    let body: Value = serde_json::from_slice(&requests[0].body).unwrap();
99    assert_eq!(body["questions"].as_object().unwrap().len(), 2, "the repeated question is asked once");
100}
101
102#[tokio::test(flavor = "multi_thread")]
103async fn a_rescanned_scan_judges_each_outer_row() {
104    let mock = MockJev::start(|req| answer(req, |id, _| id as f64 / 10.0)).await;
105    let (_pg, client) = jev_instance(&mock).await;
106    tickets(&client, 3).await;
107
108    let rows = client
109        .query(
110            "SELECT v.x, s.p FROM (VALUES (1), (3)) v(x)
111             CROSS JOIN LATERAL (SELECT jev_prob(t, 'q') AS p FROM tickets t WHERE t.id = v.x) s
112             ORDER BY v.x",
113            &[],
114        )
115        .await
116        .unwrap();
117    let got: Vec<(i32, f64)> = rows.iter().map(|r| (r.get(0), r.get(1))).collect();
118    assert_eq!(got, [(1, 0.1), (3, 0.3)]);
119    assert_eq!(mock.requests().len(), 2);
120}
121
122#[tokio::test(flavor = "multi_thread")]
123async fn a_call_outside_the_scan_is_refused_before_spending() {
124    let mock = MockJev::start(|req| answer(req, |_, _| 0.5)).await;
125    let (_pg, client) = jev_instance(&mock).await;
126
127    // A call over two relations is evaluated by the join, which no scan
128    // can sit under.
129    let err = client
130        .query("SELECT jev_prob((t.id, u.id), 'q') FROM tickets t JOIN tickets u ON u.id = t.id", &[])
131        .await
132        .expect_err("refused");
133    assert_eq!(err.as_db_error().unwrap().code(), &SqlState::FEATURE_NOT_SUPPORTED, "{err:?}");
134    assert!(mock.requests().is_empty());
135}
136
137/// Answers a Noul with `p(row id)` and a Choice with option `id % n`.
138fn answer_any(req: &Recorded, p: impl Fn(i64) -> f64) -> Reply {
139    let body: Value = serde_json::from_slice(&req.body).unwrap();
140    let id = body["state"]["id"].as_i64().unwrap();
141    let answers: serde_json::Map<String, Value> = body["questions"]
142        .as_object()
143        .unwrap()
144        .iter()
145        .map(|(key, q)| {
146            let answer = match q["criteria"].as_object() {
147                Some(criteria) => {
148                    let labels: Vec<&String> = criteria.keys().collect();
149                    let pick = labels[id as usize % labels.len()].clone();
150                    let probabilities: serde_json::Map<_, _> =
151                        labels.iter().map(|l| (l.to_string(), json!(if **l == pick { 1.0 } else { 0.0 }))).collect();
152                    json!({ "type": "choice", "choice": pick, "confidence": 1.0, "probabilities": probabilities })
153                }
154                None => json!({ "type": "noul", "noul": p(id) }),
155            };
156            (key.clone(), answer)
157        })
158        .collect();
159    Reply::json(200, json!({ "model": "jev-1.13.0", "answers": answers, "usage": { "input_tokens": 1, "output_tokens": 1 } }))
160}
161
162/// Runs `query` over 20 tickets with every request held until all 20
163/// arrive, so it hangs unless the calls are batched; then checks it
164/// against `reference`, which reads the same judgments per row (from the
165/// cache, sending nothing).
166/// EXPLAIN's `Candidate Rows` is every row judged, before the jev
167/// conditions.
168async fn batched_like_per_row(query: &str, reference: &str) {
169    let mock = MockJev::start(|req| answer_any(req, |id| id as f64 / 100.0).after_requests(20)).await;
170    let (_pg, client) = jev_instance(&mock).await;
171    tickets(&client, 20).await;
172
173    let plan: Vec<String> = client.query(&format!("EXPLAIN {query}"), &[]).await.unwrap().iter().map(|r| r.get(0)).collect();
174    let plan = plan.join("\n");
175    assert!(plan.contains("Custom Scan (JevScan)"), "{plan}");
176    assert!(plan.contains("Candidate Rows: 20"), "{plan}");
177    assert!(plan.contains("Worst-Case Cost"), "{plan}");
178    assert!(mock.requests().is_empty(), "EXPLAIN sent nothing");
179
180    let got: Vec<String> = client.simple_query(query).await.unwrap().iter().filter_map(row_text).collect();
181    assert_eq!(mock.requests().len(), 20, "one request per row");
182    assert_eq!(mock.peak_in_flight(), 20, "every row in flight at once");
183    let want: Vec<String> = client.simple_query(reference).await.unwrap().iter().filter_map(row_text).collect();
184    assert_eq!(mock.requests().len(), 20, "the reference is served from the cache");
185    assert_eq!(got, want);
186}
187
188fn row_text(message: &tokio_postgres::SimpleQueryMessage) -> Option<String> {
189    match message {
190        tokio_postgres::SimpleQueryMessage::Row(row) => {
191            Some((0..row.len()).map(|i| row.get(i).unwrap_or("NULL").to_string()).collect::<Vec<_>>().join("|"))
192        }
193        _ => None,
194    }
195}
196
197#[tokio::test(flavor = "multi_thread")]
198async fn an_aggregate_over_calls_is_batched() {
199    batched_like_per_row(
200        "SELECT sum(jev_prob(t, 'q')) FROM tickets t",
201        "SELECT sum(p) FROM (SELECT jev_prob(t, 'q') AS p FROM tickets t OFFSET 0) s",
202    )
203    .await;
204}
205
206#[tokio::test(flavor = "multi_thread")]
207async fn having_over_calls_is_batched() {
208    batched_like_per_row(
209        "SELECT id % 3 AS g, round(avg(jev_prob(t, 'q'))::numeric, 4) FROM tickets t GROUP BY 1 HAVING avg(jev_prob(t, 'q')) > 0.1 ORDER BY 1",
210        "SELECT g, round(avg(p)::numeric, 4) FROM (SELECT id % 3 AS g, jev_prob(t, 'q') AS p FROM tickets t OFFSET 0) s GROUP BY 1 HAVING avg(p) > 0.1 ORDER BY 1",
211    )
212    .await;
213}
214
215#[tokio::test(flavor = "multi_thread")]
216async fn distinct_over_calls_is_batched() {
217    batched_like_per_row(
218        "SELECT DISTINCT jev_choice(t, 'q', ARRAY['a', 'b', 'c']) AS c FROM tickets t ORDER BY 1",
219        "SELECT DISTINCT c FROM (SELECT jev_choice(t, 'q', ARRAY['a', 'b', 'c']) AS c FROM tickets t OFFSET 0) s ORDER BY 1",
220    )
221    .await;
222}
223
224#[tokio::test(flavor = "multi_thread")]
225async fn a_window_over_calls_is_batched() {
226    batched_like_per_row(
227        "SELECT id, rank() OVER (ORDER BY jev_prob(t, 'q') DESC) FROM tickets t ORDER BY id",
228        "SELECT id, rank() OVER (ORDER BY p DESC) FROM (SELECT id, jev_prob(t, 'q') AS p FROM tickets t OFFSET 0) s ORDER BY id",
229    )
230    .await;
231}
232
233async fn explain_lines(client: &tokio_postgres::Client, query: &str) -> String {
234    let plan: Vec<String> = client.query(query, &[]).await.unwrap().iter().map(|r| r.get(0)).collect();
235    plan.join("\n")
236}
237
238#[tokio::test(flavor = "multi_thread")]
239async fn explain_verbose_deparses_a_lifted_scan() {
240    // VERBOSE prints every node's Output, so the parent's rewritten
241    // OUTER_VAR must resolve through the scan's tlist, and the scan's own
242    // expressions through its child. The child is printed once, though it
243    // is both `lefttree` and the scan's custom plan.
244    let mock = MockJev::start(|req| answer(req, |id, _| id as f64 / 100.0)).await;
245    let (_pg, client) = jev_instance(&mock).await;
246    tickets(&client, 20).await;
247
248    for (query, parent) in [
249        ("SELECT sum(jev_prob(t, 'q')) FROM tickets t", "Aggregate"),
250        ("SELECT id, rank() OVER (ORDER BY jev_prob(t, 'q') DESC) FROM tickets t", "WindowAgg"),
251    ] {
252        let plan = explain_lines(&client, &format!("EXPLAIN (VERBOSE, COSTS OFF) {query}")).await;
253        assert!(plan.contains(parent), "{plan}");
254        assert!(plan.contains("Custom Scan (JevScan)"), "{plan}");
255        assert!(plan.contains("jev_prob(t.*, 'q'::text)"), "the call is deparsed: {plan}");
256        assert_eq!(plan.matches("Seq Scan on public.tickets").count(), 1, "{plan}");
257        assert!(mock.requests().is_empty(), "EXPLAIN sent nothing");
258    }
259
260    let plan = explain_lines(
261        &client,
262        "EXPLAIN (ANALYZE, VERBOSE, COSTS OFF, TIMING OFF, SUMMARY OFF, BUFFERS OFF) SELECT sum(jev_prob(t, 'q')) FROM tickets t",
263    )
264    .await;
265    assert!(plan.contains("Cache Misses: 20"), "{plan}");
266    assert!(plan.contains("Requests: 20"), "{plan}");
267    assert_eq!(plan.matches("Seq Scan on public.tickets").count(), 1, "{plan}");
268    assert_eq!(mock.requests().len(), 20);
269
270    // A generic plan is a copy; it must deparse and run as the original.
271    client
272        .batch_execute("SET plan_cache_mode = force_generic_plan; PREPARE s AS SELECT sum(jev_prob(t, 'q')) FROM tickets t")
273        .await
274        .unwrap();
275    let plan = explain_lines(&client, "EXPLAIN (VERBOSE, COSTS OFF) EXECUTE s").await;
276    assert!(plan.contains("Custom Scan (JevScan)"), "{plan}");
277    let sum: f64 = client.query_one("EXECUTE s", &[]).await.unwrap().get(0);
278    assert!((sum - 2.1).abs() < 1e-9, "{sum}");
279    assert_eq!(mock.requests().len(), 20, "served from the cache");
280}
281
282#[tokio::test(flavor = "multi_thread")]
283async fn max_cost_refuses_an_aggregate_before_sending() {
284    let mock = MockJev::start(|req| answer(req, |_, _| 0.5)).await;
285    let (_pg, client) = jev_instance(&mock).await;
286    tickets(&client, 20).await;
287    client.batch_execute("SET jev.max_cost = 0").await.unwrap();
288
289    let err = client.query_one("SELECT sum(jev_prob(t, 'q')) FROM tickets t", &[]).await.expect_err("refused");
290    assert_eq!(err.as_db_error().unwrap().code(), &SqlState::PROGRAM_LIMIT_EXCEEDED, "{err:?}");
291    assert!(mock.requests().is_empty());
292}
293
294#[tokio::test(flavor = "multi_thread")]
295async fn explain_shows_the_scan_without_spending() {
296    let mock = MockJev::start(|req| answer(req, |_, _| 0.5)).await;
297    let (_pg, client) = jev_instance(&mock).await;
298
299    let plan: Vec<String> = client
300        .query("EXPLAIN SELECT id FROM tickets t WHERE jev_prob(t, 'q') > 0.5", &[])
301        .await
302        .unwrap()
303        .iter()
304        .map(|r| r.get(0))
305        .collect();
306    assert!(plan[0].contains("Custom Scan (JevScan)"), "{plan:#?}");
307    assert!(mock.requests().is_empty());
308}
309
310#[tokio::test(flavor = "multi_thread")]
311async fn a_plan_with_any_unbatched_call_spends_nothing() {
312    let mock = MockJev::start(|req| answer(req, |_, _| 0.9)).await;
313    let (_pg, client) = jev_instance(&mock).await;
314    tickets(&client, 20).await;
315
316    // The WHERE call is the scan's; the one over two relations is the
317    // join's. Without a check before execution, the scan would judge a
318    // window first.
319    let err = client
320        .query(
321            "SELECT jev_prob((t.id, u.id), 'a') FROM tickets t JOIN tickets u ON u.id = t.id WHERE jev_prob(t, 'b') > 0.5",
322            &[],
323        )
324        .await
325        .expect_err("refused");
326    assert_eq!(err.as_db_error().unwrap().code(), &SqlState::FEATURE_NOT_SUPPORTED, "{err:?}");
327    assert!(mock.requests().is_empty(), "{} requests sent before the refusal", mock.requests().len());
328}
329
330#[tokio::test(flavor = "multi_thread")]
331async fn the_window_is_clamped_to_the_servers_stream_limit() {
332    // hyper queues streams past the peer's limit silently, so the scan
333    // must not pull rows it cannot send. A sequence in the child's quals
334    // counts the rows pulled.
335    const STREAMS: u32 = 3;
336    // Row 50 warms the connection alone; the rest wait for the whole
337    // window (1 + STREAMS requests in all), so a smaller one would hang.
338    let mock = MockJev::start_limited(STREAMS, |req| {
339        let reply = answer(req, |_, _| 0.9);
340        if state_id(req) == 50 { reply } else { reply.after_requests(1 + STREAMS as usize) }
341    })
342    .await;
343    let (_pg, client) = jev_instance(&mock).await;
344    tickets(&client, 50).await;
345    client.batch_execute("CREATE SEQUENCE pulled; SET jev.concurrency = 40").await.unwrap();
346
347    // The first statement opens the connection and reads the server's
348    // SETTINGS; the second runs on it.
349    client.query("SELECT jev_prob(t, 'q') FROM tickets t WHERE id = 50", &[]).await.unwrap();
350    let before = mock.requests().len();
351    assert_eq!(before, 1);
352    let rows = client
353        .query("SELECT id FROM tickets t WHERE nextval('pulled') > 0 AND jev_prob(t, 'q') > 0.5 LIMIT 1", &[])
354        .await
355        .unwrap();
356    assert_eq!(rows.len(), 1);
357    let pulled: i64 = client.query_one("SELECT last_value FROM pulled", &[]).await.unwrap().get(0);
358    assert_eq!(pulled, i64::from(STREAMS), "rows pulled ahead: the window is the server's limit, not jev.concurrency");
359    assert_eq!(mock.requests().len() - before, STREAMS as usize);
360}
361
362async fn requests_sent(client: &tokio_postgres::Client) -> i64 {
363    client.query_one("SELECT requests FROM jev_stats()", &[]).await.unwrap().get(0)
364}
365
366#[tokio::test(flavor = "multi_thread")]
367async fn an_or_arm_postgres_skips_is_never_judged() {
368    // The NULL-id row is answered 0.1, so only its request shows it.
369    let mock = MockJev::start(|req| {
370        let body: Value = serde_json::from_slice(&req.body).unwrap();
371        let p = match body["state"]["id"].as_i64() {
372            Some(id) if id % 2 == 0 => 0.9,
373            _ => 0.1,
374        };
375        let answers: serde_json::Map<String, Value> =
376            body["questions"].as_object().unwrap().keys().map(|k| (k.clone(), json!({ "type": "noul", "noul": p }))).collect();
377        Reply::json(200, json!({ "model": "jev-1.13.0", "answers": answers, "usage": { "input_tokens": 1, "output_tokens": 1 } }))
378    })
379    .await;
380    let (_pg, client) = jev_instance(&mock).await;
381    tickets(&client, 20).await;
382    // A NULL id makes `id <= 5` NULL, not false, so Postgres still asks.
383    client.batch_execute("INSERT INTO tickets VALUES (NULL, 'no id')").await.unwrap();
384    let before = requests_sent(&client).await;
385
386    let rows = client
387        .query("SELECT id FROM tickets t WHERE id <= 5 OR jev(t, 'q') ORDER BY id NULLS LAST", &[])
388        .await
389        .unwrap();
390    let ids: Vec<Option<i32>> = rows.iter().map(|r| r.get(0)).collect();
391    let want: Vec<Option<i32>> = (1..=5).chain((6..=20).filter(|id| id % 2 == 0)).map(Some).collect();
392    assert_eq!(ids, want);
393
394    let judged: Vec<Value> = mock
395        .requests()
396        .iter()
397        .map(|r| serde_json::from_slice::<Value>(&r.body).unwrap()["state"]["id"].clone())
398        .collect();
399    let mut ids: Vec<i64> = judged.iter().filter_map(Value::as_i64).collect();
400    ids.sort();
401    assert_eq!(ids, (6..=20).collect::<Vec<_>>(), "only the rows the OR reaches");
402    assert_eq!(judged.iter().filter(|v| v.is_null()).count(), 1, "the NULL id is reached");
403    assert_eq!(requests_sent(&client).await - before, 16, "jev_stats counts the same");
404}
405
406#[tokio::test(flavor = "multi_thread")]
407async fn and_and_case_judge_only_what_postgres_reads() {
408    // Each row's probability is fixed by its id, so the reference below
409    // is what Postgres computes calling jev_prob on exactly these rows.
410    let p = |id: i64| id as f64 / 100.0;
411    let mock = MockJev::start(move |req| answer(req, |id, _| p(id))).await;
412    let (_pg, client) = jev_instance(&mock).await;
413    tickets(&client, 20).await;
414    let before = requests_sent(&client).await;
415
416    let query = "SELECT id, \
417                        CASE WHEN id % 3 = 0 THEN jev_prob(t, 'a') \
418                             WHEN id % 3 = 1 THEN 1 - jev_prob(t, 'b') END, \
419                        (id > 15 AND jev(t, 'c')) \
420                 FROM tickets t ORDER BY id";
421    let got: Vec<String> = client.simple_query(query).await.unwrap().iter().filter_map(row_text).collect();
422
423    let want: Vec<String> = (1..=20i64)
424        .map(|id| {
425            let case = match id % 3 {
426                0 => p(id).to_string(),
427                1 => (1.0 - p(id)).to_string(),
428                _ => "NULL".to_string(),
429            };
430            let and = if id > 15 { if p(id) >= 0.5 { "t" } else { "f" } } else { "f" };
431            format!("{id}|{case}|{and}")
432        })
433        .collect();
434    assert_eq!(got, want);
435
436    let mut asked: Vec<(i64, String)> = mock
437        .requests()
438        .iter()
439        .flat_map(|r| {
440            let body: Value = serde_json::from_slice(&r.body).unwrap();
441            let id = body["state"]["id"].as_i64().unwrap();
442            body["questions"]
443                .as_object()
444                .unwrap()
445                .values()
446                .map(|q| (id, q["instructions"].as_str().unwrap().to_string()))
447                .collect::<Vec<_>>()
448        })
449        .collect();
450    asked.sort();
451    let mut expected: Vec<(i64, String)> = (1..=20i64)
452        .flat_map(|id| {
453            let mut q = Vec::new();
454            match id % 3 {
455                0 => q.push((id, "a".to_string())),
456                1 => q.push((id, "b".to_string())),
457                _ => {}
458            }
459            if id > 15 {
460                q.push((id, "c".to_string()));
461            }
462            q
463        })
464        .collect();
465    expected.sort();
466    assert_eq!(asked, expected, "only the questions Postgres reaches");
467    let rows_sent = expected.iter().map(|&(id, _)| id).collect::<std::collections::BTreeSet<_>>().len() as i64;
468    assert_eq!(requests_sent(&client).await - before, rows_sent, "one request per reached row");
469}
470
471#[tokio::test(flavor = "multi_thread")]
472async fn a_materialized_cte_is_batched() {
473    batched_like_per_row(
474        "WITH s AS MATERIALIZED (SELECT id, jev_prob(t, 'q') AS p FROM tickets t) SELECT id, p FROM s ORDER BY id",
475        "SELECT id, p FROM (SELECT id, jev_prob(t, 'q') AS p FROM tickets t OFFSET 0) s ORDER BY id",
476    )
477    .await;
478}
479
480#[tokio::test(flavor = "multi_thread")]
481async fn a_cte_read_twice_is_judged_once() {
482    batched_like_per_row(
483        "WITH s AS (SELECT id, jev_prob(t, 'q') AS p FROM tickets t) SELECT a.id, a.p, b.p FROM s a JOIN s b ON b.id = a.id ORDER BY a.id",
484        "SELECT id, p, p FROM (SELECT id, jev_prob(t, 'q') AS p FROM tickets t OFFSET 0) s ORDER BY id",
485    )
486    .await;
487}
488
489#[tokio::test(flavor = "multi_thread")]
490async fn a_scalar_initplan_is_batched() {
491    batched_like_per_row(
492        "SELECT (SELECT round(sum(jev_prob(t, 'q'))::numeric, 4) FROM tickets t)",
493        "SELECT round(sum(p)::numeric, 4) FROM (SELECT jev_prob(t, 'q') AS p FROM tickets t OFFSET 0) s",
494    )
495    .await;
496}
497
498/// Postgres keeps an EXISTS whose WHERE is volatile as a correlated
499/// SubPlan (`convert_EXISTS_sublink_to_join`), so each outer row runs it:
500/// every run's scan judges that run's rows, as per row.
501#[tokio::test(flavor = "multi_thread")]
502async fn an_exists_sublink_judges_each_run_like_per_row() {
503    let mock = MockJev::start(|req| answer(req, |id, _| id as f64 / 100.0)).await;
504    let (_pg, client) = jev_instance(&mock).await;
505    tickets(&client, 20).await;
506    let before = requests_sent(&client).await;
507
508    let query = "SELECT id FROM tickets o \
509                 WHERE EXISTS (SELECT 1 FROM tickets t WHERE t.id = o.id AND jev_prob(t, 'q') > 0.1) ORDER BY id";
510    let plan = explain_lines(&client, &format!("EXPLAIN {query}")).await;
511    assert!(plan.contains("SubPlan"), "{plan}");
512    assert!(plan.contains("Custom Scan (JevScan)"), "{plan}");
513    let got: Vec<String> = client.simple_query(query).await.unwrap().iter().filter_map(row_text).collect();
514    let reference = "SELECT id FROM (SELECT id, jev_prob(t, 'q') AS p FROM tickets t OFFSET 0) s WHERE p > 0.1 ORDER BY id";
515    let want: Vec<String> = client.simple_query(reference).await.unwrap().iter().filter_map(row_text).collect();
516    assert_eq!(got, want);
517    assert_eq!(got, (11..=20).map(|id| id.to_string()).collect::<Vec<_>>());
518    assert_eq!(requests_sent(&client).await - before, 20, "each row once; the reference is served from the cache");
519}
520
521#[tokio::test(flavor = "multi_thread")]
522async fn an_in_sublink_is_batched() {
523    batched_like_per_row(
524        "SELECT id FROM tickets o WHERE id IN (SELECT id FROM tickets t WHERE jev(t, 'q', 0.1)) ORDER BY id",
525        "SELECT id FROM (SELECT id, jev_prob(t, 'q') AS p FROM tickets t OFFSET 0) s WHERE p >= 0.1 ORDER BY id",
526    )
527    .await;
528}
529
530#[tokio::test(flavor = "multi_thread")]
531async fn a_hashed_not_in_subplan_is_batched() {
532    batched_like_per_row(
533        "SELECT id FROM tickets o WHERE id NOT IN (SELECT id FROM tickets t WHERE jev_prob(t, 'q') > 0.1) ORDER BY id",
534        "SELECT id FROM (SELECT id, jev_prob(t, 'q') AS p FROM tickets t OFFSET 0) s WHERE NOT p > 0.1 ORDER BY id",
535    )
536    .await;
537}
538
539/// A correlated SubPlan runs once per outer row, so each run's scan
540/// judges that run's rows; the answers equal the per-row reference, and a
541/// row asked again is served from the cache rather than sent.
542#[tokio::test(flavor = "multi_thread")]
543async fn a_correlated_subplan_judges_each_run_like_per_row() {
544    let mock = MockJev::start(|req| answer(req, |id, _| id as f64 / 100.0)).await;
545    let (_pg, client) = jev_instance(&mock).await;
546    tickets(&client, 20).await;
547    let before = requests_sent(&client).await;
548
549    let query = "SELECT o.id, (SELECT jev_prob(t, 'q') FROM tickets t WHERE t.id = o.id % 5 + 1) \
550                 FROM tickets o ORDER BY o.id";
551    let plan = explain_lines(&client, &format!("EXPLAIN {query}")).await;
552    assert!(plan.contains("SubPlan"), "{plan}");
553    assert!(plan.contains("Custom Scan (JevScan)"), "{plan}");
554    let got: Vec<String> = client.simple_query(query).await.unwrap().iter().filter_map(row_text).collect();
555    let want: Vec<String> = (1..=20i64).map(|id| format!("{id}|{}", (id % 5 + 1) as f64 / 100.0)).collect();
556    assert_eq!(got, want);
557    assert_eq!(requests_sent(&client).await - before, 5, "each distinct row sent once; repeats are cache hits");
558}
559
560/// Answers a Noul with a tenth of the integer column in the state: the
561/// outer value in `(t.body, o.id)`.
562fn answer_by_outer(req: &Recorded) -> Reply {
563    let body: Value = serde_json::from_slice(&req.body).unwrap();
564    let outer = body["state"].as_object().unwrap().values().find_map(Value::as_i64).expect("an integer column");
565    let answers: serde_json::Map<String, Value> = body["questions"]
566        .as_object()
567        .unwrap()
568        .keys()
569        .map(|key| (key.clone(), json!({ "type": "noul", "noul": outer as f64 / 10.0 })))
570        .collect();
571    Reply::json(200, json!({ "model": "jev-1.13.0", "answers": answers, "usage": { "input_tokens": 1, "output_tokens": 1 } }))
572}
573
574/// A call reading an outer column gets it as a Param, which changes on
575/// every run of the SubPlan: the scan evaluates it again on each rescan,
576/// sends that run's value, and serves a repeated value from the cache.
577/// The first query's child does not depend on the Param (only the
578/// projection does), so a scan that replayed its first run would answer
579/// every outer row alike.
580#[tokio::test(flavor = "multi_thread")]
581async fn a_call_reading_an_outer_column_is_judged_per_run() {
582    let mock = MockJev::start(answer_by_outer).await;
583    let (_pg, client) = jev_instance(&mock).await;
584    tickets(&client, 3).await;
585
586    // The correlated WHERE's first run repeats the first query's pair.
587    for (inner, new) in [("t.id = 1", 3), ("t.id = o.id", 2)] {
588        let before = requests_sent(&client).await;
589        let query = format!(
590            "SELECT o.id, (SELECT jev_prob((t.body, o.id), 'q') FROM tickets t WHERE {inner}) \
591             FROM (VALUES (1, 1), (2, 2), (3, 1), (4, 3)) o(n, id) ORDER BY o.n"
592        );
593        let plan = explain_lines(&client, &format!("EXPLAIN (VERBOSE) {query}")).await;
594        assert!(plan.contains("SubPlan"), "{plan}");
595        assert!(plan.contains("Custom Scan (JevScan)"), "{plan}");
596        let got: Vec<String> = client.simple_query(&query).await.unwrap().iter().filter_map(row_text).collect();
597        assert_eq!(got, ["1|0.1", "2|0.2", "1|0.1", "3|0.3"], "{inner}");
598        assert_eq!(requests_sent(&client).await - before, new, "{inner}: each new (body, id) once; repeats are hits");
599    }
600    let sent: Vec<i64> = mock
601        .requests()
602        .iter()
603        .map(|r| {
604            let body: Value = serde_json::from_slice(&r.body).unwrap();
605            body["state"].as_object().unwrap().values().find_map(Value::as_i64).unwrap()
606        })
607        .collect();
608    assert_eq!(sent, [1, 2, 3, 2, 3], "each run sends its own outer value");
609}
610
611/// A jev condition over the worktable of a `WITH RECURSIVE`: each
612/// iteration's scan judges the rows the last one produced, and the
613/// recursion stops where the condition first fails.
614#[tokio::test(flavor = "multi_thread")]
615async fn a_condition_over_a_recursive_worktable_is_judged_per_iteration() {
616    let mock = MockJev::start(|req| answer(req, |id, _| if id < 10 { 0.9 } else { 0.1 })).await;
617    let (_pg, client) = jev_instance(&mock).await;
618    tickets(&client, 20).await;
619    let before = requests_sent(&client).await;
620
621    let query = "WITH RECURSIVE r AS ( \
622                   SELECT id, body FROM tickets WHERE id = 1 \
623                   UNION ALL \
624                   SELECT t.id, t.body FROM r JOIN tickets t ON t.id = r.id + 1 WHERE jev(r, 'q') \
625                 ) SELECT id FROM r ORDER BY id";
626    let plan = explain_lines(&client, &format!("EXPLAIN {query}")).await;
627    assert!(plan.contains("WorkTable Scan"), "{plan}");
628    assert!(plan.contains("Custom Scan (JevScan)"), "{plan}");
629    let got: Vec<String> = client.simple_query(query).await.unwrap().iter().filter_map(row_text).collect();
630    assert_eq!(got, (1..=10).map(|id| id.to_string()).collect::<Vec<_>>());
631    assert_eq!(requests_sent(&client).await - before, 10, "each worktable row judged once");
632}
633
634/// The same recursion with the condition on the table the recursive term
635/// joins: that scan is rescanned every iteration, and a row it judged in
636/// an earlier one is a cache hit.
637#[tokio::test(flavor = "multi_thread")]
638async fn a_condition_in_a_recursive_term_is_judged_once_per_row() {
639    let mock = MockJev::start(|req| answer(req, |id, _| if id <= 10 { 0.9 } else { 0.1 })).await;
640    let (_pg, client) = jev_instance(&mock).await;
641    tickets(&client, 20).await;
642    let before = requests_sent(&client).await;
643
644    let query = "WITH RECURSIVE r AS ( \
645                   SELECT id FROM tickets WHERE id = 1 \
646                   UNION ALL \
647                   SELECT t.id FROM r JOIN tickets t ON t.id = r.id + 1 WHERE jev(t, 'q') \
648                 ) SELECT id FROM r ORDER BY id";
649    let plan = explain_lines(&client, &format!("EXPLAIN {query}")).await;
650    assert!(plan.contains("Custom Scan (JevScan)"), "{plan}");
651    let got: Vec<String> = client.simple_query(query).await.unwrap().iter().filter_map(row_text).collect();
652    assert_eq!(got, (1..=10).map(|id| id.to_string()).collect::<Vec<_>>());
653    let sent = requests_sent(&client).await - before;
654    assert!((10..=20).contains(&sent), "no ticket judged twice: {sent}");
655}
656
657/// A partitioned table: each partition's scan is wrapped, with the quals
658/// and the select list translated to it, so every partition's rows are
659/// judged in one statement, as the per-row reference judges them.
660async fn partitioned(client: &tokio_postgres::Client) {
661    client
662        .batch_execute(
663            "CREATE TABLE parts (id int, body text) PARTITION BY RANGE (id);
664             CREATE TABLE parts_lo PARTITION OF parts FOR VALUES FROM (1) TO (11);
665             CREATE TABLE parts_hi PARTITION OF parts FOR VALUES FROM (11) TO (21);
666             INSERT INTO parts SELECT id, body FROM tickets; ANALYZE parts;",
667        )
668        .await
669        .unwrap();
670}
671
672/// An Append pulls its partitions one after another, so each partition's
673/// rows are in flight together, and the next partition's follow.
674#[tokio::test(flavor = "multi_thread")]
675async fn a_partitioned_table_is_batched_across_partitions() {
676    let mock = MockJev::start(|req| answer_any(req, |id| id as f64 / 100.0).after_requests(10)).await;
677    let (_pg, client) = jev_instance(&mock).await;
678    tickets(&client, 20).await;
679    partitioned(&client).await;
680
681    let query = "SELECT id, jev_prob(p, 'q') FROM parts p WHERE jev(p, 'q', 0.05) ORDER BY id";
682    let plan = explain_lines(&client, &format!("EXPLAIN {query}")).await;
683    assert!(plan.contains("Seq Scan on parts_lo") && plan.contains("Seq Scan on parts_hi"), "{plan}");
684    // One per partition, and the ORDER BY projection's over the Sort.
685    assert_eq!(plan.matches("Custom Scan (JevScan)").count(), 3, "{plan}");
686    assert!(mock.requests().is_empty(), "EXPLAIN sent nothing");
687
688    let got: Vec<String> = client.simple_query(query).await.unwrap().iter().filter_map(row_text).collect();
689    assert_eq!(mock.requests().len(), 20, "one request per row; the select list shares each row's judgment");
690    assert_eq!(mock.peak_in_flight(), 10, "a partition's rows in flight at once");
691    let reference = "SELECT id, p FROM (SELECT id, jev_prob(t, 'q') AS p FROM parts t OFFSET 0) s WHERE p >= 0.05 ORDER BY id";
692    let want: Vec<String> = client.simple_query(reference).await.unwrap().iter().filter_map(row_text).collect();
693    assert_eq!(mock.requests().len(), 20, "the reference is served from the cache");
694    assert_eq!(got, want);
695    assert_eq!(got.len(), 16);
696}
697
698#[tokio::test(flavor = "multi_thread")]
699async fn equal_judgments_in_two_partitions_share_one_request() {
700    // The column list sends only `body`, so both rows' judgments are equal.
701    let mock = MockJev::start(|req| {
702        let body: Value = serde_json::from_slice(&req.body).unwrap();
703        let answers: serde_json::Map<String, Value> =
704            body["questions"].as_object().unwrap().keys().map(|k| (k.clone(), json!({ "type": "noul", "noul": 0.9 }))).collect();
705        Reply::json(200, json!({ "model": "jev-1.13.0", "answers": answers, "usage": { "input_tokens": 1, "output_tokens": 1 } }))
706    })
707    .await;
708    let (_pg, client) = jev_instance(&mock).await;
709    client
710        .batch_execute(
711            "CREATE TABLE kinds (id int, body text) PARTITION BY LIST (id);
712             CREATE TABLE kinds_a PARTITION OF kinds FOR VALUES IN (1);
713             CREATE TABLE kinds_b PARTITION OF kinds FOR VALUES IN (2);
714             INSERT INTO kinds VALUES (1, 'same'), (2, 'same');",
715        )
716        .await
717        .unwrap();
718
719    let rows = client.query("SELECT id FROM kinds k WHERE jev((k.body, 0), 'q') ORDER BY id", &[]).await.unwrap();
720    assert_eq!(rows.len(), 2);
721    let dedupe: i64 = client.query_one("SELECT dedupe_hits FROM jev_stats()", &[]).await.unwrap().get(0);
722    assert_eq!((mock.requests().len(), dedupe), (1, 1), "the second partition waited on the first's request");
723}