Skip to main content

zellij_utils/
nested_session.rs

1use crate::data::{Direction, KeyWithModifier};
2use crate::nested_session_contract::nested_session_contract as proto;
3use base64::alphabet::STANDARD as BASE64_STANDARD_ALPHABET;
4use base64::engine::general_purpose::{
5    GeneralPurpose, GeneralPurposeConfig, STANDARD as BASE64_STANDARD,
6};
7use base64::engine::{DecodePaddingMode, Engine as _};
8use prost::Message;
9use std::str::FromStr;
10
11const BASE64_DECODER: GeneralPurpose = GeneralPurpose::new(
12    &BASE64_STANDARD_ALPHABET,
13    GeneralPurposeConfig::new().with_decode_padding_mode(DecodePaddingMode::Indifferent),
14);
15
16pub const NESTED_DCS_PARAM: u16 = 26661;
17pub const NESTED_FRAME_HEADER: &[u8] = b"\x1bP26661n";
18pub const NESTED_FRAME_TERMINATOR: &[u8] = b"\x1b\\";
19
20pub const REANNOUNCE_SILENCE_MS: u64 = 3000;
21pub const REANNOUNCE_CHECK_INTERVAL_MS: u64 = 1000;
22
23pub fn reannounce_silence_ms() -> u64 {
24    std::env::var("ZELLIJ_NESTED_REANNOUNCE_SILENCE_MS")
25        .ok()
26        .and_then(|value| value.parse().ok())
27        .unwrap_or(REANNOUNCE_SILENCE_MS)
28}
29
30pub fn reannounce_check_interval_ms() -> u64 {
31    std::env::var("ZELLIJ_NESTED_REANNOUNCE_CHECK_INTERVAL_MS")
32        .ok()
33        .and_then(|value| value.parse().ok())
34        .unwrap_or(REANNOUNCE_CHECK_INTERVAL_MS)
35}
36
37#[derive(Debug, Clone, Copy, PartialEq, Eq)]
38pub enum NestedSessionCapability {
39    NestedControl,
40}
41
42#[derive(Debug, Clone, PartialEq)]
43pub enum NestedSessionMessage {
44    Announce {
45        session_name: String,
46        capabilities: Vec<NestedSessionCapability>,
47    },
48    FocusHost {
49        direction: Option<Direction>,
50    },
51    ToggleHostFullscreen {
52        fullscreen: bool,
53    },
54    Pong,
55    Bye,
56    AnnounceAck {
57        ancestry: Vec<String>,
58        capabilities: Vec<NestedSessionCapability>,
59        descend_keys: Vec<KeyWithModifier>,
60    },
61    FocusGained {
62        from_direction: Option<Direction>,
63    },
64    FocusLost,
65    FullscreenState {
66        fullscreen: bool,
67    },
68    AncestryUpdate {
69        ancestry: Vec<String>,
70    },
71    Ping,
72    ShortcutUpdate {
73        ascend_keys: Vec<KeyWithModifier>,
74        descend_keys: Vec<KeyWithModifier>,
75    },
76}
77
78fn keys_to_proto(keys: &[KeyWithModifier]) -> Vec<String> {
79    keys.iter().map(|key| key.to_kdl()).collect()
80}
81
82fn keys_from_proto(keys: &[String]) -> Vec<KeyWithModifier> {
83    let parsed: Vec<KeyWithModifier> = keys
84        .iter()
85        .filter_map(|key| KeyWithModifier::from_str(key).ok())
86        .collect();
87    if parsed.len() == keys.len() {
88        parsed
89    } else {
90        vec![]
91    }
92}
93
94fn capabilities_to_proto(capabilities: &[NestedSessionCapability]) -> Vec<i32> {
95    capabilities
96        .iter()
97        .map(|capability| match capability {
98            NestedSessionCapability::NestedControl => proto::NestedCapability::NestedControl as i32,
99        })
100        .collect()
101}
102
103fn capabilities_from_proto(capabilities: &[i32]) -> Vec<NestedSessionCapability> {
104    capabilities
105        .iter()
106        .filter_map(
107            |capability| match proto::NestedCapability::try_from(*capability).ok() {
108                Some(proto::NestedCapability::NestedControl) => {
109                    Some(NestedSessionCapability::NestedControl)
110                },
111                _ => None,
112            },
113        )
114        .collect()
115}
116
117fn direction_to_proto(direction: Option<Direction>) -> i32 {
118    match direction {
119        Some(Direction::Left) => proto::NestedDirection::Left as i32,
120        Some(Direction::Right) => proto::NestedDirection::Right as i32,
121        Some(Direction::Up) => proto::NestedDirection::Up as i32,
122        Some(Direction::Down) => proto::NestedDirection::Down as i32,
123        None => proto::NestedDirection::Unspecified as i32,
124    }
125}
126
127fn direction_from_proto(direction: i32) -> Option<Direction> {
128    match proto::NestedDirection::try_from(direction).ok() {
129        Some(proto::NestedDirection::Left) => Some(Direction::Left),
130        Some(proto::NestedDirection::Right) => Some(Direction::Right),
131        Some(proto::NestedDirection::Up) => Some(Direction::Up),
132        Some(proto::NestedDirection::Down) => Some(Direction::Down),
133        _ => None,
134    }
135}
136
137impl From<NestedSessionMessage> for proto::NestedSessionMessage {
138    fn from(message: NestedSessionMessage) -> Self {
139        use proto::nested_session_message::Payload;
140        let payload = match message {
141            NestedSessionMessage::Announce {
142                session_name,
143                capabilities,
144            } => Payload::Announce(proto::Announce {
145                session_name,
146                capabilities: capabilities_to_proto(&capabilities),
147            }),
148            NestedSessionMessage::FocusHost { direction } => Payload::FocusHost(proto::FocusHost {
149                direction: direction_to_proto(direction),
150            }),
151            NestedSessionMessage::ToggleHostFullscreen { fullscreen } => {
152                Payload::HostFullscreen(proto::ToggleHostFullscreen { fullscreen })
153            },
154            NestedSessionMessage::Pong => Payload::Pong(proto::Pong {}),
155            NestedSessionMessage::Bye => Payload::Bye(proto::Bye {}),
156            NestedSessionMessage::AnnounceAck {
157                ancestry,
158                capabilities,
159                descend_keys,
160            } => Payload::AnnounceAck(proto::AnnounceAck {
161                ancestry,
162                capabilities: capabilities_to_proto(&capabilities),
163                descend_keys: keys_to_proto(&descend_keys),
164            }),
165            NestedSessionMessage::FocusGained { from_direction } => {
166                Payload::FocusGained(proto::FocusGained {
167                    from_direction: direction_to_proto(from_direction),
168                })
169            },
170            NestedSessionMessage::FocusLost => Payload::FocusLost(proto::FocusLost {}),
171            NestedSessionMessage::FullscreenState { fullscreen } => {
172                Payload::FullscreenState(proto::FullscreenState { fullscreen })
173            },
174            NestedSessionMessage::AncestryUpdate { ancestry } => {
175                Payload::AncestryUpdate(proto::AncestryUpdate { ancestry })
176            },
177            NestedSessionMessage::Ping => Payload::Ping(proto::Ping {}),
178            NestedSessionMessage::ShortcutUpdate {
179                ascend_keys,
180                descend_keys,
181            } => Payload::ShortcutUpdate(proto::ShortcutUpdate {
182                ascend_keys: keys_to_proto(&ascend_keys),
183                descend_keys: keys_to_proto(&descend_keys),
184            }),
185        };
186        proto::NestedSessionMessage {
187            payload: Some(payload),
188        }
189    }
190}
191
192impl TryFrom<proto::NestedSessionMessage> for NestedSessionMessage {
193    type Error = ();
194    fn try_from(message: proto::NestedSessionMessage) -> Result<Self, Self::Error> {
195        use proto::nested_session_message::Payload;
196        match message.payload {
197            Some(Payload::Announce(announce)) => Ok(NestedSessionMessage::Announce {
198                session_name: announce.session_name,
199                capabilities: capabilities_from_proto(&announce.capabilities),
200            }),
201            Some(Payload::FocusHost(focus_host)) => Ok(NestedSessionMessage::FocusHost {
202                direction: direction_from_proto(focus_host.direction),
203            }),
204            Some(Payload::HostFullscreen(host_fullscreen)) => {
205                Ok(NestedSessionMessage::ToggleHostFullscreen {
206                    fullscreen: host_fullscreen.fullscreen,
207                })
208            },
209            Some(Payload::Pong(_)) => Ok(NestedSessionMessage::Pong),
210            Some(Payload::Bye(_)) => Ok(NestedSessionMessage::Bye),
211            Some(Payload::AnnounceAck(announce_ack)) => Ok(NestedSessionMessage::AnnounceAck {
212                ancestry: announce_ack.ancestry,
213                capabilities: capabilities_from_proto(&announce_ack.capabilities),
214                descend_keys: keys_from_proto(&announce_ack.descend_keys),
215            }),
216            Some(Payload::FocusGained(focus_gained)) => Ok(NestedSessionMessage::FocusGained {
217                from_direction: direction_from_proto(focus_gained.from_direction),
218            }),
219            Some(Payload::FocusLost(_)) => Ok(NestedSessionMessage::FocusLost),
220            Some(Payload::FullscreenState(fullscreen_state)) => {
221                Ok(NestedSessionMessage::FullscreenState {
222                    fullscreen: fullscreen_state.fullscreen,
223                })
224            },
225            Some(Payload::AncestryUpdate(ancestry_update)) => {
226                Ok(NestedSessionMessage::AncestryUpdate {
227                    ancestry: ancestry_update.ancestry,
228                })
229            },
230            Some(Payload::Ping(_)) => Ok(NestedSessionMessage::Ping),
231            Some(Payload::ShortcutUpdate(shortcut_update)) => {
232                Ok(NestedSessionMessage::ShortcutUpdate {
233                    ascend_keys: keys_from_proto(&shortcut_update.ascend_keys),
234                    descend_keys: keys_from_proto(&shortcut_update.descend_keys),
235                })
236            },
237            None => Err(()),
238        }
239    }
240}
241
242pub fn encode_payload(message: &NestedSessionMessage) -> Vec<u8> {
243    let proto_message: proto::NestedSessionMessage = message.clone().into();
244    proto_message.encode_to_vec()
245}
246
247pub fn decode_payload(payload_bytes: &[u8]) -> Option<NestedSessionMessage> {
248    let proto_message = proto::NestedSessionMessage::decode(payload_bytes).ok()?;
249    NestedSessionMessage::try_from(proto_message).ok()
250}
251
252pub fn encode_frame(message: &NestedSessionMessage) -> Vec<u8> {
253    encode_frame_from_payload(&encode_payload(message))
254}
255
256pub fn encode_frame_from_payload(payload_bytes: &[u8]) -> Vec<u8> {
257    let encoded = BASE64_STANDARD.encode(payload_bytes);
258    let mut frame = Vec::with_capacity(
259        NESTED_FRAME_HEADER.len() + encoded.len() + NESTED_FRAME_TERMINATOR.len(),
260    );
261    frame.extend_from_slice(NESTED_FRAME_HEADER);
262    frame.extend_from_slice(encoded.as_bytes());
263    frame.extend_from_slice(NESTED_FRAME_TERMINATOR);
264    frame
265}
266
267pub fn decode_base64(encoded: &[u8]) -> Option<Vec<u8>> {
268    BASE64_DECODER.decode(encoded).ok()
269}
270
271const MAX_PARTIAL_FRAME_BYTES: usize = 1024 * 1024;
272
273enum FrameScanStatus {
274    Complete(usize),
275    NeedMore,
276    Diverged,
277}
278
279fn frame_scan_status(buf: &[u8]) -> FrameScanStatus {
280    if buf.len() < NESTED_FRAME_HEADER.len() {
281        return if NESTED_FRAME_HEADER.starts_with(buf) {
282            FrameScanStatus::NeedMore
283        } else {
284            FrameScanStatus::Diverged
285        };
286    }
287    if &buf[..NESTED_FRAME_HEADER.len()] != NESTED_FRAME_HEADER {
288        return FrameScanStatus::Diverged;
289    }
290    let mut i = NESTED_FRAME_HEADER.len();
291    while i < buf.len() {
292        match buf[i] {
293            0x1b => match buf.get(i + 1) {
294                Some(&b'\\') => return FrameScanStatus::Complete(i + 2),
295                Some(_) => return FrameScanStatus::Diverged,
296                None => return FrameScanStatus::NeedMore,
297            },
298            _ => i += 1,
299        }
300    }
301    FrameScanStatus::NeedMore
302}
303
304#[derive(Debug, Default)]
305pub struct NestedFrameExtractor {
306    partial_frame: Vec<u8>,
307}
308
309impl NestedFrameExtractor {
310    pub fn new() -> Self {
311        Self::default()
312    }
313
314    pub fn extract(&mut self, bytes: &[u8]) -> (Vec<u8>, Vec<Vec<u8>>) {
315        let mut decoded_payloads = Vec::new();
316        if self.partial_frame.is_empty() && !bytes.contains(&0x1b) {
317            return (bytes.to_vec(), decoded_payloads);
318        }
319        let mut working = std::mem::take(&mut self.partial_frame);
320        working.extend_from_slice(bytes);
321        let mut cleaned = Vec::with_capacity(working.len());
322        let mut i = 0;
323        while i < working.len() {
324            let rest = &working[i..];
325            if rest[0] == 0x1b {
326                match frame_scan_status(rest) {
327                    FrameScanStatus::Complete(len) => {
328                        let encoded_payload =
329                            &rest[NESTED_FRAME_HEADER.len()..len - NESTED_FRAME_TERMINATOR.len()];
330                        if let Some(payload) = decode_base64(encoded_payload) {
331                            decoded_payloads.push(payload);
332                        }
333                        i += len;
334                        continue;
335                    },
336                    FrameScanStatus::NeedMore => {
337                        let tail = rest.to_vec();
338                        if tail.len() > MAX_PARTIAL_FRAME_BYTES {
339                            cleaned.extend_from_slice(&tail);
340                        } else {
341                            self.partial_frame = tail;
342                        }
343                        return (cleaned, decoded_payloads);
344                    },
345                    FrameScanStatus::Diverged => {},
346                }
347            }
348            cleaned.push(working[i]);
349            i += 1;
350        }
351        (cleaned, decoded_payloads)
352    }
353
354    pub fn partial_bytes(&self) -> &[u8] {
355        &self.partial_frame
356    }
357
358    pub fn take_partial(&mut self) -> Vec<u8> {
359        std::mem::take(&mut self.partial_frame)
360    }
361}
362
363#[cfg(test)]
364mod tests {
365    use super::*;
366    use crate::data::BareKey;
367
368    fn all_message_arms() -> Vec<NestedSessionMessage> {
369        vec![
370            NestedSessionMessage::Announce {
371                session_name: "guest".to_owned(),
372                capabilities: vec![NestedSessionCapability::NestedControl],
373            },
374            NestedSessionMessage::FocusHost {
375                direction: Some(Direction::Left),
376            },
377            NestedSessionMessage::FocusHost { direction: None },
378            NestedSessionMessage::ToggleHostFullscreen { fullscreen: true },
379            NestedSessionMessage::Pong,
380            NestedSessionMessage::Bye,
381            NestedSessionMessage::AnnounceAck {
382                ancestry: vec!["outer".to_owned(), "middle".to_owned()],
383                capabilities: vec![NestedSessionCapability::NestedControl],
384                descend_keys: vec![
385                    KeyWithModifier::new(BareKey::Char('o')).with_ctrl_modifier(),
386                    KeyWithModifier::new(BareKey::Down),
387                ],
388            },
389            NestedSessionMessage::FocusGained {
390                from_direction: Some(Direction::Right),
391            },
392            NestedSessionMessage::FocusLost,
393            NestedSessionMessage::FullscreenState { fullscreen: true },
394            NestedSessionMessage::AncestryUpdate {
395                ancestry: vec!["outer".to_owned()],
396            },
397            NestedSessionMessage::Ping,
398            NestedSessionMessage::ShortcutUpdate {
399                ascend_keys: vec![
400                    KeyWithModifier::new(BareKey::Char('o')).with_ctrl_modifier(),
401                    KeyWithModifier::new(BareKey::Up),
402                ],
403                descend_keys: vec![],
404            },
405        ]
406    }
407
408    #[test]
409    fn payload_roundtrip_preserves_every_arm() {
410        for message in all_message_arms() {
411            let decoded = decode_payload(&encode_payload(&message));
412            assert_eq!(decoded, Some(message));
413        }
414    }
415
416    #[test]
417    fn focus_direction_roundtrip_preserves_every_direction() {
418        let directions = [
419            None,
420            Some(Direction::Left),
421            Some(Direction::Right),
422            Some(Direction::Up),
423            Some(Direction::Down),
424        ];
425        for direction in directions {
426            let focus_host = NestedSessionMessage::FocusHost { direction };
427            assert_eq!(
428                decode_payload(&encode_payload(&focus_host)),
429                Some(focus_host)
430            );
431            let focus_gained = NestedSessionMessage::FocusGained {
432                from_direction: direction,
433            };
434            assert_eq!(
435                decode_payload(&encode_payload(&focus_gained)),
436                Some(focus_gained)
437            );
438        }
439    }
440
441    #[test]
442    fn frame_roundtrip_preserves_message() {
443        let message = NestedSessionMessage::Announce {
444            session_name: "guest".to_owned(),
445            capabilities: vec![NestedSessionCapability::NestedControl],
446        };
447        let frame = encode_frame(&message);
448        assert!(frame.starts_with(NESTED_FRAME_HEADER));
449        assert!(frame.ends_with(NESTED_FRAME_TERMINATOR));
450        let encoded_payload =
451            &frame[NESTED_FRAME_HEADER.len()..frame.len() - NESTED_FRAME_TERMINATOR.len()];
452        let payload = decode_base64(encoded_payload).unwrap();
453        assert_eq!(decode_payload(&payload), Some(message));
454    }
455
456    #[test]
457    fn garbage_base64_is_rejected() {
458        assert_eq!(decode_base64(b"!!!not-base64!!!"), None);
459    }
460
461    #[test]
462    fn truncated_payload_is_rejected() {
463        let payload = encode_payload(&NestedSessionMessage::Announce {
464            session_name: "a-long-session-name-to-truncate".to_owned(),
465            capabilities: vec![],
466        });
467        assert_eq!(decode_payload(&payload[..payload.len() - 5]), None);
468    }
469
470    #[test]
471    fn unknown_oneof_arm_is_rejected() {
472        let unknown_field_bytes = [0xE2, 0x03, 0x00];
473        assert_eq!(decode_payload(&unknown_field_bytes), None);
474    }
475
476    #[test]
477    fn empty_payload_is_rejected() {
478        assert_eq!(decode_payload(&[]), None);
479    }
480}