use std::sync::Arc;
use aws_smithy_types::{BigDecimal, BigInteger, Blob, DateTime, Document, DocumentSettings};
use crate::serde::{capped_container_size, SerdeError, ShapeDeserializer};
use crate::Schema;
#[derive(Debug)]
pub struct DocumentShapeDeserializer<'a> {
cursor: &'a Document,
settings: Option<Arc<dyn DocumentSettings>>,
}
impl<'a> DocumentShapeDeserializer<'a> {
pub fn new(document: &'a Document) -> Self {
Self {
cursor: document,
settings: None,
}
}
pub fn new_with_settings(
document: &'a Document,
settings: Option<Arc<dyn DocumentSettings>>,
) -> Self {
Self {
cursor: document,
settings,
}
}
}
fn type_mismatch(expected: &str, found: &Document) -> SerdeError {
SerdeError::type_mismatch(format!(
"expected {expected} document, got {}",
kind_name(found)
))
}
fn kind_name(d: &Document) -> &'static str {
match d {
Document::Null => "null",
Document::Bool(_) => "boolean",
Document::Number(_) => "number",
Document::String(_) => "string",
Document::Blob(_) => "blob",
Document::Timestamp(_) => "timestamp",
Document::BigInteger(_) => "bigInteger",
Document::BigDecimal(_) => "bigDecimal",
Document::Array(_) => "list",
Document::Object(_) => "map",
_ => "unknown",
}
}
fn resolve_member<'s>(schema: &'s Schema<'s>, wire_name: &str) -> Option<&'s Schema<'s>> {
let idx = schema
.members()
.iter()
.position(|m| m.member_name() == Some(wire_name))?;
schema.member_schema_by_index(idx)
}
impl<'a> ShapeDeserializer for DocumentShapeDeserializer<'a> {
fn read_struct(
&mut self,
schema: &Schema<'_>,
consumer: &mut dyn FnMut(&Schema<'_>, &mut dyn ShapeDeserializer) -> Result<(), SerdeError>,
) -> Result<(), SerdeError> {
let map = self
.cursor
.as_object()
.ok_or_else(|| type_mismatch("struct (map)", self.cursor))?;
for (key, value) in map {
let Some(member_schema) = resolve_member(schema, key) else {
continue;
};
let mut sub = Self::new_with_settings(value, self.settings.clone());
consumer(member_schema, &mut sub)?;
}
Ok(())
}
fn read_list(
&mut self,
_schema: &Schema<'_>,
consumer: &mut dyn FnMut(&mut dyn ShapeDeserializer) -> Result<(), SerdeError>,
) -> Result<(), SerdeError> {
let items = self
.cursor
.as_array()
.ok_or_else(|| type_mismatch("list", self.cursor))?;
for item in items {
let mut sub = Self::new_with_settings(item, self.settings.clone());
consumer(&mut sub)?;
}
Ok(())
}
fn read_map(
&mut self,
_schema: &Schema<'_>,
consumer: &mut dyn FnMut(String, &mut dyn ShapeDeserializer) -> Result<(), SerdeError>,
) -> Result<(), SerdeError> {
let entries = self
.cursor
.as_object()
.ok_or_else(|| type_mismatch("map", self.cursor))?;
for (key, value) in entries {
let mut sub = Self::new_with_settings(value, self.settings.clone());
consumer(key.clone(), &mut sub)?;
}
Ok(())
}
fn read_boolean(&mut self, _schema: &Schema<'_>) -> Result<bool, SerdeError> {
self.cursor
.as_bool()
.ok_or_else(|| type_mismatch("boolean", self.cursor))
}
fn read_byte(&mut self, _schema: &Schema<'_>) -> Result<i8, SerdeError> {
Ok(self.cursor.as_byte()?)
}
fn read_short(&mut self, _schema: &Schema<'_>) -> Result<i16, SerdeError> {
Ok(self.cursor.as_short()?)
}
fn read_integer(&mut self, _schema: &Schema<'_>) -> Result<i32, SerdeError> {
Ok(self.cursor.as_integer()?)
}
fn read_long(&mut self, _schema: &Schema<'_>) -> Result<i64, SerdeError> {
Ok(self.cursor.as_long()?)
}
fn read_float(&mut self, _schema: &Schema<'_>) -> Result<f32, SerdeError> {
Ok(self.cursor.as_float()?)
}
fn read_double(&mut self, _schema: &Schema<'_>) -> Result<f64, SerdeError> {
Ok(self.cursor.as_double()?)
}
fn read_big_integer(&mut self, _schema: &Schema<'_>) -> Result<BigInteger, SerdeError> {
Ok(self.cursor.coerce_big_integer()?)
}
fn read_big_decimal(&mut self, _schema: &Schema<'_>) -> Result<BigDecimal, SerdeError> {
Ok(self.cursor.coerce_big_decimal()?)
}
fn read_string(&mut self, _schema: &Schema<'_>) -> Result<String, SerdeError> {
match self.cursor {
Document::String(s) => Ok(s.clone()),
other => Err(type_mismatch("string", other)),
}
}
fn read_blob(&mut self, _schema: &Schema<'_>) -> Result<Blob, SerdeError> {
if let Document::Blob(b) = self.cursor {
return Ok(Blob::new(b.clone()));
}
if let (Document::String(s), Some(settings)) = (self.cursor, &self.settings) {
return settings
.coerce_string_to_blob(s)
.map(Blob::new)
.map_err(SerdeError::from);
}
Err(type_mismatch("blob", self.cursor))
}
fn read_timestamp(&mut self, _schema: &Schema<'_>) -> Result<DateTime, SerdeError> {
if let Document::Timestamp(ts) = self.cursor {
return Ok(*ts);
}
if let Some(settings) = &self.settings {
match self.cursor {
Document::String(s) => {
return settings
.coerce_string_to_timestamp(s)
.map_err(SerdeError::from);
}
Document::Number(n) => {
return settings
.coerce_number_to_timestamp(n)
.map_err(SerdeError::from);
}
_ => {}
}
}
Err(type_mismatch("timestamp", self.cursor))
}
fn read_document(&mut self, _schema: &Schema<'_>) -> Result<Document, SerdeError> {
Ok(self.cursor.clone())
}
fn is_null(&self) -> bool {
matches!(self.cursor, Document::Null)
}
fn container_size(&self) -> Option<usize> {
let raw = match self.cursor {
Document::Array(items) => items.len(),
Document::Object(entries) => entries.len(),
_ => return None,
};
Some(capped_container_size(raw))
}
}
#[cfg(test)]
mod tests {
use std::collections::HashMap;
use super::*;
use crate::document::DocumentShapeSerializer;
use crate::serde::{SerdeError, SerializableStruct, ShapeDeserializer, ShapeSerializer};
use crate::{prelude, shape_id, Schema, ShapeId, ShapeType};
use aws_smithy_types::document::DocumentObject;
use aws_smithy_types::Number;
const PERSON_ID: ShapeId<'static> = shape_id!("smithy.example", "Person");
const PERSON_NAME_ID: ShapeId<'static> = shape_id!("smithy.example", "Person", "name");
const PERSON_AGE_ID: ShapeId<'static> = shape_id!("smithy.example", "Person", "age");
static PERSON_NAME_MEMBER: Schema<'static> =
Schema::new_member(PERSON_NAME_ID, ShapeType::String, "name", 0);
static PERSON_AGE_MEMBER: Schema<'static> =
Schema::new_member(PERSON_AGE_ID, ShapeType::Integer, "age", 1);
static PERSON_SCHEMA: Schema<'static> = Schema::new_struct(
PERSON_ID,
ShapeType::Structure,
&[&PERSON_NAME_MEMBER, &PERSON_AGE_MEMBER],
);
#[derive(Debug, Default, PartialEq)]
struct Person {
name: Option<String>,
age: Option<i32>,
}
impl SerializableStruct for Person {
fn serialize_members(&self, ser: &mut dyn ShapeSerializer) -> Result<(), SerdeError> {
if let Some(n) = &self.name {
ser.write_string(&PERSON_NAME_MEMBER, n)?;
}
if let Some(a) = self.age {
ser.write_integer(&PERSON_AGE_MEMBER, a)?;
}
Ok(())
}
}
fn deserialize_person(deser: &mut dyn ShapeDeserializer) -> Result<Person, SerdeError> {
let mut out = Person::default();
deser.read_struct(&PERSON_SCHEMA, &mut |member, sub| {
match member.member_index() {
Some(0) => out.name = Some(sub.read_string(member)?),
Some(1) => out.age = Some(sub.read_integer(member)?),
_ => {}
}
Ok(())
})?;
Ok(out)
}
#[test]
fn read_string_returns_value() {
let doc = Document::String("hello".to_string());
let mut deser = DocumentShapeDeserializer::new(&doc);
assert_eq!(deser.read_string(&prelude::STRING).unwrap(), "hello");
}
#[test]
fn read_string_on_non_string_errors() {
let doc = Document::Number(Number::PosInt(1));
let mut deser = DocumentShapeDeserializer::new(&doc);
let err = deser.read_string(&prelude::STRING).unwrap_err();
assert!(matches!(err, SerdeError::TypeMismatch { .. }));
}
#[test]
fn read_integer_with_coercion() {
let doc = Document::Number(Number::PosInt(42));
let mut deser = DocumentShapeDeserializer::new(&doc);
assert_eq!(deser.read_integer(&prelude::INTEGER).unwrap(), 42);
}
#[test]
fn read_integer_overflow_errors() {
let doc = Document::Number(Number::PosInt(u64::MAX));
let mut deser = DocumentShapeDeserializer::new(&doc);
let err = deser.read_integer(&prelude::INTEGER).unwrap_err();
assert!(matches!(err, SerdeError::NumericCoercionOverflow { .. }));
}
#[test]
fn read_boolean_returns_value() {
let doc = Document::Bool(true);
let mut deser = DocumentShapeDeserializer::new(&doc);
assert!(deser.read_boolean(&prelude::BOOLEAN).unwrap());
}
#[test]
fn read_blob_returns_value_for_native_blob() {
let doc = Document::Blob(vec![1, 2, 3]);
let mut deser = DocumentShapeDeserializer::new(&doc);
let blob = deser.read_blob(&prelude::BLOB).unwrap();
assert_eq!(blob.as_ref(), &[1u8, 2, 3]);
}
#[test]
fn read_blob_on_string_errors_without_settings() {
let doc = Document::String("YWJjZA==".to_string());
let mut deser = DocumentShapeDeserializer::new(&doc);
let err = deser.read_blob(&prelude::BLOB).unwrap_err();
assert!(matches!(err, SerdeError::TypeMismatch { .. }));
}
#[test]
fn is_null_for_null_document() {
let doc = Document::Null;
let deser = DocumentShapeDeserializer::new(&doc);
assert!(deser.is_null());
}
#[test]
fn is_null_false_for_non_null() {
let doc = Document::String("not-null".to_string());
let deser = DocumentShapeDeserializer::new(&doc);
assert!(!deser.is_null());
}
#[test]
fn read_list_iterates_elements() {
let doc = Document::Array(vec![
Document::String("a".to_string()),
Document::String("b".to_string()),
Document::String("c".to_string()),
]);
let mut deser = DocumentShapeDeserializer::new(&doc);
let mut collected = Vec::new();
deser
.read_list(&prelude::DOCUMENT, &mut |sub| {
collected.push(sub.read_string(&prelude::STRING)?);
Ok(())
})
.unwrap();
assert_eq!(collected, ["a", "b", "c"]);
}
#[test]
fn read_map_iterates_entries() {
let doc = Document::Object(DocumentObject::from([
("k1".to_string(), Document::String("v1".to_string())),
("k2".to_string(), Document::String("v2".to_string())),
]));
let mut deser = DocumentShapeDeserializer::new(&doc);
let mut collected = HashMap::new();
deser
.read_map(&prelude::DOCUMENT, &mut |key, sub| {
collected.insert(key, sub.read_string(&prelude::STRING)?);
Ok(())
})
.unwrap();
assert_eq!(collected.get("k1").map(String::as_str), Some("v1"));
assert_eq!(collected.get("k2").map(String::as_str), Some("v2"));
}
#[test]
fn container_size_on_list() {
let doc = Document::Array(vec![Document::Null; 5]);
let deser = DocumentShapeDeserializer::new(&doc);
assert_eq!(deser.container_size(), Some(5));
}
#[test]
fn container_size_on_scalar_is_none() {
let doc = Document::String("foo".to_string());
let deser = DocumentShapeDeserializer::new(&doc);
assert!(deser.container_size().is_none());
}
#[test]
fn read_struct_with_consumer_dispatch() {
let doc = Document::Object(DocumentObject::from([
("name".to_string(), Document::String("Alex".to_string())),
("age".to_string(), Document::Number(Number::PosInt(30))),
]));
let mut deser = DocumentShapeDeserializer::new(&doc);
let person = deserialize_person(&mut deser).unwrap();
assert_eq!(
person,
Person {
name: Some("Alex".into()),
age: Some(30),
}
);
}
#[test]
fn read_struct_with_missing_optional_member() {
let doc = Document::Object(DocumentObject::from([(
"name".to_string(),
Document::String("Sam".to_string()),
)]));
let mut deser = DocumentShapeDeserializer::new(&doc);
let person = deserialize_person(&mut deser).unwrap();
assert_eq!(
person,
Person {
name: Some("Sam".into()),
age: None,
}
);
}
#[test]
fn read_struct_ignores_unknown_members() {
let doc = Document::Object(DocumentObject::from([
("name".to_string(), Document::String("Joe".to_string())),
(
"unknown_field".to_string(),
Document::String("ignored".to_string()),
),
]));
let mut deser = DocumentShapeDeserializer::new(&doc);
let person = deserialize_person(&mut deser).unwrap();
assert_eq!(person.name.as_deref(), Some("Joe"));
}
#[test]
fn read_struct_on_non_map_errors() {
let doc = Document::String("not-a-struct".to_string());
let mut deser = DocumentShapeDeserializer::new(&doc);
let err = deserialize_person(&mut deser).unwrap_err();
assert!(matches!(err, SerdeError::TypeMismatch { .. }));
}
#[test]
fn round_trip_through_document() {
let original = Person {
name: Some("Iago".into()),
age: Some(7),
};
let mut ser = DocumentShapeSerializer::new();
ser.write_struct(&PERSON_SCHEMA, &original).unwrap();
let doc = ser.finish().unwrap();
let mut deser = DocumentShapeDeserializer::new(doc.document());
let restored = deserialize_person(&mut deser).unwrap();
assert_eq!(restored, original);
assert_eq!(doc.discriminator(), Some("smithy.example#Person"));
}
}