postjevsql.git / tests / choice.rs
1//! `jev_choice(row, q, options)` (contract *SQL surface*): the top label
2//! of a Choice. Options are sent in the caller's order, which is part of
3//! the question and so of the cache key; a Choice takes 2 to 255 options,
4//! and a single option is the answer without asking.
5
6use serde_json::{Value, json};
7use tokio_postgres::error::SqlState;
8use support::jev_instance;
9use support::mock_jev::{MockJev, Reply};
10
11/// 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}
30
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}
105
106/// A label is palloc'd per row in per-tuple memory, which the parent
107/// resets as it pulls the next row. Palloc'd in the query's context
108/// instead (ExecutorState), every row's label stayed until the statement
109/// 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}