use std::ffi::OsString;
use std::path::PathBuf;
use anyhow::Result;
use clap::{Args, CommandFactory, FromArgMatches, Parser, Subcommand, ValueEnum};
use super::trace_cli::{self, QueryLabelFilter, QueryOptions, TraceQueryKind, QUERY_MAX_LIMIT};
#[derive(Parser, Debug)]
#[command(
name = "candle-graph",
version,
about = "Import and visualize Candle execution traces",
long_about = "Capability-qualified evidence and atomic bundles from candle-graph/trace/10 runs."
)]
pub struct Cli {
#[command(subcommand)]
command: Command,
}
impl Cli {
pub fn run(self) -> Result<()> {
match self.command {
Command::Import(import) => {
trace_cli::run_import(&import.trace, import.output.as_deref())
}
#[cfg(feature = "visualizer")]
Command::View(view) => {
trace_cli::run_view(&view.trace, &view.output, view.nsight_dir.as_deref())
}
Command::Summary(summary) => trace_cli::run_summary(
&summary.input,
summary.output.as_deref(),
summary.require_valid,
),
Command::Query(query) => {
let options = query.options();
trace_cli::run_query(
&query.input,
query.kind.into(),
&options,
query.output.as_deref(),
)
}
Command::Overview(overview) => {
trace_cli::run_overview(&overview.input, overview.output.as_deref())
}
Command::Compare(compare) => trace_cli::run_compare(
&compare.baseline,
&compare.candidate,
compare.unverified_traces,
compare.require_eligible,
compare.output.as_deref(),
),
Command::Report(report) => trace_cli::run_report(
&report.trace,
report.nsight_dir.as_deref(),
&report.bundle,
report.output.as_deref(),
),
Command::Verify(verify) => {
trace_cli::run_verify(&verify.bundle, verify.semantic, verify.output.as_deref())
}
Command::Protocol(protocol) => trace_cli::run_protocol(protocol.output.as_deref()),
Command::CampaignStatus(status) => {
trace_cli::run_campaign_status(&status.manifest, status.output.as_deref())
}
Command::Series(series) => trace_cli::run_series(
series.manifest.as_deref(),
&series.bundle,
series.label_prefix.as_deref(),
series.output.as_deref(),
),
}
}
pub fn parse_as_cargo_subcommand() -> Self {
let mut argv: Vec<OsString> = std::env::args_os().collect();
if argv
.get(1)
.is_some_and(|argument| argument == "candle-graph")
{
argv.remove(1);
}
let command = Self::command()
.name("cargo-candle-graph")
.bin_name("cargo candle-graph")
.about("Import and analyze candle-graph execution trace files");
let matches = command.get_matches_from(argv);
match Self::from_arg_matches(&matches) {
Ok(cli) => cli,
Err(error) => error.exit(),
}
}
}
pub fn command_catalog() -> Vec<serde_json::Value> {
Cli::command()
.get_subcommands()
.map(|subcommand| {
serde_json::json!({
"name": subcommand.get_name(),
"about": subcommand.get_about().map(|about| about.to_string()),
})
})
.collect()
}
#[derive(Subcommand, Debug)]
enum Command {
Import(ImportArgs),
#[cfg(feature = "visualizer")]
View(ViewArgs),
Summary(SummaryArgs),
Query(QueryArgs),
Overview(OverviewArgs),
Compare(CompareArgs),
Report(ReportArgs),
Verify(VerifyArgs),
Protocol(ProtocolArgs),
CampaignStatus(CampaignStatusArgs),
Series(SeriesArgs),
}
#[derive(Args, Debug)]
struct ImportArgs {
#[arg(value_name = "INPUT")]
trace: PathBuf,
#[arg(long, short, value_name = "FILE")]
output: Option<PathBuf>,
}
#[cfg(feature = "visualizer")]
#[derive(Args, Debug)]
struct ViewArgs {
#[arg(value_name = "TRACE")]
trace: PathBuf,
#[arg(long, value_name = "FILE")]
output: PathBuf,
#[arg(long, value_name = "DIR")]
nsight_dir: Option<PathBuf>,
}
#[derive(Args, Debug)]
struct SummaryArgs {
#[arg(value_name = "INPUT")]
input: PathBuf,
#[arg(long)]
require_valid: bool,
#[arg(long, short, value_name = "FILE")]
output: Option<PathBuf>,
}
#[derive(Args, Debug)]
struct QueryArgs {
#[arg(value_name = "INPUT")]
input: PathBuf,
#[arg(long, value_enum)]
kind: CliTraceQueryKind,
#[arg(long, value_name = "S", conflicts_with = "label_prefix")]
label: Option<String>,
#[arg(long, value_name = "S")]
label_prefix: Option<String>,
#[arg(
long,
value_name = "N",
value_parser = parse_query_limit,
conflicts_with = "all"
)]
limit: Option<usize>,
#[arg(long, value_name = "N", conflicts_with = "all")]
offset: Option<usize>,
#[arg(long, conflicts_with_all = ["limit", "offset"])]
all: bool,
#[arg(long, short, value_name = "FILE")]
output: Option<PathBuf>,
}
impl QueryArgs {
fn filter(&self) -> Option<QueryLabelFilter> {
self.label
.clone()
.map(QueryLabelFilter::Exact)
.or_else(|| self.label_prefix.clone().map(QueryLabelFilter::Prefix))
}
fn options(&self) -> QueryOptions {
QueryOptions {
filter: self.filter(),
limit: self.limit,
offset: self.offset,
all: self.all,
}
}
}
fn parse_query_limit(value: &str) -> std::result::Result<usize, String> {
let limit = value
.parse::<usize>()
.map_err(|_| format!("limit must be an integer in 1..={QUERY_MAX_LIMIT}"))?;
if !(1..=QUERY_MAX_LIMIT).contains(&limit) {
return Err(format!("limit must be in 1..={QUERY_MAX_LIMIT}"));
}
Ok(limit)
}
#[derive(Args, Debug)]
struct OverviewArgs {
#[arg(value_name = "INPUT")]
input: PathBuf,
#[arg(long, short, value_name = "FILE")]
output: Option<PathBuf>,
}
#[derive(Args, Debug)]
struct CompareArgs {
#[arg(long, required = true, num_args = 1.., value_name = "BUNDLE")]
baseline: Vec<PathBuf>,
#[arg(long, required = true, num_args = 1.., value_name = "BUNDLE")]
candidate: Vec<PathBuf>,
#[arg(long)]
unverified_traces: bool,
#[arg(long)]
require_eligible: bool,
#[arg(long, short, value_name = "FILE")]
output: Option<PathBuf>,
}
#[derive(Args, Debug)]
struct ReportArgs {
#[arg(value_name = "TRACE")]
trace: PathBuf,
#[arg(long, value_name = "DIR")]
nsight_dir: Option<PathBuf>,
#[arg(long, value_name = "DIR")]
bundle: PathBuf,
#[arg(long, short, value_name = "FILE")]
output: Option<PathBuf>,
}
#[derive(Args, Debug)]
struct VerifyArgs {
#[arg(value_name = "BUNDLE")]
bundle: PathBuf,
#[arg(long)]
semantic: bool,
#[arg(long, short, value_name = "FILE")]
output: Option<PathBuf>,
}
#[derive(Args, Debug)]
struct ProtocolArgs {
#[arg(long, short, value_name = "FILE")]
output: Option<PathBuf>,
}
#[derive(Args, Debug)]
struct CampaignStatusArgs {
#[arg(long, value_name = "FILE")]
manifest: PathBuf,
#[arg(long, short, value_name = "FILE")]
output: Option<PathBuf>,
}
#[derive(Args, Debug)]
struct SeriesArgs {
#[arg(
long,
value_name = "FILE",
conflicts_with = "bundle",
required_unless_present = "bundle"
)]
manifest: Option<PathBuf>,
#[arg(long, value_name = "DIR", num_args = 1..)]
bundle: Vec<PathBuf>,
#[arg(long, value_name = "P")]
label_prefix: Option<String>,
#[arg(long, short, value_name = "FILE")]
output: Option<PathBuf>,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, ValueEnum)]
enum CliTraceQueryKind {
Labels,
SlowestHost,
SlowestDevice,
Activations,
Heaviest,
Memory,
Spans,
Tensors,
TensorStats,
Gradients,
Capabilities,
GpuStatus,
GpuCorrelation,
GpuPhases,
GpuKernels,
GpuAttributionGaps,
}
impl From<CliTraceQueryKind> for TraceQueryKind {
fn from(kind: CliTraceQueryKind) -> Self {
match kind {
CliTraceQueryKind::Labels => Self::Labels,
CliTraceQueryKind::SlowestHost => Self::SlowestHost,
CliTraceQueryKind::SlowestDevice => Self::SlowestDevice,
CliTraceQueryKind::Activations => Self::Activations,
CliTraceQueryKind::Heaviest => Self::Heaviest,
CliTraceQueryKind::Memory => Self::Memory,
CliTraceQueryKind::Spans => Self::Spans,
CliTraceQueryKind::Tensors => Self::Tensors,
CliTraceQueryKind::TensorStats => Self::TensorStats,
CliTraceQueryKind::Gradients => Self::Gradients,
CliTraceQueryKind::Capabilities => Self::Capabilities,
CliTraceQueryKind::GpuStatus => Self::GpuStatus,
CliTraceQueryKind::GpuCorrelation => Self::GpuCorrelation,
CliTraceQueryKind::GpuPhases => Self::GpuPhases,
CliTraceQueryKind::GpuKernels => Self::GpuKernels,
CliTraceQueryKind::GpuAttributionGaps => Self::GpuAttributionGaps,
}
}
}