use std::collections::BTreeMap;
use std::io::{self, Read, Write};
use serde::{Deserialize, Serialize};
use serde_json::Value;
const MAX_FRAME: u32 = 256 * 1024 * 1024;
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct Frame {
#[serde(default)]
pub req: u64,
#[serde(flatten)]
pub msg: Msg,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(tag = "t", rename_all = "snake_case")]
pub enum Msg {
Hello {
client_id: String,
group: String,
session: bool,
kind: String, },
Nodes,
ClusterResources,
AvailablePerNode,
CreatePg {
bundles: Vec<f64>,
strategy: String,
},
PgTable {
pg_id: String,
},
RemovePg {
pg_id: String,
},
CreateActor {
name: String,
num_gpus: f64,
pg_id: String,
bundle_index: usize,
env: BTreeMap<String, String>,
},
Call {
actor_id: String,
method: String,
},
Get {
ref_id: String,
timeout_ms: Option<u64>,
},
Wait {
ref_ids: Vec<String>,
num_returns: usize,
timeout_ms: Option<u64>,
},
KillActor {
actor_id: String,
},
Status {
group: Option<String>,
},
StopAll {
group: Option<String>,
},
Ok0,
Err {
error: String,
},
HelloOk {
node_id: String,
node_ip: String,
gcs_address: String,
head_node_id: String,
},
NodesOk {
nodes: Vec<Value>,
},
ResourcesOk {
resources: BTreeMap<String, f64>,
},
AvailOk {
nodes: BTreeMap<String, BTreeMap<String, f64>>,
},
CreatePgOk {
pg_id: String,
ready_ref: String,
},
PgTableOk {
table: Value,
},
CreateActorOk {
actor_id: String,
node_id: String,
gpu_ids: Vec<u32>,
},
CallOk {
ref_id: String,
},
GetOk {
status: String,
#[serde(default)]
reason: String,
},
WaitOk {
ready: Vec<String>,
},
StatusOk {
data: Value,
},
AgentRegister {
agent_id: String,
group: String,
node_ip: String,
gpus: Vec<u32>,
#[serde(default = "default_gpu_vendor")]
gpu_vendor: String,
cpus: u32,
container: String,
pid: u32,
#[serde(default)]
services: BTreeMap<String, String>,
#[serde(default)]
services_ports: BTreeMap<String, ServicePort>,
#[serde(default)]
service_notes: BTreeMap<String, String>,
resume: Vec<ResumeActor>,
#[serde(default)]
unacked_refs: Vec<String>,
},
AgentRegisterOk {
node_id: String,
},
Spawn {
actor_id: String,
name: String,
env: BTreeMap<String, String>,
gpu_ids: Vec<u32>,
node_id: String,
gcs_address: String,
},
SpawnResult {
actor_id: String,
ok: bool,
#[serde(default)]
error: String,
#[serde(default)]
pid: u32,
},
CallActor {
actor_id: String,
ref_id: String,
method: String,
},
ActorResult {
ref_id: String,
ok: bool,
#[serde(default)]
error: String,
},
ActorExit {
actor_id: String,
exit_code: Option<i32>,
signal: Option<i32>,
},
Kill {
actor_id: String,
},
ServiceNote {
service: String,
note: String,
},
Ping,
Pong,
PeerHello {
node_id: String,
node_ip: String,
control_addr: String,
http_port: u16,
#[serde(default)]
addrs: Vec<String>,
#[serde(default)]
addr_tags: BTreeMap<String, Vec<String>>,
#[serde(default)]
probes: bool,
},
PeerHelloOk {
node_id: String,
node_ip: String,
#[serde(default)]
control_addr: String,
#[serde(default)]
http_port: u16,
#[serde(default)]
addrs: Vec<String>,
#[serde(default)]
addr_tags: BTreeMap<String, Vec<String>>,
#[serde(default)]
probes: bool,
},
Probe {
node_id: String,
local_addr: String,
},
ProbeOk {
node_id: String,
},
PeerStatus {
data: Value,
},
PeerEvent {
origin: String,
line: String,
},
HostHello {
actor_id: String,
},
Ctor, CtorOk,
CtorErr {
#[serde(default)]
error: String,
},
HostCall {
ref_id: String,
method: String,
},
HostResult {
ref_id: String,
ok: bool,
},
}
pub fn default_gpu_vendor() -> String {
"nvidia".to_string()
}
pub fn write_frame<W: Write>(w: &mut W, frame: &Frame, payload: &[u8]) -> io::Result<()> {
let header = serde_json::to_vec(frame)?;
w.write_all(&(header.len() as u32).to_le_bytes())?;
w.write_all(&(payload.len() as u32).to_le_bytes())?;
w.write_all(&header)?;
w.write_all(payload)?;
w.flush()
}
pub fn read_frame<R: Read>(r: &mut R) -> io::Result<Option<(Frame, Vec<u8>)>> {
let mut lens = [0u8; 8];
if !read_exact_or_eof(r, &mut lens)? {
return Ok(None);
}
let hlen = u32::from_le_bytes(lens[0..4].try_into().unwrap());
let plen = u32::from_le_bytes(lens[4..8].try_into().unwrap());
if hlen > MAX_FRAME || plen > MAX_FRAME {
return Err(io::Error::new(
io::ErrorKind::InvalidData,
format!("frame too large: header={hlen} payload={plen}"),
));
}
let mut header = vec![0u8; hlen as usize];
r.read_exact(&mut header)?;
let mut payload = vec![0u8; plen as usize];
r.read_exact(&mut payload)?;
let frame: Frame = serde_json::from_slice(&header).map_err(|e| {
io::Error::new(
io::ErrorKind::InvalidData,
format!(
"bad frame header: {e}: {}",
String::from_utf8_lossy(&header)
),
)
})?;
Ok(Some((frame, payload)))
}
fn read_exact_or_eof<R: Read>(r: &mut R, buf: &mut [u8]) -> io::Result<bool> {
let mut filled = 0;
while filled < buf.len() {
match r.read(&mut buf[filled..]) {
Ok(0) if filled == 0 => return Ok(false),
Ok(0) => {
return Err(io::Error::new(
io::ErrorKind::UnexpectedEof,
"EOF mid-frame",
))
}
Ok(n) => filled += n,
Err(e) if e.kind() == io::ErrorKind::Interrupted => continue,
Err(e) => return Err(e),
}
}
Ok(true)
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ServicePort {
pub port: u16,
#[serde(default)]
pub path: String,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ResumeActor {
pub actor_id: String,
pub name: String,
pub gpu_ids: Vec<u32>,
pub pid: u32,
pub pending_refs: Vec<String>,
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn roundtrip() {
let mut buf: Vec<u8> = Vec::new();
let f = Frame {
req: 7,
msg: Msg::Call {
actor_id: "a1".into(),
method: "run".into(),
},
};
write_frame(&mut buf, &f, b"PAYLOAD").unwrap();
let mut cur = std::io::Cursor::new(buf);
let (g, p) = read_frame(&mut cur).unwrap().unwrap();
assert_eq!(g.req, 7);
assert_eq!(p, b"PAYLOAD");
match g.msg {
Msg::Call { actor_id, method } => {
assert_eq!(actor_id, "a1");
assert_eq!(method, "run");
}
other => panic!("wrong variant: {other:?}"),
}
assert!(read_frame(&mut cur).unwrap().is_none());
}
#[test]
fn eof_mid_frame_is_an_error() {
let mut buf: Vec<u8> = Vec::new();
let f = Frame {
req: 1,
msg: Msg::Ping,
};
write_frame(&mut buf, &f, b"").unwrap();
buf.truncate(buf.len() - 2);
let mut cur = std::io::Cursor::new(buf);
assert!(read_frame(&mut cur).is_err());
}
}