1//! Prices the four release gates' recording runs (contract *Testing*, 2//! "Release gates") before anything is spent: for each run, the requests 3//! it sends and their worst case at `jev.price_per_mtok`, every request 4//! billed three times, as `jev.max_cost` budgets it. 5//! 6//! - **Layout parity**: the same rows, one question each, sent both ways: 7//! row as state (one request per row) and row in instructions (every 8//! row a question in one request, the state empty). Row in instructions 9//! is built here from jev-protocol directly, since the extension's 10//! planner cannot reach it until this gate passes. 11//! - **Unsure band**: one question about one row, asked repeatedly. 12//! - **Label tree**: the first round of the label-tree search, lookahead 13//! included, with the "does any label fit" Noul it carries. 14//! - **Ranking**: `ORDER BY jev_prob … LIMIT` over two rows, one request 15//! each. 16//! 17//! `--send` records all four gates as replay fixtures, 18//! `<out-dir>/<gate>.json`, in the format `tests/replay.rs` reads. The 19//! layout-parity and unsure-band requests are sent as built here, each 20//! distinct request once (the unsure band's repeats are one request, so 21//! its fixture holds one answer). The label-tree and ranking runs go 22//! through a throwaway Postgres serving the built extension, which asks a 23//! jev-mock that forwards every request to the real endpoint with the key 24//! and records the exchange. The label tree runs its 25//! whole search, not only the round priced above; `--max-cost` (set as 26//! `jev.max_cost`) bounds what the two runs may spend. 27//! 28//! A gate is recorded once and replayed from then on, so `--send` refuses, 29//! before any key is read or request sent, when any of the four fixtures 30//! already exists in `--out-dir`; `--rerecord` is the deliberate way to 31//! spend again and overwrite them. Without it each fixture is also 32//! written only if it does not exist, so two concurrent runs cannot both 33//! write one. 34 35#![forbid(unsafe_code)] 36 37use std::process::ExitCode; 38 39mod send; 40 41use jev_protocol::{Json, ModelId, Noul, Questions, request_bytes}; 42use postjevsql_core::label_tree::{Ask, Branch, Describe, Params, Search, Step, Tree}; 43use postjevsql_core::{Planner, TokenRatio}; 44 45const USAGE: &str = "usage: record-gates --model <versioned id> [--price-per-mtok <$>] \ 46[--repeats <n>] [--send [--api-key-file <path>] [--max-cost <$>] [--out-dir <dir>] [--rerecord]]"; 47 48pub(crate) use support::gates::{QUESTION, ROWS, TREE_PATHS}; 49 50pub(crate) struct Args { 51 pub(crate) model: ModelId, 52 pub(crate) price_per_mtok: f64, 53 ratio: TokenRatio, 54 repeats: usize, 55 send: bool, 56 api_key_file: Option<String>, 57 pub(crate) max_cost: f64, 58 pub(crate) out_dir: String, 59 pub(crate) rerecord: bool, 60} 61 62fn parse(mut argv: impl Iterator<Item = String>) -> Result<Args, String> { 63 let (mut model, mut price, mut repeats) = (None, 0.042, 30); 64 let (mut send, mut api_key_file, mut max_cost, mut out_dir) = (false, None, 0.01, "tests/fixtures".to_owned()); 65 let mut rerecord = false; 66 while let Some(arg) = argv.next() { 67 let mut value = || argv.next().ok_or(format!("{arg} needs a value")); 68 match arg.as_str() { 69 "--model" => model = Some(ModelId::pinned(&value()?).map_err(|e| e.to_string())?), 70 "--price-per-mtok" => price = number(&arg, &value()?)?, 71 "--repeats" => repeats = value()?.parse().map_err(|_| "--repeats takes a count".to_owned())?, 72 "--api-key-file" => api_key_file = Some(value()?), 73 "--send" => send = true, 74 "--max-cost" => max_cost = number(&arg, &value()?)?, 75 "--out-dir" => out_dir = value()?, 76 "--rerecord" => rerecord = true, 77 _ => return Err(format!("unknown argument {arg:?}")), 78 } 79 } 80 // No cache to learn from, so the conservative fallback, which can 81 // only overstate the price. 82 let ratio = TokenRatio::FALLBACK; 83 if repeats < 2 { 84 return Err("--repeats must be at least 2, or there is no spread to measure".into()); 85 } 86 let model = model.ok_or("--model is required: the gates are recorded against a pinned version")?; 87 if rerecord && !send { 88 return Err("--rerecord only applies to --send".into()); 89 } 90 Ok(Args { model, price_per_mtok: price, ratio, repeats, send, api_key_file, max_cost, out_dir, rerecord }) 91} 92 93fn number(flag: &str, v: &str) -> Result<f64, String> { 94 v.parse::<f64>().ok().filter(|n| n.is_finite() && *n > 0.0).ok_or(format!("{flag} takes a positive number")) 95} 96 97/// One run's requests, as the bytes that would be sent. 98pub(crate) struct Run { 99 name: &'static str, 100 pub(crate) requests: Vec<Vec<u8>>, 101} 102 103fn noul(instructions: Json) -> Questions { 104 let mut q = Questions::new(); 105 q.noul("q", Noul::new(instructions)).expect("one Noul is a valid request"); 106 q 107} 108 109/// A row sent as the state, as the planner places it. 110fn row_as_state(model: &ModelId, row: &str) -> Vec<u8> { 111 let planner = Planner::new(); 112 let placed = planner.place(planner.layout(1), row).expect("the built-in rows are JSON"); 113 request_bytes(model, placed.state(), &noul(Json::text(QUESTION))).expect("a valid request") 114} 115 116pub(crate) fn layout_parity(model: &ModelId) -> Vec<Run> { 117 let state = ROWS.iter().map(|row| row_as_state(model, row)).collect(); 118 let mut questions = Questions::new(); 119 for (i, row) in ROWS.iter().enumerate() { 120 let instructions = serde_json::json!({ "question": QUESTION, "row": serde_json::from_str::<serde_json::Value>(row).unwrap() }); 121 let instructions = Json::canonical(&instructions.to_string()).expect("canonical JSON"); 122 questions.noul(&format!("r{i}"), Noul::new(instructions)).expect("distinct ids"); 123 } 124 let empty = Json::canonical("{}").expect("canonical JSON"); 125 vec![ 126 Run { name: "layout parity: row as state", requests: state }, 127 Run { 128 name: "layout parity: row in instructions", 129 requests: vec![request_bytes(model, &empty, &questions).expect("a valid request")], 130 }, 131 ] 132} 133 134pub(crate) fn unsure_band(model: &ModelId, repeats: usize) -> Run { 135 Run { name: "unsure band", requests: vec![row_as_state(model, ROWS[0]); repeats] } 136} 137 138fn label_tree(model: &ModelId, ratio: TokenRatio) -> Run { 139 let tree = Tree::new(Branch::from_paths(TREE_PATHS.iter().map(|p| p.split(" > "))).expect("a valid tree")) 140 .expect("a valid tree"); 141 let describe = Describe::default(); 142 let choice = |ask: Ask| tree.choice(ask.node(), ask.order(), QUESTION, describe).expect("asked nodes branch"); 143 let mut search = Search::new(&tree, Params::default()); 144 let Step::Ask(asks) = search.step(|ask| { 145 let bytes = jev_protocol::question_bytes(&jev_protocol::Question::Choice(choice(ask))).expect("a valid question"); 146 ratio.tokens(0.0, bytes.len() as f64).ceil() as usize 147 }) else { 148 unreachable!("a tree with more than one leaf asks before it is done") 149 }; 150 let mut questions = Questions::new(); 151 for (i, ask) in asks.into_iter().enumerate() { 152 questions.choice(&format!("c{i}"), choice(ask)).expect("distinct ids"); 153 } 154 questions.noul("fit", tree.gate(QUESTION, describe).expect("more than one leaf")).expect("distinct ids"); 155 let planner = Planner::new(); 156 let placed = planner.place(planner.layout(questions.len()), ROWS[0]).expect("the built-in rows are JSON"); 157 Run { name: "label tree (round 1)", requests: vec![request_bytes(model, placed.state(), &questions).expect("a valid request")] } 158} 159 160fn ranking(model: &ModelId) -> Run { 161 Run { name: "ranking", requests: ROWS[..2].iter().map(|row| row_as_state(model, row)).collect() } 162} 163 164/// Exactly one of `TYPESAFE_API_KEY` and `--api-key-file`, as the 165/// extension reads its key: both set is refused rather than ranked, 166/// which would silently bill the wrong account. 167fn api_key(file: Option<&str>) -> Result<String, String> { 168 match (std::env::var("TYPESAFE_API_KEY").ok().filter(|k| !k.is_empty()), file) { 169 (Some(_), Some(_)) => Err("both TYPESAFE_API_KEY and --api-key-file are set; set exactly one".into()), 170 (None, None) => Err("--send needs an API key: set TYPESAFE_API_KEY or --api-key-file".into()), 171 (Some(key), None) => Ok(key), 172 (None, Some(path)) => std::fs::read_to_string(path) 173 .map(|k| k.trim().to_owned()) 174 .map_err(|e| format!("cannot read --api-key-file: {e}")), 175 } 176} 177 178fn main() -> ExitCode { 179 let args = match parse(std::env::args().skip(1)) { 180 Ok(a) => a, 181 Err(e) => { 182 eprintln!("record-gates: {e}\n{USAGE}"); 183 return ExitCode::from(2); 184 } 185 }; 186 let mut runs = layout_parity(&args.model); 187 runs.extend([unsure_band(&args.model, args.repeats), label_tree(&args.model, args.ratio), ranking(&args.model)]); 188 189 println!( 190 "model {}, ${} per Mtok, {} chars per token, {} billed attempts per request", 191 args.model.as_str(), 192 args.price_per_mtok, 193 args.ratio.chars_per_token(), 194 postjevsql_core::token_ratio::BILLED_ATTEMPTS, 195 ); 196 println!("{:<36} {:>8} {:>12} {:>14}", "run", "requests", "input tokens", "worst case $"); 197 let (mut requests, mut chars) = (0.0, 0.0); 198 for run in &runs { 199 let (n, c) = (run.requests.len() as f64, run.requests.iter().map(Vec::len).sum::<usize>() as f64); 200 requests += n; 201 chars += c; 202 println!( 203 "{:<36} {:>8} {:>12.0} {:>14.6}", 204 run.name, 205 n, 206 args.ratio.tokens(n, c), 207 args.ratio.worst_case(n, c, args.price_per_mtok) 208 ); 209 } 210 println!( 211 "{:<36} {:>8} {:>12.0} {:>14.6}", 212 "total", 213 requests, 214 args.ratio.tokens(requests, chars), 215 args.ratio.worst_case(requests, chars, args.price_per_mtok) 216 ); 217 218 if args.send { 219 if let Err(e) = send::refuse_recorded(&args) { 220 eprintln!("record-gates: {e}"); 221 return ExitCode::from(2); 222 } 223 let key = match api_key(args.api_key_file.as_deref()) { 224 Ok(k) => k, 225 Err(e) => { 226 eprintln!("record-gates: {e}"); 227 return ExitCode::from(2); 228 } 229 }; 230 return match send::record(&args, &key) { 231 Ok(()) => ExitCode::SUCCESS, 232 Err(e) => { 233 eprintln!("record-gates: {e}"); 234 ExitCode::FAILURE 235 } 236 }; 237 } 238 ExitCode::SUCCESS 239}