#[cfg(test)]
mod tests;
use serde::de::{Deserialize, Deserializer, IgnoredAny, MapAccess, Visitor};
use serde_json::error::Category;
use std::{collections::TryReserveError, fmt};
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub enum ResponseFormat {
Json,
Hex,
LabeledHex,
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub struct ResponseLimits {
pub input_bytes: usize,
pub decoded_bytes: usize,
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub enum JsonErrorKind {
Syntax,
Data,
EndOfInput,
Io,
}
#[derive(Debug)]
pub enum ResponseError {
InputLimit {
limit: usize,
},
InvalidUtf8 {
offset: usize,
},
Json {
kind: JsonErrorKind,
line: usize,
column: usize,
},
MissingResponseBytes,
MissingHexLabel,
EmptyHex,
InvalidHex {
offset: usize,
},
OddHexLength,
DecodedLimit {
limit: usize,
},
Allocation(TryReserveError),
}
impl fmt::Display for ResponseError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::InputLimit { limit } => write!(f, "response input exceeds {limit} bytes"),
Self::InvalidUtf8 { offset } => {
write!(f, "response JSON is not UTF-8 at byte {offset}")
}
Self::Json { kind, line, column } => {
write!(f, "response JSON failed ({kind:?}) at {line}:{column}")
}
Self::MissingResponseBytes => f.write_str("response JSON lacks response_bytes"),
Self::MissingHexLabel => f.write_str("response hex label is missing"),
Self::EmptyHex => f.write_str("response hex is empty"),
Self::InvalidHex { offset } => write!(f, "invalid response hex at byte {offset}"),
Self::OddHexLength => f.write_str("response hex has an odd digit count"),
Self::DecodedLimit { limit } => write!(f, "decoded response exceeds {limit} bytes"),
Self::Allocation(source) => write!(f, "response allocation failed: {source}"),
}
}
}
impl std::error::Error for ResponseError {
fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
match self {
Self::Allocation(source) => Some(source),
_ => None,
}
}
}
pub fn decode(
input: &[u8],
format: ResponseFormat,
limits: ResponseLimits,
) -> Result<Vec<u8>, ResponseError> {
if input.len() > limits.input_bytes {
return Err(ResponseError::InputLimit {
limit: limits.input_bytes,
});
}
match format {
ResponseFormat::Json => {
std::str::from_utf8(input).map_err(|error| ResponseError::InvalidUtf8 {
offset: error.valid_up_to(),
})?;
let envelope: Envelope =
serde_json::from_slice(input).map_err(|error| json_error(&error))?;
let hex = envelope.0.ok_or(ResponseError::MissingResponseBytes)?;
decode_hex(hex.as_bytes(), false, limits.decoded_bytes)
}
ResponseFormat::Hex => decode_text_hex(input, limits.decoded_bytes),
ResponseFormat::LabeledHex => {
let input = input.trim_ascii_start();
let hex = input
.strip_prefix(b"response (hex):")
.ok_or(ResponseError::MissingHexLabel)?;
decode_text_hex(hex, limits.decoded_bytes)
}
}
}
fn decode_text_hex(hex: &[u8], limit: usize) -> Result<Vec<u8>, ResponseError> {
if hex.iter().all(u8::is_ascii_whitespace) {
return Err(ResponseError::EmptyHex);
}
decode_hex(hex, true, limit)
}
fn decode_hex(hex: &[u8], whitespace: bool, limit: usize) -> Result<Vec<u8>, ResponseError> {
let mut digits = 0;
for (offset, &byte) in hex.iter().enumerate() {
if byte.is_ascii_hexdigit() {
digits += 1;
} else if !(whitespace && byte.is_ascii_whitespace()) {
return Err(ResponseError::InvalidHex { offset });
}
}
if digits % 2 != 0 {
return Err(ResponseError::OddHexLength);
}
if digits / 2 > limit {
return Err(ResponseError::DecodedLimit { limit });
}
let mut bytes = Vec::new();
bytes
.try_reserve_exact(digits / 2)
.map_err(ResponseError::Allocation)?;
let mut high = None;
for byte in hex.iter().copied().filter(u8::is_ascii_hexdigit) {
let nibble = if byte <= b'9' {
byte - b'0'
} else {
byte.to_ascii_lowercase() - b'a' + 10
};
if let Some(previous) = high.take() {
bytes.push((previous << 4) | nibble);
} else {
high = Some(nibble);
}
}
Ok(bytes)
}
fn json_error(error: &serde_json::Error) -> ResponseError {
let kind = match error.classify() {
Category::Io => JsonErrorKind::Io,
Category::Syntax => JsonErrorKind::Syntax,
Category::Data => JsonErrorKind::Data,
Category::Eof => JsonErrorKind::EndOfInput,
};
ResponseError::Json {
kind,
line: error.line(),
column: error.column(),
}
}
struct Envelope(Option<String>);
impl<'de> Deserialize<'de> for Envelope {
fn deserialize<D: Deserializer<'de>>(deserializer: D) -> Result<Self, D::Error> {
struct EnvelopeVisitor;
impl<'de> Visitor<'de> for EnvelopeVisitor {
type Value = Envelope;
fn expecting(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str("an ICP response object")
}
fn visit_map<M: MapAccess<'de>>(self, mut map: M) -> Result<Envelope, M::Error> {
let mut response = None;
let mut seen = false;
while let Some(key) = map.next_key::<String>()? {
if key == "response_bytes" {
if seen {
return Err(serde::de::Error::duplicate_field("response_bytes"));
}
seen = true;
response = map.next_value::<Option<String>>()?;
} else {
map.next_value::<IgnoredAny>()?;
}
}
Ok(Envelope(response))
}
}
deserializer.deserialize_map(EnvelopeVisitor)
}
}