use std::env;
use crate::tensor::{Result, TensorError};
use super::ENV_FULL_CLUSTER_JSON;
#[derive(Debug, Clone, Default, PartialEq, Eq)]
pub struct SshConfig {
pub target: Option<String>,
pub port: Option<u16>,
pub user: Option<String>,
pub identity_file: Option<String>,
pub options: Vec<String>,
}
#[derive(Debug, Clone)]
pub struct FullCluster {
pub controller: FullController,
pub workers: Vec<FullWorker>,
pub salt: crate::distributed::wire::SessionSalt,
pub env: std::collections::BTreeMap<String, String>,
}
#[derive(Debug, Clone)]
pub struct FullController {
pub host: String,
pub port: u16,
pub path: String,
pub docker: Option<String>,
pub arch: Option<String>,
pub data_path: Option<String>,
pub join: Option<JoinKnobs>,
}
#[derive(Debug, Clone, Default, PartialEq, Eq)]
pub struct JoinKnobs {
pub min_rank_start: Option<usize>,
pub join_timeout_secs: Option<u64>,
pub target_ranks: Option<usize>,
pub max_join_timeout_secs: Option<u64>,
pub open_admission: Option<bool>,
}
impl FullCluster {
pub fn with_session_salt(mut self, salt: crate::distributed::wire::SessionSalt) -> Self {
self.salt = salt;
self
}
}
#[derive(Debug, Clone)]
pub struct FullWorker {
pub host: String,
pub ranks: Vec<usize>,
pub local_devices: Option<Vec<u8>>,
pub nccl_socket_ifname: String,
pub path: String,
pub arch: Option<String>,
pub ssh: Option<SshConfig>,
pub tunnel: bool,
pub env: std::collections::BTreeMap<String, String>,
}
impl FullWorker {
pub fn ssh_target(&self) -> &str {
self.ssh
.as_ref()
.and_then(|s| s.target.as_deref())
.unwrap_or(&self.host)
}
}
impl FullCluster {
pub fn from_env() -> Result<Self> {
let raw = env::var(ENV_FULL_CLUSTER_JSON).map_err(|e| {
TensorError::new(&format!(
"cluster launcher: reading {ENV_FULL_CLUSTER_JSON} failed: {e}"
))
})?;
let bytes = crate::distributed::cluster::hex_decode(raw.trim()).map_err(|e| {
TensorError::new(&format!(
"cluster launcher: {ENV_FULL_CLUSTER_JSON} hex-decode failed: {e}"
))
})?;
let val: serde_json::Value = serde_json::from_slice(&bytes).map_err(|e| {
TensorError::new(&format!(
"cluster launcher: {ENV_FULL_CLUSTER_JSON} JSON parse failed: {e}"
))
})?;
Self::from_value(&val)
}
pub fn from_value(val: &serde_json::Value) -> Result<Self> {
let obj = val.as_object().ok_or_else(|| {
TensorError::new("cluster launcher: top-level JSON must be an object")
})?;
let controller_val = obj
.get("controller")
.and_then(|v| v.as_object())
.ok_or_else(|| {
TensorError::new("cluster launcher: controller (object) required")
})?;
let controller_host = controller_val
.get("host")
.and_then(|v| v.as_str())
.ok_or_else(|| {
TensorError::new("cluster launcher: controller.host (string) required")
})?
.to_string();
if controller_host.trim().is_empty() {
return Err(TensorError::new(
"cluster launcher: controller.host must be non-empty",
));
}
let controller_port_u64 = controller_val
.get("port")
.and_then(|v| v.as_u64())
.ok_or_else(|| {
TensorError::new("cluster launcher: controller.port (u16) required")
})?;
let controller_port = u16::try_from(controller_port_u64).map_err(|_| {
TensorError::new(&format!(
"cluster launcher: controller.port must fit in u16 (got {controller_port_u64})"
))
})?;
let controller_path = controller_val
.get("path")
.and_then(|v| v.as_str())
.ok_or_else(|| {
TensorError::new("cluster launcher: controller.path (string) required")
})?
.to_string();
let controller_docker = controller_val
.get("docker")
.and_then(|v| v.as_str())
.map(String::from);
let controller_arch = controller_val
.get("arch")
.and_then(|v| v.as_str())
.map(String::from);
let controller_data_path = controller_val
.get("data_path")
.and_then(|v| v.as_str())
.map(String::from);
let controller_join = parse_join_block(controller_val.get("join"))?;
let workers_val = obj
.get("workers")
.and_then(|v| v.as_array())
.ok_or_else(|| TensorError::new("cluster launcher: workers (array) required"))?;
if workers_val.is_empty() {
return Err(TensorError::new(
"cluster launcher: workers must be non-empty",
));
}
let workers: Vec<FullWorker> = workers_val
.iter()
.enumerate()
.map(|(i, w)| parse_full_worker(w, i))
.collect::<Result<_>>()?;
let mut all: Vec<usize> = workers.iter().flat_map(|w| w.ranks.iter().copied()).collect();
let ws = all.len();
all.sort_unstable();
let expected: Vec<usize> = (0..ws).collect();
if all != expected {
return Err(TensorError::new(&format!(
"cluster launcher: ranks across workers must be exactly 0..{ws} \
with no duplicates or gaps, got sorted-unique sequence {all:?}"
)));
}
let env = parse_env_block(obj.get("env"), "cluster.env")?;
Ok(FullCluster {
controller: FullController {
host: controller_host,
port: controller_port,
path: controller_path,
docker: controller_docker,
arch: controller_arch,
data_path: controller_data_path,
join: controller_join,
},
workers,
salt: [0u8; crate::distributed::wire::SESSION_SALT_BYTES],
env,
})
}
pub fn world_size(&self) -> usize {
self.workers.iter().map(|w| w.ranks.len()).sum()
}
pub fn spans_multiple_workers(&self) -> bool {
self.workers.len() > 1
}
pub fn to_json(&self) -> serde_json::Value {
let workers: Vec<serde_json::Value> = self
.workers
.iter()
.map(|h| {
let mut o = serde_json::Map::new();
o.insert("host".into(), serde_json::Value::String(h.host.clone()));
o.insert(
"ranks".into(),
serde_json::Value::Array(
h.ranks.iter().map(|r| serde_json::Value::from(*r)).collect(),
),
);
let ld = match &h.local_devices {
None => serde_json::Value::String("all".into()),
Some(v) => serde_json::Value::Array(
v.iter().map(|d| serde_json::Value::from(*d)).collect(),
),
};
o.insert("local_devices".into(), ld);
o.insert(
"nccl_socket_ifname".into(),
serde_json::Value::String(h.nccl_socket_ifname.clone()),
);
o.insert("path".into(), serde_json::Value::String(h.path.clone()));
if let Some(a) = &h.arch {
o.insert("arch".into(), serde_json::Value::String(a.clone()));
}
if let Some(s) = &h.ssh {
let mut ssh_obj = serde_json::Map::new();
if let Some(t) = &s.target {
ssh_obj.insert("target".into(), serde_json::Value::String(t.clone()));
}
if let Some(p) = s.port {
ssh_obj.insert("port".into(), serde_json::Value::from(p));
}
if let Some(u) = &s.user {
ssh_obj.insert("user".into(), serde_json::Value::String(u.clone()));
}
if let Some(i) = &s.identity_file {
ssh_obj.insert(
"identity_file".into(),
serde_json::Value::String(i.clone()),
);
}
if !s.options.is_empty() {
ssh_obj.insert(
"options".into(),
serde_json::Value::Array(
s.options
.iter()
.map(|opt| serde_json::Value::String(opt.clone()))
.collect(),
),
);
}
o.insert("ssh".into(), serde_json::Value::Object(ssh_obj));
}
if h.tunnel {
o.insert("tunnel".into(), serde_json::Value::Bool(true));
}
if !h.env.is_empty() {
let mut env_obj = serde_json::Map::new();
for (k, v) in &h.env {
env_obj.insert(k.clone(), serde_json::Value::String(v.clone()));
}
o.insert("env".into(), serde_json::Value::Object(env_obj));
}
serde_json::Value::Object(o)
})
.collect();
let mut top = serde_json::Map::new();
let mut controller_obj = serde_json::Map::new();
controller_obj.insert(
"host".into(),
serde_json::Value::String(self.controller.host.clone()),
);
controller_obj.insert(
"port".into(),
serde_json::Value::from(self.controller.port),
);
controller_obj.insert(
"path".into(),
serde_json::Value::String(self.controller.path.clone()),
);
if let Some(s) = &self.controller.docker {
controller_obj.insert("docker".into(), serde_json::Value::String(s.clone()));
}
if let Some(s) = &self.controller.arch {
controller_obj.insert("arch".into(), serde_json::Value::String(s.clone()));
}
if let Some(s) = &self.controller.data_path {
controller_obj.insert("data_path".into(), serde_json::Value::String(s.clone()));
}
if let Some(j) = &self.controller.join {
let mut join_obj = serde_json::Map::new();
if let Some(n) = j.min_rank_start {
join_obj.insert("min_rank_start".into(), serde_json::Value::from(n));
}
if let Some(n) = j.join_timeout_secs {
join_obj.insert("join_timeout".into(), serde_json::Value::from(n));
}
if let Some(n) = j.target_ranks {
join_obj.insert("target_ranks".into(), serde_json::Value::from(n));
}
if let Some(n) = j.max_join_timeout_secs {
join_obj.insert("max_join_timeout".into(), serde_json::Value::from(n));
}
if let Some(b) = j.open_admission {
join_obj.insert("open_admission".into(), serde_json::Value::Bool(b));
}
if !join_obj.is_empty() {
controller_obj.insert("join".into(), serde_json::Value::Object(join_obj));
}
}
top.insert("controller".into(), serde_json::Value::Object(controller_obj));
top.insert("workers".into(), serde_json::Value::Array(workers));
if !self.env.is_empty() {
let mut env_obj = serde_json::Map::new();
for (k, v) in &self.env {
env_obj.insert(k.clone(), serde_json::Value::String(v.clone()));
}
top.insert("env".into(), serde_json::Value::Object(env_obj));
}
serde_json::Value::Object(top)
}
}
fn parse_full_worker(v: &serde_json::Value, i: usize) -> Result<FullWorker> {
let obj = v.as_object().ok_or_else(|| {
TensorError::new(&format!("cluster launcher: workers[{i}] must be an object"))
})?;
let host = obj
.get("host")
.and_then(|v| v.as_str())
.ok_or_else(|| {
TensorError::new(&format!(
"cluster launcher: workers[{i}].host (string) required"
))
})?
.to_string();
if host.trim().is_empty() {
return Err(TensorError::new(&format!(
"cluster launcher: workers[{i}].host must be non-empty"
)));
}
let name = host;
let ranks_arr = obj
.get("ranks")
.and_then(|v| v.as_array())
.ok_or_else(|| {
TensorError::new(&format!(
"cluster launcher: workers[{i}] ({name:?}): ranks (array) required"
))
})?;
let ranks: Vec<usize> = ranks_arr
.iter()
.enumerate()
.map(|(j, e)| {
let n = e.as_u64().ok_or_else(|| {
TensorError::new(&format!(
"cluster launcher: workers[{i}].ranks[{j}]: non-integer entry"
))
})?;
usize::try_from(n).map_err(|_| {
TensorError::new(&format!(
"cluster launcher: workers[{i}].ranks[{j}]: value {n} out of range"
))
})
})
.collect::<Result<_>>()?;
let local_devices = match obj.get("local_devices") {
None => {
return Err(TensorError::new(&format!(
"cluster launcher: workers[{i}] ({name:?}): local_devices required"
)));
}
Some(serde_json::Value::String(s)) if s == "all" => None,
Some(serde_json::Value::String(s)) => {
return Err(TensorError::new(&format!(
"cluster launcher: workers[{i}] ({name:?}): local_devices: \
expected \"all\" or array, got string {s:?}"
)));
}
Some(serde_json::Value::Array(arr)) => {
let v: Vec<u8> = arr
.iter()
.enumerate()
.map(|(j, e)| {
let n = e.as_u64().ok_or_else(|| {
TensorError::new(&format!(
"cluster launcher: workers[{i}].local_devices[{j}]: \
non-integer entry"
))
})?;
u8::try_from(n).map_err(|_| {
TensorError::new(&format!(
"cluster launcher: workers[{i}].local_devices[{j}]: \
value {n} does not fit in u8"
))
})
})
.collect::<Result<_>>()?;
if v.len() != ranks.len() {
return Err(TensorError::new(&format!(
"cluster launcher: workers[{i}] ({name:?}): ranks ({}) and \
local_devices ({}) length mismatch",
ranks.len(),
v.len()
)));
}
Some(v)
}
Some(other) => {
return Err(TensorError::new(&format!(
"cluster launcher: workers[{i}] ({name:?}): local_devices: \
expected \"all\" or array, got {other}"
)));
}
};
let nccl_socket_ifname = obj
.get("nccl_socket_ifname")
.and_then(|v| v.as_str())
.ok_or_else(|| {
TensorError::new(&format!(
"cluster launcher: workers[{i}] ({name:?}): nccl_socket_ifname (string) required"
))
})?
.to_string();
let path = obj
.get("path")
.and_then(|v| v.as_str())
.ok_or_else(|| {
TensorError::new(&format!(
"cluster launcher: workers[{i}] ({name:?}): path (string) required"
))
})?
.to_string();
if path.trim().is_empty() {
return Err(TensorError::new(&format!(
"cluster launcher: workers[{i}] ({name:?}): path must be non-empty"
)));
}
let arch = obj
.get("arch")
.and_then(|v| v.as_str())
.map(String::from);
let ssh = parse_ssh_block(
obj.get("ssh"),
&format!("workers[{i}] ({name:?})"),
)?;
let tunnel = match obj.get("tunnel") {
None | Some(serde_json::Value::Null) => false,
Some(serde_json::Value::Bool(b)) => *b,
Some(other) => {
return Err(TensorError::new(&format!(
"cluster launcher: workers[{i}] ({name:?}): tunnel must be a \
boolean, got {other}"
)));
}
};
let env = parse_env_block(
obj.get("env"),
&format!("workers[{i}] ({name:?}).env"),
)?;
Ok(FullWorker {
host: name,
ranks,
local_devices,
nccl_socket_ifname,
path,
arch,
ssh,
tunnel,
env,
})
}
fn parse_join_block(v: Option<&serde_json::Value>) -> Result<Option<JoinKnobs>> {
let obj = match v {
None | Some(serde_json::Value::Null) => return Ok(None),
Some(serde_json::Value::Object(m)) => m,
Some(other) => {
return Err(TensorError::new(&format!(
"cluster launcher: controller.join must be a map \
(min_rank_start, join_timeout, target_ranks, \
max_join_timeout, open_admission), got {other}"
)));
}
};
let uint = |key: &str| -> Result<Option<u64>> {
match obj.get(key) {
None | Some(serde_json::Value::Null) => Ok(None),
Some(v) => match v.as_u64() {
Some(n) => Ok(Some(n)),
None => Err(TensorError::new(&format!(
"cluster launcher: controller.join.{key} must be a \
non-negative integer, got {v}"
))),
},
}
};
let open_admission = match obj.get("open_admission") {
None | Some(serde_json::Value::Null) => None,
Some(serde_json::Value::Bool(b)) => Some(*b),
Some(other) => {
return Err(TensorError::new(&format!(
"cluster launcher: controller.join.open_admission must be a \
boolean, got {other}"
)));
}
};
const KNOWN: [&str; 5] = [
"min_rank_start",
"join_timeout",
"target_ranks",
"max_join_timeout",
"open_admission",
];
for k in obj.keys() {
if !KNOWN.contains(&k.as_str()) {
return Err(TensorError::new(&format!(
"cluster launcher: controller.join.{k}: unknown field \
(expected one of {KNOWN:?})"
)));
}
}
Ok(Some(JoinKnobs {
min_rank_start: uint("min_rank_start")?.map(|n| n as usize),
join_timeout_secs: uint("join_timeout")?,
target_ranks: uint("target_ranks")?.map(|n| n as usize),
max_join_timeout_secs: uint("max_join_timeout")?,
open_admission,
}))
}
fn parse_ssh_block(
v: Option<&serde_json::Value>,
label: &str,
) -> Result<Option<SshConfig>> {
let obj = match v {
None | Some(serde_json::Value::Null) => return Ok(None),
Some(serde_json::Value::Object(m)) => m,
Some(other) => {
return Err(TensorError::new(&format!(
"cluster launcher: {label}.ssh must be a map (target, port, \
user, identity_file, options), got {other}"
)));
}
};
let target = obj
.get("target")
.and_then(|v| v.as_str())
.map(String::from);
let port = match obj.get("port") {
None | Some(serde_json::Value::Null) => None,
Some(v) => {
let n = v.as_u64().ok_or_else(|| {
TensorError::new(&format!(
"cluster launcher: {label}.ssh.port must be integer"
))
})?;
Some(u16::try_from(n).map_err(|_| {
TensorError::new(&format!(
"cluster launcher: {label}.ssh.port {n} does not fit in u16"
))
})?)
}
};
let user = obj
.get("user")
.and_then(|v| v.as_str())
.map(String::from);
let identity_file = obj
.get("identity_file")
.and_then(|v| v.as_str())
.map(String::from);
let options: Vec<String> = match obj.get("options") {
None | Some(serde_json::Value::Null) => Vec::new(),
Some(serde_json::Value::Array(arr)) => arr
.iter()
.enumerate()
.map(|(j, e)| {
e.as_str().map(String::from).ok_or_else(|| {
TensorError::new(&format!(
"cluster launcher: {label}.ssh.options[{j}]: must be string"
))
})
})
.collect::<Result<_>>()?,
Some(other) => {
return Err(TensorError::new(&format!(
"cluster launcher: {label}.ssh.options must be array of strings, got {other}"
)));
}
};
Ok(Some(SshConfig { target, port, user, identity_file, options }))
}
fn parse_env_block(
v: Option<&serde_json::Value>,
label: &str,
) -> Result<std::collections::BTreeMap<String, String>> {
use std::collections::BTreeMap;
match v {
None | Some(serde_json::Value::Null) => Ok(BTreeMap::new()),
Some(serde_json::Value::Object(map)) => {
let mut out = BTreeMap::new();
for (k, val) in map {
if crate::distributed::cluster::is_reserved_cluster_env_key(k) {
return Err(TensorError::new(&format!(
"cluster launcher: {label}[{k:?}] is reserved \
(launcher-owned rank identity). GPU scoping belongs \
in `local_devices:`; FLODL_INTERNAL_* vars are set \
by the launcher itself."
)));
}
let valid_key = !k.is_empty()
&& k.chars()
.next()
.is_some_and(|c| c.is_ascii_alphabetic() || c == '_')
&& k.chars().all(|c| c.is_ascii_alphanumeric() || c == '_');
if !valid_key {
return Err(TensorError::new(&format!(
"cluster launcher: {label}[{k:?}] is not a valid env var \
name ([A-Za-z_][A-Za-z0-9_]*)"
)));
}
let s = val.as_str().ok_or_else(|| {
TensorError::new(&format!(
"cluster launcher: {label}[{k:?}] must be a string, got {val}"
))
})?;
out.insert(k.clone(), s.to_string());
}
Ok(out)
}
Some(other) => Err(TensorError::new(&format!(
"cluster launcher: {label} must be an object (NAME → string VALUE), \
got {other}"
))),
}
}