kcode_rust_bins/
protocol.rs1use 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}