1use serde::{Deserialize, Serialize};
4
5use crate::{MediaSource, ModelCallContext, ModelError, ModelFuture, ModelRef};
6
7#[derive(Clone, Copy, Debug, Default, Deserialize, Eq, PartialEq, Serialize)]
9#[serde(rename_all = "snake_case")]
10#[non_exhaustive]
11pub enum ImageFormat {
12 #[default]
14 Png,
15 Webp,
17 Jpeg,
19}
20
21impl ImageFormat {
22 #[must_use]
24 pub const fn media_type(self) -> &'static str {
25 match self {
26 Self::Png => "image/png",
27 Self::Webp => "image/webp",
28 Self::Jpeg => "image/jpeg",
29 }
30 }
31}
32
33#[derive(Clone, Debug, Deserialize, PartialEq, Serialize)]
35pub struct ImageGenerationRequest {
36 pub model: ModelRef,
38 pub prompt: String,
40 pub count: u8,
42 pub size: Option<String>,
44 pub quality: Option<String>,
46 pub format: ImageFormat,
48 pub transparent: bool,
50}
51
52#[derive(Clone, Debug, Deserialize, Eq, PartialEq, Serialize)]
54pub struct GeneratedImage {
55 pub source: MediaSource,
57 pub revised_prompt: Option<String>,
59}
60
61#[derive(Clone, Debug, Deserialize, Eq, PartialEq, Serialize)]
63pub struct ImageGenerationResponse {
64 pub images: Vec<GeneratedImage>,
66}
67
68pub trait ImageGenerationModel: Send + Sync {
70 fn generate_image(
72 &self,
73 request: ImageGenerationRequest,
74 context: ModelCallContext,
75 ) -> ModelFuture<'_, Result<ImageGenerationResponse, ModelError>>;
76}
77
78#[derive(Clone, Copy, Debug, Default, Deserialize, Eq, PartialEq, Serialize)]
80#[serde(rename_all = "snake_case")]
81#[non_exhaustive]
82pub enum SpeechFormat {
83 #[default]
85 Mp3,
86 Opus,
88 Aac,
90 Flac,
92 Wav,
94 Pcm,
96}
97
98impl SpeechFormat {
99 #[must_use]
101 pub const fn media_type(self) -> &'static str {
102 match self {
103 Self::Mp3 => "audio/mpeg",
104 Self::Opus => "audio/opus",
105 Self::Aac => "audio/aac",
106 Self::Flac => "audio/flac",
107 Self::Wav => "audio/wav",
108 Self::Pcm => "audio/L16",
109 }
110 }
111}
112
113#[derive(Clone, Debug, Deserialize, PartialEq, Serialize)]
115pub struct SpeechRequest {
116 pub model: ModelRef,
118 pub input: String,
120 pub voice: String,
122 pub instructions: Option<String>,
124 pub format: SpeechFormat,
126 pub speed: Option<f32>,
128}
129
130#[derive(Clone, Debug, Deserialize, Eq, PartialEq, Serialize)]
132pub struct SpeechResponse {
133 pub media_type: String,
135 #[serde(with = "byte_serde")]
137 pub bytes: Vec<u8>,
138}
139
140pub trait SpeechModel: Send + Sync {
142 fn synthesize_speech(
144 &self,
145 request: SpeechRequest,
146 context: ModelCallContext,
147 ) -> ModelFuture<'_, Result<SpeechResponse, ModelError>>;
148}
149
150#[derive(Clone, Debug, Deserialize, PartialEq, Serialize)]
152pub struct TranscriptionRequest {
153 pub model: ModelRef,
155 pub file_name: String,
157 pub media_type: String,
159 #[serde(with = "byte_serde")]
161 pub bytes: Vec<u8>,
162 pub language: Option<String>,
164 pub prompt: Option<String>,
166}
167
168#[derive(Clone, Debug, Deserialize, PartialEq, Serialize)]
170pub struct TranscriptionResponse {
171 pub text: String,
173 pub language: Option<String>,
175 pub duration_seconds: Option<f64>,
177}
178
179pub trait TranscriptionModel: Send + Sync {
181 fn transcribe(
183 &self,
184 request: TranscriptionRequest,
185 context: ModelCallContext,
186 ) -> ModelFuture<'_, Result<TranscriptionResponse, ModelError>>;
187}
188
189mod byte_serde {
190 use serde::{Deserialize, Deserializer, Serializer};
191
192 pub fn serialize<S>(value: &[u8], serializer: S) -> Result<S::Ok, S::Error>
193 where
194 S: Serializer,
195 {
196 serializer.serialize_bytes(value)
197 }
198
199 pub fn deserialize<'de, D>(deserializer: D) -> Result<Vec<u8>, D::Error>
200 where
201 D: Deserializer<'de>,
202 {
203 Vec::<u8>::deserialize(deserializer)
204 }
205}