pub mod capabilities;
pub mod commands;
pub mod config;
pub mod dependency;
pub mod launchers;
pub mod models;
pub mod proxy;
pub mod registry;
pub mod utils;
pub mod version {
include!(concat!(env!("OUT_DIR"), "/version.rs"));
}
pub mod providers;
use alog::{MessageLevel, alog};
use clap::{Parser, Subcommand};
use commands::{
CapabilityCommands, HardwareCommands, LauncherCommands, ModelCommands, ProviderCommands,
};
use utils::ui::{UI_REGISTRY, Ui, run_interactive_tui};
extern crate paste;
#[derive(Parser, Debug)]
#[command(name = "granite-cli")]
#[command(about = "Universal Model Adapter with Capabilities", long_about = None)]
struct Cli {
#[command(subcommand)]
command: Option<Commands>,
#[arg(
short,
long,
global = true,
default_value = "warning",
env = "LOG_LEVEL"
)]
log_level: String,
#[arg(long, global = true, default_value = "", env = "LOG_FILTERS")]
log_filters: String,
#[arg(long, global = true, env = "LOG_JSON")]
log_json: bool,
#[arg(long, global = true, env = "LOG_THREAD_ID")]
log_thread_id: bool,
}
#[derive(clap::Args, Debug)]
struct ModelWithOutput {
#[arg(short, long, global = true, default_value = "terminal")]
output: String,
#[command(subcommand)]
subcommand: ModelSubcommands,
}
#[derive(clap::Args, Debug)]
struct CapabilityWithOutput {
#[arg(short, long, global = true, default_value = "terminal")]
output: String,
#[command(subcommand)]
subcommand: CapabilitySubcommands,
}
#[derive(clap::Args, Debug)]
struct ProviderWithOutput {
#[arg(short, long, global = true, default_value = "terminal")]
output: String,
#[command(subcommand)]
subcommand: ProviderSubcommands,
}
#[derive(clap::Args, Debug)]
struct LaunchWithOutput {
#[arg(short, long, global = true, default_value = "terminal")]
output: String,
launcher_id: String,
#[arg(long)]
dry_run: bool,
#[arg(short = 'u', long = "usage-tracking")]
usage_tracking: bool,
#[arg(trailing_var_arg = true)]
args: Vec<String>,
}
#[derive(clap::Args, Debug)]
struct LauncherWithOutput {
#[arg(short, long, global = true, default_value = "terminal")]
output: String,
#[command(subcommand)]
subcommand: LauncherSubcommands,
}
#[derive(Subcommand, Debug)]
enum Commands {
Model(ModelWithOutput),
Capability(CapabilityWithOutput),
Provider(ProviderWithOutput),
Launcher(LauncherWithOutput),
Hardware,
Launch(LaunchWithOutput),
Version,
}
#[derive(Subcommand, Debug)]
enum ModelSubcommands {
Catalog {
#[arg(short, long)]
r#type: Option<String>,
},
List {
#[arg(short, long)]
r#type: Option<String>,
},
Search {
query: String,
},
Recommend {
#[arg(short, long)]
r#type: Option<String>,
#[arg(short = 'p', long = "providers", value_delimiter = ',')]
providers: Vec<String>,
#[arg(long)]
wide: bool,
},
Info {
model_id: String,
},
Setup {
model_id: String,
},
Pull {
model_id: String,
},
Remove {
model_id: String,
},
}
#[derive(Subcommand, Debug)]
enum CapabilitySubcommands {
Catalog,
List,
Info {
capability_id: String,
},
Setup {
capability_type: String,
#[arg(long = "id")]
instance_id: Option<String>,
},
Remove {
capability_id: String,
},
}
#[derive(Subcommand, Debug)]
enum ProviderSubcommands {
Catalog {
#[arg(long)]
wide: bool,
},
List,
Setup {
provider_type: String,
#[arg(long = "id")]
instance_id: Option<String>,
},
Health {
provider_id: Option<String>,
},
Remove {
provider_id: String,
},
}
#[derive(Subcommand, Debug)]
enum LauncherSubcommands {
Catalog,
List,
Setup {
launcher_type: String,
#[arg(long = "id")]
instance_id: Option<String>,
},
Remove {
launcher_id: String,
},
}
pub struct AppContext {
pub config: config::Config,
pub ui: std::sync::Arc<dyn Ui>,
}
fn construct_ui(output: &str) -> Box<dyn Ui> {
UI_REGISTRY
.construct(
output,
output,
&serde_json::json!({}),
&crate::config::Config::default(),
)
.unwrap_or_else(|_| {
eprintln!("Unknown output format '{output}'. Valid: terminal, plain, json, markdown");
std::process::exit(1);
})
}
fn construct_context(
output: &str,
log_level: &str,
log_filters: &str,
log_json: bool,
log_thread_id: bool,
) -> AppContext {
let ui = construct_ui(output);
let ui: std::sync::Arc<dyn Ui> = std::sync::Arc::from(ui);
let formatter_kind = if log_json {
alog::FormatterKind::Json
} else {
alog::FormatterKind::Pretty
};
let ui_arc_clone = Arc::clone(&ui);
let ui_writer = UiWriter {
ui: Arc::clone(&ui),
};
alog::configure(alog::Config {
default_level: log_level.parse().unwrap(),
filters: alog::Filters::Spec(log_filters.to_string()),
formatter: alog::FormatterKind::Custom(Box::new(UiFormatter::new(
formatter_kind,
ui_arc_clone,
))),
writer: alog::Writer::Custom(Box::new(ui_writer)),
thread_id: log_thread_id,
});
alog!("MAIN", MessageLevel::Debug, "Welcome to granite-cli!");
let config = config::Config::new().unwrap_or_else(|e| {
ui.error(&format!("Failed to load config: {e}"));
std::process::exit(1);
});
AppContext { config, ui }
}
use std::io::{self, Write};
use std::sync::Arc;
struct UiWriter {
ui: Arc<dyn Ui>,
}
impl Write for UiWriter {
fn write(&mut self, buf: &[u8]) -> io::Result<usize> {
let text = std::str::from_utf8(buf).unwrap_or("");
for line in text.split('\n') {
if line.is_empty() {
continue;
}
self.ui.info(line);
}
Ok(buf.len())
}
fn flush(&mut self) -> io::Result<()> {
Ok(())
}
}
struct UiFormatter {
inner: Box<dyn alog::Formatter>,
ui: Arc<dyn Ui>,
}
impl UiFormatter {
fn new(kind: alog::FormatterKind, ui: Arc<dyn Ui>) -> Self {
Self {
inner: match kind {
alog::FormatterKind::Pretty => Box::new(alog::PrettyFormatter::default()),
alog::FormatterKind::Json => Box::new(alog::JsonFormatter),
alog::FormatterKind::Custom(c) => c,
},
ui,
}
}
}
impl alog::Formatter for UiFormatter {
fn format(&self, record: &alog::LogRecord<'_>) -> String {
let formatted = self.inner.format(record).trim_end_matches('\n').to_string();
match record.level {
MessageLevel::Fatal | MessageLevel::Error => self.ui.error_mark(&formatted),
MessageLevel::Warning => self.ui.warn_mark(&formatted),
MessageLevel::Info => formatted,
_ => self.ui.detail_mark(&formatted),
}
}
}
#[tokio::main]
async fn main() {
let cli = Cli::parse();
let log_level = cli.log_level.clone();
let log_filters = cli.log_filters.clone();
let log_json = cli.log_json;
let log_thread_id = cli.log_thread_id;
let command = cli.command;
let result: Result<(), ()> = match command {
Some(Commands::Model(wrapper)) => {
let mut ctx = construct_context(
&wrapper.output,
&log_level,
&log_filters,
log_json,
log_thread_id,
);
run_model_command(&mut ctx, wrapper.subcommand)
.await
.map_err(|e| ctx.ui.error(&e.to_string()))
}
Some(Commands::Capability(wrapper)) => {
let mut ctx = construct_context(
&wrapper.output,
&log_level,
&log_filters,
log_json,
log_thread_id,
);
run_capability_command(&mut ctx, wrapper.subcommand)
.await
.map_err(|e| ctx.ui.error(&e.to_string()))
}
Some(Commands::Provider(wrapper)) => {
let mut ctx = construct_context(
&wrapper.output,
&log_level,
&log_filters,
log_json,
log_thread_id,
);
run_provider_command(&mut ctx, wrapper.subcommand)
.await
.map_err(|e| ctx.ui.error(&e.to_string()))
}
Some(Commands::Hardware) => {
let ctx = construct_context(
"terminal",
&log_level,
&log_filters,
log_json,
log_thread_id,
);
HardwareCommands::show(&ctx).map_err(|e| ctx.ui.error(&e.to_string()))
}
Some(Commands::Launcher(wrapper)) => {
let mut ctx = construct_context(
&wrapper.output,
&log_level,
&log_filters,
log_json,
log_thread_id,
);
run_launcher_command(&mut ctx, wrapper.subcommand)
.await
.map_err(|e| ctx.ui.error(&e.to_string()))
}
Some(Commands::Launch(wrapper)) => {
let ctx = construct_context(
&wrapper.output,
&log_level,
&log_filters,
log_json,
log_thread_id,
);
run_launch(
&*ctx.ui,
&wrapper.launcher_id,
&wrapper.args,
wrapper.dry_run,
wrapper.usage_tracking,
)
.await
.map_err(|e| ctx.ui.error(&e.to_string()))
}
Some(Commands::Version) => {
let _ctx =
construct_context("warning", &log_level, &log_filters, log_json, log_thread_id);
println!("{}", version::version_string());
Ok(())
}
None => {
let ctx = construct_context(
"terminal",
&log_level,
&log_filters,
log_json,
log_thread_id,
);
run_interactive_tui(ctx)
.await
.map_err(|e| eprintln!("Error: {e}"))
}
};
if result.is_err() {
std::process::exit(1);
}
}
async fn run_model_command(ctx: &mut AppContext, subcmd: ModelSubcommands) -> anyhow::Result<()> {
match subcmd {
ModelSubcommands::Catalog { r#type } => {
let filter = match r#type.as_deref() {
Some("text") => Some(models::ModelType::Text),
Some("vision") => Some(models::ModelType::Vision),
Some("speech") => Some(models::ModelType::Speech),
Some("embedding") => Some(models::ModelType::Embedding),
Some(t) => {
anyhow::bail!(
"Unknown model type: {t}. Valid types: text, vision, speech, embedding"
);
}
None => None,
};
ModelCommands::catalog(ctx, filter)
}
ModelSubcommands::List { r#type } => {
let filter = match r#type.as_deref() {
Some("text") => Some(models::ModelType::Text),
Some("vision") => Some(models::ModelType::Vision),
Some("speech") => Some(models::ModelType::Speech),
Some("embedding") => Some(models::ModelType::Embedding),
Some(t) => {
anyhow::bail!(
"Unknown model type: {t}. Valid types: text, vision, speech, embedding"
);
}
None => None,
};
ModelCommands::list(ctx, filter)
}
ModelSubcommands::Search { query } => ModelCommands::search(ctx, &query),
ModelSubcommands::Recommend {
r#type,
providers,
wide,
} => {
let filter = match r#type.as_deref() {
Some("text") => Some(models::ModelType::Text),
Some("vision") => Some(models::ModelType::Vision),
Some("speech") => Some(models::ModelType::Speech),
Some("embedding") => Some(models::ModelType::Embedding),
Some(t) => {
anyhow::bail!(
"Unknown model type: {t}. Valid types: text, vision, speech, embedding"
);
}
None => None,
};
ModelCommands::recommend(ctx, filter, &providers, wide)
}
ModelSubcommands::Info { model_id } => ModelCommands::info(ctx, &model_id),
ModelSubcommands::Setup { model_id } => ModelCommands::setup(ctx, &model_id).await,
ModelSubcommands::Pull { model_id } => ModelCommands::pull(ctx, &model_id).await,
ModelSubcommands::Remove { model_id } => ModelCommands::remove(ctx, &model_id),
}
}
async fn run_capability_command(
ctx: &mut AppContext,
subcmd: CapabilitySubcommands,
) -> anyhow::Result<()> {
match subcmd {
CapabilitySubcommands::Catalog => CapabilityCommands::catalog(ctx),
CapabilitySubcommands::List => CapabilityCommands::list(ctx),
CapabilitySubcommands::Info { capability_id } => {
CapabilityCommands::info(ctx, &capability_id)
}
CapabilitySubcommands::Setup {
capability_type,
instance_id,
} => CapabilityCommands::setup(ctx, &capability_type, instance_id.as_deref()).await,
CapabilitySubcommands::Remove { capability_id } => {
CapabilityCommands::remove(ctx, &capability_id)
}
}
}
async fn run_provider_command(
ctx: &mut AppContext,
subcmd: ProviderSubcommands,
) -> anyhow::Result<()> {
match subcmd {
ProviderSubcommands::Catalog { wide } => ProviderCommands::catalog(ctx, wide),
ProviderSubcommands::List => ProviderCommands::list(ctx),
ProviderSubcommands::Setup {
provider_type,
instance_id,
} => ProviderCommands::setup(ctx, &provider_type, instance_id.as_deref()).await,
ProviderSubcommands::Health { provider_id } => {
ProviderCommands::health(ctx, provider_id.as_deref()).await
}
ProviderSubcommands::Remove { provider_id } => ProviderCommands::remove(ctx, &provider_id),
}
}
async fn run_launcher_command(
ctx: &mut AppContext,
subcmd: LauncherSubcommands,
) -> anyhow::Result<()> {
match subcmd {
LauncherSubcommands::Catalog => LauncherCommands::catalog(ctx),
LauncherSubcommands::List => LauncherCommands::list(ctx),
LauncherSubcommands::Setup {
launcher_type,
instance_id,
} => LauncherCommands::setup(ctx, &launcher_type, instance_id.as_deref()).await,
LauncherSubcommands::Remove { launcher_id } => LauncherCommands::remove(ctx, &launcher_id),
}
}
async fn run_launch(
ui: &dyn Ui,
launcher_id: &str,
args: &[String],
dry_run: bool,
usage_tracking: bool,
) -> anyhow::Result<()> {
use crate::capabilities::CAPABILITY_REGISTRY;
use crate::launchers::LAUNCHER_REGISTRY;
use crate::launchers::LaunchContext;
use crate::proxy::{ProxyServer, UsageTracker, UsageTrackingContext};
use std::sync::Mutex;
let mut config = crate::config::Config::new()?;
let lc = config
.get_launcher(launcher_id)
.ok_or_else(|| {
anyhow::anyhow!(
"No launcher configured with id '{launcher_id}'. \
Run `granite-cli launcher setup {launcher_id}` first."
)
})?
.clone();
let track_usage = usage_tracking && !dry_run;
let tracker = Arc::new(UsageTracker::new());
let proxy_servers: Arc<Mutex<Vec<ProxyServer>>> = Arc::new(Mutex::new(Vec::new()));
if track_usage {
config.usage_tracking = Some(UsageTrackingContext {
tracker: Arc::clone(&tracker),
servers: Arc::clone(&proxy_servers),
});
}
let mut launcher = LAUNCHER_REGISTRY
.construct(&lc.launcher_type, &lc.launcher_id, &lc.config, &config)
.map_err(|e| anyhow::anyhow!("Failed to construct launcher: {e}"))?;
let launch_ctx = LaunchContext {
launcher_id: launcher_id.to_string(),
working_dir: std::env::current_dir()?,
base_env: std::collections::HashMap::new(),
dry_run,
};
let mut bound_capabilities: Vec<Box<dyn crate::capabilities::Capability>> = Vec::new();
for cap_id in &lc.enabled_capabilities {
let cap_cfg = config.get_capability(cap_id).ok_or_else(|| {
anyhow::anyhow!(
"Launcher '{launcher_id}' references capability '{cap_id}' \
which is not configured. Run `granite-cli capability setup` first."
)
})?;
let capability = CAPABILITY_REGISTRY
.construct(
&cap_cfg.capability_type,
&cap_cfg.capability_id,
&cap_cfg.config,
&config,
)
.map_err(|e| anyhow::anyhow!("Failed to construct capability '{cap_id}': {e}"))?;
capability.on_setup().await?;
launcher.bind_capability(capability.as_ref()).await?;
bound_capabilities.push(capability);
}
for capability in &bound_capabilities {
capability.on_pre_launch(&launch_ctx).await?;
}
let launch_result = launcher.launch(args, &launch_ctx, ui).await;
for capability in bound_capabilities.iter().rev() {
if let Err(e) = capability.on_post_launch(&launch_ctx).await {
ui.warn(&format!(
"on_post_launch failed for capability '{}': {e}",
capability.instance_id()
));
}
if let Err(e) = capability.on_shutdown(&launch_ctx).await {
ui.warn(&format!(
"on_shutdown failed for capability '{}': {e}",
capability.instance_id()
));
}
}
let status = launch_result?;
if track_usage {
let started: Vec<ProxyServer> = proxy_servers.lock().unwrap().drain(..).collect();
for server in started {
server.shutdown().await;
}
print_usage_summary(ui, &tracker);
}
if !status.success() {
anyhow::bail!(
"'{}' exited with status {}",
launcher_id,
status.code().unwrap_or(-1)
);
}
Ok(())
}
fn print_usage_summary(ui: &dyn Ui, tracker: &proxy::UsageTracker) {
let snapshot = tracker.snapshot();
if snapshot.is_empty() {
return;
}
let mut rows: Vec<Vec<String>> = snapshot
.iter()
.map(|(label, s)| {
vec![
label.clone(),
s.requests.to_string(),
s.input_tokens.to_string(),
s.output_tokens.to_string(),
s.cache_creation_tokens.to_string(),
s.cache_read_tokens.to_string(),
]
})
.collect();
rows.sort_by(|a, b| a[0].cmp(&b[0]));
let total = snapshot
.values()
.fold(proxy::UsageStats::default(), |mut acc, s| {
acc.requests += s.requests;
acc.input_tokens += s.input_tokens;
acc.output_tokens += s.output_tokens;
acc.cache_creation_tokens += s.cache_creation_tokens;
acc.cache_read_tokens += s.cache_read_tokens;
acc
});
rows.push(vec![
"Total".to_string(),
total.requests.to_string(),
total.input_tokens.to_string(),
total.output_tokens.to_string(),
total.cache_creation_tokens.to_string(),
total.cache_read_tokens.to_string(),
]);
ui.table(
"Usage",
&[
"Binding",
"Requests",
"Input Tokens",
"Output Tokens",
"Cache Write",
"Cache Read",
],
&rows,
);
}