Skip to main content

ctl_task_proto/
lib.rs

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
169/// Writes one length-prefixed task protocol message.
170///
171/// # Errors
172///
173/// Returns an error when serialization fails, the encoded frame is oversized,
174/// or the transport cannot be written.
175pub 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
197/// Reads one length-prefixed task protocol message.
198///
199/// A clean end of stream before the next frame returns `Ok(None)`.
200///
201/// # Errors
202///
203/// Returns an error when the transport fails mid-frame, the declared frame is
204/// oversized, or the payload cannot be decoded.
205pub 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}