1use std::future::Future;
4
5use bytes::Bytes;
6use chrono::DateTime;
7use chrono::Utc;
8use serde::Deserialize;
9use serde::Serialize;
10use tokio_util::sync::CancellationToken;
11
12use crate::dynamic::BoxStream;
13use crate::error::ProviderError;
14use crate::json::JsonValue;
15use crate::language_model::RequestMetadata;
16use crate::language_model::ResponseMetadata;
17use crate::language_model::StreamError;
18use crate::shared::AudioFormat;
19use crate::shared::Headers;
20use crate::shared::ModelId;
21use crate::shared::ProviderId;
22use crate::shared::ProviderMetadata;
23use crate::shared::ProviderOptions;
24use crate::shared::Warning;
25use crate::shared::base64_bytes;
26
27pub trait SpeechTranslationModel: Send + Sync + 'static {
29 fn provider(&self) -> &ProviderId;
31
32 fn model_id(&self) -> &ModelId;
34
35 fn do_stream(
37 &self,
38 options: SpeechTranslationStreamOptions,
39 ) -> impl Future<Output = Result<SpeechTranslationStreamResult, ProviderError>> + Send;
40}
41
42pub struct SpeechTranslationStreamOptions {
44 pub audio: BoxStream<'static, Bytes>,
46 pub input_audio_format: AudioFormat,
48 pub target_language: String,
50 pub source_language: Option<String>,
52 pub output_audio_format: Option<AudioFormat>,
54 pub provider_options: ProviderOptions,
56 pub headers: Headers,
58 pub include_raw_chunks: bool,
60 pub cancellation: CancellationToken,
62}
63
64impl std::fmt::Debug for SpeechTranslationStreamOptions {
65 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
66 f.debug_struct("SpeechTranslationStreamOptions")
67 .field("audio", &"<stream>")
68 .field("input_audio_format", &self.input_audio_format)
69 .field("target_language", &self.target_language)
70 .field("source_language", &self.source_language)
71 .field("output_audio_format", &self.output_audio_format)
72 .field("provider_options", &self.provider_options)
73 .field("headers", &self.headers)
74 .field("include_raw_chunks", &self.include_raw_chunks)
75 .finish_non_exhaustive()
76 }
77}
78
79#[derive(Debug, Clone, Copy, Default, PartialEq, Serialize, Deserialize)]
81pub struct SpeechTranslationUsage {
82 #[serde(default, skip_serializing_if = "Option::is_none")]
84 pub input_audio_seconds: Option<f64>,
85 #[serde(default, skip_serializing_if = "Option::is_none")]
87 pub input_audio_tokens: Option<u64>,
88 #[serde(default, skip_serializing_if = "Option::is_none")]
90 pub output_audio_tokens: Option<u64>,
91 #[serde(default, skip_serializing_if = "Option::is_none")]
93 pub input_text_tokens: Option<u64>,
94 #[serde(default, skip_serializing_if = "Option::is_none")]
96 pub output_text_tokens: Option<u64>,
97}
98
99#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
101#[serde(tag = "type", rename_all = "kebab-case")]
102#[non_exhaustive]
103pub enum SpeechTranslationStreamPart {
104 StreamStart {
106 #[serde(default)]
108 warnings: Vec<Warning>,
109 },
110 Audio {
112 #[serde(default, skip_serializing_if = "Option::is_none")]
114 id: Option<String>,
115 #[serde(with = "base64_bytes")]
117 audio: Bytes,
118 #[serde(default, skip_serializing_if = "Option::is_none")]
120 provider_metadata: Option<ProviderMetadata>,
121 },
122 OutputTextDelta {
124 #[serde(default, skip_serializing_if = "Option::is_none")]
126 id: Option<String>,
127 delta: String,
129 #[serde(default, skip_serializing_if = "Option::is_none")]
131 provider_metadata: Option<ProviderMetadata>,
132 },
133 OutputTextFinal {
135 #[serde(default, skip_serializing_if = "Option::is_none")]
137 id: Option<String>,
138 text: String,
140 #[serde(default, skip_serializing_if = "Option::is_none")]
142 provider_metadata: Option<ProviderMetadata>,
143 },
144 SourceTranscriptDelta {
146 #[serde(default, skip_serializing_if = "Option::is_none")]
148 id: Option<String>,
149 delta: String,
151 #[serde(default, skip_serializing_if = "Option::is_none")]
153 provider_metadata: Option<ProviderMetadata>,
154 },
155 SourceTranscriptPartial {
157 #[serde(default, skip_serializing_if = "Option::is_none")]
159 id: Option<String>,
160 text: String,
162 #[serde(default, skip_serializing_if = "Option::is_none")]
164 start_second: Option<f64>,
165 #[serde(default, skip_serializing_if = "Option::is_none")]
167 end_second: Option<f64>,
168 #[serde(default, skip_serializing_if = "Option::is_none")]
170 channel_index: Option<u32>,
171 #[serde(default, skip_serializing_if = "Option::is_none")]
173 provider_metadata: Option<ProviderMetadata>,
174 },
175 SourceTranscriptFinal {
177 #[serde(default, skip_serializing_if = "Option::is_none")]
179 id: Option<String>,
180 text: String,
182 #[serde(default, skip_serializing_if = "Option::is_none")]
184 start_second: Option<f64>,
185 #[serde(default, skip_serializing_if = "Option::is_none")]
187 end_second: Option<f64>,
188 #[serde(default, skip_serializing_if = "Option::is_none")]
190 channel_index: Option<u32>,
191 #[serde(default, skip_serializing_if = "Option::is_none")]
193 provider_metadata: Option<ProviderMetadata>,
194 },
195 ResponseMetadata {
197 #[serde(default, skip_serializing_if = "Option::is_none")]
199 timestamp: Option<DateTime<Utc>>,
200 #[serde(default, skip_serializing_if = "Option::is_none")]
202 model_id: Option<ModelId>,
203 #[serde(default, skip_serializing_if = "Option::is_none")]
205 headers: Option<Headers>,
206 #[serde(default, skip_serializing_if = "Option::is_none")]
208 body: Option<JsonValue>,
209 },
210 Finish {
212 source_text: String,
214 output_text: String,
216 #[serde(default, skip_serializing_if = "Option::is_none")]
218 duration_in_seconds: Option<f64>,
219 #[serde(default, skip_serializing_if = "Option::is_none")]
221 usage: Option<SpeechTranslationUsage>,
222 #[serde(default, skip_serializing_if = "Option::is_none")]
224 provider_metadata: Option<ProviderMetadata>,
225 },
226 Raw {
228 raw_value: JsonValue,
230 },
231 Error {
233 error: StreamError,
235 },
236}
237
238pub struct SpeechTranslationStreamResult {
240 pub stream: BoxStream<'static, SpeechTranslationStreamPart>,
242 pub request: RequestMetadata,
244 pub response: ResponseMetadata,
246}
247
248impl std::fmt::Debug for SpeechTranslationStreamResult {
249 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
250 f.debug_struct("SpeechTranslationStreamResult")
251 .field("stream", &"<stream>")
252 .field("request", &self.request)
253 .field("response", &self.response)
254 .finish()
255 }
256}