main.rsannotatedmain.rssource239 lines · 10.3 KB · raw

Prices the four release gates' recording runs (contract Testing, "Release gates") before anything is spent: for each run, the requests it sends and their worst case at jev.price_per_mtok, every request billed three times, as jev.max_cost budgets it.

  • Layout parity: the same rows, one question each, sent both ways: row as state (one request per row) and row in instructions (every row a question in one request, the state empty). Row in instructions is built here from jev-protocol directly, since the extension's planner cannot reach it until this gate passes.
  • Unsure band: one question about one row, asked repeatedly.
  • Label tree: the first round of the label-tree search, lookahead included, with the "does any label fit" Noul it carries.
  • Ranking: ORDER BY jev_prob … LIMIT over two rows, one request each.

--send records all four gates as replay fixtures, <out-dir>/<gate>.json, in the format tests/replay.rs reads. The layout-parity and unsure-band requests are sent as built here, each distinct request once (the unsure band's repeats are one request, so its fixture holds one answer). The label-tree and ranking runs go through a throwaway Postgres serving the built extension, which asks a jev-mock that forwards every request to the real endpoint with the key and records the exchange. The label tree runs its whole search, not only the round priced above; --max-cost (set as jev.max_cost) bounds what the two runs may spend.

A gate is recorded once and replayed from then on, so --send refuses, before any key is read or request sent, when any of the four fixtures already exists in --out-dir; --rerecord is the deliberate way to spend again and overwrite them. Without it each fixture is also written only if it does not exist, so two concurrent runs cannot both write one.

35#![forbid(unsafe_code)]
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}

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

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

Exactly one of TYPESAFE_API_KEY and --api-key-file, as the extension reads its key: both set is refused rather than ranked, 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}
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}