use std::io::{Read, Write};
use std::net::{TcpStream, ToSocketAddrs};
use std::time::Duration;
use crate::config::{self, DEFAULT_CONTROLLER_PORT};
use crate::context::Context;
use crate::style;
const HTTP_TIMEOUT: Duration = Duration::from_secs(5);
pub fn run(json: bool, addr_override: Option<&str>) -> i32 {
let (candidates, origin) = resolve_candidates(addr_override);
let mut last_err = String::new();
for addr in &candidates {
match fetch_state(addr) {
Ok(body) => {
if json {
println!("{body}");
} else {
print_status(addr, &body);
}
return 0;
}
Err(e) => last_err = e,
}
}
crate::cli_error!(
"no cluster run listening at {} ({origin}): {last_err}\n\
The status endpoint lives on the controller's training port for \
exactly as long as the run does — a refused connection usually \
just means no run is up.",
candidates.join(" / "),
);
1
}
fn resolve_candidates(addr_override: Option<&str>) -> (Vec<String>, String) {
if let Some(addr) = addr_override {
let addr = if addr.contains(':') {
addr.to_string()
} else {
format!("{addr}:{DEFAULT_CONTROLLER_PORT}")
};
return (vec![addr], "--addr".to_string());
}
if let Ok(env_name) = std::env::var("FDL_ENV") {
if let Some(cluster) = load_cluster_for_env(&env_name) {
let host = cluster.controller.host.clone();
let port = cluster.controller.port;
let mut candidates = vec![format!("{host}:{port}")];
if host != "127.0.0.1" && host != "localhost" {
candidates.push(format!("127.0.0.1:{port}"));
}
return (candidates, format!("fdl.{env_name}.yml controller"));
}
}
eprintln!(
"{}",
style::dim(&format!(
"fdl status: no cluster env active; trying \
127.0.0.1:{DEFAULT_CONTROLLER_PORT} (pass --addr or use \
`fdl @<env> status` to target a specific controller)"
)),
);
(
vec![format!("127.0.0.1:{DEFAULT_CONTROLLER_PORT}")],
"convention default".to_string(),
)
}
fn load_cluster_for_env(env_name: &str) -> Option<config::ClusterConfig> {
let ctx = Context::resolve();
let config_path = config::find_config(&ctx.root)?;
let project = config::load_project_with_env(&config_path, Some(env_name)).ok()?;
project.cluster
}
fn fetch_state(addr: &str) -> Result<String, String> {
let sock_addr = addr
.to_socket_addrs()
.map_err(|e| format!("cannot resolve {addr}: {e}"))?
.next()
.ok_or_else(|| format!("cannot resolve {addr}"))?;
let mut stream = TcpStream::connect_timeout(&sock_addr, HTTP_TIMEOUT)
.map_err(|e| format!("connect: {e}"))?;
stream
.set_read_timeout(Some(HTTP_TIMEOUT))
.and_then(|()| stream.set_write_timeout(Some(HTTP_TIMEOUT)))
.map_err(|e| format!("socket setup: {e}"))?;
stream
.write_all(
format!(
"GET /state.json HTTP/1.1\r\nHost: {addr}\r\n\
Connection: close\r\n\r\n"
)
.as_bytes(),
)
.map_err(|e| format!("send request: {e}"))?;
let mut response = String::new();
stream
.read_to_string(&mut response)
.map_err(|e| format!("read response: {e}"))?;
let (head, body) = response
.split_once("\r\n\r\n")
.ok_or_else(|| "malformed HTTP response".to_string())?;
let status_line = head.lines().next().unwrap_or_default();
if !status_line.contains(" 200 ") {
return Err(format!(
"endpoint answered {} — {}",
status_line.trim_start_matches("HTTP/1.1 "),
body.trim(),
));
}
Ok(body.trim().to_string())
}
fn print_status(addr: &str, body: &str) {
let state: serde_json::Value = match serde_json::from_str(body) {
Ok(v) => v,
Err(_) => {
println!("{body}");
return;
}
};
let phase = state["phase"].as_str().unwrap_or("unknown");
let painted_phase = match phase {
"training" | "done" => style::green(phase),
"waiting" | "forming" => style::yellow(phase),
"failed" => style::red(phase),
other => other.to_string(),
};
println!(
"cluster run @ {addr} — {}",
style::bold(&painted_phase),
);
let joined_ranks = state["joined_ranks"].as_u64().unwrap_or(0);
let joined_hosts = state["joined_hosts"].as_u64().unwrap_or(0);
let quorum = state["min_rank_start"].as_u64().unwrap_or(0);
let target = state["target_ranks"]
.as_u64()
.map(|t| t.to_string())
.unwrap_or_else(|| "none".to_string());
println!(
" ranks: {joined_ranks} joined across {joined_hosts} host(s) \
(quorum {quorum}, target {target})",
);
if phase == "waiting" {
let fmt_remaining = |v: &serde_json::Value| match v.as_u64() {
Some(s) => format!("{s}s left"),
None => "expired".to_string(),
};
println!(
" window: {} hard cap: {}",
fmt_remaining(&state["window_remaining_secs"]),
fmt_remaining(&state["cap_remaining_secs"]),
);
}
let Some(members) = state["members"].as_array() else {
return;
};
if members.is_empty() {
println!(" hosts: none joined yet");
return;
}
println!(" hosts:");
let host_width = members
.iter()
.filter_map(|m| m["host"].as_str())
.map(str::len)
.max()
.unwrap_or(0);
for m in members {
let host = m["host"].as_str().unwrap_or("?");
let ranks: Vec<String> = m["ranks"]
.as_array()
.map(|a| a.iter().filter_map(|r| r.as_u64()).map(|r| r.to_string()).collect())
.unwrap_or_default();
let joined_at = m["joined_at_secs"].as_u64().unwrap_or(0);
let libtorch = m["libtorch"].as_str().unwrap_or("?");
let padded_host = format!("{host:<host_width$}");
println!(
" {} ranks [{}] {} libtorch {} joined +{joined_at}s",
style::bold(&padded_host),
ranks.join(", "),
summarize_gpus(&m["gpus"]),
libtorch,
);
}
}
fn summarize_gpus(gpus: &serde_json::Value) -> String {
let names: Vec<&str> = gpus
.as_array()
.map(|a| a.iter().filter_map(|g| g.as_str()).collect())
.unwrap_or_default();
if names.is_empty() {
return "no GPUs listed".to_string();
}
if names.iter().all(|n| *n == names[0]) {
return format!("{}x {}", names.len(), names[0]);
}
names.join(", ")
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn addr_override_gets_default_port_when_bare() {
let (candidates, origin) = resolve_candidates(Some("10.0.0.7"));
assert_eq!(candidates, vec!["10.0.0.7:1337".to_string()]);
assert_eq!(origin, "--addr");
let (candidates, _) = resolve_candidates(Some("10.0.0.7:9000"));
assert_eq!(candidates, vec!["10.0.0.7:9000".to_string()]);
}
#[test]
fn gpu_summary_groups_identical_names() {
let gpus = serde_json::json!(["GP106", "GP106"]);
assert_eq!(summarize_gpus(&gpus), "2x GP106");
let gpus = serde_json::json!(["GP106", "RTX 5060 Ti"]);
assert_eq!(summarize_gpus(&gpus), "GP106, RTX 5060 Ti");
assert_eq!(summarize_gpus(&serde_json::json!([])), "no GPUs listed");
}
#[test]
fn fetch_state_reports_non_200_with_body() {
let listener = std::net::TcpListener::bind("127.0.0.1:0").unwrap();
let addr = listener.local_addr().unwrap();
let server = std::thread::spawn(move || {
let (mut stream, _) = listener.accept().unwrap();
let mut buf = [0u8; 512];
let _ = stream.read(&mut buf);
let body = r#"{"error":"no membership state published yet"}"#;
let _ = stream.write_all(
format!(
"HTTP/1.1 503 Service Unavailable\r\n\
Content-Type: application/json\r\n\
Connection: close\r\n\
Content-Length: {}\r\n\r\n{body}",
body.len(),
)
.as_bytes(),
);
});
let err = fetch_state(&addr.to_string()).unwrap_err();
assert!(err.contains("503"), "{err}");
assert!(err.contains("no membership state"), "{err}");
server.join().unwrap();
}
#[test]
fn fetch_state_round_trips_200_body() {
let listener = std::net::TcpListener::bind("127.0.0.1:0").unwrap();
let addr = listener.local_addr().unwrap();
let server = std::thread::spawn(move || {
let (mut stream, _) = listener.accept().unwrap();
let mut buf = [0u8; 512];
let _ = stream.read(&mut buf);
let body = r#"{"phase":"training","joined_ranks":3}"#;
let _ = stream.write_all(
format!(
"HTTP/1.1 200 OK\r\n\
Content-Type: application/json\r\n\
Connection: close\r\n\
Content-Length: {}\r\n\r\n{body}",
body.len(),
)
.as_bytes(),
);
});
let body = fetch_state(&addr.to_string()).unwrap();
assert_eq!(body, r#"{"phase":"training","joined_ranks":3}"#);
server.join().unwrap();
}
}