use std::collections::BTreeMap;
use std::path::{Path, PathBuf};
use serde::{Deserialize, Serialize};
use super::cluster::{ClusterConfig, DdpConfig, OutputConfig, TrainingConfig};
#[derive(Debug, Default, Deserialize)]
#[serde(deny_unknown_fields)]
pub struct ProjectConfig {
#[serde(default)]
pub description: Option<String>,
#[serde(default)]
pub commands: BTreeMap<String, CommandSpec>,
#[serde(default)]
pub cluster: Option<ClusterConfig>,
}
#[derive(Debug, Default, Deserialize)]
#[serde(deny_unknown_fields)]
pub struct CommandConfig {
#[serde(default)]
pub description: Option<String>,
#[serde(default)]
pub entry: Option<String>,
#[serde(default)]
pub docker: Option<String>,
#[serde(default)]
pub ddp: Option<DdpConfig>,
#[serde(default)]
pub training: Option<TrainingConfig>,
#[serde(default)]
pub output: Option<OutputConfig>,
#[serde(default)]
pub commands: BTreeMap<String, CommandSpec>,
#[serde(default, rename = "arg-name")]
pub arg_name: Option<String>,
#[serde(default)]
pub schema: Option<Schema>,
#[serde(default)]
pub compile: Option<bool>,
}
#[derive(Debug, Default, Clone)]
pub struct CommandSpec {
pub description: Option<String>,
pub run: Option<String>,
pub append: Option<String>,
pub path: Option<String>,
pub docker: Option<String>,
pub ddp: Option<DdpConfig>,
pub training: Option<TrainingConfig>,
pub output: Option<OutputConfig>,
pub options: BTreeMap<String, serde_json::Value>,
pub cluster: Option<bool>,
pub load_error: Option<String>,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum CommandKind {
Run,
Path,
Preset,
}
impl CommandSpec {
pub fn kind(&self) -> Result<CommandKind, String> {
if let Some(e) = &self.load_error {
return Err(e.clone());
}
if self.docker.is_some() && self.run.is_none() {
return Err(
"command declares `docker:` without `run:`; \
`docker:` only wraps inline run-scripts"
.to_string(),
);
}
if self.append.is_some() && self.run.is_none() {
return Err(
"command declares `append:` without `run:`; \
`append:` only forwards trailing tokens for inline run-scripts"
.to_string(),
);
}
match (self.run.as_deref(), self.path.as_deref()) {
(Some(_), Some(_)) => Err(
"command declares both `run:` and `path:`; \
only one is allowed"
.to_string(),
),
(Some(_), None) => Ok(CommandKind::Run),
(None, Some(_)) => Ok(CommandKind::Path),
(None, None) => {
if self.ddp.is_some()
|| self.training.is_some()
|| self.output.is_some()
|| !self.options.is_empty()
{
Ok(CommandKind::Preset)
} else {
Ok(CommandKind::Path)
}
}
}
}
pub fn resolve_path(&self, name: &str, parent_dir: &Path) -> PathBuf {
match &self.path {
Some(p) => parent_dir.join(p),
None => parent_dir.join(name),
}
}
}
impl<'de> Deserialize<'de> for CommandSpec {
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
where
D: serde::Deserializer<'de>,
{
#[derive(Deserialize)]
#[serde(deny_unknown_fields)]
struct Inner {
#[serde(default)]
description: Option<String>,
#[serde(default)]
run: Option<String>,
#[serde(default)]
append: Option<String>,
#[serde(default)]
path: Option<String>,
#[serde(default)]
docker: Option<String>,
#[serde(default)]
ddp: Option<DdpConfig>,
#[serde(default)]
training: Option<TrainingConfig>,
#[serde(default)]
output: Option<OutputConfig>,
#[serde(default)]
options: BTreeMap<String, serde_json::Value>,
#[serde(default)]
cluster: Option<bool>,
}
let raw = serde_yaml_ng::Value::deserialize(deserializer)?;
if matches!(raw, serde_yaml_ng::Value::Null) {
return Ok(Self::default());
}
let inner: Inner = match serde_yaml_ng::from_value(raw) {
Ok(inner) => inner,
Err(e) => {
return Ok(Self {
load_error: Some(e.to_string()),
..Self::default()
});
}
};
Ok(Self {
description: inner.description,
run: inner.run,
append: inner.append,
path: inner.path,
docker: inner.docker,
ddp: inner.ddp,
training: inner.training,
output: inner.output,
options: inner.options,
cluster: inner.cluster,
load_error: None,
})
}
}
#[derive(Debug, Clone, Default, Deserialize, Serialize)]
#[serde(deny_unknown_fields)]
pub struct Schema {
#[serde(default, skip_serializing_if = "Vec::is_empty")]
pub args: Vec<ArgSpec>,
#[serde(default, skip_serializing_if = "BTreeMap::is_empty")]
pub options: BTreeMap<String, OptionSpec>,
#[serde(default, skip_serializing_if = "is_false")]
pub strict: bool,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub description: Option<String>,
#[serde(default, skip_serializing_if = "BTreeMap::is_empty")]
pub commands: BTreeMap<String, Schema>,
}
#[derive(Debug, Clone, Deserialize, Serialize)]
#[serde(deny_unknown_fields)]
pub struct OptionSpec {
#[serde(rename = "type")]
pub ty: String,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub description: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub default: Option<serde_json::Value>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub choices: Option<Vec<serde_json::Value>>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub short: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub env: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
#[allow(dead_code)]
pub completer: Option<String>,
}
#[derive(Debug, Clone, Deserialize, Serialize)]
#[serde(deny_unknown_fields)]
pub struct ArgSpec {
pub name: String,
#[serde(rename = "type")]
pub ty: String,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub description: Option<String>,
#[serde(default = "default_required")]
pub required: bool,
#[serde(default, skip_serializing_if = "is_false")]
pub variadic: bool,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub default: Option<serde_json::Value>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub choices: Option<Vec<serde_json::Value>>,
#[serde(default, skip_serializing_if = "Option::is_none")]
#[allow(dead_code)]
pub completer: Option<String>,
}
fn is_false(b: &bool) -> bool {
!*b
}
fn default_required() -> bool {
true
}
const RESERVED_LONGS: &[&str] = &[
"help", "version", "quiet", "env",
];
const RESERVED_SHORTS: &[&str] = &[
"h", "V", "q", "v", "e",
];
const VALID_TYPES: &[&str] = &[
"string", "int", "float", "bool", "path",
"list[string]", "list[int]", "list[float]", "list[path]",
];
pub fn validate_schema(schema: &Schema) -> Result<(), String> {
if !schema.commands.is_empty() {
if !schema.args.is_empty() || !schema.options.is_empty() {
return Err(
"schema declares both `commands` (a subcommand tree) and \
top-level `args`/`options`; a node is either a leaf or a \
branch, not both — move the flags onto the subcommands"
.to_string(),
);
}
for (name, child) in &schema.commands {
if name.trim().is_empty() {
return Err("schema `commands` has an empty subcommand name".to_string());
}
validate_schema(child).map_err(|e| format!("subcommand `{name}`: {e}"))?;
}
return Ok(());
}
let mut short_seen: BTreeMap<String, String> = BTreeMap::new();
for (long, spec) in &schema.options {
if !VALID_TYPES.contains(&spec.ty.as_str()) {
return Err(format!(
"option --{}: unknown type '{}' (valid: {})",
long,
spec.ty,
VALID_TYPES.join(", ")
));
}
if RESERVED_LONGS.contains(&long.as_str()) {
return Err(format!(
"option --{long} shadows a reserved fdl-level flag"
));
}
if let Some(s) = &spec.short {
if s.chars().count() != 1 {
return Err(format!(
"option --{long}: `short: \"{s}\"` must be a single character"
));
}
if RESERVED_SHORTS.contains(&s.as_str()) {
return Err(format!(
"option --{long}: short -{s} shadows a reserved fdl-level flag"
));
}
if let Some(prev) = short_seen.insert(s.clone(), long.clone()) {
return Err(format!(
"options --{prev} and --{long} both declare short -{s}"
));
}
}
}
let mut seen_optional = false;
let mut name_seen: BTreeMap<String, ()> = BTreeMap::new();
for (i, arg) in schema.args.iter().enumerate() {
if !VALID_TYPES.contains(&arg.ty.as_str()) {
return Err(format!(
"arg <{}>: unknown type '{}' (valid: {})",
arg.name,
arg.ty,
VALID_TYPES.join(", ")
));
}
if name_seen.insert(arg.name.clone(), ()).is_some() {
return Err(format!("duplicate positional name <{}>", arg.name));
}
if arg.variadic && i != schema.args.len() - 1 {
return Err(format!(
"arg <{}>: variadic positional must be the last one",
arg.name
));
}
let is_optional = !arg.required || arg.default.is_some();
if arg.required && arg.default.is_some() {
return Err(format!(
"arg <{}>: `required: true` with a default is a contradiction",
arg.name
));
}
if seen_optional && arg.required && arg.default.is_none() {
return Err(format!(
"arg <{}>: required positional cannot follow an optional one",
arg.name
));
}
if is_optional {
seen_optional = true;
}
}
Ok(())
}