1use clap::{ArgAction, Parser, ValueEnum};
2use serde::Deserialize;
3
4const DEFAULT_INTERVAL_SECS: u64 = 1;
5const RENICE_MIN: i64 = -20;
6const RENICE_MAX: i64 = 19;
7const DEFAULT_CONFIG_FILE: &str = "resource-tracker.toml";
8
9#[derive(Debug, Clone, Copy, PartialEq, ValueEnum)]
15pub enum OutputFormat {
16 Json,
18 Csv,
21}
22
23#[derive(Debug, Default, Deserialize)]
28struct TomlConfig {
29 job: Option<TomlJob>,
30 tracker: Option<TomlTracker>,
31}
32
33#[derive(Debug, Deserialize)]
34struct TomlJob {
35 name: Option<String>,
37 pid: Option<i32>,
39}
40
41#[derive(Debug, Deserialize)]
42struct TomlTracker {
43 interval_secs: Option<u64>,
45 renice: Option<i32>,
47 aggregate_cpu_steal: Option<bool>,
49}
50
51#[derive(Debug, Clone, Default)]
59pub struct JobMetadata {
60 pub project_name: Option<String>,
61 pub job_name: Option<String>,
62 pub stage_name: Option<String>,
63 pub task_name: Option<String>,
64 pub team: Option<String>,
65 pub env: Option<String>,
66 pub language: Option<String>,
67 pub orchestrator: Option<String>,
68 pub executor: Option<String>,
69 pub external_run_id: Option<String>,
70 pub container_image: Option<String>,
71 pub tags: Vec<String>,
73 pub command: Vec<String>,
76}
77
78#[derive(Debug, Parser)]
83#[command(
84 name = "resource-tracker",
85 about = "Lightweight Linux resource & GPU tracker.\n\n\
86 Shell-wrapper mode: resource-tracker [FLAGS] -- <command> [args...]\n\
87 The tracker will spawn <command>, monitor it, and exit when it exits.",
88 version
89)]
90struct Cli {
91 #[arg(short = 'p', long, value_name = "PID")]
95 pid: Option<i32>,
96
97 #[arg(short = 'i', long, value_name = "SECS")]
99 interval: Option<u64>,
100
101 #[arg(
105 short = 'r',
106 long = "renice",
107 value_name = "VALUE",
108 env = "TRACKER_RENICE",
109 num_args = 0..=1,
110 default_missing_value = "19",
111 value_parser = clap::value_parser!(i32).range(RENICE_MIN..=RENICE_MAX),
112 verbatim_doc_comment,
113 )]
114 renice: Option<i32>,
115
116 #[arg(
118 long = "aggregate-cpu-steal",
119 value_name = "AGGREGATE_CPU_STEAL",
120 env = "TRACKER_AGGREGATE_CPU_STEAL",
121 default_missing_value = "true"
122 )]
123 aggregate_cpu_steal: Option<bool>,
124
125 #[arg(short = 'c', long, value_name = "FILE", default_value = DEFAULT_CONFIG_FILE)]
127 config: String,
128
129 #[arg(short = 'f', long, value_name = "FORMAT", default_value = "json")]
131 format: OutputFormat,
132
133 #[arg(short = 'o', long, value_name = "FILE", env = "TRACKER_OUTPUT")]
136 output: Option<String>,
137
138 #[arg(long, env = "TRACKER_QUIET")]
141 quiet: bool,
142
143 #[arg(long, value_name = "NAME", env = "TRACKER_PROJECT_NAME")]
146 project_name: Option<String>,
147
148 #[arg(short = 'n', long, value_name = "NAME", env = "TRACKER_JOB_NAME")]
150 job_name: Option<String>,
151
152 #[arg(long, value_name = "NAME", env = "TRACKER_STAGE_NAME")]
154 stage_name: Option<String>,
155
156 #[arg(long, value_name = "NAME", env = "TRACKER_TASK_NAME")]
158 task_name: Option<String>,
159
160 #[arg(long, value_name = "NAME", env = "TRACKER_TEAM")]
162 team: Option<String>,
163
164 #[arg(long, value_name = "ENV", env = "TRACKER_ENV")]
166 env: Option<String>,
167
168 #[arg(long, value_name = "LANG", env = "TRACKER_LANGUAGE")]
170 language: Option<String>,
171
172 #[arg(long, value_name = "NAME", env = "TRACKER_ORCHESTRATOR")]
174 orchestrator: Option<String>,
175
176 #[arg(long, value_name = "NAME", env = "TRACKER_EXECUTOR")]
178 executor: Option<String>,
179
180 #[arg(long, value_name = "ID", env = "TRACKER_EXTERNAL_RUN_ID")]
182 external_run_id: Option<String>,
183
184 #[arg(long, value_name = "IMAGE", env = "TRACKER_CONTAINER_IMAGE")]
186 container_image: Option<String>,
187
188 #[arg(long = "tag", value_name = "KEY=VALUE", action = ArgAction::Append)]
190 tags: Vec<String>,
191
192 #[arg(
196 trailing_var_arg = true,
197 allow_hyphen_values = true,
198 value_name = "COMMAND"
199 )]
200 command: Vec<String>,
201}
202
203#[derive(Debug, Clone)]
209pub struct Config {
210 pub pid: Option<i32>,
213 pub interval_secs: u64,
215 pub renice: Option<i32>,
217 pub aggregate_cpu_steal: bool,
219 pub format: OutputFormat,
221 pub output_file: Option<String>,
224 pub quiet: bool,
226 pub metadata: JobMetadata,
228 pub command: Vec<String>,
230}
231
232impl Config {
233 pub fn load() -> Self {
236 let cli = Cli::parse();
237
238 let toml: TomlConfig = std::fs::read_to_string(&cli.config)
240 .ok()
241 .and_then(|s| toml::from_str(&s).ok())
242 .unwrap_or_default();
243
244 let interval_secs = cli
245 .interval
246 .or_else(|| toml.tracker.as_ref().and_then(|t| t.interval_secs))
247 .unwrap_or(DEFAULT_INTERVAL_SECS);
248
249 if interval_secs == 0 {
250 eprintln!("error: --interval must be >= 1 (got 0)");
251 std::process::exit(1);
252 }
253
254 let renice = cli
255 .renice
256 .or_else(|| toml.tracker.as_ref().and_then(|t| t.renice));
257
258 let aggregate_cpu_steal = cli
259 .aggregate_cpu_steal
260 .or_else(|| toml.tracker.as_ref().and_then(|t| t.aggregate_cpu_steal))
261 .unwrap_or(true);
262
263 let pid = cli.pid.or_else(|| toml.job.as_ref().and_then(|j| j.pid));
264
265 let metadata = JobMetadata {
266 project_name: cli.project_name,
267 job_name: cli
268 .job_name
269 .or_else(|| toml.job.as_ref().and_then(|j| j.name.clone())),
270 stage_name: cli.stage_name,
271 task_name: cli.task_name,
272 team: cli.team,
273 env: cli.env,
274 language: cli.language,
275 orchestrator: cli.orchestrator,
276 executor: cli.executor,
277 external_run_id: cli.external_run_id,
278 container_image: cli.container_image,
279 tags: cli.tags,
280 command: cli.command.clone(),
281 };
282
283 Config {
284 pid,
285 interval_secs,
286 renice,
287 aggregate_cpu_steal,
288 format: cli.format,
289 output_file: cli.output,
290 quiet: cli.quiet,
291 metadata,
292 command: cli.command,
293 }
294 }
295}
296
297#[cfg(test)]
302mod tests {
303 use super::*;
304
305 #[test]
307 fn test_toml_config_deserializes() {
308 let toml_str = r#"
309[job]
310name = "benchmark"
311pid = 12345
312
313[tracker]
314interval_secs = 5
315"#;
316 let cfg: TomlConfig = toml::from_str(toml_str).expect("TOML parse failed");
317 let job = cfg.job.as_ref().expect("job section missing");
318 assert_eq!(job.name.as_deref(), Some("benchmark"));
319 assert_eq!(job.pid, Some(12345));
320 let tracker = cfg.tracker.as_ref().expect("tracker section missing");
321 assert_eq!(tracker.interval_secs, Some(5));
322 }
323
324 #[test]
326 fn test_toml_config_default_is_all_none() {
327 let cfg = TomlConfig::default();
328 assert!(cfg.job.is_none(), "job must be None in default TomlConfig");
329 assert!(
330 cfg.tracker.is_none(),
331 "tracker must be None in default TomlConfig"
332 );
333 }
334
335 #[test]
337 fn test_job_metadata_default_all_none() {
338 let m = JobMetadata::default();
339 assert!(m.project_name.is_none());
340 assert!(m.job_name.is_none());
341 assert!(m.stage_name.is_none());
342 assert!(m.task_name.is_none());
343 assert!(m.team.is_none());
344 assert!(m.env.is_none());
345 assert!(m.language.is_none());
346 assert!(m.orchestrator.is_none());
347 assert!(m.executor.is_none());
348 assert!(m.external_run_id.is_none());
349 assert!(m.container_image.is_none());
350 assert!(
351 m.tags.is_empty(),
352 "tags must be empty in default JobMetadata"
353 );
354 }
355
356 #[test]
358 fn test_output_format_equality() {
359 assert_eq!(OutputFormat::Json, OutputFormat::Json);
360 assert_eq!(OutputFormat::Csv, OutputFormat::Csv);
361 assert_ne!(OutputFormat::Json, OutputFormat::Csv);
362 }
363
364 #[test]
366 fn test_toml_config_ignores_unknown_keys() {
367 let toml_str = r#"
368[job]
369name = "run1"
370unknown_field = "ignored"
371"#;
372 let result: Result<TomlConfig, _> = toml::from_str(toml_str);
374 let _ = result; }
379}