Skip to main content

kcode_rust_bins/
protocol.rs

1use std::io::{Read, Write};
2
3use crate::model::{Error, Result, RustBinInput, RustBinOutput};
4
5const MAGIC: [u8; 8] = *b"KRBIN\x02\0\0";
6const MAX_TEXT_BYTES: usize = 16 * 1024 * 1024;
7const MAX_OBJECTS: usize = 1024;
8const MAX_OBJECT_BYTES: usize = 512 * 1024 * 1024;
9const MAX_TOTAL_BYTES: usize = 1024 * 1024 * 1024;
10
11pub fn read_input() -> Result<RustBinInput> {
12    let stdin = std::io::stdin();
13    let mut input = stdin.lock();
14    let message = read_message(&mut input)?;
15    Ok(RustBinInput {
16        text: message.0,
17        objects: message.1,
18    })
19}
20
21pub fn write_output(output: &RustBinOutput) -> Result<()> {
22    let stdout = std::io::stdout();
23    let mut writer = stdout.lock();
24    write_message(&mut writer, &output.text, &output.objects)?;
25    writer
26        .flush()
27        .map_err(|error| Error::new("protocol", format!("flush stdout: {error}")))
28}
29
30pub(crate) fn write_message(
31    writer: &mut impl Write,
32    text: &str,
33    objects: &[Vec<u8>],
34) -> Result<()> {
35    validate_parts(text, objects)?;
36    writer
37        .write_all(&MAGIC)
38        .and_then(|()| writer.write_all(&(text.len() as u64).to_le_bytes()))
39        .and_then(|()| writer.write_all(text.as_bytes()))
40        .and_then(|()| writer.write_all(&(objects.len() as u32).to_le_bytes()))
41        .map_err(|error| Error::new("protocol", format!("write frame header: {error}")))?;
42    for object in objects {
43        writer
44            .write_all(&(object.len() as u64).to_le_bytes())
45            .and_then(|()| writer.write_all(object))
46            .map_err(|error| Error::new("protocol", format!("write frame object: {error}")))?;
47    }
48    Ok(())
49}
50
51pub(crate) fn read_output(reader: &mut impl Read) -> Result<RustBinOutput> {
52    let (text, objects) = read_message(reader)?;
53    Ok(RustBinOutput { text, objects })
54}
55
56fn read_message(reader: &mut impl Read) -> Result<(String, Vec<Vec<u8>>)> {
57    let mut magic = [0_u8; 8];
58    read_exact(reader, &mut magic, "magic")?;
59    if magic != MAGIC {
60        return Err(Error::new("protocol", "invalid frame magic or revision"));
61    }
62
63    let text_length = read_u64(reader, "text length")?;
64    let text_length = usize::try_from(text_length)
65        .map_err(|_| Error::new("protocol", "text length overflows usize"))?;
66    if text_length > MAX_TEXT_BYTES {
67        return Err(Error::new("protocol", "text exceeds 16 MiB"));
68    }
69    let mut text = vec![0_u8; text_length];
70    read_exact(reader, &mut text, "text")?;
71    let text = String::from_utf8(text)
72        .map_err(|error| Error::new("protocol", format!("text is not UTF-8: {error}")))?;
73
74    let count = read_u32(reader, "object count")? as usize;
75    if count > MAX_OBJECTS {
76        return Err(Error::new("protocol", "object count exceeds 1,024"));
77    }
78    let mut total = text_length;
79    let mut objects = Vec::with_capacity(count);
80    for _ in 0..count {
81        let length = read_u64(reader, "object length")?;
82        let length = usize::try_from(length)
83            .map_err(|_| Error::new("protocol", "object length overflows usize"))?;
84        if length > MAX_OBJECT_BYTES {
85            return Err(Error::new("protocol", "object exceeds 512 MiB"));
86        }
87        total = total
88            .checked_add(length)
89            .ok_or_else(|| Error::new("protocol", "aggregate payload length overflow"))?;
90        if total > MAX_TOTAL_BYTES {
91            return Err(Error::new("protocol", "aggregate payload exceeds 1 GiB"));
92        }
93        let mut object = vec![0_u8; length];
94        read_exact(reader, &mut object, "object bytes")?;
95        objects.push(object);
96    }
97
98    let mut trailing = [0_u8; 1];
99    match reader.read(&mut trailing) {
100        Ok(0) => Ok((text, objects)),
101        Ok(_) => Err(Error::new("protocol", "trailing bytes after frame")),
102        Err(error) => Err(Error::new(
103            "protocol",
104            format!("read frame terminator: {error}"),
105        )),
106    }
107}
108
109fn validate_parts(text: &str, objects: &[Vec<u8>]) -> Result<()> {
110    if text.len() > MAX_TEXT_BYTES {
111        return Err(Error::new("protocol", "text exceeds 16 MiB"));
112    }
113    if objects.len() > MAX_OBJECTS {
114        return Err(Error::new("protocol", "object count exceeds 1,024"));
115    }
116    let mut total = text.len();
117    for object in objects {
118        if object.len() > MAX_OBJECT_BYTES {
119            return Err(Error::new("protocol", "object exceeds 512 MiB"));
120        }
121        total = total
122            .checked_add(object.len())
123            .ok_or_else(|| Error::new("protocol", "aggregate payload length overflow"))?;
124        if total > MAX_TOTAL_BYTES {
125            return Err(Error::new("protocol", "aggregate payload exceeds 1 GiB"));
126        }
127    }
128    Ok(())
129}
130
131fn read_u64(reader: &mut impl Read, label: &str) -> Result<u64> {
132    let mut bytes = [0_u8; 8];
133    read_exact(reader, &mut bytes, label)?;
134    Ok(u64::from_le_bytes(bytes))
135}
136
137fn read_u32(reader: &mut impl Read, label: &str) -> Result<u32> {
138    let mut bytes = [0_u8; 4];
139    read_exact(reader, &mut bytes, label)?;
140    Ok(u32::from_le_bytes(bytes))
141}
142
143fn read_exact(reader: &mut impl Read, bytes: &mut [u8], label: &str) -> Result<()> {
144    reader
145        .read_exact(bytes)
146        .map_err(|error| Error::new("protocol", format!("read {label}: {error}")))
147}
148
149#[cfg(test)]
150mod tests {
151    use super::{read_output, write_message};
152
153    #[test]
154    fn round_trips_text_and_ordered_objects() {
155        let mut bytes = Vec::new();
156        write_message(
157            &mut bytes,
158            "answer: 42",
159            &[b"first".to_vec(), b"second".to_vec()],
160        )
161        .unwrap();
162        let output = read_output(&mut bytes.as_slice()).unwrap();
163        assert_eq!(output.text, "answer: 42");
164        assert_eq!(output.objects, [b"first".to_vec(), b"second".to_vec()]);
165    }
166
167    #[test]
168    fn rejects_trailing_bytes() {
169        let mut bytes = Vec::new();
170        write_message(&mut bytes, "plain text", &[]).unwrap();
171        bytes.push(0);
172        assert!(read_output(&mut bytes.as_slice()).is_err());
173    }
174}