use ctl_component_info::{ComponentBuildInfo, ComponentInfo};
use serde::{Deserialize, Serialize, de::DeserializeOwned};
use std::io;
use tokio::io::{AsyncRead, AsyncReadExt as _, AsyncWrite, AsyncWriteExt as _};
pub const PROTOCOL_VERSION: u16 = 2;
const MAX_FRAME_BYTES: usize = 16 * 1024;
#[derive(Debug, Serialize, Deserialize)]
#[serde(tag = "type", rename_all = "snake_case", deny_unknown_fields)]
pub enum ClientMessage {
PrepareCtmuxRestart {
protocol_version: u16,
expected_remote_id: String,
},
Confirm {},
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct RunningCtmux {
pub build: Option<ComponentBuildInfo>,
pub protocol_version: Option<u16>,
pub control_protocol_version: u16,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct CtmuxPreparation {
pub remote_id: String,
pub running: RunningCtmux,
pub available: ComponentInfo,
}
#[derive(Debug, Serialize, Deserialize)]
pub struct CtmuxRestartCompleted {
pub after: ComponentInfo,
pub terminated_sessions: u32,
}
#[derive(Debug, Serialize, Deserialize)]
#[serde(tag = "type", rename_all = "snake_case", deny_unknown_fields)]
pub enum ServerMessage {
Prepared {
info: CtmuxPreparation,
},
Completed {
result: CtmuxRestartCompleted,
},
Error {
code: String,
message: String,
may_have_stopped: bool,
},
}
pub async fn read<R: AsyncRead + Unpin, T: DeserializeOwned>(reader: &mut R) -> io::Result<T> {
let size = reader.read_u32().await? as usize;
if size == 0 || size > MAX_FRAME_BYTES {
return Err(io::Error::new(
io::ErrorKind::InvalidData,
"Invalid maintenance frame size.",
));
}
let mut bytes = vec![0; size];
reader.read_exact(&mut bytes).await?;
serde_json::from_slice(&bytes).map_err(|error| io::Error::new(io::ErrorKind::InvalidData, error))
}
pub async fn write<W: AsyncWrite + Unpin, T: Serialize>(
writer: &mut W,
value: &T,
) -> io::Result<()> {
let bytes = serde_json::to_vec(value).map_err(io::Error::other)?;
if bytes.len() > MAX_FRAME_BYTES {
return Err(io::Error::new(
io::ErrorKind::InvalidData,
"Maintenance frame is too large.",
));
}
writer
.write_u32(u32::try_from(bytes.len()).map_err(io::Error::other)?)
.await?;
writer.write_all(&bytes).await?;
writer.flush().await
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn rejects_unbounded_frames_and_arbitrary_operations() {
assert!(
read::<_, ClientMessage>(&mut u32::MAX.to_be_bytes().as_slice())
.await
.is_err()
);
assert!(serde_json::from_str::<ClientMessage>(r#"{"type":"confirm","command":"sh"}"#).is_err());
assert!(serde_json::from_str::<ClientMessage>(r#"{"type":"exec"}"#).is_err());
}
#[tokio::test]
async fn preparation_and_confirmation_are_separate_frames() {
let mut bytes = Vec::new();
write(
&mut bytes,
&ClientMessage::PrepareCtmuxRestart {
protocol_version: PROTOCOL_VERSION,
expected_remote_id: "test-identity".into(),
},
)
.await
.unwrap();
write(&mut bytes, &ClientMessage::Confirm {}).await.unwrap();
let mut reader = bytes.as_slice();
assert!(matches!(
read::<_, ClientMessage>(&mut reader).await.unwrap(),
ClientMessage::PrepareCtmuxRestart { .. }
));
assert_ne!(reader, &[] as &[u8]);
assert!(matches!(
read::<_, ClientMessage>(&mut reader).await.unwrap(),
ClientMessage::Confirm {}
));
}
}