use std::path::PathBuf;
use serde::{Deserialize, Serialize};
use url::Url;
use crate::{
audio::AudioResponseFormat,
errors::OapiError,
rest::post::{Post, PostNoStream},
};
#[derive(Debug, Serialize, Default, Clone)]
pub struct TranscriptionRequest {
#[serde(skip_serializing)]
pub file: PathBuf,
pub model: String,
#[serde(skip_serializing_if = "Option::is_none")]
pub language: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub prompt: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub response_format: Option<AudioResponseFormat>,
#[serde(skip_serializing_if = "Option::is_none")]
pub temperature: Option<f32>,
#[serde(skip_serializing_if = "Option::is_none")]
pub timestamp_granularities: Option<Vec<TimestampGranularity>>,
#[serde(skip_serializing)]
pub chunking_strategy: Option<ChunkingStrategy>,
#[serde(skip_serializing_if = "Option::is_none")]
pub include: Option<Vec<Include>>,
}
#[derive(Debug, Serialize, Clone, Copy)]
#[serde(rename_all = "snake_case")]
pub enum TimestampGranularity {
Word,
Segment,
}
#[derive(Debug, Serialize, Clone, Copy)]
#[serde(rename_all = "snake_case")]
pub enum Include {
Logprobs,
}
#[derive(Debug, Clone)]
pub enum ChunkingStrategy {
Auto,
ServerVad(ServerVadConfig),
}
#[derive(Debug, Clone, Default)]
pub struct ServerVadConfig {
pub prefix_padding_ms: Option<u32>,
pub silence_duration_ms: Option<u32>,
pub threshold: Option<f32>,
}
#[derive(Debug, Deserialize, Clone)]
#[serde(untagged)]
pub enum TranscriptionResponse {
Verbose(TranscriptionVerbose),
Plain(Transcription),
}
#[derive(Debug, Deserialize, Clone)]
pub struct Transcription {
pub text: String,
pub languages: Option<Vec<TranscriptionLanguage>>,
pub logprobs: Option<Vec<TranscriptionLogprob>>,
pub usage: Option<TranscriptionUsage>,
}
#[derive(Debug, Deserialize, Clone, PartialEq)]
pub struct TranscriptionLanguage {
pub code: String,
}
#[derive(Debug, Deserialize, Clone, PartialEq)]
pub struct TranscriptionLogprob {
pub token: Option<String>,
pub bytes: Option<Vec<f32>>,
pub logprob: Option<f32>,
}
#[derive(Debug, Deserialize, Clone, PartialEq)]
#[serde(tag = "type", rename_all = "snake_case")]
pub enum TranscriptionUsage {
Tokens {
input_tokens: u64,
output_tokens: u64,
total_tokens: u64,
input_token_details: Option<UsageTokensInputTokenDetails>,
},
Duration {
seconds: f64,
},
}
#[derive(Debug, Deserialize, Clone, PartialEq)]
pub struct UsageTokensInputTokenDetails {
pub audio_tokens: Option<u64>,
pub text_tokens: Option<u64>,
}
#[derive(Debug, Deserialize, Clone)]
pub struct TranscriptionVerbose {
pub duration: f64,
pub language: String,
pub text: String,
pub segments: Option<Vec<TranscriptionSegment>>,
pub usage: Option<TranscriptionVerboseUsage>,
pub words: Option<Vec<TranscriptionWord>>,
}
#[derive(Debug, Deserialize, Clone)]
pub struct TranscriptionVerboseUsage {
pub seconds: f64,
}
#[derive(Debug, Deserialize, Clone)]
pub struct TranscriptionSegment {
pub id: u64,
pub avg_logprob: f64,
pub compression_ratio: f64,
pub end: f64,
pub no_speech_prob: f64,
pub seek: u64,
pub start: f64,
pub temperature: f64,
pub text: String,
pub tokens: Vec<u64>,
}
#[derive(Debug, Deserialize, Clone)]
pub struct TranscriptionWord {
pub end: f64,
pub start: f64,
pub word: String,
}
crate::impl_from_str!(TranscriptionResponse);
impl Post for TranscriptionRequest {
#[inline]
fn is_streaming(&self) -> bool {
false
}
fn build_url(&self, base_url: &str) -> Result<String, OapiError> {
let mut url = Url::parse(base_url.trim_end_matches('/')).map_err(OapiError::UrlError)?;
url.path_segments_mut()
.map_err(|_| OapiError::UrlCannotBeBase(base_url.to_string()))?
.push("audio")
.push("transcriptions");
Ok(url.to_string())
}
}
impl PostNoStream for TranscriptionRequest {
type Response = TranscriptionResponse;
async fn get_response_string(
&self,
client: &reqwest::Client,
url: &str,
key: &str,
) -> Result<String, OapiError> {
if !self.file.exists() {
return Err(OapiError::FileNotFoundError(self.file.clone()));
}
let content = tokio::fs::read(&self.file).await?;
let file_name = self
.file
.file_name()
.and_then(|name| name.to_str())
.ok_or_else(|| OapiError::ResponseError("Invalid file name".to_string()))?
.to_string();
let file_part = reqwest::multipart::Part::bytes(content).file_name(file_name);
let mut form = reqwest::multipart::Form::new().part("file", file_part);
form = form.text("model", self.model.clone());
if let Some(language) = &self.language {
form = form.text("language", language.clone());
}
if let Some(prompt) = &self.prompt {
form = form.text("prompt", prompt.clone());
}
if let Some(response_format) = self.response_format {
let literal = crate::audio::enum_to_literal(&response_format)?;
form = form.text("response_format", literal);
}
if let Some(temperature) = self.temperature {
form = form.text("temperature", temperature.to_string());
}
if let Some(granularities) = &self.timestamp_granularities {
for granularity in granularities {
let literal = crate::audio::enum_to_literal(granularity)?;
form = form.text("timestamp_granularities[]", literal);
}
}
if let Some(include) = &self.include {
for item in include {
let literal = crate::audio::enum_to_literal(item)?;
form = form.text("include[]", literal);
}
}
if let Some(chunking_strategy) = &self.chunking_strategy {
let value = match chunking_strategy {
ChunkingStrategy::Auto => "auto".to_string(),
ChunkingStrategy::ServerVad(config) => {
let mut map = serde_json::Map::new();
map.insert("type".to_string(), "server_vad".into());
if let Some(v) = config.prefix_padding_ms {
map.insert("prefix_padding_ms".to_string(), v.into());
}
if let Some(v) = config.silence_duration_ms {
map.insert("silence_duration_ms".to_string(), v.into());
}
if let Some(v) = config.threshold {
map.insert("threshold".to_string(), v.into());
}
serde_json::to_string(&map).map_err(|e| {
OapiError::ResponseError(format!(
"Failed to serialize chunking_strategy: {e}"
))
})?
}
};
form = form.text("chunking_strategy", value);
}
let response = client
.post(url)
.header("Accept", "application/json")
.bearer_auth(key)
.multipart(form)
.send()
.await?;
crate::rest::response_text_checked(response).await
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_build_url() {
let request = TranscriptionRequest::default();
let url = request.build_url("https://api.openai.com/v1/").unwrap();
assert_eq!(url, "https://api.openai.com/v1/audio/transcriptions");
}
#[test]
fn enum_literals() {
assert_eq!(
crate::audio::enum_to_literal(&AudioResponseFormat::VerboseJson).unwrap(),
"verbose_json"
);
assert_eq!(
crate::audio::enum_to_literal(&TimestampGranularity::Word).unwrap(),
"word"
);
assert_eq!(
crate::audio::enum_to_literal(&Include::Logprobs).unwrap(),
"logprobs"
);
}
#[test]
fn parse_plain_response() {
let content = r#"{
"text": "The quick brown fox jumped over the lazy dog."
}"#;
let response: TranscriptionResponse = content.parse().unwrap();
let TranscriptionResponse::Plain(transcription) = response else {
panic!("expected plain transcription");
};
assert_eq!(
transcription.text,
"The quick brown fox jumped over the lazy dog."
);
assert_eq!(transcription.languages, None);
assert_eq!(transcription.logprobs, None);
assert_eq!(transcription.usage, None);
}
#[test]
fn parse_verbose_response() {
let content = r#"{
"duration": 8.47,
"language": "english",
"text": "The quick brown fox jumped over the lazy dog.",
"segments": [
{
"id": 0,
"avg_logprob": -0.2365,
"compression_ratio": 1.7174,
"end": 3.48,
"no_speech_prob": 0.01485,
"seek": 0,
"start": 0.0,
"temperature": 0.0,
"text": " The quick brown fox jumped over the lazy dog.",
"tokens": [464, 2069, 7586, 21831, 18045, 625, 262, 16931, 3290, 13]
}
],
"words": [
{
"end": 0.36,
"start": 0.06,
"word": "The"
}
],
"usage": {
"type": "duration",
"seconds": 8.47
}
}"#;
let response: TranscriptionResponse = content.parse().unwrap();
let TranscriptionResponse::Verbose(verbose) = response else {
panic!("expected verbose transcription");
};
assert_eq!(verbose.duration, 8.47);
assert_eq!(verbose.language, "english");
assert_eq!(
verbose.text,
"The quick brown fox jumped over the lazy dog."
);
let segments = verbose.segments.unwrap();
assert_eq!(segments.len(), 1);
assert_eq!(segments[0].id, 0);
assert_eq!(segments[0].start, 0.0);
assert_eq!(segments[0].end, 3.48);
assert_eq!(segments[0].tokens.len(), 10);
let words = verbose.words.unwrap();
assert_eq!(words[0].word, "The");
assert_eq!(words[0].start, 0.06);
let usage = verbose.usage.unwrap();
assert_eq!(usage.seconds, 8.47);
}
#[test]
fn parse_tokens_usage() {
let content = r#"{
"text": "Hello.",
"usage": {
"type": "tokens",
"input_tokens": 76,
"output_tokens": 13,
"total_tokens": 89,
"input_token_details": {
"audio_tokens": 76,
"text_tokens": 0
}
}
}"#;
let response: TranscriptionResponse = content.parse().unwrap();
let TranscriptionResponse::Plain(transcription) = response else {
panic!("expected plain transcription");
};
let TranscriptionUsage::Tokens {
input_tokens,
output_tokens,
total_tokens,
input_token_details,
} = transcription.usage.unwrap()
else {
panic!("expected tokens usage");
};
assert_eq!(input_tokens, 76);
assert_eq!(output_tokens, 13);
assert_eq!(total_tokens, 89);
let details = input_token_details.unwrap();
assert_eq!(details.audio_tokens, Some(76));
assert_eq!(details.text_tokens, Some(0));
}
}