postjevsql.git / tests / batching.rs

The batch scan (contract Execution): every jev* call in a query is evaluated by one plan node, which judges rows concurrently as streams on the backend's one HTTP/2 connection, and each row alone.

5use serde_json::{Value, json};
6use support::jev_instance;
7use support::mock_jev::{MockJev, Recorded, Reply};
8use tokio_postgres::error::SqlState;

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}
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}

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}

Runs query over 20 tickets with every request held until all 20 arrive, so it hangs unless the calls are batched; then checks it against reference, which reads the same judgments per row (from the cache, sending nothing). EXPLAIN's Candidate Rows is every row judged, before the jev 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}
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}

Postgres keeps an EXISTS whose WHERE is volatile as a correlated SubPlan (convert_EXISTS_sublink_to_join), so each outer row runs it: 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}
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}

A correlated SubPlan runs once per outer row, so each run's scan judges that run's rows; the answers equal the per-row reference, and a 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}

Answers a Noul with a tenth of the integer column in the state: the 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}

A call reading an outer column gets it as a Param, which changes on every run of the SubPlan: the scan evaluates it again on each rescan, sends that run's value, and serves a repeated value from the cache. The first query's child does not depend on the Param (only the projection does), so a scan that replayed its first run would answer 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}

A jev condition over the worktable of a WITH RECURSIVE: each iteration's scan judges the rows the last one produced, and the 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}

The same recursion with the condition on the table the recursive term joins: that scan is rescanned every iteration, and a row it judged in 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}

A partitioned table: each partition's scan is wrapped, with the quals and the select list translated to it, so every partition's rows are 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}

An Append pulls its partitions one after another, so each partition's 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}
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}