Skip to main content

gproxy_protocol/protocol/openai/generate_content/
response_items.rs

1use std::collections::BTreeMap;
2
3use serde::{Deserialize, Serialize, de};
4use serde_json::Value;
5
6use super::super::common::*;
7
8mod actions;
9mod content;
10mod message;
11mod typed;
12
13pub use actions::*;
14pub use content::*;
15pub use message::*;
16pub use typed::*;
17
18#[derive(Debug, Clone, PartialEq, Serialize)]
19#[serde(untagged)]
20pub enum ResponseItem {
21    Message(ResponseMessageItem),
22    Typed(TypedResponseItem),
23    Unknown(UnknownResponseItem),
24}
25
26impl<'de> Deserialize<'de> for ResponseItem {
27    fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
28    where
29        D: serde::Deserializer<'de>,
30    {
31        let value = Value::deserialize(deserializer)?;
32        let type_name = value.get("type").and_then(Value::as_str);
33
34        let Some(type_name) = type_name else {
35            if let Ok(message) = serde_json::from_value::<ResponseMessageItem>(value.clone()) {
36                return Ok(Self::Message(message));
37            }
38
39            if let Some(item_reference) = item_reference_without_type(&value) {
40                return Ok(Self::Typed(item_reference));
41            }
42
43            return serde_json::from_value(value)
44                .map(Self::Unknown)
45                .map_err(de::Error::custom);
46        };
47
48        let item_type =
49            serde_json::from_value::<ResponseItemType>(Value::String(type_name.to_owned()))
50                .map_err(de::Error::custom)?;
51
52        match item_type {
53            ResponseItemType::Known(ResponseItemTypeKnown::Message) => {
54                serde_json::from_value(value)
55                    .map(Self::Message)
56                    .map_err(de::Error::custom)
57            }
58            ResponseItemType::Known(_) => serde_json::from_value(value)
59                .map(Self::Typed)
60                .map_err(de::Error::custom),
61            ResponseItemType::Unknown(_) => serde_json::from_value(value)
62                .map(Self::Unknown)
63                .map_err(de::Error::custom),
64        }
65    }
66}
67
68fn item_reference_without_type(value: &Value) -> Option<TypedResponseItem> {
69    let object = value.as_object()?;
70    let id = object.get("id")?.as_str()?.to_owned();
71    let mut extra = Extra::new();
72
73    for (key, value) in object {
74        if key != "id" {
75            extra.insert(key.clone(), value.clone());
76        }
77    }
78
79    Some(TypedResponseItem::ItemReference { id, extra })
80}
81
82#[derive(Debug, Clone, PartialEq, Serialize)]
83#[serde(transparent)]
84pub struct ResponseOutputItem(pub ResponseItem);
85
86impl<'de> Deserialize<'de> for ResponseOutputItem {
87    fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
88    where
89        D: serde::Deserializer<'de>,
90    {
91        let item = ResponseItem::deserialize(deserializer)?;
92        validate_response_output_item(&item).map_err(de::Error::custom)?;
93        Ok(Self(item))
94    }
95}
96
97fn validate_response_output_item(item: &ResponseItem) -> Result<(), &'static str> {
98    let ResponseItem::Typed(typed) = item else {
99        return Ok(());
100    };
101
102    match typed {
103        TypedResponseItem::ComputerCallOutput { id, status, .. } => {
104            require_some(id, "computer_call_output.id")?;
105            require_some(status, "computer_call_output.status")?;
106        }
107        TypedResponseItem::FunctionCallOutput { id, status, .. } => {
108            require_some(id, "function_call_output.id")?;
109            require_some(status, "function_call_output.status")?;
110        }
111        TypedResponseItem::ToolSearchCall {
112            id,
113            call_id,
114            execution,
115            status,
116            ..
117        } => {
118            require_some(id, "tool_search_call.id")?;
119            require_some(call_id, "tool_search_call.call_id")?;
120            require_some(execution, "tool_search_call.execution")?;
121            require_some(status, "tool_search_call.status")?;
122        }
123        TypedResponseItem::ToolSearchOutput {
124            id,
125            call_id,
126            execution,
127            status,
128            ..
129        } => {
130            require_some(id, "tool_search_output.id")?;
131            require_some(call_id, "tool_search_output.call_id")?;
132            require_some(execution, "tool_search_output.execution")?;
133            require_some(status, "tool_search_output.status")?;
134        }
135        TypedResponseItem::AdditionalTools { id, .. } => {
136            require_some(id, "additional_tools.id")?;
137        }
138        TypedResponseItem::ShellCall {
139            id,
140            environment,
141            status,
142            ..
143        } => {
144            require_some(id, "shell_call.id")?;
145            require_some(environment, "shell_call.environment")?;
146            require_some(status, "shell_call.status")?;
147        }
148        TypedResponseItem::ShellCallOutput {
149            id,
150            max_output_length,
151            status,
152            ..
153        } => {
154            require_some(id, "shell_call_output.id")?;
155            require_some(max_output_length, "shell_call_output.max_output_length")?;
156            require_some(status, "shell_call_output.status")?;
157        }
158        _ => {}
159    }
160
161    Ok(())
162}
163
164fn require_some<T>(value: &Option<T>, field: &'static str) -> Result<(), &'static str> {
165    value.as_ref().map(|_| ()).ok_or(field)
166}
167
168#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
169pub struct UnknownResponseItem {
170    #[serde(rename = "type", skip_serializing_if = "Option::is_none")]
171    pub type_: Option<ResponseItemType>,
172    #[serde(default, flatten, skip_serializing_if = "BTreeMap::is_empty")]
173    pub extra: Extra,
174}
175
176#[cfg(test)]
177mod tests {
178    use super::*;
179
180    /// Regression: `ResponseItem`/`ResponseMessageItem` must serialize flat,
181    /// matching their hand-written flat `Deserialize` implementations.
182    #[test]
183    fn input_message_serializes_flat() {
184        let flat = serde_json::json!({"type": "message", "role": "user", "content": "hi"});
185        let item: ResponseItem = serde_json::from_value(flat.clone()).unwrap();
186        let back = serde_json::to_value(&item).unwrap();
187        assert!(
188            back.get("Message").is_none() && back.get("EasyInput").is_none(),
189            "must not be externally tagged: {back}"
190        );
191        assert_eq!(back["role"], "user", "{back}");
192        assert_eq!(back, flat);
193    }
194
195    /// Regression (#146): Codex CLI replays assistant turns without `id` / `status`,
196    /// carrying `output_text` parts. They must decode as easy-input `OutputParts`
197    /// and round-trip unchanged instead of failing the whole request.
198    #[test]
199    fn replayed_assistant_history_decodes_as_easy_input_output_parts() {
200        let replayed = serde_json::json!({
201            "type": "message",
202            "role": "assistant",
203            "content": [{"type": "output_text", "text": "hello"}]
204        });
205        let item: ResponseItem = serde_json::from_value(replayed.clone()).unwrap();
206        let ResponseItem::Message(ResponseMessageItem::EasyInput(message)) = &item else {
207            panic!("expected EasyInput, got: {item:?}");
208        };
209        assert!(
210            matches!(&message.content, ResponseEasyInputContent::OutputParts(parts) if parts.len() == 1),
211            "expected OutputParts: {:?}",
212            message.content
213        );
214        assert_eq!(serde_json::to_value(&item).unwrap(), replayed);
215    }
216
217    /// Guard: ordinary input parts must keep matching `Parts`, not be swallowed
218    /// by the `OutputParts` arm added after it.
219    #[test]
220    fn easy_input_text_parts_still_decode_as_input_parts() {
221        let body = serde_json::json!({
222            "type": "message",
223            "role": "assistant",
224            "content": [{"type": "input_text", "text": "hi"}]
225        });
226        let item: ResponseItem = serde_json::from_value(body).unwrap();
227        let ResponseItem::Message(ResponseMessageItem::EasyInput(message)) = &item else {
228            panic!("expected EasyInput, got: {item:?}");
229        };
230        assert!(
231            matches!(&message.content, ResponseEasyInputContent::Parts(parts) if parts.len() == 1),
232            "expected Parts: {:?}",
233            message.content
234        );
235    }
236}