gproxy_protocol/protocol/openai/generate_content/
response_items.rs1use 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 #[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 #[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 #[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}