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::MediaType;
use crate::shared::ModelId;
use crate::shared::ProviderId;
use crate::shared::ProviderMetadata;
use crate::shared::ProviderOptions;
use crate::shared::Warning;
pub trait TranscriptionModel: Send + Sync + 'static {
fn provider(&self) -> &ProviderId;
fn model_id(&self) -> &ModelId;
fn do_generate(
&self,
options: TranscriptionOptions,
) -> impl Future<Output = Result<TranscriptionResult, ProviderError>> + Send;
fn supports_stream(&self) -> bool {
false
}
fn do_stream(
&self,
options: TranscriptionStreamOptions,
) -> impl Future<Output = Result<TranscriptionStreamResult, ProviderError>> + Send {
let _ = options;
std::future::ready(Err(ProviderError::unsupported("transcription streaming")))
}
}
#[derive(Debug, Clone)]
pub struct TranscriptionOptions {
pub audio: Bytes,
pub media_type: MediaType,
pub provider_options: ProviderOptions,
pub headers: Headers,
pub cancellation: CancellationToken,
}
impl TranscriptionOptions {
#[must_use]
pub fn new(audio: Bytes, media_type: impl Into<MediaType>) -> Self {
Self {
audio,
media_type: media_type.into(),
provider_options: ProviderOptions::new(),
headers: Headers::new(),
cancellation: CancellationToken::new(),
}
}
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct TranscriptionSegment {
pub text: String,
pub start_second: f64,
pub end_second: f64,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct TranscriptionResult {
pub text: String,
#[serde(default)]
pub segments: Vec<TranscriptionSegment>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub language: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub duration_in_seconds: Option<f64>,
#[serde(default)]
pub warnings: Vec<Warning>,
#[serde(default)]
pub request: RequestMetadata,
#[serde(default)]
pub response: ResponseMetadata,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub provider_metadata: Option<ProviderMetadata>,
}
pub struct TranscriptionStreamOptions {
pub audio: BoxStream<'static, Bytes>,
pub input_audio_format: AudioFormat,
pub provider_options: ProviderOptions,
pub headers: Headers,
pub include_raw_chunks: bool,
pub cancellation: CancellationToken,
}
impl std::fmt::Debug for TranscriptionStreamOptions {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("TranscriptionStreamOptions")
.field("audio", &"<stream>")
.field("input_audio_format", &self.input_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, PartialEq, Serialize, Deserialize)]
#[serde(tag = "type", rename_all = "kebab-case")]
#[non_exhaustive]
pub enum TranscriptionStreamPart {
StreamStart {
#[serde(default)]
warnings: Vec<Warning>,
},
TranscriptDelta {
#[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>,
},
TranscriptPartial {
#[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")]
duration_in_seconds: 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>,
},
TranscriptFinal {
#[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 {
text: String,
#[serde(default)]
segments: Vec<TranscriptionSegment>,
#[serde(default, skip_serializing_if = "Option::is_none")]
language: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
duration_in_seconds: Option<f64>,
#[serde(default, skip_serializing_if = "Option::is_none")]
provider_metadata: Option<ProviderMetadata>,
},
Raw {
raw_value: JsonValue,
},
Error {
error: StreamError,
},
}
pub struct TranscriptionStreamResult {
pub stream: BoxStream<'static, TranscriptionStreamPart>,
pub request: RequestMetadata,
pub response: ResponseMetadata,
}
impl std::fmt::Debug for TranscriptionStreamResult {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("TranscriptionStreamResult")
.field("stream", &"<stream>")
.field("request", &self.request)
.field("response", &self.response)
.finish()
}
}