1//! `jev_choice_tree(row, q, paths [, tau])` (contract *SQL surface*, 2//! label tree): a best-first search over a label tree, one Choice per 3//! node. Each round's Choices about a row go in one request, and each 4//! node's Choice is its own `jev_cache` row, so a repeat sends nothing. 5 6use serde_json::{Value, json}; 7use tokio_postgres::error::SqlState; 8use support::jev_instance; 9use support::mock_jev::{MockJev, Reply}; 10 11/// The labels on the path the mock favours. 12const FAVOURED: [&str; 3] = ["support", "billing", "refunds"]; 13 14/// What the mock answers the "does any label fit" Noul. 15const FIT: f64 = 0.35; 16 17/// Gives a favoured option 0.9 and the rest share 0.1; a Choice with no 18/// favoured option is uniform. A Noul, the gate, is [`FIT`]. 19async fn favouring() -> MockJev { 20 MockJev::start(|req| { 21 let body: Value = serde_json::from_slice(&req.body).unwrap(); 22 let mut answers = serde_json::Map::new(); 23 for (qid, q) in body["questions"].as_object().unwrap() { 24 if q["type"] == "noul" { 25 answers.insert(qid.clone(), json!({ "type": "noul", "noul": FIT })); 26 continue; 27 } 28 let labels: Vec<&String> = q["criteria"].as_object().unwrap().keys().collect(); 29 let n = labels.len() as f64; 30 let top = labels.iter().find(|l| FAVOURED.contains(&l.as_str())); 31 let p = |l: &String| match top { 32 Some(t) if *t == l => 0.9, 33 Some(_) => 0.1 / (n - 1.0), 34 None => 1.0 / n, 35 }; 36 let choice = top.copied().unwrap_or(labels[0]).clone(); 37 let probabilities: serde_json::Map<_, _> = labels.iter().map(|l| (l.to_string(), json!(p(l)))).collect(); 38 answers.insert( 39 qid.clone(), 40 json!({ "type": "choice", "choice": choice, "confidence": 0.8, "probabilities": probabilities }), 41 ); 42 } 43 Reply::json(200, json!({ "model": "jev-1.13.0", "answers": answers, "usage": { "input_tokens": 300, "output_tokens": 20 } })) 44 }) 45 .await 46} 47 48const PATHS: &str = "ARRAY['support > billing > refunds', 'support > billing > invoices', 'support > tech', 49 'sales > new', 'sales > renewals']"; 50 51#[tokio::test(flavor = "multi_thread")] 52async fn finds_the_leaf_one_request_per_round_and_caches_each_node() { 53 let mock = favouring().await; 54 let (_pg, client) = jev_instance(&mock).await; 55 56 let sql = format!("SELECT jev_choice_tree(t, 'Which team?', {PATHS}) FROM tickets t"); 57 let path: String = client.query_one(&sql, &[]).await.unwrap().get(0); 58 assert_eq!(path, "support > billing > refunds"); 59 60 // Every request carries one row's Choices, and a node below the root 61 // names its parent path. 62 let requests = mock.requests(); 63 assert!(!requests.is_empty()); 64 let mut choices = 0; 65 for r in &requests { 66 let body: Value = serde_json::from_slice(&r.body).unwrap(); 67 choices += body["questions"].as_object().unwrap().len(); 68 } 69 let sent: Vec<String> = requests.iter().map(|r| String::from_utf8(r.body.to_vec()).unwrap()).collect(); 70 assert!(sent.iter().any(|b| b.contains("Within: support > billing")), "no node named its parent path: {sent:?}"); 71 72 // One cache row per node Choice. 73 let cached: i64 = client.query_one("SELECT count(*) FROM jev_cache", &[]).await.unwrap().get(0); 74 assert_eq!(cached, choices as i64); 75 76 // A repeat is answered from the cache, round by round. 77 let path: String = client.query_one(&sql, &[]).await.unwrap().get(0); 78 assert_eq!(path, "support > billing > refunds"); 79 assert_eq!(mock.requests().len(), requests.len(), "served from the cache"); 80} 81 82#[tokio::test(flavor = "multi_thread")] 83async fn tau_stops_at_the_deepest_node_that_likely() { 84 let mock = favouring().await; 85 let (_pg, client) = jev_instance(&mock).await; 86 87 // support 0.9, billing 0.81, refunds 0.729. 88 let path: String = client 89 .query_one(&format!("SELECT jev_choice_tree(t, 'Which team?', {PATHS}, 0.75) FROM tickets t"), &[]) 90 .await 91 .unwrap() 92 .get(0); 93 assert_eq!(path, "support > billing"); 94 95 // Nothing is that likely: NULL. 96 let none: Option<String> = client 97 .query_one(&format!("SELECT jev_choice_tree(t, 'Which team?', {PATHS}, 0.95) FROM tickets t"), &[]) 98 .await 99 .unwrap() 100 .get(0); 101 assert_eq!(none, None); 102} 103 104/// `jev_choice_tree_full`'s columns: path, probability, separation, 105/// depth, leaf. `fit` is read by [`fit`]. 106type TreeResult = (Option<String>, f64, Option<f64>, i32, bool); 107 108/// `jev_choice_tree_full`'s `fit` column. 109async fn fit(client: &tokio_postgres::Client, paths: &str) -> Option<f64> { 110 let sql = format!("SELECT (jev_choice_tree_full(t, 'Which team?', {paths})).fit FROM tickets t"); 111 client.query_one(&sql, &[]).await.unwrap().get(0) 112} 113 114/// The Nouls a request asked, by instructions. 115fn nouls(body: &[u8]) -> Vec<String> { 116 let body: Value = serde_json::from_slice(body).unwrap(); 117 body["questions"] 118 .as_object() 119 .unwrap() 120 .values() 121 .filter(|q| q["type"] == "noul") 122 .map(|q| q["instructions"].as_str().unwrap().to_owned()) 123 .collect() 124} 125 126#[tokio::test(flavor = "multi_thread")] 127async fn full_asks_whether_any_label_fits_in_the_first_request() { 128 let mock = favouring().await; 129 let (pg, client) = jev_instance(&mock).await; 130 131 // The scalar function returns no fit, so it asks no gate. 132 let sql = format!("SELECT jev_choice_tree(t, 'Which team?', {PATHS}) FROM tickets t"); 133 client.query_one(&sql, &[]).await.unwrap(); 134 assert!(mock.requests().iter().all(|r| nouls(&r.body).is_empty()), "the scalar asked a gate"); 135 let before = mock.requests().len(); 136 137 // The node Choices are cached, so the full search's one request is the 138 // gate, which shows the tree as the root describes it. 139 assert_eq!(fit(&client, PATHS).await, Some(FIT)); 140 let requests = mock.requests(); 141 assert_eq!(requests.len(), before + 1); 142 let gate = nouls(&requests[before].body); 143 assert_eq!( 144 gate, 145 ["Which team?\n\nDoes any of these labels fit? Contains: support, sales. For example: refunds, invoices, tech, new, renewals"] 146 ); 147 148 // On a fresh cache the gate rides in the first round's request, with 149 // the root's Choice. The first instance is stopped first: rebinding 150 // would keep it alive, and a thread may hold one slot at a time. 151 drop((client, pg)); 152 let (_pg, client) = jev_instance(&mock).await; 153 let before = mock.requests().len(); 154 assert_eq!(fit(&client, PATHS).await, Some(FIT)); 155 let requests = &mock.requests()[before..]; 156 assert_eq!(nouls(&requests[0].body).len(), 1); 157 let first: Value = serde_json::from_slice(&requests[0].body).unwrap(); 158 let root = first["questions"].as_object().unwrap().values().any(|q| { 159 q["type"] == "choice" && q["criteria"].as_object().unwrap().keys().eq(["support", "sales"]) 160 }); 161 assert!(root, "the gate went without the root's Choice: {first}"); 162 assert!(requests[1..].iter().all(|r| nouls(&r.body).is_empty()), "the gate was asked again"); 163 164 // A tree of one leaf sends nothing, so no fit. 165 let before = mock.requests().len(); 166 assert_eq!(fit(&client, "ARRAY['support > billing']").await, None); 167 assert_eq!(mock.requests().len(), before); 168} 169 170async fn tree_full(client: &tokio_postgres::Client, tau: Option<f64>) -> TreeResult { 171 let tau = tau.map_or(String::new(), |t| format!(", {t}")); 172 let sql = format!("SELECT to_jsonb(jev_choice_tree_full(t, 'Which team?', {PATHS}{tau}))::text FROM tickets t"); 173 let text: String = client.query_one(&sql, &[]).await.unwrap().get(0); 174 let r: Value = serde_json::from_str(&text).unwrap(); 175 let depth = i32::try_from(r["depth"].as_i64().unwrap()).unwrap(); 176 (r["path"].as_str().map(str::to_owned), r["probability"].as_f64().unwrap(), r["separation"].as_f64(), depth, r["leaf"].as_bool().unwrap()) 177} 178 179fn close(a: f64, b: f64) -> bool { 180 (a - b).abs() < 1e-9 181} 182 183#[tokio::test(flavor = "multi_thread")] 184async fn full_reports_probability_separation_depth_and_leaf() { 185 let mock = favouring().await; 186 let (_pg, client) = jev_instance(&mock).await; 187 188 // refunds 0.9³ = 0.729; the runner-up leaf is tech, 0.9 × 0.1 = 0.09. 189 let (path, p, separation, depth, leaf) = tree_full(&client, None).await; 190 assert_eq!(path.as_deref(), Some("support > billing > refunds")); 191 assert!(close(p, 0.729), "{p}"); 192 assert!(close(separation.unwrap(), 0.729 / 0.09), "{separation:?}"); 193 assert_eq!((depth, leaf), (3, true)); 194 let sent = mock.requests().len(); 195 196 // Under tau a stop proves no runner-up, so separation is NULL. 197 let (path, p, separation, depth, leaf) = tree_full(&client, Some(0.75)).await; 198 assert_eq!(path.as_deref(), Some("support > billing")); 199 assert!(close(p, 0.81), "{p}"); 200 assert_eq!((separation, depth, leaf), (None, 2, false)); 201 202 let (path, p, separation, depth, leaf) = tree_full(&client, Some(0.5)).await; 203 assert_eq!(path.as_deref(), Some("support > billing > refunds")); 204 assert!(close(p, 0.729), "{p}"); 205 assert_eq!((separation, depth, leaf), (None, 3, true)); 206 207 // Nothing that likely: a stop at the root, whose path is NULL. 208 let (path, p, separation, depth, leaf) = tree_full(&client, Some(0.95)).await; 209 assert_eq!(path, None); 210 assert!(close(p, 1.0), "{p}"); 211 assert_eq!((separation, depth, leaf), (None, 0, false)); 212 213 // Every node's judgment is shared with the first search's. 214 assert_eq!(mock.requests().len(), sent, "served from the cache"); 215} 216 217#[tokio::test(flavor = "multi_thread")] 218async fn a_bad_tree_is_refused_before_sending() { 219 let mock = favouring().await; 220 let (_pg, client) = jev_instance(&mock).await; 221 222 for paths in [ 223 "ARRAY[]::text[]", 224 "ARRAY['a > b', NULL]", 225 "ARRAY['a > b', 'a > b']", 226 "ARRAY['a', 'a > b']", 227 ] { 228 let err = client 229 .query(&format!("SELECT jev_choice_tree(t, 'Which team?', {paths}) FROM tickets t"), &[]) 230 .await 231 .expect_err("refused"); 232 assert_eq!(err.as_db_error().expect("an ERROR").code(), &SqlState::INVALID_PARAMETER_VALUE, "{paths}"); 233 } 234 let err = client 235 .query(&format!("SELECT jev_choice_tree(t, 'Which team?', {PATHS}, 1.5) FROM tickets t"), &[]) 236 .await 237 .expect_err("refused"); 238 assert_eq!(err.as_db_error().expect("an ERROR").code(), &SqlState::INVALID_PARAMETER_VALUE); 239 assert!(mock.requests().is_empty(), "refused before sending"); 240} 241 242#[tokio::test(flavor = "multi_thread")] 243async fn lookahead_is_priced_at_the_learned_ratio() { 244 let mock = favouring().await; 245 let (_pg, client) = jev_instance(&mock).await; 246 let sql = format!("SELECT jev_choice_tree(t, 'Which team?', {PATHS}) FROM tickets t"); 247 248 let path: String = client.query_one(&sql, &[]).await.unwrap().get(0); 249 assert_eq!(path, "support > billing > refunds"); 250 let with_lookahead = mock.requests().len(); 251 252 // Learned at a tiny fraction of a character per token, no child's 253 // Choice fits the lookahead budget, so every level is its own round. 254 client 255 .execute( 256 "INSERT INTO jev_cache (key, model, answered_model, namespace, layout, state, question, 257 input_tokens, output_tokens, answer) 258 VALUES ('\\x00', 'jev-1.13.0', 'jev-1.13.0', 'seed', 'row_as_state', '', '', 1000000267, 0, '{}')", 259 &[], 260 ) 261 .await 262 .unwrap(); 263 client.batch_execute("SET jev.cache_namespace = 'again'").await.unwrap(); 264 let path: String = client.query_one(&sql, &[]).await.unwrap().get(0); 265 assert_eq!(path, "support > billing > refunds"); 266 let without = mock.requests().len() - with_lookahead; 267 assert!(without > with_lookahead, "{without} rounds unpriced against {with_lookahead} with lookahead"); 268}