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