1#[cfg(test)]
8mod tests;
9
10use serde::de::{Deserialize, Deserializer, IgnoredAny, MapAccess, Visitor};
11use serde_json::error::Category;
12use std::{collections::TryReserveError, fmt};
13
14#[derive(Clone, Copy, Debug, Eq, PartialEq)]
16pub enum ResponseFormat {
17 Json,
21 Hex,
24 LabeledHex,
28}
29
30#[derive(Clone, Copy, Debug, Eq, PartialEq)]
32pub struct ResponseLimits {
33 pub input_bytes: usize,
35 pub decoded_bytes: usize,
37}
38
39#[derive(Clone, Copy, Debug, Eq, PartialEq)]
41pub enum JsonErrorKind {
42 Syntax,
44 Data,
46 EndOfInput,
48 Io,
50}
51
52#[derive(Debug)]
54pub enum ResponseError {
55 InputLimit {
57 limit: usize,
59 },
60 InvalidUtf8 {
62 offset: usize,
64 },
65 Json {
67 kind: JsonErrorKind,
69 line: usize,
71 column: usize,
73 },
74 MissingResponseBytes,
76 MissingHexLabel,
78 EmptyHex,
80 InvalidHex {
82 offset: usize,
85 },
86 OddHexLength,
88 DecodedLimit {
90 limit: usize,
92 },
93 Allocation(TryReserveError),
95}
96
97impl fmt::Display for ResponseError {
98 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
99 match self {
100 Self::InputLimit { limit } => write!(f, "response input exceeds {limit} bytes"),
101 Self::InvalidUtf8 { offset } => {
102 write!(f, "response JSON is not UTF-8 at byte {offset}")
103 }
104 Self::Json { kind, line, column } => {
105 write!(f, "response JSON failed ({kind:?}) at {line}:{column}")
106 }
107 Self::MissingResponseBytes => f.write_str("response JSON lacks response_bytes"),
108 Self::MissingHexLabel => f.write_str("response hex label is missing"),
109 Self::EmptyHex => f.write_str("response hex is empty"),
110 Self::InvalidHex { offset } => write!(f, "invalid response hex at byte {offset}"),
111 Self::OddHexLength => f.write_str("response hex has an odd digit count"),
112 Self::DecodedLimit { limit } => write!(f, "decoded response exceeds {limit} bytes"),
113 Self::Allocation(source) => write!(f, "response allocation failed: {source}"),
114 }
115 }
116}
117
118impl std::error::Error for ResponseError {
119 fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
120 match self {
121 Self::Allocation(source) => Some(source),
122 _ => None,
123 }
124 }
125}
126
127pub fn decode(
140 input: &[u8],
141 format: ResponseFormat,
142 limits: ResponseLimits,
143) -> Result<Vec<u8>, ResponseError> {
144 if input.len() > limits.input_bytes {
145 return Err(ResponseError::InputLimit {
146 limit: limits.input_bytes,
147 });
148 }
149 match format {
150 ResponseFormat::Json => {
151 std::str::from_utf8(input).map_err(|error| ResponseError::InvalidUtf8 {
152 offset: error.valid_up_to(),
153 })?;
154 let envelope: Envelope =
155 serde_json::from_slice(input).map_err(|error| json_error(&error))?;
156 let hex = envelope.0.ok_or(ResponseError::MissingResponseBytes)?;
157 decode_hex(hex.as_bytes(), false, limits.decoded_bytes)
158 }
159 ResponseFormat::Hex => decode_text_hex(input, limits.decoded_bytes),
160 ResponseFormat::LabeledHex => {
161 let input = input.trim_ascii_start();
162 let hex = input
163 .strip_prefix(b"response (hex):")
164 .ok_or(ResponseError::MissingHexLabel)?;
165 decode_text_hex(hex, limits.decoded_bytes)
166 }
167 }
168}
169
170fn decode_text_hex(hex: &[u8], limit: usize) -> Result<Vec<u8>, ResponseError> {
171 if hex.iter().all(u8::is_ascii_whitespace) {
172 return Err(ResponseError::EmptyHex);
173 }
174 decode_hex(hex, true, limit)
175}
176
177fn decode_hex(hex: &[u8], whitespace: bool, limit: usize) -> Result<Vec<u8>, ResponseError> {
178 let mut digits = 0;
179 for (offset, &byte) in hex.iter().enumerate() {
180 if byte.is_ascii_hexdigit() {
181 digits += 1;
182 } else if !(whitespace && byte.is_ascii_whitespace()) {
183 return Err(ResponseError::InvalidHex { offset });
184 }
185 }
186 if digits % 2 != 0 {
187 return Err(ResponseError::OddHexLength);
188 }
189 if digits / 2 > limit {
190 return Err(ResponseError::DecodedLimit { limit });
191 }
192 let mut bytes = Vec::new();
193 bytes
194 .try_reserve_exact(digits / 2)
195 .map_err(ResponseError::Allocation)?;
196 let mut high = None;
197 for byte in hex.iter().copied().filter(u8::is_ascii_hexdigit) {
198 let nibble = if byte <= b'9' {
200 byte - b'0'
201 } else {
202 byte.to_ascii_lowercase() - b'a' + 10
203 };
204 if let Some(previous) = high.take() {
205 bytes.push((previous << 4) | nibble);
206 } else {
207 high = Some(nibble);
208 }
209 }
210 Ok(bytes)
211}
212
213fn json_error(error: &serde_json::Error) -> ResponseError {
214 let kind = match error.classify() {
215 Category::Io => JsonErrorKind::Io,
216 Category::Syntax => JsonErrorKind::Syntax,
217 Category::Data => JsonErrorKind::Data,
218 Category::Eof => JsonErrorKind::EndOfInput,
219 };
220 ResponseError::Json {
221 kind,
222 line: error.line(),
223 column: error.column(),
224 }
225}
226
227struct Envelope(Option<String>);
228
229impl<'de> Deserialize<'de> for Envelope {
230 fn deserialize<D: Deserializer<'de>>(deserializer: D) -> Result<Self, D::Error> {
231 struct EnvelopeVisitor;
232 impl<'de> Visitor<'de> for EnvelopeVisitor {
233 type Value = Envelope;
234
235 fn expecting(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
236 f.write_str("an ICP response object")
237 }
238
239 fn visit_map<M: MapAccess<'de>>(self, mut map: M) -> Result<Envelope, M::Error> {
240 let mut response = None;
241 let mut seen = false;
242 while let Some(key) = map.next_key::<String>()? {
243 if key == "response_bytes" {
244 if seen {
245 return Err(serde::de::Error::duplicate_field("response_bytes"));
246 }
247 seen = true;
248 response = map.next_value::<Option<String>>()?;
249 } else {
250 map.next_value::<IgnoredAny>()?;
251 }
252 }
253 Ok(Envelope(response))
254 }
255 }
256 deserializer.deserialize_map(EnvelopeVisitor)
257 }
258}