main.rsannotatedmain.rssource239 lines · 10.3 KB · raw
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}