Skip to main content

gproxy_protocol/protocol/openai/images/
stream.rs

1use 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}