use std::future::Future;
use bytes::Bytes;
use chrono::DateTime;
use chrono::Utc;
use serde::Deserialize;
use serde::Serialize;
use tokio_util::sync::CancellationToken;
use crate::dynamic::BoxStream;
use crate::error::ProviderError;
use crate::json::JsonValue;
use crate::language_model::RequestMetadata;
use crate::language_model::ResponseMetadata;
use crate::language_model::StreamError;
use crate::shared::AudioFormat;
use crate::shared::Headers;
use crate::shared::ModelId;
use crate::shared::ProviderId;
use crate::shared::ProviderMetadata;
use crate::shared::ProviderOptions;
use crate::shared::Warning;
use crate::shared::base64_bytes;
pub trait SpeechTranslationModel: Send + Sync + 'static {
fn provider(&self) -> &ProviderId;
fn model_id(&self) -> &ModelId;
fn do_stream(
&self,
options: SpeechTranslationStreamOptions,
) -> impl Future<Output = Result<SpeechTranslationStreamResult, ProviderError>> + Send;
}
pub struct SpeechTranslationStreamOptions {
pub audio: BoxStream<'static, Bytes>,
pub input_audio_format: AudioFormat,
pub target_language: String,
pub source_language: Option<String>,
pub output_audio_format: Option<AudioFormat>,
pub provider_options: ProviderOptions,
pub headers: Headers,
pub include_raw_chunks: bool,
pub cancellation: CancellationToken,
}
impl std::fmt::Debug for SpeechTranslationStreamOptions {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("SpeechTranslationStreamOptions")
.field("audio", &"<stream>")
.field("input_audio_format", &self.input_audio_format)
.field("target_language", &self.target_language)
.field("source_language", &self.source_language)
.field("output_audio_format", &self.output_audio_format)
.field("provider_options", &self.provider_options)
.field("headers", &self.headers)
.field("include_raw_chunks", &self.include_raw_chunks)
.finish_non_exhaustive()
}
}
#[derive(Debug, Clone, Copy, Default, PartialEq, Serialize, Deserialize)]
pub struct SpeechTranslationUsage {
#[serde(default, skip_serializing_if = "Option::is_none")]
pub input_audio_seconds: Option<f64>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub input_audio_tokens: Option<u64>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub output_audio_tokens: Option<u64>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub input_text_tokens: Option<u64>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub output_text_tokens: Option<u64>,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
#[serde(tag = "type", rename_all = "kebab-case")]
#[non_exhaustive]
pub enum SpeechTranslationStreamPart {
StreamStart {
#[serde(default)]
warnings: Vec<Warning>,
},
Audio {
#[serde(default, skip_serializing_if = "Option::is_none")]
id: Option<String>,
#[serde(with = "base64_bytes")]
audio: Bytes,
#[serde(default, skip_serializing_if = "Option::is_none")]
provider_metadata: Option<ProviderMetadata>,
},
OutputTextDelta {
#[serde(default, skip_serializing_if = "Option::is_none")]
id: Option<String>,
delta: String,
#[serde(default, skip_serializing_if = "Option::is_none")]
provider_metadata: Option<ProviderMetadata>,
},
OutputTextFinal {
#[serde(default, skip_serializing_if = "Option::is_none")]
id: Option<String>,
text: String,
#[serde(default, skip_serializing_if = "Option::is_none")]
provider_metadata: Option<ProviderMetadata>,
},
SourceTranscriptDelta {
#[serde(default, skip_serializing_if = "Option::is_none")]
id: Option<String>,
delta: String,
#[serde(default, skip_serializing_if = "Option::is_none")]
provider_metadata: Option<ProviderMetadata>,
},
SourceTranscriptPartial {
#[serde(default, skip_serializing_if = "Option::is_none")]
id: Option<String>,
text: String,
#[serde(default, skip_serializing_if = "Option::is_none")]
start_second: Option<f64>,
#[serde(default, skip_serializing_if = "Option::is_none")]
end_second: Option<f64>,
#[serde(default, skip_serializing_if = "Option::is_none")]
channel_index: Option<u32>,
#[serde(default, skip_serializing_if = "Option::is_none")]
provider_metadata: Option<ProviderMetadata>,
},
SourceTranscriptFinal {
#[serde(default, skip_serializing_if = "Option::is_none")]
id: Option<String>,
text: String,
#[serde(default, skip_serializing_if = "Option::is_none")]
start_second: Option<f64>,
#[serde(default, skip_serializing_if = "Option::is_none")]
end_second: Option<f64>,
#[serde(default, skip_serializing_if = "Option::is_none")]
channel_index: Option<u32>,
#[serde(default, skip_serializing_if = "Option::is_none")]
provider_metadata: Option<ProviderMetadata>,
},
ResponseMetadata {
#[serde(default, skip_serializing_if = "Option::is_none")]
timestamp: Option<DateTime<Utc>>,
#[serde(default, skip_serializing_if = "Option::is_none")]
model_id: Option<ModelId>,
#[serde(default, skip_serializing_if = "Option::is_none")]
headers: Option<Headers>,
#[serde(default, skip_serializing_if = "Option::is_none")]
body: Option<JsonValue>,
},
Finish {
source_text: String,
output_text: String,
#[serde(default, skip_serializing_if = "Option::is_none")]
duration_in_seconds: Option<f64>,
#[serde(default, skip_serializing_if = "Option::is_none")]
usage: Option<SpeechTranslationUsage>,
#[serde(default, skip_serializing_if = "Option::is_none")]
provider_metadata: Option<ProviderMetadata>,
},
Raw {
raw_value: JsonValue,
},
Error {
error: StreamError,
},
}
pub struct SpeechTranslationStreamResult {
pub stream: BoxStream<'static, SpeechTranslationStreamPart>,
pub request: RequestMetadata,
pub response: ResponseMetadata,
}
impl std::fmt::Debug for SpeechTranslationStreamResult {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("SpeechTranslationStreamResult")
.field("stream", &"<stream>")
.field("request", &self.request)
.field("response", &self.response)
.finish()
}
}