use std::path::Path;
use anyhow::Context as _;
use tritium_nn::ModelRunner;
pub(crate) fn read_token_file(path: &Path) -> anyhow::Result<Vec<u32>> {
let text = std::fs::read_to_string(path)
.with_context(|| format!("failed to read token file `{}`", path.display()))?;
let raw: Vec<i64> = serde_json::from_str(&text).with_context(|| {
format!(
"failed to parse `{}` as a JSON array of ints",
path.display()
)
})?;
raw.into_iter()
.map(|v| {
u32::try_from(v)
.with_context(|| format!("token id {v} is out of range for u32 (0..=4294967295)"))
})
.collect()
}
#[must_use]
pub(crate) fn render_output(tokens: &[u32]) -> String {
use std::fmt::Write as _;
let mut out = String::new();
let json: Vec<String> = tokens.iter().map(u32::to_string).collect();
let _ = writeln!(out, "[{}]", json.join(", "));
for t in tokens {
let _ = writeln!(out, "{t}");
}
out
}
pub(crate) fn run(
model_path: &Path,
tokens: &[u32],
max_new: usize,
greedy: bool,
eos: u32,
) -> anyhow::Result<()> {
if !greedy {
eprintln!(
"note: only greedy decoding is implemented in v0.20; \
`--greedy=false` still decodes greedily"
);
}
let bytes = std::fs::read(model_path)
.with_context(|| format!("failed to read model `{}`", model_path.display()))?;
let mut runner = ModelRunner::load_cpu(&bytes)
.with_context(|| format!("failed to load model `{}`", model_path.display()))?;
let generated = runner
.generate(tokens, max_new, eos)
.context("generation failed")?;
print!("{}", render_output(&generated));
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn render_output_is_json_then_lines() {
let report = render_output(&[1, 2, 3]);
assert!(report.starts_with("[1, 2, 3]\n"), "{report}");
assert!(report.contains("\n1\n2\n3\n"), "{report}");
}
#[test]
fn render_output_empty_is_empty_array() {
let report = render_output(&[]);
assert_eq!(report, "[]\n");
}
#[test]
fn read_token_file_parses_array() {
let mut tmp = std::env::temp_dir();
tmp.push(format!("tritium-cli-tokens-{}.json", std::process::id()));
std::fs::write(&tmp, b"[1, 128000, 9906]").expect("write temp");
let toks = read_token_file(&tmp).expect("parse tokens");
let _ = std::fs::remove_file(&tmp);
assert_eq!(toks, vec![1, 128_000, 9906]);
}
#[test]
fn read_token_file_rejects_negative() {
let mut tmp = std::env::temp_dir();
tmp.push(format!(
"tritium-cli-tokens-neg-{}.json",
std::process::id()
));
std::fs::write(&tmp, b"[1, -5]").expect("write temp");
let err = read_token_file(&tmp).expect_err("negative must error");
let _ = std::fs::remove_file(&tmp);
assert!(format!("{err:#}").contains("out of range"), "{err:#}");
}
#[test]
fn read_token_file_rejects_non_json() {
let mut tmp = std::env::temp_dir();
tmp.push(format!(
"tritium-cli-tokens-bad-{}.json",
std::process::id()
));
std::fs::write(&tmp, b"not json").expect("write temp");
let err = read_token_file(&tmp).expect_err("bad json must error");
let _ = std::fs::remove_file(&tmp);
assert!(format!("{err:#}").contains("failed to parse"), "{err:#}");
}
#[test]
fn run_on_missing_model_errors_cleanly() {
let err = run(
Path::new("/nonexistent/model.gguf"),
&[1, 2, 3],
4,
true,
128_001,
)
.expect_err("missing model must error");
assert!(
format!("{err:#}").contains("failed to read model"),
"{err:#}"
);
}
}