gproxy_protocol/protocol/openai/images/
stream.rs1use std::collections::BTreeMap;
2
3use serde::{Deserialize, Serialize, de};
4use serde_json::Value;
5
6use super::super::common::*;
7use super::responses::ImageUsage;
8
9#[derive(Debug, Clone, PartialEq, Serialize)]
10#[non_exhaustive]
11pub enum ImageStreamEvent {
12 Known(KnownImageStreamEvent),
13 Unknown(UnknownImageStreamEvent),
14}
15
16impl<'de> Deserialize<'de> for ImageStreamEvent {
17 fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
18 where
19 D: serde::Deserializer<'de>,
20 {
21 let value = Value::deserialize(deserializer)?;
22 match image_stream_event_type::<D::Error>(&value)? {
23 Some(ImageStreamEventType::Known(_)) => serde_json::from_value(value)
24 .map(Self::Known)
25 .map_err(de::Error::custom),
26 Some(ImageStreamEventType::Unknown(_)) | None => serde_json::from_value(value)
27 .map(Self::Unknown)
28 .map_err(de::Error::custom),
29 }
30 }
31}
32
33#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
34#[serde(tag = "type")]
35#[non_exhaustive]
36pub enum KnownImageStreamEvent {
37 #[serde(rename = "image_generation.partial_image")]
38 ImageGenerationPartialImage {
39 b64_json: String,
40 partial_image_index: u32,
41 #[serde(default, flatten, skip_serializing_if = "BTreeMap::is_empty")]
42 extra: Extra,
43 },
44 #[serde(rename = "image_generation.completed")]
45 ImageGenerationCompleted {
46 b64_json: String,
47 #[serde(skip_serializing_if = "Option::is_none")]
48 usage: Option<ImageUsage>,
49 #[serde(default, flatten, skip_serializing_if = "BTreeMap::is_empty")]
50 extra: Extra,
51 },
52 #[serde(rename = "image_edit.partial_image")]
53 ImageEditPartialImage {
54 b64_json: String,
55 partial_image_index: u32,
56 #[serde(default, flatten, skip_serializing_if = "BTreeMap::is_empty")]
57 extra: Extra,
58 },
59 #[serde(rename = "image_edit.completed")]
60 ImageEditCompleted {
61 b64_json: String,
62 #[serde(skip_serializing_if = "Option::is_none")]
63 usage: Option<ImageUsage>,
64 #[serde(default, flatten, skip_serializing_if = "BTreeMap::is_empty")]
65 extra: Extra,
66 },
67}
68
69#[derive(Debug, Clone, PartialEq, Serialize)]
70#[non_exhaustive]
71pub enum ImageGenerationStreamEvent {
72 Known(KnownImageGenerationStreamEvent),
73 Unknown(UnknownImageStreamEvent),
74}
75
76impl<'de> Deserialize<'de> for ImageGenerationStreamEvent {
77 fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
78 where
79 D: serde::Deserializer<'de>,
80 {
81 let value = Value::deserialize(deserializer)?;
82 match image_stream_event_type::<D::Error>(&value)? {
83 Some(ImageStreamEventType::Known(
84 ImageStreamEventTypeKnown::ImageGenerationPartialImage
85 | ImageStreamEventTypeKnown::ImageGenerationCompleted,
86 )) => serde_json::from_value(value)
87 .map(Self::Known)
88 .map_err(de::Error::custom),
89 Some(ImageStreamEventType::Known(_)) => Err(de::Error::custom(
90 "known image edit stream event cannot deserialize as image generation stream event",
91 )),
92 Some(ImageStreamEventType::Unknown(_)) | None => serde_json::from_value(value)
93 .map(Self::Unknown)
94 .map_err(de::Error::custom),
95 }
96 }
97}
98
99#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
100#[serde(tag = "type")]
101#[non_exhaustive]
102pub enum KnownImageGenerationStreamEvent {
103 #[serde(rename = "image_generation.partial_image")]
104 PartialImage {
105 b64_json: String,
106 partial_image_index: u32,
107 #[serde(default, flatten, skip_serializing_if = "BTreeMap::is_empty")]
108 extra: Extra,
109 },
110 #[serde(rename = "image_generation.completed")]
111 Completed {
112 b64_json: String,
113 #[serde(skip_serializing_if = "Option::is_none")]
114 usage: Option<ImageUsage>,
115 #[serde(default, flatten, skip_serializing_if = "BTreeMap::is_empty")]
116 extra: Extra,
117 },
118}
119
120#[derive(Debug, Clone, PartialEq, Serialize)]
121#[non_exhaustive]
122pub enum ImageEditStreamEvent {
123 Known(KnownImageEditStreamEvent),
124 Unknown(UnknownImageStreamEvent),
125}
126
127impl<'de> Deserialize<'de> for ImageEditStreamEvent {
128 fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
129 where
130 D: serde::Deserializer<'de>,
131 {
132 let value = Value::deserialize(deserializer)?;
133 match image_stream_event_type::<D::Error>(&value)? {
134 Some(ImageStreamEventType::Known(
135 ImageStreamEventTypeKnown::ImageEditPartialImage
136 | ImageStreamEventTypeKnown::ImageEditCompleted,
137 )) => serde_json::from_value(value)
138 .map(Self::Known)
139 .map_err(de::Error::custom),
140 Some(ImageStreamEventType::Known(_)) => Err(de::Error::custom(
141 "known image generation stream event cannot deserialize as image edit stream event",
142 )),
143 Some(ImageStreamEventType::Unknown(_)) | None => serde_json::from_value(value)
144 .map(Self::Unknown)
145 .map_err(de::Error::custom),
146 }
147 }
148}
149
150#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
151#[serde(tag = "type")]
152#[non_exhaustive]
153pub enum KnownImageEditStreamEvent {
154 #[serde(rename = "image_edit.partial_image")]
155 PartialImage {
156 b64_json: String,
157 partial_image_index: u32,
158 #[serde(default, flatten, skip_serializing_if = "BTreeMap::is_empty")]
159 extra: Extra,
160 },
161 #[serde(rename = "image_edit.completed")]
162 Completed {
163 b64_json: String,
164 #[serde(skip_serializing_if = "Option::is_none")]
165 usage: Option<ImageUsage>,
166 #[serde(default, flatten, skip_serializing_if = "BTreeMap::is_empty")]
167 extra: Extra,
168 },
169}
170
171#[derive(Debug, Clone, PartialEq, Serialize, Deserialize, gproxy_protocol_macros::WireBuilder)]
172#[non_exhaustive]
173pub struct UnknownImageStreamEvent {
174 #[serde(rename = "type", skip_serializing_if = "Option::is_none")]
175 pub type_: Option<ImageStreamEventType>,
176 #[serde(default, flatten, skip_serializing_if = "BTreeMap::is_empty")]
177 pub extra: Extra,
178}
179
180fn image_stream_event_type<E>(value: &Value) -> Result<Option<ImageStreamEventType>, E>
181where
182 E: de::Error,
183{
184 let Some(type_name) = value.get("type").and_then(Value::as_str) else {
185 return Ok(None);
186 };
187
188 serde_json::from_value(Value::String(type_name.to_owned()))
189 .map(Some)
190 .map_err(de::Error::custom)
191}