1pub mod control;
2
3use serde::{Deserialize, Serialize, de::DeserializeOwned};
4use std::io;
5use thiserror::Error;
6use tokio::io::{AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt};
7
8pub const PROTOCOL_VERSION: u16 = 4;
9pub const MAX_FRAME_SIZE: usize = 8 * 1024 * 1024;
10
11#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
12#[serde(rename_all = "snake_case")]
13pub enum ExecutionMode {
14 Interactive,
15 Background,
16}
17
18#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
19pub struct TaskDefinition {
20 pub name: String,
21 pub program: String,
22 pub arguments: Vec<String>,
23 pub working_directory: Option<String>,
24 pub execution_mode: ExecutionMode,
25}
26
27#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
28#[serde(rename_all = "snake_case")]
29pub enum DesiredState {
30 Stopped,
31 Running,
32}
33
34#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
35#[serde(rename_all = "snake_case")]
36pub enum RunState {
37 Starting,
38 Unknown,
39 Running,
40 Completed,
41 Failed,
42 Stopped,
43}
44
45#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
46pub struct InteractiveRun {
47 #[serde(default)]
48 pub released: bool,
49 pub ctmux_socket: std::path::PathBuf,
50 pub instance_id: String,
51 pub session_id: Option<String>,
52}
53
54#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
55pub struct RunInfo {
56 #[serde(default, skip_serializing_if = "Option::is_none")]
57 pub definition: Option<TaskDefinition>,
58 #[serde(default, skip_serializing_if = "Option::is_none")]
59 pub interactive: Option<InteractiveRun>,
60 pub run_id: String,
61 pub state: RunState,
62 pub started_at_ms: u64,
63 pub ended_at_ms: Option<u64>,
64 pub exit_code: Option<i32>,
65}
66
67#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
68pub struct TaskInfo {
69 pub task_id: String,
70 pub definition: TaskDefinition,
71 pub desired_state: DesiredState,
72 pub active_run: Option<RunInfo>,
73 pub last_run: Option<RunInfo>,
74}
75
76#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
77#[serde(rename_all = "snake_case")]
78pub enum LogStream {
79 Stdout,
80 Stderr,
81}
82
83#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
84pub struct LogEvent {
85 pub run_id: String,
86 pub sequence: u64,
87 pub stream: LogStream,
88 pub data: Vec<u8>,
89}
90
91#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
92#[serde(tag = "type", rename_all = "snake_case")]
93pub enum ClientMessage {
94 Handshake {
95 protocol_version: u16,
96 client_name: String,
97 },
98 CreateTask {
99 definition: TaskDefinition,
100 },
101 RegisterTask {
102 task_id: String,
103 definition: TaskDefinition,
104 },
105 UpdateTask {
106 task: String,
107 definition: TaskDefinition,
108 },
109 ListTasks,
110 ShowTask {
111 task: String,
112 },
113 StartTask {
114 task: String,
115 },
116 StopTask {
117 task: String,
118 },
119 RestartTask {
120 task: String,
121 },
122 RemoveTask {
123 task: String,
124 },
125 ReadLogs {
126 task: String,
127 after_sequence: Option<u64>,
128 follow: bool,
129 },
130}
131
132#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
133#[serde(tag = "type", rename_all = "snake_case")]
134pub enum ServerMessage {
135 HandshakeAccepted { protocol_version: u16 },
136 TaskCreated { task: TaskInfo },
137 TaskList { tasks: Vec<TaskInfo> },
138 TaskStatus { task: TaskInfo },
139 TaskRemoved { task_id: String },
140 Log { event: LogEvent },
141 LogsFinished,
142 Error { code: ErrorCode, message: String },
143}
144
145#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
146#[serde(rename_all = "snake_case")]
147pub enum ErrorCode {
148 InvalidRequest,
149 ProtocolVersionMismatch,
150 InvalidDefinition,
151 TaskNotFound,
152 NameConflict,
153 AlreadyRunning,
154 NotRunning,
155 UnsupportedExecutionMode,
156 Internal,
157}
158
159#[derive(Debug, Error)]
160pub enum CodecError {
161 #[error("I/O error: {0}")]
162 Io(#[from] io::Error),
163 #[error("frame length {actual} exceeds the maximum of {maximum} bytes")]
164 FrameTooLarge { actual: usize, maximum: usize },
165 #[error("invalid task protocol JSON frame: {0}")]
166 Json(#[from] serde_json::Error),
167}
168
169pub async fn write_frame<W, T>(writer: &mut W, message: &T) -> Result<(), CodecError>
176where
177 W: AsyncWrite + Unpin,
178 T: Serialize,
179{
180 let payload = serde_json::to_vec(message)?;
181 if payload.len() > MAX_FRAME_SIZE {
182 return Err(CodecError::FrameTooLarge {
183 actual: payload.len(),
184 maximum: MAX_FRAME_SIZE,
185 });
186 }
187 let length = u32::try_from(payload.len()).map_err(|_| CodecError::FrameTooLarge {
188 actual: payload.len(),
189 maximum: MAX_FRAME_SIZE,
190 })?;
191 writer.write_all(&length.to_be_bytes()).await?;
192 writer.write_all(&payload).await?;
193 writer.flush().await?;
194 Ok(())
195}
196
197pub async fn read_frame<R, T>(reader: &mut R) -> Result<Option<T>, CodecError>
206where
207 R: AsyncRead + Unpin,
208 T: DeserializeOwned,
209{
210 let mut length_bytes = [0_u8; 4];
211 match reader.read(&mut length_bytes[..1]).await {
212 Ok(0) => return Ok(None),
213 Ok(_) => {
214 reader.read_exact(&mut length_bytes[1..]).await?;
215 }
216 Err(error) => return Err(error.into()),
217 }
218 let length = u32::from_be_bytes(length_bytes) as usize;
219 if length > MAX_FRAME_SIZE {
220 return Err(CodecError::FrameTooLarge {
221 actual: length,
222 maximum: MAX_FRAME_SIZE,
223 });
224 }
225 let mut payload = vec![0; length];
226 reader.read_exact(&mut payload).await?;
227 Ok(Some(serde_json::from_slice(&payload)?))
228}
229
230#[cfg(test)]
231mod tests {
232 use super::*;
233
234 #[test]
235 fn legacy_background_runs_remain_readable() {
236 let run: RunInfo = serde_json::from_str(
237 r#"{"run_id":"old-run","state":"completed","started_at_ms":1,"ended_at_ms":2,"exit_code":0}"#,
238 )
239 .unwrap();
240 assert_eq!(run.state, RunState::Completed);
241 assert!(run.interactive.is_none());
242 }
243
244 #[tokio::test]
245 async fn frames_round_trip() {
246 let message = ClientMessage::ListTasks;
247 let (mut client, mut server) = tokio::io::duplex(1024);
248 let write = write_frame(&mut client, &message);
249 let read = read_frame::<_, ClientMessage>(&mut server);
250 let (written, received) = tokio::join!(write, read);
251 written.unwrap();
252 assert_eq!(received.unwrap(), Some(message));
253 }
254}