postjevsql.git / tests / budget.rs
1//! Spend guards (contract *Cost and safety*): `jev.max_rows` and
2//! `jev.max_cost` refuse a statement over budget before its first request,
3//! from the planner's estimate, and stop one whose estimate was wrong
4//! before the request that would cross the line. The running count is
5//! the statement's, across its scans, and the running cost is what the
6//! answers' `usage` reports plus the worst case of what is in flight.
7
8use support::mock_jev::MockJev;
9use support::mock_jev::Reply;
10use support::{jev_instance, noul};
11use tokio_postgres::error::SqlState;
12
13const QUERY: &str = "SELECT jev_prob(t, 'Urgent?') FROM tickets t";
14
15/// An instance whose `tickets` holds `rows` rows, analyzed so the
16/// planner's estimate is exact.
17async fn with_rows(mock: &MockJev, rows: i32) -> (support::postgres::Instance, tokio_postgres::Client) {
18    let (pg, client) = jev_instance(mock).await;
19    client
20        .batch_execute(&format!(
21            "INSERT INTO tickets SELECT g, 'more' FROM generate_series(2, {rows}) g; ANALYZE tickets;"
22        ))
23        .await
24        .expect("rows");
25    (pg, client)
26}
27
28#[tokio::test(flavor = "multi_thread")]
29async fn over_max_rows_is_refused_before_anything_is_sent() {
30    let mock = MockJev::start(|_| Reply::json(200, noul(0.5))).await;
31    let (_pg, client) = with_rows(&mock, 5).await;
32    client.batch_execute("SET jev.max_rows = 3").await.unwrap();
33
34    let err = client.query(QUERY, &[]).await.expect_err("refused");
35    let db = err.as_db_error().expect("an ERROR");
36    assert_eq!(db.code(), &SqlState::PROGRAM_LIMIT_EXCEEDED, "{err:?}");
37    assert_eq!(db.message(), "this statement is estimated to send 5 rows, over jev.max_rows (3)");
38    assert!(mock.requests().is_empty());
39}
40
41/// A jev condition in WHERE judges every row that reaches it, so its
42/// guessed selectivity does not shrink the estimate.
43#[tokio::test(flavor = "multi_thread")]
44async fn a_jev_condition_does_not_shrink_the_estimate() {
45    let mock = MockJev::start(|_| Reply::json(200, noul(0.5))).await;
46    let (_pg, client) = with_rows(&mock, 5).await;
47    client.batch_execute("SET jev.max_rows = 3").await.unwrap();
48
49    let err = client.query("SELECT id FROM tickets t WHERE jev(t, 'Urgent?')", &[]).await.expect_err("refused");
50    let db = err.as_db_error().expect("an ERROR");
51    assert_eq!(db.code(), &SqlState::PROGRAM_LIMIT_EXCEEDED, "{err:?}");
52    assert_eq!(db.message(), "this statement is estimated to send 5 rows, over jev.max_rows (3)");
53    assert!(mock.requests().is_empty());
54}
55
56#[tokio::test(flavor = "multi_thread")]
57async fn over_max_cost_is_refused_before_anything_is_sent() {
58    let mock = MockJev::start(|_| Reply::json(200, noul(0.5))).await;
59    let (_pg, client) = with_rows(&mock, 5).await;
60    // 5 requests of at least 267 tokens, 3 attempts each, at $1 per token.
61    client.batch_execute("SET jev.price_per_mtok = 1000000; SET jev.max_cost = 1000").await.unwrap();
62
63    let err = client.query(QUERY, &[]).await.expect_err("refused");
64    let db = err.as_db_error().expect("an ERROR");
65    assert_eq!(db.code(), &SqlState::PROGRAM_LIMIT_EXCEEDED, "{err:?}");
66    assert!(db.message().starts_with("this statement is estimated to spend up to $"), "{}", db.message());
67    assert!(db.message().ends_with("over jev.max_cost ($1000)"), "{}", db.message());
68    assert!(mock.requests().is_empty());
69}
70
71#[tokio::test(flavor = "multi_thread")]
72async fn within_budget_runs() {
73    let mock = MockJev::start(|_| Reply::json(200, noul(0.5))).await;
74    let (_pg, client) = with_rows(&mock, 5).await;
75    client.batch_execute("SET jev.max_rows = 5; SET jev.max_cost = 1").await.unwrap();
76
77    assert_eq!(client.query(QUERY, &[]).await.expect("within budget").len(), 5);
78    assert_eq!(mock.requests().len(), 5);
79}
80
81#[tokio::test(flavor = "multi_thread")]
82async fn a_low_estimate_still_stops_at_the_limit() {
83    let mock = MockJev::start(|_| Reply::json(200, noul(0.5))).await;
84    // Analyzed at one row, then grown: every row still fits the one page,
85    // so the planner keeps estimating one.
86    let (_pg, client) = with_rows(&mock, 1).await;
87    client
88        .batch_execute(
89            "INSERT INTO tickets SELECT g, 'more' FROM generate_series(2, 6) g;
90             SET jev.max_rows = 2; SET jev.concurrency = 1",
91        )
92        .await
93        .unwrap();
94
95    let err = client.query(QUERY, &[]).await.expect_err("stopped");
96    let db = err.as_db_error().expect("an ERROR");
97    assert_eq!(db.code(), &SqlState::PROGRAM_LIMIT_EXCEEDED, "{err:?}");
98    assert_eq!(db.message(), "this statement would send more than jev.max_rows (2) rows");
99    assert_eq!(mock.requests().len(), 2);
100}
101
102#[tokio::test(flavor = "multi_thread")]
103async fn no_limit_by_default() {
104    let mock = MockJev::start(|_| Reply::json(200, noul(0.5))).await;
105    let (_pg, client) = jev_instance(&mock).await;
106    let row = client.query_one("SHOW jev.max_rows", &[]).await.unwrap();
107    assert_eq!(row.get::<_, String>(0), "-1");
108    client.query_one(QUERY, &[]).await.expect("unlimited");
109}
110
111/// An answer reporting `input_tokens` billed.
112fn billed(input_tokens: u64) -> serde_json::Value {
113    let mut answer = noul(0.5);
114    answer["usage"]["input_tokens"] = input_tokens.into();
115    answer
116}
117
118/// Six rows the planner estimates as one, judged one at a time, at $1
119/// per token: every row's worst case is under 900 tokens (3 attempts of
120/// 267 plus its characters over 2.6).
121async fn underestimated(mock: &MockJev, max_cost: u32) -> (support::postgres::Instance, tokio_postgres::Client) {
122    let (pg, client) = with_rows(mock, 1).await;
123    client
124        .batch_execute(&format!(
125            "INSERT INTO tickets SELECT g, 'more' FROM generate_series(2, 6) g;
126             SET jev.concurrency = 1; SET jev.price_per_mtok = 1000000; SET jev.max_cost = {max_cost}"
127        ))
128        .await
129        .unwrap();
130    (pg, client)
131}
132
133#[tokio::test(flavor = "multi_thread")]
134async fn the_running_cost_is_what_answers_report() {
135    // Two worst cases would pass $1,500; six answers of 10 tokens do not.
136    let mock = MockJev::start(|_| Reply::json(200, billed(10))).await;
137    let (_pg, client) = underestimated(&mock, 1500).await;
138
139    assert_eq!(client.query(QUERY, &[]).await.expect("within what was billed").len(), 6);
140    assert_eq!(mock.requests().len(), 6);
141}
142
143#[tokio::test(flavor = "multi_thread")]
144async fn an_answer_billed_over_the_estimate_stops_the_statement() {
145    // One answer billed 3,000 tokens is over $2,000 on its own, though
146    // two worst cases would fit.
147    let mock = MockJev::start(|_| Reply::json(200, billed(3000))).await;
148    let (_pg, client) = underestimated(&mock, 2000).await;
149
150    let err = client.query(QUERY, &[]).await.expect_err("stopped");
151    let db = err.as_db_error().expect("an ERROR");
152    assert_eq!(db.code(), &SqlState::PROGRAM_LIMIT_EXCEEDED, "{err:?}");
153    assert_eq!(db.message(), "this statement would spend more than jev.max_cost ($2000)");
154    assert_eq!(mock.requests().len(), 1);
155}
156
157#[tokio::test(flavor = "multi_thread")]
158async fn the_row_limit_counts_every_scan_of_the_statement() {
159    let mock = MockJev::start(|_| Reply::json(200, noul(0.5))).await;
160    // Estimated at one row per scan, and three in each.
161    let (_pg, client) = with_rows(&mock, 1).await;
162    client
163        .batch_execute(
164            "INSERT INTO tickets SELECT g, 'more' FROM generate_series(2, 3) g;
165             SET jev.max_rows = 4; SET jev.concurrency = 1",
166        )
167        .await
168        .unwrap();
169
170    let err = client
171        .query(
172            "SELECT id FROM tickets t WHERE jev_prob(t, 'Urgent?') > 0
173             UNION ALL
174             SELECT id FROM tickets t WHERE jev_prob(t, 'Angry?') > 0",
175            &[],
176        )
177        .await
178        .expect_err("stopped");
179    let db = err.as_db_error().expect("an ERROR");
180    assert_eq!(db.code(), &SqlState::PROGRAM_LIMIT_EXCEEDED, "{err:?}");
181    assert_eq!(db.message(), "this statement would send more than jev.max_rows (4) rows");
182    assert_eq!(mock.requests().len(), 4);
183}
184
185/// EXPLAIN's properties by label.
186async fn explain(client: &tokio_postgres::Client) -> std::collections::HashMap<String, String> {
187    let rows = client.query(&format!("EXPLAIN {QUERY}"), &[]).await.expect("explained");
188    rows.iter()
189        .filter_map(|r| {
190            let line: String = r.get(0);
191            let (label, value) = line.trim().split_once(": ")?;
192            Some((label.to_owned(), value.to_owned()))
193        })
194        .collect()
195}
196
197#[tokio::test(flavor = "multi_thread")]
198async fn the_token_estimate_learns_from_cached_usage() {
199    // Every answer reports far more tokens than 2.6 characters each would.
200    let mock = MockJev::start(|_| Reply::json(200, billed(5000))).await;
201    let (_pg, client) = with_rows(&mock, 5).await;
202    let tokens = |plan: &std::collections::HashMap<String, String>| -> i64 {
203        plan.get("Estimated Input Tokens").expect("estimated").parse().expect("an integer")
204    };
205
206    let before = tokens(&explain(&client).await);
207    assert_eq!(client.query(QUERY, &[]).await.expect("judged").len(), 5);
208    let after = tokens(&explain(&client).await);
209    assert!(after > 5 * before, "learned {after} tokens against the fallback's {before}");
210
211    // A statement priced at the fallback is refused once the ratio is
212    // learned, before anything more is sent.
213    let sent = mock.requests().len();
214    client
215        .batch_execute(&format!(
216            "SET jev.cache_namespace = 'again'; SET jev.price_per_mtok = 1000000; SET jev.max_cost = {}",
217            3 * 2 * before
218        ))
219        .await
220        .unwrap();
221    let err = client.query(QUERY, &[]).await.expect_err("refused");
222    assert_eq!(err.as_db_error().expect("an ERROR").code(), &SqlState::PROGRAM_LIMIT_EXCEEDED, "{err:?}");
223    assert_eq!(mock.requests().len(), sent);
224}