rig_core/
audio_generation.rs1use crate::markers::{Missing, Provided};
4use crate::wasm_compat::{WasmCompatSend, WasmCompatSync};
5use serde_json::Value;
6
7crate::provider_response::provider_error_enum!(
8 AudioGenerationError, "audio generation" {
14 #[error("RequestError: {0}")]
16 RequestError(#[from] Box<dyn std::error::Error + Send + Sync + 'static>),
17 }
18);
19
20pub struct AudioGenerationResponse<T> {
21 pub audio: Vec<u8>,
22 pub response: T,
23}
24
25pub trait AudioGenerationModel: Sized + Clone + WasmCompatSend + WasmCompatSync {
26 type Response: WasmCompatSend + WasmCompatSync;
27
28 type Client;
29
30 fn make(client: &Self::Client, model: impl Into<String>) -> Self;
31
32 fn audio_generation(
33 &self,
34 request: AudioGenerationRequest,
35 ) -> impl std::future::Future<
36 Output = Result<AudioGenerationResponse<Self::Response>, AudioGenerationError>,
37 > + WasmCompatSend;
38
39 fn audio_generation_request(&self) -> AudioGenerationRequestBuilder<Self, Missing, Missing> {
40 AudioGenerationRequestBuilder::new(self.clone())
41 }
42}
43pub struct AudioGenerationRequest {
44 pub text: String,
45 pub voice: String,
46 pub speed: f32,
47 pub additional_params: Option<Value>,
48}
49
50pub struct AudioGenerationRequestBuilder<M, T = Missing, V = Missing>
51where
52 M: AudioGenerationModel,
53{
54 model: M,
55 text: T,
56 voice: V,
57 speed: f32,
58 additional_params: Option<Value>,
59}
60
61impl<M> AudioGenerationRequestBuilder<M, Missing, Missing>
62where
63 M: AudioGenerationModel,
64{
65 pub fn new(model: M) -> Self {
66 Self {
67 model,
68 text: Missing,
69 voice: Missing,
70 speed: 1.0,
71 additional_params: None,
72 }
73 }
74}
75
76impl<M, T, V> AudioGenerationRequestBuilder<M, T, V>
77where
78 M: AudioGenerationModel,
79{
80 pub fn text(self, text: &str) -> AudioGenerationRequestBuilder<M, Provided<String>, V> {
82 AudioGenerationRequestBuilder {
83 model: self.model,
84 text: Provided(text.to_string()),
85 voice: self.voice,
86 speed: self.speed,
87 additional_params: self.additional_params,
88 }
89 }
90
91 pub fn voice(self, voice: &str) -> AudioGenerationRequestBuilder<M, T, Provided<String>> {
93 AudioGenerationRequestBuilder {
94 model: self.model,
95 text: self.text,
96 voice: Provided(voice.to_string()),
97 speed: self.speed,
98 additional_params: self.additional_params,
99 }
100 }
101
102 pub fn speed(mut self, speed: f32) -> Self {
104 self.speed = speed;
105 self
106 }
107
108 pub fn additional_params(mut self, params: Value) -> Self {
110 self.additional_params = Some(params);
111 self
112 }
113}
114
115impl<M> AudioGenerationRequestBuilder<M, Provided<String>, Provided<String>>
116where
117 M: AudioGenerationModel,
118{
119 pub fn build(self) -> AudioGenerationRequest {
120 AudioGenerationRequest {
121 text: self.text.0,
122 voice: self.voice.0,
123 speed: self.speed,
124 additional_params: self.additional_params,
125 }
126 }
127
128 pub async fn send(self) -> Result<AudioGenerationResponse<M::Response>, AudioGenerationError> {
129 let model = self.model.clone();
130
131 model.audio_generation(self.build()).await
132 }
133}
134
135#[cfg(test)]
136mod provider_response_tests {
137 use super::*;
138 use crate::{http_client, provider_response};
139 use http::StatusCode;
140
141 #[test]
142 fn audio_generation_error_provider_response_helpers_with_preserved_json_body() {
143 let body = r#"{"error":{"message":"invalid voice"}}"#;
144 let error = AudioGenerationError::ProviderResponse(
145 provider_response::ProviderResponseError::without_status(body.to_string()),
146 );
147
148 assert_eq!(error.provider_response_body(), Some(body));
149 assert_eq!(error.provider_response_status(), None);
150 assert_eq!(
151 error.provider_response_json().expect("valid JSON"),
152 Some(serde_json::json!({ "error": { "message": "invalid voice" } }))
153 );
154 }
155
156 #[test]
157 fn audio_generation_error_provider_response_helpers_with_http_non_success() {
158 let body = r#"{"error":{"message":"bad request"}}"#;
159 let error =
160 AudioGenerationError::HttpError(http_client::Error::InvalidStatusCodeWithMessage(
161 StatusCode::BAD_REQUEST,
162 body.to_string(),
163 ));
164
165 assert_eq!(error.provider_response_body(), Some(body));
166 assert_eq!(
167 error.provider_response_status(),
168 Some(StatusCode::BAD_REQUEST)
169 );
170 assert_eq!(
171 error.provider_response_json().expect("valid JSON"),
172 Some(serde_json::json!({ "error": { "message": "bad request" } }))
173 );
174 }
175
176 #[test]
177 fn audio_generation_error_provider_error_is_not_a_provider_response() {
178 let error = AudioGenerationError::ProviderError("internal diagnostic".to_string());
179
180 assert_eq!(error.provider_response_body(), None);
181 assert_eq!(error.provider_response_status(), None);
182 assert_eq!(error.provider_response_json().expect("no body"), None);
183 }
184
185 #[test]
186 fn audio_generation_error_provider_response_helpers_with_unrelated_variant() {
187 let error = AudioGenerationError::ResponseError("parse failed".to_string());
188
189 assert_eq!(error.provider_response_body(), None);
190 assert_eq!(error.provider_response_status(), None);
191 assert_eq!(error.provider_response_json().expect("no body"), None);
192 }
193}