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}