Skip to main content

kcode_rust_bins/
protocol.rs

1use std::io::{Read, Write};
2
3use crate::model::{Error, Result, RustBinInput, RustBinOutput, validate_json};
4
5const MAGIC: [u8; 8] = *b"KRBIN\x01\0\0";
6const MAX_JSON_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        json: 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.json, &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    json: &str,
33    objects: &[Vec<u8>],
34) -> Result<()> {
35    validate_parts(json, objects)?;
36    writer
37        .write_all(&MAGIC)
38        .and_then(|()| writer.write_all(&(json.len() as u64).to_le_bytes()))
39        .and_then(|()| writer.write_all(json.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 (json, objects) = read_message(reader)?;
53    Ok(RustBinOutput { json, 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 json_length = read_u64(reader, "JSON length")?;
64    let json_length = usize::try_from(json_length)
65        .map_err(|_| Error::new("protocol", "JSON length overflows usize"))?;
66    if json_length > MAX_JSON_BYTES {
67        return Err(Error::new("protocol", "JSON exceeds 16 MiB"));
68    }
69    let mut json = vec![0_u8; json_length];
70    read_exact(reader, &mut json, "JSON")?;
71    let json = String::from_utf8(json)
72        .map_err(|error| Error::new("protocol", format!("JSON is not UTF-8: {error}")))?;
73    validate_json(&json)?;
74
75    let count = read_u32(reader, "object count")? as usize;
76    if count > MAX_OBJECTS {
77        return Err(Error::new("protocol", "object count exceeds 1,024"));
78    }
79    let mut total = json_length;
80    let mut objects = Vec::with_capacity(count);
81    for _ in 0..count {
82        let length = read_u64(reader, "object length")?;
83        let length = usize::try_from(length)
84            .map_err(|_| Error::new("protocol", "object length overflows usize"))?;
85        if length > MAX_OBJECT_BYTES {
86            return Err(Error::new("protocol", "object exceeds 512 MiB"));
87        }
88        total = total
89            .checked_add(length)
90            .ok_or_else(|| Error::new("protocol", "aggregate payload length overflow"))?;
91        if total > MAX_TOTAL_BYTES {
92            return Err(Error::new("protocol", "aggregate payload exceeds 1 GiB"));
93        }
94        let mut object = vec![0_u8; length];
95        read_exact(reader, &mut object, "object bytes")?;
96        objects.push(object);
97    }
98
99    let mut trailing = [0_u8; 1];
100    match reader.read(&mut trailing) {
101        Ok(0) => Ok((json, objects)),
102        Ok(_) => Err(Error::new("protocol", "trailing bytes after frame")),
103        Err(error) => Err(Error::new(
104            "protocol",
105            format!("read frame terminator: {error}"),
106        )),
107    }
108}
109
110fn validate_parts(json: &str, objects: &[Vec<u8>]) -> Result<()> {
111    validate_json(json)?;
112    if json.len() > MAX_JSON_BYTES {
113        return Err(Error::new("protocol", "JSON exceeds 16 MiB"));
114    }
115    if objects.len() > MAX_OBJECTS {
116        return Err(Error::new("protocol", "object count exceeds 1,024"));
117    }
118    let mut total = json.len();
119    for object in objects {
120        if object.len() > MAX_OBJECT_BYTES {
121            return Err(Error::new("protocol", "object exceeds 512 MiB"));
122        }
123        total = total
124            .checked_add(object.len())
125            .ok_or_else(|| Error::new("protocol", "aggregate payload length overflow"))?;
126        if total > MAX_TOTAL_BYTES {
127            return Err(Error::new("protocol", "aggregate payload exceeds 1 GiB"));
128        }
129    }
130    Ok(())
131}
132
133fn read_u64(reader: &mut impl Read, label: &str) -> Result<u64> {
134    let mut bytes = [0_u8; 8];
135    read_exact(reader, &mut bytes, label)?;
136    Ok(u64::from_le_bytes(bytes))
137}
138
139fn read_u32(reader: &mut impl Read, label: &str) -> Result<u32> {
140    let mut bytes = [0_u8; 4];
141    read_exact(reader, &mut bytes, label)?;
142    Ok(u32::from_le_bytes(bytes))
143}
144
145fn read_exact(reader: &mut impl Read, bytes: &mut [u8], label: &str) -> Result<()> {
146    reader
147        .read_exact(bytes)
148        .map_err(|error| Error::new("protocol", format!("read {label}: {error}")))
149}
150
151#[cfg(test)]
152mod tests {
153    use super::{read_output, write_message};
154
155    #[test]
156    fn round_trips_one_json_value_and_ordered_objects() {
157        let mut bytes = Vec::new();
158        write_message(
159            &mut bytes,
160            "{\"answer\":42}",
161            &[b"first".to_vec(), b"second".to_vec()],
162        )
163        .unwrap();
164        let output = read_output(&mut bytes.as_slice()).unwrap();
165        assert_eq!(output.json, "{\"answer\":42}");
166        assert_eq!(output.objects, [b"first".to_vec(), b"second".to_vec()]);
167    }
168
169    #[test]
170    fn rejects_invalid_json_and_trailing_bytes() {
171        let mut bytes = Vec::new();
172        assert!(write_message(&mut bytes, "not-json", &[]).is_err());
173        write_message(&mut bytes, "null", &[]).unwrap();
174        bytes.push(0);
175        assert!(read_output(&mut bytes.as_slice()).is_err());
176    }
177}