postjevsql.git / tests / budget.rs

Spend guards (contract Cost and safety): jev.max_rows and jev.max_cost refuse a statement over budget before its first request, from the planner's estimate, and stop one whose estimate was wrong before the request that would cross the line. The running count is the statement's, across its scans, and the running cost is what the answers' usage reports plus the worst case of what is in flight.

8use support::mock_jev::MockJev;
9use support::mock_jev::Reply;
10use support::{jev_instance, noul};
11use tokio_postgres::error::SqlState;
13const QUERY: &str = "SELECT jev_prob(t, 'Urgent?') FROM tickets t";

An instance whose tickets holds rows rows, analyzed so the 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}
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}

A jev condition in WHERE judges every row that reaches it, so its 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}
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}

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}

Six rows the planner estimates as one, judged one at a time, at $1 per token: every row's worst case is under 900 tokens (3 attempts of 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}
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}

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}
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}