Skip to main content

openkind_core/
state.rs

1//! `state` field — the content to evaluate.
2//!
3//! From the spec:
4//! > `state`: string | object | array (required)
5//! > The content to evaluate. A plain string for text, or structured data
6//! > (object/array) for things like chat logs, records, or the current
7//! > state of your application.
8
9use schemars::JsonSchema;
10use serde::de::{MapAccess, SeqAccess, Visitor};
11use serde::{Deserialize, Deserializer, Serialize};
12
13/// State to evaluate. Permissive JSON shape: any string, object, or array.
14#[derive(Debug, Clone, PartialEq, Serialize, JsonSchema)]
15#[serde(untagged)]
16pub enum State {
17    /// Plain text.
18    Text(String),
19    /// Structured object (e.g. `{ "user": {...}, "logs": [...], "diff": "..." }`).
20    Object(serde_json::Map<String, serde_json::Value>),
21    /// Ordered list (e.g. chat log entries).
22    Array(Vec<serde_json::Value>),
23}
24
25// `#[serde(untagged)]` deserialization buffers the whole value into serde's
26// internal `Content` type before picking a variant, allocating on every
27// request parse. The visitor below streams the same decision directly from
28// the payload shape with identical accepted inputs: strings, objects, and
29// arrays deserialize to the same variants, and any other shape is rejected.
30impl<'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        // Each variant serializes back to the same JSON shape it came from.
109        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}