use std::io::{Read, Write};
use crate::canonical::{CanonicalError, ExitClass};
use crate::cli::{parse_args, Args};
use crate::config::{config_path, defaults, partial_from_env, read_config_file, OutMode};
use crate::pipeline::{open_input, read_files, read_request};
use crate::store::{Clock, CredStore, ModelCache};
use crate::transport::Transport;
mod fetch;
use fetch::fetch_count;
pub struct CountIo<'a> {
pub stdout: &'a mut dyn Write,
pub stderr: &'a mut dyn Write,
pub transport: &'a dyn Transport,
pub store: &'a dyn CredStore,
pub cache: &'a dyn ModelCache,
pub clock: &'a dyn Clock,
}
pub fn count_tokens(args: &Args, reader: &mut dyn Read, io: &mut CountIo) -> u8 {
match run_count(args, reader, io) {
Ok(code) => code,
Err(e) => {
let _ = writeln!(io.stderr, "{}", e.message);
e.exit_code()
}
}
}
fn run_count(args: &Args, reader: &mut dyn Read, io: &mut CountIo) -> Result<u8, CanonicalError> {
let flags = parse_args(&args.argv)?;
if flags.help {
return Ok(super::emit(io.stdout, super::HELP));
}
if flags.skill {
return Ok(super::emit(io.stdout, super::SKILL));
}
if flags.version {
return Ok(super::emit(io.stdout, super::VERSION_LINE));
}
let file = read_config_file(&config_path(flags.config_path, &args.env))?;
let env = partial_from_env(&args.env)?;
let merged = flags.config.or(env).or(file).or(defaults());
let json = merged.output == Some(OutMode::Ndjson);
let file_parts = match read_files(&flags.files) {
Ok(parts) => parts,
Err((path, e)) => {
let _ = writeln!(io.stderr, "cannot read --file `{}`: {e}", path.display());
return Ok(ExitClass::NoInput.code());
}
};
let mut input_file;
let rdr: &mut dyn Read = match &flags.input {
Some(path) => match open_input(Some(path)) {
Ok(f) => {
input_file = f;
&mut *input_file
}
Err(_) => {
let _ = writeln!(io.stderr, "cannot open --input file `{}`", path.display());
return Ok(ExitClass::NoInput.code());
}
},
None => reader,
};
let request = read_request(flags.prompt.as_deref(), file_parts, rdr)?;
let req_model = (!request.model.is_empty()).then(|| request.model.clone());
let cfg = merged.into_resolved(req_model.as_deref(), Some(io.cache))?;
let n = fetch_count(request, cfg, io)?;
print_count(io.stdout, n, json).map_err(write_failed)?;
Ok(0)
}
fn print_count(out: &mut dyn Write, n: u32, json: bool) -> std::io::Result<()> {
if json {
writeln!(out, "{}", serde_json::json!({ "input_tokens": n }))
} else {
writeln!(out, "{n}")
}
}
fn write_failed(e: std::io::Error) -> CanonicalError {
CanonicalError {
kind: crate::canonical::ErrorKind::Transport,
message: format!("failed to write token count: {e}"),
provider_detail: None,
retry_after_seconds: None,
}
}