jev_choice(row, q, options) (contract SQL surface): the top label
of a Choice. Options are sent in the caller's order, which is part of
the question and so of the cache key; a Choice takes 2 to 255 options,
and a single option is the answer without asking.
Answers every Choice with its last option.
12async fn last_option() -> MockJev { 13 MockJev::start(|req| { 14 let body: Value = serde_json::from_slice(&req.body).unwrap(); 15 let mut answers = serde_json::Map::new(); 16 for (qid, q) in body["questions"].as_object().unwrap() { 17 let labels: Vec<&String> = q["criteria"].as_object().unwrap().keys().collect(); 18 let last = labels.last().unwrap().to_string(); 19 let probabilities: serde_json::Map<_, _> = 20 labels.iter().map(|l| (l.to_string(), json!(if **l == last { 1.0 } else { 0.0 }))).collect(); 21 answers.insert( 22 qid.clone(), 23 json!({ "type": "choice", "choice": last, "confidence": 0.9, "probabilities": probabilities }), 24 ); 25 } 26 Reply::json(200, json!({ "model": "jev-1.13.0", "answers": answers, "usage": { "input_tokens": 296, "output_tokens": 20 } })) 27 }) 28 .await 29}
31const OPTIONS: &str = "ARRAY['sales', 'billing', 'other']"; 32 33#[tokio::test(flavor = "multi_thread")] 34async fn top_label_in_caller_order_and_cached() { 35 let mock = last_option().await; 36 let (_pg, client) = jev_instance(&mock).await; 37 client.batch_execute("INSERT INTO tickets VALUES (2, 'b')").await.unwrap(); 38 39 let sql = format!("SELECT id, jev_choice(t, 'Which team?', {OPTIONS}) FROM tickets t ORDER BY id"); 40 let rows = client.query(&sql, &[]).await.unwrap(); 41 let got: Vec<(i32, String)> = rows.iter().map(|r| (r.get(0), r.get(1))).collect(); 42 assert_eq!(got, [(1, "other".to_owned()), (2, "other".to_owned())]); 43 44 // The options went in the caller's order, not sorted. 45 let requests = mock.requests(); 46 assert_eq!(requests.len(), 2); 47 let sent = String::from_utf8(requests[0].body.to_vec()).unwrap(); 48 assert!(sent.contains(r#""criteria":{"sales":null,"billing":null,"other":null}"#), "options out of order in {sent}"); 49 50 // A repeat is a cache hit, and sends nothing. 51 let before: i64 = client.query_one("SELECT cache_hits FROM jev_stats()", &[]).await.unwrap().get(0); 52 let rows = client.query(&sql, &[]).await.unwrap(); 53 assert_eq!(rows.iter().map(|r| r.get::<_, String>(1)).collect::<Vec<_>>(), ["other", "other"]); 54 let after: i64 = client.query_one("SELECT cache_hits FROM jev_stats()", &[]).await.unwrap().get(0); 55 assert_eq!(after - before, 2); 56 assert_eq!(mock.requests().len(), 2, "served from the cache"); 57 58 // Reordering the options is another question, with its own key. 59 let label: String = client 60 .query_one("SELECT jev_choice(t, 'Which team?', ARRAY['other', 'billing', 'sales']) FROM tickets t WHERE id = 1", &[]) 61 .await 62 .unwrap() 63 .get(0); 64 assert_eq!(label, "sales"); 65 assert_eq!(mock.requests().len(), 3); 66 let cached: i64 = client.query_one("SELECT count(*) FROM jev_cache", &[]).await.unwrap().get(0); 67 assert_eq!(cached, 3); 68} 69 70#[tokio::test(flavor = "multi_thread")] 71async fn one_option_is_answered_locally() { 72 let mock = last_option().await; 73 let (_pg, client) = jev_instance(&mock).await; 74 75 let label: String = 76 client.query_one("SELECT jev_choice(t, 'Which team?', ARRAY['sales']) FROM tickets t", &[]).await.unwrap().get(0); 77 assert_eq!(label, "sales"); 78 assert!(mock.requests().is_empty(), "nothing sent"); 79 let cached: i64 = client.query_one("SELECT count(*) FROM jev_cache", &[]).await.unwrap().get(0); 80 assert_eq!(cached, 0, "nothing judged, nothing cached"); 81} 82 83#[tokio::test(flavor = "multi_thread")] 84async fn zero_or_more_than_255_options_are_refused() { 85 let mock = last_option().await; 86 let (_pg, client) = jev_instance(&mock).await; 87 88 let many = format!("ARRAY[{}]", (0..256).map(|i| format!("'o{i}'")).collect::<Vec<_>>().join(",")); 89 for options in ["ARRAY[]::text[]", many.as_str(), "ARRAY['a', NULL]", "ARRAY['a', 'a']"] { 90 let err = client 91 .query(&format!("SELECT jev_choice(t, 'Which team?', {options}) FROM tickets t"), &[]) 92 .await 93 .expect_err("refused"); 94 assert_eq!(err.as_db_error().expect("an ERROR").code(), &SqlState::INVALID_PARAMETER_VALUE, "{options}"); 95 } 96 assert!(mock.requests().is_empty(), "refused before sending"); 97 98 // 255 options are sent. 99 let most = format!("ARRAY[{}]", (0..255).map(|i| format!("'o{i}'")).collect::<Vec<_>>().join(",")); 100 let label: String = 101 client.query_one(&format!("SELECT jev_choice(t, 'Which team?', {most}) FROM tickets t"), &[]).await.unwrap().get(0); 102 assert_eq!(label, "o254"); 103 assert_eq!(mock.requests().len(), 1); 104}
A label is palloc'd per row in per-tuple memory, which the parent resets as it pulls the next row. Palloc'd in the query's context instead (ExecutorState), every row's label stayed until the statement ended, so a long scan with large labels grew without bound.
110#[tokio::test(flavor = "multi_thread")] 111async fn labels_do_not_accumulate_in_executor_state() { 112 const ROWS: usize = 1000; 113 const LABEL: usize = 4000; 114 let mock = last_option().await; 115 let (_pg, client) = jev_instance(&mock).await; 116 client 117 .batch_execute(&format!( 118 "INSERT INTO tickets SELECT g, 'ticket ' || g FROM generate_series(2, {ROWS}) g; ANALYZE tickets; 119 -- The statement's own ExecutorState is the shallowest; the one 120 -- this function's query runs in sits below it. 121 CREATE FUNCTION executor_bytes() RETURNS bigint VOLATILE LANGUAGE sql AS $$ 122 SELECT total_bytes FROM pg_backend_memory_contexts 123 WHERE name = 'ExecutorState' ORDER BY level LIMIT 1 124 $$;" 125 )) 126 .await 127 .unwrap(); 128 129 let sql = format!( 130 "SELECT jev_choice(t, 'Which team?', ARRAY[repeat('a', {LABEL}), repeat('b', {LABEL})]), executor_bytes() 131 FROM tickets t" 132 ); 133 let rows = client.query(&sql, &[]).await.unwrap(); 134 assert_eq!(rows.len(), ROWS); 135 assert!(rows.iter().all(|r| r.get::<_, String>(0) == "b".repeat(LABEL)), "the last option, in full"); 136 137 // Past the first in-flight window, the context stays flat. Kept per 138 // row, the labels alone would add ROWS × LABEL bytes (about 4 MB). 139 let bytes: Vec<i64> = rows.iter().map(|r| r.get(1)).collect(); 140 let settled = bytes[200]; 141 let peak = *bytes.iter().max().unwrap(); 142 assert!( 143 peak - settled < 512 * 1024, 144 "ExecutorState grew from {settled} to {peak} bytes over the statement: {:?}", 145 bytes.iter().step_by(100).collect::<Vec<_>>() 146 ); 147}