use gwk_domain::protocol::KernelErrorCode;
use serde::Deserialize;
use serde::de::{DeserializeOwned, MapAccess, SeqAccess, Visitor};
use super::WireError;
const DUPLICATE_KEY: &str = "duplicate object key";
struct NoDuplicateKeys(serde_json::Value);
impl<'de> Deserialize<'de> for NoDuplicateKeys {
fn deserialize<D: serde::Deserializer<'de>>(d: D) -> Result<Self, D::Error> {
d.deserialize_any(StrictValue).map(NoDuplicateKeys)
}
}
struct StrictValue;
impl<'de> Visitor<'de> for StrictValue {
type Value = serde_json::Value;
fn expecting(&self, f: &mut std::fmt::Formatter) -> std::fmt::Result {
f.write_str("a JSON value with no duplicate object key")
}
fn visit_unit<E>(self) -> Result<Self::Value, E> {
Ok(serde_json::Value::Null)
}
fn visit_none<E>(self) -> Result<Self::Value, E> {
Ok(serde_json::Value::Null)
}
fn visit_bool<E>(self, v: bool) -> Result<Self::Value, E> {
Ok(serde_json::Value::Bool(v))
}
fn visit_i64<E>(self, v: i64) -> Result<Self::Value, E> {
Ok(serde_json::Value::Number(v.into()))
}
fn visit_u64<E>(self, v: u64) -> Result<Self::Value, E> {
Ok(serde_json::Value::Number(v.into()))
}
fn visit_f64<E: serde::de::Error>(self, v: f64) -> Result<Self::Value, E> {
serde_json::Number::from_f64(v)
.map(serde_json::Value::Number)
.ok_or_else(|| E::custom("a non-finite number is not JSON"))
}
fn visit_str<E>(self, v: &str) -> Result<Self::Value, E> {
Ok(serde_json::Value::String(v.to_owned()))
}
fn visit_seq<A: SeqAccess<'de>>(self, mut seq: A) -> Result<Self::Value, A::Error> {
let mut out = Vec::new();
while let Some(NoDuplicateKeys(value)) = seq.next_element()? {
out.push(value);
}
Ok(serde_json::Value::Array(out))
}
fn visit_map<A: MapAccess<'de>>(self, mut map: A) -> Result<Self::Value, A::Error> {
let mut out = serde_json::Map::new();
while let Some(key) = map.next_key::<String>()? {
let NoDuplicateKeys(value) = map.next_value()?;
if out.insert(key.clone(), value).is_some() {
return Err(serde::de::Error::custom(format!("{DUPLICATE_KEY} {key:?}")));
}
}
Ok(serde_json::Value::Object(out))
}
}
pub fn decode<T: DeserializeOwned>(body: &[u8]) -> Result<T, WireError> {
let text = std::str::from_utf8(body).map_err(|e| {
WireError::new(
KernelErrorCode::Schema,
format!("frame body is not UTF-8: {e}"),
)
})?;
let mut stream = serde_json::Deserializer::from_str(text).into_iter::<NoDuplicateKeys>();
let NoDuplicateKeys(value) = stream
.next()
.transpose()
.map_err(|e| {
let text = e.to_string();
let code = if text.starts_with(DUPLICATE_KEY) {
KernelErrorCode::DuplicateKey
} else {
KernelErrorCode::Schema
};
WireError::new(code, format!("frame body: {text}"))
})?
.ok_or_else(|| WireError::new(KernelErrorCode::Schema, "frame body holds no JSON value"))?;
let rest = text[stream.byte_offset()..].trim();
if !rest.is_empty() {
return Err(WireError::new(
KernelErrorCode::Schema,
format!("{} trailing bytes after the JSON value", rest.len()),
));
}
serde_json::from_value(value)
.map_err(|e| WireError::new(KernelErrorCode::Validation, format!("frame body: {e}")))
}
#[cfg(test)]
mod tests {
use gwk_domain::protocol::{ClientControl, KernelRequest};
use super::*;
fn control(raw: &str) -> Result<ClientControl, WireError> {
decode::<ClientControl>(raw.as_bytes())
}
#[test]
fn a_well_formed_request_decodes() {
let ok = control(r#"{"type":"request","request_id":"r-1","request":{"type":"health"}}"#)
.expect("decode");
match ok {
ClientControl::Request { request, .. } => {
assert_eq!(request, KernelRequest::Health {});
}
other => panic!("{other:?}"),
}
}
#[test]
fn a_field_less_request_still_refuses_an_extra_field() {
let error = control(
r#"{"type":"request","request_id":"r-1","request":{"type":"health","limit":500}}"#,
)
.expect_err("extra field accepted");
assert_eq!(error.code, KernelErrorCode::Validation);
assert!(error.message.contains("limit"), "{error}");
}
#[test]
fn an_unknown_field_on_a_known_request_is_refused() {
let error = control(
r#"{"type":"request","request_id":"r-1","request":{"type":"read_events","limit":5,"depth":2}}"#,
)
.expect_err("unknown field accepted");
assert_eq!(error.code, KernelErrorCode::Validation);
assert!(error.message.contains("depth"), "{error}");
}
#[test]
fn a_duplicate_key_is_refused_at_every_depth() {
for raw in [
r#"{"type":"request","type":"request","request_id":"r-1","request":{"type":"health"}}"#,
r#"{"type":"request","request_id":"r-1","request":{"type":"submit_command","envelope":{"command_id":"c","project_id":"p","command_type":"ingest_record","schema_version":1,"issued_at":"t","actor":{"kind":"kernel"},"origin":{"system":"gw"},"idempotency_key":"k","payload":{"deep":{"cost":1,"cost":2}}}}}"#,
] {
let error = control(raw).expect_err("duplicate key accepted");
assert_eq!(error.code, KernelErrorCode::DuplicateKey);
assert!(error.message.contains(DUPLICATE_KEY), "{error}");
}
}
#[test]
fn a_legitimate_repeat_in_two_different_objects_is_not_a_duplicate() {
control(
r#"{"type":"request","request_id":"r-1","request":{"type":"submit_command","envelope":{"command_id":"c","project_id":"p","command_type":"ingest_record","schema_version":1,"issued_at":"t","actor":{"kind":"kernel"},"origin":{"system":"gw"},"idempotency_key":"k","payload":{"rows":[{"kind":"a"},{"kind":"b"}]}}}}"#,
)
.expect("a repeated name in sibling objects was refused");
}
#[test]
fn bytes_that_are_not_utf8_are_refused_before_parsing() {
let error = decode::<ClientControl>(&[0x7b, 0xff, 0x7d]).expect_err("invalid UTF-8");
assert_eq!(error.code, KernelErrorCode::Schema);
assert!(error.message.contains("UTF-8"), "{error}");
}
#[test]
fn a_second_value_in_one_frame_is_trailing_garbage() {
let error = control(
r#"{"type":"request","request_id":"r-1","request":{"type":"health"}}{"type":"request","request_id":"r-2","request":{"type":"status"}}"#,
)
.expect_err("two values in one frame accepted");
assert_eq!(error.code, KernelErrorCode::Schema);
assert!(error.message.contains("trailing"), "{error}");
}
#[test]
fn an_empty_body_is_a_refusal_and_not_a_default() {
let error = decode::<ClientControl>(b" ").expect_err("empty body accepted");
assert_eq!(error.code, KernelErrorCode::Schema);
}
}