use clap::{Parser, Subcommand};
mod chat;
mod doctor;
mod prune_score;
mod serve;
#[derive(Parser)]
#[command(name = "lattice", about = "Pure-Rust transformer inference engine")]
struct Cli {
#[command(subcommand)]
command: Command,
}
#[derive(Subcommand)]
enum Command {
Chat {
#[arg(long)]
model: String,
#[arg(long, default_value = "256")]
max_tokens: usize,
#[arg(long, default_value = "0.7")]
temperature: f32,
#[arg(long)]
tokenizer_dir: Option<String>,
},
Serve {
#[arg(long)]
model: String,
#[arg(long, default_value = "127.0.0.1")]
host: String,
#[arg(long, default_value = "8080")]
port: u16,
#[arg(long, default_value = "256")]
max_tokens: usize,
#[arg(long)]
model_id: Option<String>,
#[arg(long)]
tokenizer_dir: Option<String>,
#[arg(
long,
default_value = "32",
value_parser = clap::builder::RangedU64ValueParser::<usize>::new()
.range(1..=(tokio::sync::Semaphore::MAX_PERMITS as u64))
)]
max_pending: usize,
#[arg(long)]
preload_vision: bool,
},
Doctor {
#[arg(long)]
model: String,
#[arg(long)]
context: Option<usize>,
#[arg(long)]
tokenizer_dir: Option<String>,
},
PruneScore {
#[command(flatten)]
args: prune_score::Args,
},
}
use lattice_inference::model_format as backend;
#[tokio::main]
async fn main() {
let cli = Cli::parse();
match cli.command {
Command::Chat {
model,
max_tokens,
temperature,
tokenizer_dir,
} => {
chat::run_chat(&model, max_tokens, temperature, tokenizer_dir.as_deref());
}
Command::Serve {
model,
host,
port,
max_tokens,
model_id,
tokenizer_dir,
max_pending,
preload_vision,
} => {
use std::path::Path;
use std::sync::Arc;
use std::sync::atomic::AtomicU64;
let served_model_id = model_id.unwrap_or_else(|| {
Path::new(&model)
.file_name()
.and_then(|n| n.to_str())
.unwrap_or("lattice")
.to_string()
});
let model_path = Path::new(&model);
let format = backend::detect_format(model_path);
eprintln!("Loading model from {model}...");
let model_backend: serve::ModelBackend = match format {
backend::ModelFormat::Safetensors => {
match lattice_inference::model::qwen35::Qwen35Model::from_safetensors(
model_path,
) {
Ok(m) => serve::ModelBackend::Cpu(Arc::new(m)),
Err(e) => {
eprintln!("Error: failed to load model: {e}");
std::process::exit(1);
}
}
}
backend::ModelFormat::Q4 => {
#[cfg(feature = "metal-gpu")]
{
let tokenizer_dir_path =
tokenizer_dir.as_ref().map(std::path::PathBuf::from);
match serve::ModelBackend::spawn_metal(
model_path.to_path_buf(),
tokenizer_dir_path,
max_pending,
preload_vision,
) {
Ok((backend, _max_context)) => backend,
Err(e) => {
eprintln!("Error: failed to load Q4 model: {e}");
std::process::exit(1);
}
}
}
#[cfg(not(feature = "metal-gpu"))]
{
let _ = &tokenizer_dir;
let _ = max_pending;
let _ = preload_vision;
eprintln!("Error: {}", backend::metal_gpu_required_message(model_path));
std::process::exit(1);
}
}
backend::ModelFormat::Unknown => {
eprintln!(
"Error: {}",
backend::unrecognized_format_message(model_path)
);
std::process::exit(1);
}
_ => {
eprintln!(
"Error: {}",
backend::unrecognized_format_message(model_path)
);
std::process::exit(1);
}
};
eprintln!("Model loaded. Serving as '{served_model_id}'.");
let embedding_model =
match lattice_inference::serve::embeddings::EmbeddingModel::from_directory(
model_path,
) {
Ok(embedding_model) => {
eprintln!(
"Embeddings enabled: pooled {}-dim vectors from {model}.",
embedding_model.dimensions()
);
Some(Arc::new(embedding_model))
}
Err(err) => {
eprintln!("Embeddings disabled ({model}): {err}");
None
}
};
let state = serve::AppState {
model: model_backend,
default_max_tokens: max_tokens,
max_tokens_cap: 4096,
model_id: served_model_id.clone(),
request_counter: Arc::new(AtomicU64::new(0)),
embedding_model,
};
let app = serve::router(state);
let addr = format!("{host}:{port}");
let listener = match tokio::net::TcpListener::bind(&addr).await {
Ok(l) => l,
Err(e) => {
drop(app);
eprintln!("Error: failed to bind to {addr}: {e}");
std::process::exit(1);
}
};
eprintln!(
"Listening on {addr} (model: {served_model_id}, max_tokens default: {max_tokens})"
);
eprintln!(" POST /v1/chat/completions");
eprintln!(" GET /health");
if let Err(e) = lattice_inference::serve::serve_until_shutdown(listener, app).await {
eprintln!("Server error: {e}");
std::process::exit(1);
}
}
Command::Doctor {
model,
context,
tokenizer_dir,
} => {
use std::path::Path;
let model_path = Path::new(&model);
let tokenizer_dir_path = tokenizer_dir.as_deref().map(Path::new);
match doctor::build_report(model_path, tokenizer_dir_path, context, None) {
Ok(report) => {
println!("{report}");
if !report.is_ready() {
eprintln!("doctor: model is NOT usable as configured (see reasons above)");
std::process::exit(1);
}
}
Err(e) => {
eprintln!("Error: {e}");
std::process::exit(1);
}
}
}
Command::PruneScore { args } => match prune_score::run(&args) {
Ok(true) => {}
Ok(false) => std::process::exit(1),
Err(e) => {
eprintln!("Error: {e}");
std::process::exit(1);
}
},
}
}
#[cfg(test)]
mod max_pending_cli_tests {
use super::*;
fn parse_max_pending(args: &[&str]) -> Result<usize, clap::Error> {
let mut full = vec!["lattice", "serve", "--model", "/tmp/model"];
full.extend_from_slice(args);
match Cli::try_parse_from(full)?.command {
Command::Serve { max_pending, .. } => Ok(max_pending),
_ => panic!("expected Command::Serve, got a different Command variant"),
}
}
#[test]
fn max_pending_omitted_defaults_to_32() {
assert_eq!(parse_max_pending(&[]).expect("no --max-pending"), 32);
}
#[test]
fn max_pending_zero_is_rejected() {
parse_max_pending(&["--max-pending", "0"])
.expect_err("0 admits nothing and must be rejected, not silently accepted");
}
#[test]
fn max_pending_one_above_max_permits_is_rejected() {
let too_big = (tokio::sync::Semaphore::MAX_PERMITS as u128 + 1).to_string();
parse_max_pending(&["--max-pending", &too_big]).expect_err(
"Semaphore::MAX_PERMITS + 1 must be rejected before it can panic Semaphore::new",
);
}
#[test]
fn max_pending_at_max_permits_is_accepted() {
let at_max = tokio::sync::Semaphore::MAX_PERMITS.to_string();
assert_eq!(
parse_max_pending(&["--max-pending", &at_max])
.expect("Semaphore::MAX_PERMITS itself is the inclusive upper bound"),
tokio::sync::Semaphore::MAX_PERMITS
);
}
#[test]
fn max_pending_negative_is_rejected() {
parse_max_pending(&["--max-pending", "-1"])
.expect_err("a negative value must be rejected, not silently defaulted");
}
#[test]
fn max_pending_malformed_is_rejected() {
parse_max_pending(&["--max-pending", "not-a-number"])
.expect_err("a non-numeric value must be rejected, not silently defaulted");
}
#[test]
fn max_pending_valid_override_is_accepted() {
assert_eq!(
parse_max_pending(&["--max-pending", "8"]).expect("8 is a valid cap"),
8
);
}
}
#[cfg(test)]
mod preload_vision_cli_tests {
use super::*;
fn parse_preload_vision(args: &[&str]) -> bool {
let mut full = vec!["lattice", "serve", "--model", "/tmp/model"];
full.extend_from_slice(args);
match Cli::try_parse_from(full)
.expect("fixed --model arg always parses")
.command
{
Command::Serve { preload_vision, .. } => preload_vision,
_ => panic!("expected Command::Serve, got a different Command variant"),
}
}
#[test]
fn preload_vision_omitted_defaults_to_false() {
assert!(!parse_preload_vision(&[]));
}
#[test]
fn preload_vision_flag_present_is_true() {
assert!(parse_preload_vision(&["--preload-vision"]));
}
}