postjevsql.git / tests / label_tree.rs
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}