1use ctl_component_info::{ComponentBuildInfo, ComponentInfo};
3use serde::{Deserialize, Serialize, de::DeserializeOwned};
4use std::io;
5use tokio::io::{AsyncRead, AsyncReadExt as _, AsyncWrite, AsyncWriteExt as _};
6
7pub const PROTOCOL_VERSION: u16 = 2;
8const MAX_FRAME_BYTES: usize = 16 * 1024;
9
10#[derive(Debug, Serialize, Deserialize)]
11#[serde(tag = "type", rename_all = "snake_case", deny_unknown_fields)]
12pub enum ClientMessage {
13 PrepareCtmuxRestart {
14 protocol_version: u16,
15 expected_remote_id: String,
16 },
17 Confirm {},
18}
19
20#[derive(Debug, Clone, Serialize, Deserialize)]
21pub struct RunningCtmux {
22 pub build: Option<ComponentBuildInfo>,
23 pub protocol_version: Option<u16>,
24 pub control_protocol_version: u16,
25}
26
27#[derive(Debug, Clone, Serialize, Deserialize)]
28pub struct CtmuxPreparation {
29 pub remote_id: String,
30 pub running: RunningCtmux,
31 pub available: ComponentInfo,
32}
33
34#[derive(Debug, Serialize, Deserialize)]
35pub struct CtmuxRestartCompleted {
36 pub after: ComponentInfo,
37 pub terminated_sessions: u32,
38}
39
40#[derive(Debug, Serialize, Deserialize)]
41#[serde(tag = "type", rename_all = "snake_case", deny_unknown_fields)]
42pub enum ServerMessage {
43 Prepared {
44 info: CtmuxPreparation,
45 },
46 Completed {
47 result: CtmuxRestartCompleted,
48 },
49 Error {
50 code: String,
51 message: String,
52 may_have_stopped: bool,
53 },
54}
55
56pub async fn read<R: AsyncRead + Unpin, T: DeserializeOwned>(reader: &mut R) -> io::Result<T> {
61 let size = reader.read_u32().await? as usize;
62 if size == 0 || size > MAX_FRAME_BYTES {
63 return Err(io::Error::new(
64 io::ErrorKind::InvalidData,
65 "Invalid maintenance frame size.",
66 ));
67 }
68 let mut bytes = vec![0; size];
69 reader.read_exact(&mut bytes).await?;
70 serde_json::from_slice(&bytes).map_err(|error| io::Error::new(io::ErrorKind::InvalidData, error))
71}
72
73pub async fn write<W: AsyncWrite + Unpin, T: Serialize>(
78 writer: &mut W,
79 value: &T,
80) -> io::Result<()> {
81 let bytes = serde_json::to_vec(value).map_err(io::Error::other)?;
82 if bytes.len() > MAX_FRAME_BYTES {
83 return Err(io::Error::new(
84 io::ErrorKind::InvalidData,
85 "Maintenance frame is too large.",
86 ));
87 }
88 writer
89 .write_u32(u32::try_from(bytes.len()).map_err(io::Error::other)?)
90 .await?;
91 writer.write_all(&bytes).await?;
92 writer.flush().await
93}
94
95#[cfg(test)]
96mod tests {
97 use super::*;
98
99 #[tokio::test]
100 async fn rejects_unbounded_frames_and_arbitrary_operations() {
101 assert!(
102 read::<_, ClientMessage>(&mut u32::MAX.to_be_bytes().as_slice())
103 .await
104 .is_err()
105 );
106 assert!(serde_json::from_str::<ClientMessage>(r#"{"type":"confirm","command":"sh"}"#).is_err());
107 assert!(serde_json::from_str::<ClientMessage>(r#"{"type":"exec"}"#).is_err());
108 }
109
110 #[tokio::test]
111 async fn preparation_and_confirmation_are_separate_frames() {
112 let mut bytes = Vec::new();
113 write(
114 &mut bytes,
115 &ClientMessage::PrepareCtmuxRestart {
116 protocol_version: PROTOCOL_VERSION,
117 expected_remote_id: "test-identity".into(),
118 },
119 )
120 .await
121 .unwrap();
122 write(&mut bytes, &ClientMessage::Confirm {}).await.unwrap();
123 let mut reader = bytes.as_slice();
124 assert!(matches!(
125 read::<_, ClientMessage>(&mut reader).await.unwrap(),
126 ClientMessage::PrepareCtmuxRestart { .. }
127 ));
128 assert_ne!(reader, &[] as &[u8]);
129 assert!(matches!(
130 read::<_, ClientMessage>(&mut reader).await.unwrap(),
131 ClientMessage::Confirm {}
132 ));
133 }
134}