Skip to main content

rig_core/
audio_generation.rs

1//! Everything related to audio generation (ie, Text To Speech).
2//! Rig abstracts over a number of different providers using the [AudioGenerationModel] trait.
3use crate::markers::{Missing, Provided};
4use crate::wasm_compat::{WasmCompatSend, WasmCompatSync};
5use serde_json::Value;
6
7crate::provider_response::provider_error_enum!(
8    ///
9    /// HTTP audio failures preserve the provider's status and body: a non-success
10    /// response surfaces as [`Self::HttpError`], and a provider error envelope
11    /// returned with a 2xx status surfaces as [`Self::ProviderResponse`] (for
12    /// example the Hyperbolic audio path). Both are read by the helpers.
13    AudioGenerationError, "audio generation" {
14        /// Error building the audio generation request
15        #[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    /// Sets the text for the audio generation request
81    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    /// The voice of the generated audio
92    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    /// The speed of the generated audio
103    pub fn speed(mut self, speed: f32) -> Self {
104        self.speed = speed;
105        self
106    }
107
108    /// Adds additional parameters to the audio generation request.
109    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}