1use schemars::JsonSchema;
10use serde::de::{MapAccess, SeqAccess, Visitor};
11use serde::{Deserialize, Deserializer, Serialize};
12
13#[derive(Debug, Clone, PartialEq, Serialize, JsonSchema)]
15#[serde(untagged)]
16pub enum State {
17 Text(String),
19 Object(serde_json::Map<String, serde_json::Value>),
21 Array(Vec<serde_json::Value>),
23}
24
25impl<'de> Deserialize<'de> for State {
31 fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
32 where
33 D: Deserializer<'de>,
34 {
35 struct StateVisitor;
36
37 impl<'de> Visitor<'de> for StateVisitor {
38 type Value = State;
39
40 fn expecting(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
41 formatter.write_str("a string, object, or array")
42 }
43
44 fn visit_str<E>(self, value: &str) -> Result<Self::Value, E>
45 where
46 E: serde::de::Error,
47 {
48 Ok(State::Text(value.to_owned()))
49 }
50
51 fn visit_map<A>(self, mut map: A) -> Result<Self::Value, A::Error>
52 where
53 A: MapAccess<'de>,
54 {
55 let mut object = serde_json::Map::new();
56 while let Some(key) = map.next_key::<String>()? {
57 object.insert(key, map.next_value()?);
58 }
59 Ok(State::Object(object))
60 }
61
62 fn visit_seq<A>(self, mut seq: A) -> Result<Self::Value, A::Error>
63 where
64 A: SeqAccess<'de>,
65 {
66 let mut array = Vec::new();
67 while let Some(item) = seq.next_element::<serde_json::Value>()? {
68 array.push(item);
69 }
70 Ok(State::Array(array))
71 }
72 }
73
74 deserializer.deserialize_any(StateVisitor)
75 }
76}
77
78#[cfg(test)]
79mod tests {
80 use super::*;
81
82 #[test]
83 fn state_deserializes_each_json_shape_to_matching_variant() {
84 let text: State = serde_json::from_str("\"hello\"").unwrap();
85 assert_eq!(text, State::Text("hello".into()));
86
87 let object: State = serde_json::from_str(r#"{"user": {"id": 7}, "logs": []}"#).unwrap();
88 match object {
89 State::Object(map) => {
90 assert_eq!(map.len(), 2);
91 assert_eq!(map["user"]["id"], serde_json::json!(7));
92 assert_eq!(map["logs"], serde_json::json!([]));
93 }
94 other => panic!("expected object variant, got {other:?}"),
95 }
96
97 let array: State = serde_json::from_str(r#"["b", "a", {"k": 1}]"#).unwrap();
98 match array {
99 State::Array(items) => {
100 assert_eq!(items.len(), 3);
101 assert_eq!(items[0], serde_json::json!("b"));
102 assert_eq!(items[1], serde_json::json!("a"));
103 assert_eq!(items[2]["k"], serde_json::json!(1));
104 }
105 other => panic!("expected array variant, got {other:?}"),
106 }
107
108 assert_eq!(
110 serde_json::to_value(State::Text("hello".into())).unwrap(),
111 serde_json::json!("hello")
112 );
113 let object: State = serde_json::from_str(r#"{"a": [1, true, null]}"#).unwrap();
114 assert_eq!(
115 serde_json::to_value(object).unwrap(),
116 serde_json::json!({"a": [1, true, null]})
117 );
118 }
119
120 #[test]
121 fn state_rejects_number_bool_and_null() {
122 for raw in ["7", "3.5", "true", "false", "null"] {
123 let err = match serde_json::from_str::<State>(raw) {
124 Ok(state) => panic!("`{raw}` must not deserialize into a State, got {state:?}"),
125 Err(err) => err,
126 };
127 assert!(
128 err.to_string().contains("a string, object, or array"),
129 "unexpected error for `{raw}`: {err}"
130 );
131 }
132 }
133}