postjevsql.git / tests / cache.rs
cache.rsannotatedcache.rssource259 lines · 11.2 KB · raw
1//! The durable cache (contract *Cache*): an answer is stored under the
2//! hash of what was sent, a repeat is served without a request, and
3//! anything that changes what would be sent misses.
4
5use std::sync::Arc;
6use std::sync::atomic::{AtomicUsize, Ordering};
7
8use support::mock_jev::{MockJev, Reply};
9use support::{jev_instance, jev_instance_with, noul};
10
11const ASK: &str = "SELECT jev_prob(t, 'Urgent?') FROM tickets t";
12
13async fn counting(p: f64) -> (MockJev, Arc<AtomicUsize>) {
14    let sent = Arc::new(AtomicUsize::new(0));
15    let seen = sent.clone();
16    let mock = MockJev::start(move |_| {
17        seen.fetch_add(1, Ordering::SeqCst);
18        Reply::json(200, noul(p))
19    })
20    .await;
21    (mock, sent)
22}
23
24async fn prob(client: &tokio_postgres::Client, sql: &str) -> f64 {
25    client.query_one(sql, &[]).await.expect("answered").get(0)
26}
27
28#[tokio::test(flavor = "multi_thread")]
29async fn a_miss_is_sent_and_stored_with_its_receipt() {
30    let (mock, sent) = counting(0.83).await;
31    let (_pg, client) = jev_instance(&mock).await;
32
33    assert_eq!(prob(&client, ASK).await, 0.83);
34    assert_eq!(sent.load(Ordering::SeqCst), 1);
35
36    let row = client
37        .query_one(
38            "SELECT length(key), model, answered_model, namespace, input_tokens, output_tokens,
39                    (answer->>'noul')::float8, state <> '' AND length(question) > 0
40             FROM jev_cache",
41            &[],
42        )
43        .await
44        .expect("one stored answer");
45    assert_eq!(row.get::<_, i32>(0), 32);
46    assert_eq!(row.get::<_, String>(1), "jev-1.13.0");
47    assert_eq!(row.get::<_, String>(2), "jev-1.13.0");
48    assert_eq!(row.get::<_, String>(3), "");
49    assert_eq!(row.get::<_, i64>(4), 296);
50    assert_eq!(row.get::<_, i64>(5), 20);
51    assert_eq!(row.get::<_, f64>(6), 0.83);
52    assert!(row.get::<_, bool>(7));
53}
54
55#[tokio::test(flavor = "multi_thread")]
56async fn a_hit_is_served_without_a_request() {
57    let (mock, sent) = counting(0.83).await;
58    let (_pg, client) = jev_instance(&mock).await;
59
60    assert_eq!(prob(&client, ASK).await, 0.83);
61    assert_eq!(prob(&client, ASK).await, 0.83);
62    assert_eq!(sent.load(Ordering::SeqCst), 1, "the repeat was served from the cache");
63
64    let r = client.query_one("SELECT cache_hits, cache_misses FROM jev_stats()", &[]).await.unwrap();
65    assert_eq!((r.get::<_, i64>(0), r.get::<_, i64>(1)), (1, 1));
66}
67
68#[tokio::test(flavor = "multi_thread")]
69async fn a_changed_namespace_question_or_row_misses() {
70    let (mock, sent) = counting(0.83).await;
71    let (_pg, client) = jev_instance(&mock).await;
72
73    prob(&client, ASK).await;
74    client.batch_execute("SET jev.cache_namespace = 'rejudge'").await.unwrap();
75    prob(&client, ASK).await;
76    assert_eq!(sent.load(Ordering::SeqCst), 2, "a new namespace re-judges");
77
78    prob(&client, "SELECT jev_prob(t, 'Spam?') FROM tickets t").await;
79    assert_eq!(sent.load(Ordering::SeqCst), 3, "a new question is asked");
80
81    client.batch_execute("UPDATE tickets SET body = body || '!'").await.unwrap();
82    prob(&client, "SELECT jev_prob(t, 'Spam?') FROM tickets t").await;
83    assert_eq!(sent.load(Ordering::SeqCst), 4, "a changed row is asked again");
84
85    client.batch_execute("RESET jev.cache_namespace; DELETE FROM jev_cache").await.unwrap();
86    prob(&client, ASK).await;
87    assert_eq!(sent.load(Ordering::SeqCst), 5, "a deleted answer is asked again");
88}
89
90#[tokio::test(flavor = "multi_thread")]
91async fn equal_rows_in_one_statement_share_one_request() {
92    let (mock, sent) = counting(0.83).await;
93    let (_pg, client) = jev_instance(&mock).await;
94    client
95        .batch_execute("CREATE TABLE twins (body text); INSERT INTO twins VALUES ('same'), ('same'), ('other')")
96        .await
97        .unwrap();
98
99    let rows = client.query("SELECT jev_prob(t, 'Urgent?') FROM twins t", &[]).await.expect("answered");
100    assert!(rows.iter().all(|r| r.get::<_, f64>(0) == 0.83), "every row gets the shared answer");
101    assert_eq!(sent.load(Ordering::SeqCst), 2, "the equal rows were sent once");
102
103    let r = client.query_one("SELECT cache_hits, cache_misses, dedupe_hits FROM jev_stats()", &[]).await.unwrap();
104    assert_eq!((r.get::<_, i64>(0), r.get::<_, i64>(1), r.get::<_, i64>(2)), (0, 2, 1));
105}
106
107#[tokio::test(flavor = "multi_thread")]
108async fn equal_judgments_in_two_scans_of_one_statement_share_one_request() {
109    // 'slow' is still in flight in whichever scan asks first when the
110    // other scan reaches it; 'fast' is answered and stored by then.
111    let sent = Arc::new(AtomicUsize::new(0));
112    let seen = sent.clone();
113    let mock = MockJev::start(move |r| {
114        seen.fetch_add(1, Ordering::SeqCst);
115        let reply = Reply::json(200, noul(0.83));
116        if String::from_utf8_lossy(&r.body).contains("slow") {
117            reply.after(std::time::Duration::from_millis(500))
118        } else {
119            reply
120        }
121    })
122    .await;
123    let (_pg, client) = jev_instance(&mock).await;
124    client
125        .batch_execute(
126            "CREATE TABLE a (body text); INSERT INTO a VALUES ('fast'), ('slow');
127             CREATE TABLE b (body text); INSERT INTO b VALUES ('fast'), ('slow');
128             SET enable_hashjoin = off; SET enable_mergejoin = off",
129        )
130        .await
131        .unwrap();
132
133    let rows = client
134        .query("SELECT 1 FROM a, b WHERE jev_prob(a, 'Urgent?') > 0 AND jev_prob(b, 'Urgent?') > 0", &[])
135        .await
136        .expect("answered");
137    assert_eq!(rows.len(), 4);
138    assert_eq!(sent.load(Ordering::SeqCst), 2, "each distinct judgment was sent once across both scans");
139}
140
141#[tokio::test(flavor = "multi_thread")]
142async fn a_call_waiting_on_a_request_a_rescan_drops_asks_again() {
143    // With a window of two rows, the plan runs:
144    // 1. the outer scan sends 'one' and 'two', and emits 'one';
145    // 2. the inner scan, for 'one', sends 'fast' and 'slow', and LIMIT 1
146    //    stops it at 'fast' with 'slow' still in flight;
147    // 3. the outer scan pulls its own 'slow', equal to the inner one's,
148    //    and waits on that request; it emits 'two';
149    // 4. the inner scan is rescanned for 'two', which drops its 'slow'.
150    //    Its filter no longer includes 'slow', so no one else asks it:
151    //    the outer call must send it itself.
152    let slow = Arc::new(AtomicUsize::new(0));
153    let seen = slow.clone();
154    let mock = MockJev::start(move |r| {
155        let reply = Reply::json(200, noul(0.83));
156        if String::from_utf8_lossy(&r.body).contains("slow") {
157            seen.fetch_add(1, Ordering::SeqCst);
158            reply.after(std::time::Duration::from_millis(1500))
159        } else {
160            reply
161        }
162    })
163    .await;
164    let (_pg, client) = jev_instance_with(&mock, &[("jev.concurrency", "2")]).await;
165    client
166        .batch_execute(
167            "CREATE TABLE a (id int, body text); INSERT INTO a VALUES (1, 'one'), (2, 'two'), (3, 'slow');
168             CREATE TABLE b (id int, body text); INSERT INTO b VALUES (0, 'fast'), (3, 'slow');",
169        )
170        .await
171        .unwrap();
172
173    let rows = client
174        .query(
175            "SELECT o.id, o.p FROM (SELECT id, jev_prob(a, 'Urgent?') AS p FROM a OFFSET 0) o
176             CROSS JOIN LATERAL (
177                 SELECT jev_prob(b, 'Urgent?') FROM b WHERE b.id = 0 OR o.id = 1 LIMIT 1
178             ) s
179             ORDER BY o.id",
180            &[],
181        )
182        .await
183        .expect("the outer call is answered, not failed with its shared request");
184    let got: Vec<(i32, f64)> = rows.iter().map(|r| (r.get(0), r.get(1))).collect();
185    assert_eq!(got, [(1, 0.83), (2, 0.83), (3, 0.83)]);
186    assert_eq!(mock.abandoned(), 1, "the rescan dropped the inner 'slow' in flight");
187    assert_eq!(slow.load(Ordering::SeqCst), 2, "'slow' was sent again once, after the drop");
188
189    let r = client.query_one("SELECT dedupe_hits FROM jev_stats()", &[]).await.unwrap();
190    // The outer 'slow', and the inner 'fast' on the two rescans.
191    assert_eq!(r.get::<_, i64>(0), 3, "the outer 'slow' first waited on the inner one");
192}
193
194/// The `label: value` lines of a text EXPLAIN.
195async fn explain(client: &tokio_postgres::Client, sql: &str) -> std::collections::HashMap<String, String> {
196    let rows = client.query(&format!("EXPLAIN {sql}"), &[]).await.expect("explained");
197    rows.iter()
198        .filter_map(|r| {
199            let line: String = r.get(0);
200            let (label, value) = line.trim().split_once(": ")?;
201            Some((label.to_owned(), value.to_owned()))
202        })
203        .collect()
204}
205
206#[tokio::test(flavor = "multi_thread")]
207async fn explain_prices_the_scan_and_analyze_counts_its_hits() {
208    let (mock, sent) = counting(0.83).await;
209    let (_pg, client) = jev_instance(&mock).await;
210    client.batch_execute("ANALYZE tickets").await.unwrap();
211
212    // Plain EXPLAIN runs no rows, so it cannot know the hits: it prices
213    // every candidate as a miss, the bound the spend guards use.
214    let plan = explain(&client, ASK).await;
215    assert_eq!(plan.get("Candidate Rows").map(String::as_str), Some("1"), "{plan:#?}");
216    assert_eq!(plan.get("Estimated Requests").map(String::as_str), Some("1"), "{plan:#?}");
217    assert!(plan.contains_key("Estimated Input Tokens"), "{plan:#?}");
218    assert!(plan.get("Worst-Case Cost").is_some_and(|c| c.ends_with(" USD")), "{plan:#?}");
219    assert!(!plan.contains_key("Cache Hits"), "{plan:#?}");
220    assert_eq!(sent.load(Ordering::SeqCst), 0, "EXPLAIN sends nothing");
221
222    let first = explain(&client, &format!("(ANALYZE) {ASK}")).await;
223    assert_eq!(first.get("Cache Hits").map(String::as_str), Some("0"), "{first:#?}");
224    assert_eq!(first.get("Cache Misses").map(String::as_str), Some("1"), "{first:#?}");
225    assert_eq!(first.get("Requests").map(String::as_str), Some("1"), "{first:#?}");
226    assert_eq!(first.get("Input Tokens").map(String::as_str), Some("296"), "{first:#?}");
227    assert_eq!(first.get("Cost").map(String::as_str), Some("0.000012 USD"), "{first:#?}");
228
229    let again = explain(&client, &format!("(ANALYZE) {ASK}")).await;
230    assert_eq!(again.get("Cache Hits").map(String::as_str), Some("1"), "{again:#?}");
231    assert_eq!(again.get("Cache Misses").map(String::as_str), Some("0"), "{again:#?}");
232    assert_eq!(again.get("Requests").map(String::as_str), Some("0"), "{again:#?}");
233    assert_eq!(again.get("Cost").map(String::as_str), Some("0.000000 USD"), "{again:#?}");
234    assert_eq!(sent.load(Ordering::SeqCst), 1);
235
236    // COSTS OFF hides the estimate, as it hides the planner's.
237    let bare = explain(&client, &format!("(COSTS OFF) {ASK}")).await;
238    assert!(!bare.contains_key("Candidate Rows"), "{bare:#?}");
239}
240
241#[tokio::test(flavor = "multi_thread")]
242async fn cache_only_serves_hits_and_fails_misses_without_sending() {
243    let (mock, sent) = counting(0.83).await;
244    let (_pg, client) = jev_instance(&mock).await;
245    prob(&client, ASK).await;
246
247    // Nothing is sent, so a zero budget does not refuse it.
248    client.batch_execute("SET jev.cache_only = on; SET jev.max_cost = 0").await.unwrap();
249    assert_eq!(prob(&client, ASK).await, 0.83, "a hit is served");
250
251    let err = client
252        .query_one("SELECT jev_prob(t, 'Spam?') FROM tickets t", &[])
253        .await
254        .expect_err("a miss fails");
255    let db = err.as_db_error().expect("a server error");
256    assert_eq!(db.code().code(), "55000", "{db:?}");
257    assert!(db.message().contains("jev.cache_only"), "{db:?}");
258    assert_eq!(sent.load(Ordering::SeqCst), 1, "nothing was sent in cache_only mode");
259}