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}