1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
//! Audio transcription requests, normalized responses, and model interfaces.
//!
//! ```no_run
//! use rig_core::DynModel;
//! use rig_core::operation::Transcription;
//! use rig_core::transcription::TranscriptionRequestBuilder;
//!
//! # async fn example(model: DynModel<Transcription>) -> Result<(), Box<dyn std::error::Error>> {
//! let request = TranscriptionRequestBuilder::from_file("audio.wav")?.build();
//! let response = model.call(request).await?;
//! # let _ = response;
//! # Ok(())
//! # }
//! ```
use crate::completion::Usage;
use crate::json_utils;
use serde::{Deserialize, Serialize};
use std::io;
use std::{fs, path::Path};
/// Transcript and normalized provider metadata, with provider-specific data
/// available through [`Self::raw`].
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct TranscriptionResponse {
/// The transcribed text.
pub text: String,
/// Provider-reported token usage. Unreported counters remain `None`;
/// this field does not contain audio duration.
#[serde(default)]
pub usage: Usage,
/// Stable descriptor name of the provider that produced this response,
/// for example `"openai"`. Always populated.
pub provider: String,
/// Provider-reported model identifier, when the wire response named one.
/// This is the model the provider says answered, not the model requested.
#[serde(default)]
pub model: Option<String>,
/// Provider-assigned response-scoped identifier, when reported.
#[serde(default, skip_serializing_if = "Option::is_none")]
pub response_id: Option<String>,
/// Transport request ID from HTTP headers, or `None` when unreported.
#[serde(default, skip_serializing_if = "Option::is_none")]
pub provider_request_id: Option<String>,
/// Provider response document. Defaults to null until populated.
#[serde(default, skip_serializing_if = "serde_json::Value::is_null")]
pub raw: serde_json::Value,
}
impl TranscriptionResponse {
/// A response carrying `text`. The driver writes the provider, the
/// transport request id and the reply document; decoders set what the
/// provider reported.
pub fn new(text: impl Into<String>) -> Self {
Self {
text: text.into(),
usage: Usage::default(),
provider: String::new(),
model: None,
response_id: None,
provider_request_id: None,
raw: serde_json::Value::Null,
}
}
}
/// Struct representing a general transcription request that can be sent to a transcription model provider.
pub struct TranscriptionRequest {
/// The file data to be sent to the transcription model provider
pub data: Vec<u8>,
/// The file name to be used in the request
pub filename: String,
/// The language used in the response from the transcription model provider
pub language: Option<String>,
/// The prompt to be sent to the transcription model provider
pub prompt: Option<String>,
/// The temperature sent to the transcription model provider
pub temperature: Option<f64>,
/// Additional parameters to be sent to the transcription model provider
pub additional_params: Option<serde_json::Value>,
}
/// The filename a request carries until the caller or its file names it.
const DEFAULT_FILENAME: &str = "file";
/// Builds a transcription request over the supplied audio. The audio is not
/// validated for format or nonemptiness.
pub struct TranscriptionRequestBuilder {
request: TranscriptionRequest,
}
impl TranscriptionRequestBuilder {
/// A request over `data`, named `"file"` until [`Self::filename`] names it.
pub fn new(data: Vec<u8>) -> Self {
Self {
request: TranscriptionRequest {
data,
filename: DEFAULT_FILENAME.to_owned(),
language: None,
prompt: None,
temperature: None,
additional_params: None,
},
}
}
/// A request over the file at `path`, named after its base name. Reads the
/// file synchronously and returns I/O errors unchanged.
pub fn from_file(path: impl AsRef<Path>) -> io::Result<Self> {
let path = path.as_ref();
let filename = path
.file_name()
.map(|name| name.to_string_lossy().into_owned());
Ok(Self::new(fs::read(path)?).filename(filename))
}
/// Names the audio file; `None` restores the default name.
pub fn filename(mut self, filename: impl Into<Option<String>>) -> Self {
self.request.filename = filename
.into()
.unwrap_or_else(|| DEFAULT_FILENAME.to_owned());
self
}
/// Sets the output language.
pub fn language(mut self, language: String) -> Self {
self.request.language = Some(language);
self
}
/// Sets the prompt sent with the audio.
pub fn prompt(mut self, prompt: String) -> Self {
self.request.prompt = Some(prompt);
self
}
/// Sets the sampling temperature.
pub fn temperature(mut self, temperature: f64) -> Self {
self.request.temperature = Some(temperature);
self
}
/// Merges provider-specific parameters over earlier ones, key by key for
/// JSON objects; `None` clears existing parameters.
pub fn additional_params(mut self, params: impl Into<Option<serde_json::Value>>) -> Self {
self.request.additional_params =
json_utils::merge_params(self.request.additional_params.take(), params.into());
self
}
/// Builds the transcription request.
pub fn build(self) -> TranscriptionRequest {
self.request
}
}
#[cfg(test)]
mod builder_tests;