Skip to main content

gproxy_protocol/openai/images/stream/
mod.rs

1mod edit;
2mod generation;
3
4#[cfg(test)]
5mod tests;
6
7use serde::{Deserialize, Serialize, de};
8use serde_json::Value;
9
10use crate::openai::common::{ImageStreamEventType, Rest};
11
12use super::ImageUsage;
13
14pub use edit::*;
15pub use generation::*;
16
17#[derive(Debug, Clone, PartialEq, Serialize, Deserialize, gproxy_protocol_macros::WireBuilder)]
18#[cfg_attr(not(feature = "exhaustive"), non_exhaustive)]
19pub struct ImagePartialEvent {
20    pub b64_json: String,
21    pub partial_image_index: u32,
22    #[serde(default, flatten)]
23    pub rest: Rest,
24}
25
26#[derive(Debug, Clone, PartialEq, Serialize, Deserialize, gproxy_protocol_macros::WireBuilder)]
27#[cfg_attr(not(feature = "exhaustive"), non_exhaustive)]
28pub struct ImageCompletedEvent {
29    pub b64_json: String,
30    #[serde(skip_serializing_if = "Option::is_none")]
31    pub usage: Option<ImageUsage>,
32    #[serde(default, flatten)]
33    pub rest: Rest,
34}
35
36#[derive(Debug, Clone, PartialEq, Serialize, Deserialize, gproxy_protocol_macros::WireBuilder)]
37#[cfg_attr(not(feature = "exhaustive"), non_exhaustive)]
38pub struct UnknownImageStreamEvent {
39    #[serde(rename = "type", skip_serializing_if = "Option::is_none")]
40    pub type_: Option<ImageStreamEventType>,
41    #[serde(default, flatten)]
42    pub rest: Rest,
43}
44
45#[derive(Debug, Clone, PartialEq, Serialize)]
46#[serde(untagged)]
47#[cfg_attr(not(feature = "exhaustive"), non_exhaustive)]
48pub enum ImageStreamEvent {
49    Known(KnownImageStreamEvent),
50    Unknown(UnknownImageStreamEvent),
51}
52
53impl<'de> Deserialize<'de> for ImageStreamEvent {
54    fn deserialize<D: serde::Deserializer<'de>>(deserializer: D) -> Result<Self, D::Error> {
55        let value = Value::deserialize(deserializer)?;
56        match image_stream_event_type::<D::Error>(&value)? {
57            Some(ImageStreamEventType::Known(_)) => serde_json::from_value(value)
58                .map(Self::Known)
59                .map_err(de::Error::custom),
60            Some(ImageStreamEventType::Unknown(_)) | None => serde_json::from_value(value)
61                .map(Self::Unknown)
62                .map_err(de::Error::custom),
63        }
64    }
65}
66
67#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
68#[serde(tag = "type")]
69#[cfg_attr(not(feature = "exhaustive"), non_exhaustive)]
70pub enum KnownImageStreamEvent {
71    #[serde(rename = "image_generation.partial_image")]
72    ImageGenerationPartialImage(ImagePartialEvent),
73    #[serde(rename = "image_generation.completed")]
74    ImageGenerationCompleted(ImageCompletedEvent),
75    #[serde(rename = "image_edit.partial_image")]
76    ImageEditPartialImage(ImagePartialEvent),
77    #[serde(rename = "image_edit.completed")]
78    ImageEditCompleted(ImageCompletedEvent),
79}
80
81pub(super) fn image_stream_event_type<E>(value: &Value) -> Result<Option<ImageStreamEventType>, E>
82where
83    E: de::Error,
84{
85    let Some(type_name) = value.get("type").and_then(Value::as_str) else {
86        return Ok(None);
87    };
88    serde_json::from_value(Value::String(type_name.to_owned()))
89        .map(Some)
90        .map_err(de::Error::custom)
91}