1use serde::{Deserialize, Serialize, de::DeserializeOwned};
4
5use crate::error::ProtocolResult;
6
7pub const PROTOCOL_VERSION: u8 = 7;
13
14pub const FLAG_TERMINAL: u8 = 0b0000_0001;
18
19pub const FLAG_SESSION_START: u8 = 0b0000_0010;
23
24pub const FLAG_SHUTDOWN: u8 = 0b0000_0100;
29
30pub const FRAME_HEADER_SIZE: usize = 5;
33
34#[derive(Debug, Clone, Serialize, Deserialize)]
47pub struct Message {
48 pub v: u8,
56
57 pub t: MessageType,
59
60 #[serde(skip)]
65 pub id: u32,
66
67 #[serde(skip)]
71 pub flags: u8,
72
73 #[serde(with = "serde_bytes")]
75 pub p: Vec<u8>,
76}
77
78#[derive(
86 Debug,
87 Clone,
88 Copy,
89 PartialEq,
90 Eq,
91 Hash,
92 strum::IntoStaticStr,
93 strum::EnumString,
94 strum::EnumIter,
95)]
96pub enum MessageType {
97 #[strum(serialize = "core.ready")]
99 Ready,
100
101 #[strum(serialize = "core.init.resolved")]
103 InitResolved,
104
105 #[strum(serialize = "core.init.ack")]
107 InitAck,
108
109 #[strum(serialize = "core.shutdown")]
111 Shutdown,
112
113 #[strum(serialize = "core.relay.client.disconnected")]
115 RelayClientDisconnected,
116
117 #[strum(serialize = "core.clock.sync")]
119 ClockSync,
120
121 #[strum(serialize = "core.ping")]
123 Ping,
124
125 #[strum(serialize = "core.pong")]
127 Pong,
128
129 #[strum(serialize = "core.touch")]
131 Touch,
132
133 #[strum(serialize = "core.touched")]
135 Touched,
136
137 #[strum(serialize = "core.error")]
139 CoreError,
140
141 #[strum(serialize = "core.exec.request")]
143 ExecRequest,
144
145 #[strum(serialize = "core.exec.started")]
147 ExecStarted,
148
149 #[strum(serialize = "core.exec.stdin")]
151 ExecStdin,
152
153 #[strum(serialize = "core.exec.stdin.error")]
158 ExecStdinError,
159
160 #[strum(serialize = "core.exec.stdout")]
162 ExecStdout,
163
164 #[strum(serialize = "core.exec.stderr")]
166 ExecStderr,
167
168 #[strum(serialize = "core.exec.exited")]
170 ExecExited,
171
172 #[strum(serialize = "core.exec.failed")]
176 ExecFailed,
177
178 #[strum(serialize = "core.exec.resize")]
180 ExecResize,
181
182 #[strum(serialize = "core.exec.signal")]
184 ExecSignal,
185
186 #[strum(serialize = "core.fs.request")]
188 FsRequest,
189
190 #[strum(serialize = "core.fs.response")]
192 FsResponse,
193
194 #[strum(serialize = "core.fs.data")]
196 FsData,
197
198 #[strum(serialize = "core.tcp.connect")]
200 TcpConnect,
201
202 #[strum(serialize = "core.tcp.connected")]
204 TcpConnected,
205
206 #[strum(serialize = "core.tcp.data")]
208 TcpData,
209
210 #[strum(serialize = "core.tcp.eof")]
212 TcpEof,
213
214 #[strum(serialize = "core.tcp.close")]
216 TcpClose,
217
218 #[strum(serialize = "core.tcp.closed")]
220 TcpClosed,
221
222 #[strum(serialize = "core.tcp.failed")]
224 TcpFailed,
225
226 #[strum(serialize = "core.bootstrap")]
228 Bootstrap,
229}
230
231impl Message {
236 pub fn new(t: MessageType, id: u32, p: Vec<u8>) -> Self {
238 let flags = t.flags();
239 Self {
240 v: PROTOCOL_VERSION,
241 t,
242 id,
243 flags,
244 p,
245 }
246 }
247
248 pub fn with_payload<T: Serialize>(
250 t: MessageType,
251 id: u32,
252 payload: &T,
253 ) -> ProtocolResult<Self> {
254 let mut p = Vec::new();
255 ciborium::into_writer(payload, &mut p)?;
256 let flags = t.flags();
257 Ok(Self {
258 v: PROTOCOL_VERSION,
259 t,
260 id,
261 flags,
262 p,
263 })
264 }
265
266 pub fn payload<T: DeserializeOwned>(&self) -> ProtocolResult<T> {
268 Ok(ciborium::from_reader(&self.p[..])?)
269 }
270}
271
272impl MessageType {
273 pub fn flags(&self) -> u8 {
275 match self {
276 Self::Pong
277 | Self::Touched
278 | Self::CoreError
279 | Self::ExecExited
280 | Self::ExecFailed
281 | Self::FsResponse
282 | Self::TcpClosed
283 | Self::TcpFailed => FLAG_TERMINAL,
284 Self::ExecRequest | Self::FsRequest | Self::TcpConnect => FLAG_SESSION_START,
285 Self::Shutdown => FLAG_SHUTDOWN,
286 _ => 0,
287 }
288 }
289
290 pub fn min_protocol_version(&self) -> u8 {
311 match self {
312 Self::Ready
313 | Self::InitResolved
314 | Self::InitAck
315 | Self::Shutdown
316 | Self::RelayClientDisconnected
317 | Self::ClockSync
318 | Self::ExecRequest
319 | Self::ExecStarted
320 | Self::ExecStdin
321 | Self::ExecStdinError
322 | Self::ExecStdout
323 | Self::ExecStderr
324 | Self::ExecExited
325 | Self::ExecFailed
326 | Self::ExecResize
327 | Self::ExecSignal => 1,
328 Self::FsRequest | Self::FsResponse | Self::FsData => 2,
329 Self::CoreError => 5,
330 Self::Ping | Self::Pong | Self::Touch | Self::Touched => 6,
331 Self::Bootstrap => 7,
332 Self::TcpConnect
333 | Self::TcpConnected
334 | Self::TcpData
335 | Self::TcpEof
336 | Self::TcpClose
337 | Self::TcpClosed
338 | Self::TcpFailed => 4,
339 }
340 }
341
342 pub fn is_available_at(&self, peer_generation: u8) -> bool {
351 self.min_protocol_version() <= peer_generation
352 }
353
354 pub fn as_str(&self) -> &'static str {
359 (*self).into()
360 }
361
362 pub fn from_wire_str(s: &str) -> Option<Self> {
365 s.parse().ok()
366 }
367}
368
369impl Serialize for MessageType {
374 fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
375 where
376 S: serde::Serializer,
377 {
378 serializer.serialize_str(self.as_str())
379 }
380}
381
382impl<'de> Deserialize<'de> for MessageType {
383 fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
384 where
385 D: serde::Deserializer<'de>,
386 {
387 let s = String::deserialize(deserializer)?;
388 Self::from_wire_str(&s)
389 .ok_or_else(|| serde::de::Error::custom(format!("unknown message type: {s}")))
390 }
391}
392
393#[cfg(test)]
398mod tests {
399 use super::*;
400
401 #[test]
402 fn test_message_type_roundtrip() {
403 let types = [
404 (MessageType::Bootstrap, "core.bootstrap"),
405 (MessageType::Ready, "core.ready"),
406 (MessageType::InitResolved, "core.init.resolved"),
407 (MessageType::InitAck, "core.init.ack"),
408 (MessageType::Shutdown, "core.shutdown"),
409 (
410 MessageType::RelayClientDisconnected,
411 "core.relay.client.disconnected",
412 ),
413 (MessageType::ClockSync, "core.clock.sync"),
414 (MessageType::Ping, "core.ping"),
415 (MessageType::Pong, "core.pong"),
416 (MessageType::Touch, "core.touch"),
417 (MessageType::Touched, "core.touched"),
418 (MessageType::CoreError, "core.error"),
419 (MessageType::ExecRequest, "core.exec.request"),
420 (MessageType::ExecStarted, "core.exec.started"),
421 (MessageType::ExecStdin, "core.exec.stdin"),
422 (MessageType::ExecStdinError, "core.exec.stdin.error"),
423 (MessageType::ExecStdout, "core.exec.stdout"),
424 (MessageType::ExecStderr, "core.exec.stderr"),
425 (MessageType::ExecExited, "core.exec.exited"),
426 (MessageType::ExecFailed, "core.exec.failed"),
427 (MessageType::ExecResize, "core.exec.resize"),
428 (MessageType::ExecSignal, "core.exec.signal"),
429 (MessageType::FsRequest, "core.fs.request"),
430 (MessageType::FsResponse, "core.fs.response"),
431 (MessageType::FsData, "core.fs.data"),
432 (MessageType::TcpConnect, "core.tcp.connect"),
433 (MessageType::TcpConnected, "core.tcp.connected"),
434 (MessageType::TcpData, "core.tcp.data"),
435 (MessageType::TcpEof, "core.tcp.eof"),
436 (MessageType::TcpClose, "core.tcp.close"),
437 (MessageType::TcpClosed, "core.tcp.closed"),
438 (MessageType::TcpFailed, "core.tcp.failed"),
439 ];
440
441 for (mt, expected_str) in &types {
442 assert_eq!(mt.as_str(), *expected_str);
443 assert_eq!(MessageType::from_wire_str(expected_str).unwrap(), *mt);
444 }
445 }
446
447 #[test]
448 fn test_message_type_serde_roundtrip() {
449 let types = [
450 MessageType::Bootstrap,
451 MessageType::Ready,
452 MessageType::InitResolved,
453 MessageType::InitAck,
454 MessageType::Shutdown,
455 MessageType::RelayClientDisconnected,
456 MessageType::ClockSync,
457 MessageType::Ping,
458 MessageType::Pong,
459 MessageType::Touch,
460 MessageType::Touched,
461 MessageType::CoreError,
462 MessageType::ExecRequest,
463 MessageType::ExecStarted,
464 MessageType::ExecStdin,
465 MessageType::ExecStdinError,
466 MessageType::ExecStdout,
467 MessageType::ExecStderr,
468 MessageType::ExecExited,
469 MessageType::ExecFailed,
470 MessageType::ExecResize,
471 MessageType::ExecSignal,
472 MessageType::FsRequest,
473 MessageType::FsResponse,
474 MessageType::FsData,
475 MessageType::TcpConnect,
476 MessageType::TcpConnected,
477 MessageType::TcpData,
478 MessageType::TcpEof,
479 MessageType::TcpClose,
480 MessageType::TcpClosed,
481 MessageType::TcpFailed,
482 ];
483
484 for mt in &types {
485 let mut buf = Vec::new();
486 ciborium::into_writer(mt, &mut buf).unwrap();
487 let decoded: MessageType = ciborium::from_reader(&buf[..]).unwrap();
488 assert_eq!(&decoded, mt);
489 }
490 }
491
492 #[test]
493 fn test_unknown_message_type() {
494 assert!(MessageType::from_wire_str("core.unknown").is_none());
495 }
496
497 #[test]
498 fn test_message_with_payload_roundtrip() {
499 use crate::exec::ExecExited;
500
501 let msg =
502 Message::with_payload(MessageType::ExecExited, 7, &ExecExited { code: 42 }).unwrap();
503
504 assert_eq!(msg.t, MessageType::ExecExited);
505 assert_eq!(msg.id, 7);
506 assert_eq!(msg.flags, FLAG_TERMINAL);
507
508 let payload: ExecExited = msg.payload().unwrap();
509 assert_eq!(payload.code, 42);
510 }
511
512 #[test]
513 fn test_message_type_flags() {
514 assert_eq!(MessageType::ExecExited.flags(), FLAG_TERMINAL);
515 assert_eq!(MessageType::ExecFailed.flags(), FLAG_TERMINAL);
516 assert_eq!(MessageType::FsResponse.flags(), FLAG_TERMINAL);
517 assert_eq!(MessageType::TcpClosed.flags(), FLAG_TERMINAL);
518 assert_eq!(MessageType::TcpFailed.flags(), FLAG_TERMINAL);
519 assert_eq!(MessageType::Pong.flags(), FLAG_TERMINAL);
520 assert_eq!(MessageType::Touched.flags(), FLAG_TERMINAL);
521 assert_eq!(MessageType::ExecRequest.flags(), FLAG_SESSION_START);
522 assert_eq!(MessageType::FsRequest.flags(), FLAG_SESSION_START);
523 assert_eq!(MessageType::TcpConnect.flags(), FLAG_SESSION_START);
524 assert_eq!(MessageType::Ready.flags(), 0);
525 assert_eq!(MessageType::Bootstrap.flags(), 0);
526 assert_eq!(MessageType::InitResolved.flags(), 0);
527 assert_eq!(MessageType::InitAck.flags(), 0);
528 assert_eq!(MessageType::Shutdown.flags(), FLAG_SHUTDOWN);
529 assert_eq!(MessageType::ClockSync.flags(), 0);
530 assert_eq!(MessageType::Ping.flags(), 0);
531 assert_eq!(MessageType::Touch.flags(), 0);
532 assert_eq!(MessageType::ExecStarted.flags(), 0);
533 assert_eq!(MessageType::ExecStdin.flags(), 0);
534 assert_eq!(MessageType::ExecStdout.flags(), 0);
535 assert_eq!(MessageType::ExecStderr.flags(), 0);
536 assert_eq!(MessageType::ExecResize.flags(), 0);
537 assert_eq!(MessageType::ExecSignal.flags(), 0);
538 assert_eq!(MessageType::FsData.flags(), 0);
539 assert_eq!(MessageType::TcpConnected.flags(), 0);
540 assert_eq!(MessageType::TcpData.flags(), 0);
541 assert_eq!(MessageType::TcpEof.flags(), 0);
542 assert_eq!(MessageType::TcpClose.flags(), 0);
543 }
544
545 #[test]
546 fn test_additive_fields_keep_old_and_new_compatible() {
547 use serde::{Deserialize, Serialize};
550
551 #[derive(Serialize, Deserialize)]
553 struct Old {
554 a: u32,
555 b: u32,
556 }
557
558 #[derive(Serialize, Deserialize, Debug, PartialEq)]
560 struct New {
561 a: u32,
562 b: u32,
563 #[serde(default)]
564 c: u32,
565 }
566
567 let mut new_bytes = Vec::new();
569 ciborium::into_writer(&New { a: 1, b: 2, c: 3 }, &mut new_bytes).unwrap();
570 let as_old: Old = ciborium::from_reader(&new_bytes[..]).unwrap();
571 assert_eq!((as_old.a, as_old.b), (1, 2));
572
573 let mut old_bytes = Vec::new();
575 ciborium::into_writer(&Old { a: 1, b: 2 }, &mut old_bytes).unwrap();
576 let as_new: New = ciborium::from_reader(&old_bytes[..]).unwrap();
577 assert_eq!(as_new, New { a: 1, b: 2, c: 0 });
578 }
579
580 #[test]
581 fn test_is_available_at() {
582 assert!(MessageType::ExecRequest.is_available_at(1));
584 assert!(MessageType::ExecRequest.is_available_at(2));
585 assert!(MessageType::ExecRequest.is_available_at(PROTOCOL_VERSION));
586 assert!(!MessageType::FsRequest.is_available_at(1));
588 assert!(MessageType::FsRequest.is_available_at(2));
589 assert!(MessageType::FsRequest.is_available_at(PROTOCOL_VERSION));
590 assert!(!MessageType::Ping.is_available_at(5));
592 assert!(MessageType::Ping.is_available_at(6));
593 assert!(MessageType::Ping.is_available_at(PROTOCOL_VERSION));
594 assert!(!MessageType::Bootstrap.is_available_at(6));
596 assert!(MessageType::Bootstrap.is_available_at(PROTOCOL_VERSION));
597 }
598
599 #[test]
600 fn test_min_protocol_version_per_type() {
601 let baseline = [
604 MessageType::Ready,
605 MessageType::InitResolved,
606 MessageType::InitAck,
607 MessageType::Shutdown,
608 MessageType::RelayClientDisconnected,
609 MessageType::ClockSync,
610 MessageType::ExecRequest,
611 MessageType::ExecStarted,
612 MessageType::ExecStdin,
613 MessageType::ExecStdinError,
614 MessageType::ExecStdout,
615 MessageType::ExecStderr,
616 MessageType::ExecExited,
617 MessageType::ExecFailed,
618 MessageType::ExecResize,
619 MessageType::ExecSignal,
620 ];
621 for mt in &baseline {
622 assert_eq!(mt.min_protocol_version(), 1, "{mt:?} should be v1 baseline");
623 }
624
625 for mt in [
628 MessageType::FsRequest,
629 MessageType::FsResponse,
630 MessageType::FsData,
631 ] {
632 assert_eq!(mt.min_protocol_version(), 2, "{mt:?} should require gen 2");
633 }
634
635 for mt in [
636 MessageType::Ping,
637 MessageType::Pong,
638 MessageType::Touch,
639 MessageType::Touched,
640 ] {
641 assert_eq!(mt.min_protocol_version(), 6, "{mt:?} should require gen 6");
642 }
643
644 assert_eq!(MessageType::Bootstrap.min_protocol_version(), 7);
645
646 assert!(MessageType::FsRequest.min_protocol_version() <= PROTOCOL_VERSION);
648 }
649
650 #[test]
651 fn test_message_new_computes_flags() {
652 let msg = Message::new(MessageType::ExecRequest, 1, Vec::new());
653 assert_eq!(msg.flags, FLAG_SESSION_START);
654
655 let msg = Message::new(MessageType::ExecStdout, 1, Vec::new());
656 assert_eq!(msg.flags, 0);
657 }
658}