Skip to main content

ferrin_google/
speech.rs

1//! Gemini text-to-speech model (`generateContent` with `AUDIO` modality).
2
3use base64::Engine;
4use bytes::Bytes;
5use ferrin_provider_util::http::ResponseHandlers;
6use ferrin_provider_util::http::json_response_handler;
7use ferrin_provider_util::http::post_json;
8use ferrin_spec::JsonObject;
9use ferrin_spec::JsonValue;
10use ferrin_spec::MediaType;
11use ferrin_spec::ModelId;
12use ferrin_spec::ProviderId;
13use ferrin_spec::ResponseMetadata;
14use ferrin_spec::error::InvalidArgumentError;
15use ferrin_spec::error::ProviderError;
16use ferrin_spec::language_model::RequestMetadata;
17use ferrin_spec::shared::Warning;
18use ferrin_spec::speech_model::SpeechModel;
19use ferrin_spec::speech_model::SpeechOptions;
20use ferrin_spec::speech_model::SpeechResult;
21use serde::Deserialize;
22use serde_json::json;
23
24use crate::config::SharedConfig;
25use crate::error::failed_response_handler;
26use crate::options::parse_merged;
27use crate::output::OutputMapper;
28
29/// Provider id family.
30pub const FAMILY: &str = "speech";
31
32/// Voice used when none is requested.
33pub const DEFAULT_VOICE: &str = "Kore";
34
35/// Sample rate assumed when the response media type carries none.
36pub const DEFAULT_SAMPLE_RATE: u32 = 24_000;
37
38/// Speech options (`provider_options["google"]`).
39#[derive(Debug, Clone, Default, PartialEq, Deserialize)]
40#[serde(rename_all = "camelCase")]
41pub struct GoogleSpeechOptions {
42    /// Multi-speaker voice configuration
43    /// (`{speakerVoiceConfigs: [{speaker, voiceConfig: {prebuiltVoiceConfig: {voiceName}}}]}`).
44    #[serde(default)]
45    pub multi_speaker_voice_config: Option<JsonObject>,
46}
47
48#[derive(Deserialize)]
49#[serde(rename_all = "camelCase")]
50struct SpeechResponse {
51    response_id: Option<String>,
52    candidates: Option<Vec<SpeechCandidate>>,
53}
54
55#[derive(Deserialize)]
56struct SpeechCandidate {
57    content: Option<SpeechContent>,
58}
59
60#[derive(Deserialize)]
61struct SpeechContent {
62    parts: Option<Vec<SpeechPart>>,
63}
64
65#[derive(Deserialize)]
66#[serde(rename_all = "camelCase")]
67struct SpeechPart {
68    inline_data: Option<SpeechAudio>,
69}
70
71#[derive(Deserialize)]
72#[serde(rename_all = "camelCase")]
73struct SpeechAudio {
74    data: Option<String>,
75    mime_type: Option<String>,
76}
77
78/// Wraps signed 16-bit little-endian mono PCM in a 44-byte WAV header.
79#[must_use]
80pub fn add_wav_header(pcm: &[u8], sample_rate: u32) -> Bytes {
81    let channels: u16 = 1;
82    let bits_per_sample: u16 = 16;
83    let block_align = channels * bits_per_sample / 8;
84    let byte_rate = sample_rate * u32::from(block_align);
85    let data_size = u32::try_from(pcm.len()).unwrap_or(u32::MAX);
86    let mut out = Vec::with_capacity(44 + pcm.len());
87    out.extend_from_slice(b"RIFF");
88    out.extend_from_slice(&(36u32.saturating_add(data_size)).to_le_bytes());
89    out.extend_from_slice(b"WAVE");
90    out.extend_from_slice(b"fmt ");
91    out.extend_from_slice(&16u32.to_le_bytes());
92    out.extend_from_slice(&1u16.to_le_bytes());
93    out.extend_from_slice(&channels.to_le_bytes());
94    out.extend_from_slice(&sample_rate.to_le_bytes());
95    out.extend_from_slice(&byte_rate.to_le_bytes());
96    out.extend_from_slice(&block_align.to_le_bytes());
97    out.extend_from_slice(&bits_per_sample.to_le_bytes());
98    out.extend_from_slice(b"data");
99    out.extend_from_slice(&data_size.to_le_bytes());
100    out.extend_from_slice(pcm);
101    Bytes::from(out)
102}
103
104/// Sample rate encoded in a media type such as `audio/L16;codec=pcm;rate=24000`.
105#[must_use]
106pub fn parse_sample_rate(media_type: &str) -> Option<u32> {
107    media_type
108        .split(';')
109        .map(str::trim)
110        .find_map(|parameter| parameter.strip_prefix("rate="))
111        .and_then(|rate| rate.parse().ok())
112}
113
114/// Text-to-speech model backed by the Gemini TTS models.
115#[derive(Debug, Clone)]
116pub struct GoogleSpeechModel {
117    config: SharedConfig,
118    provider: ProviderId,
119    model_id: ModelId,
120}
121
122/// A prepared speech request.
123#[derive(Debug, Clone)]
124pub struct PreparedSpeechRequest {
125    /// Request body.
126    pub body: JsonValue,
127    /// Warnings.
128    pub warnings: Vec<Warning>,
129    /// Whether raw PCM is returned instead of WAV.
130    pub raw_pcm: bool,
131}
132
133impl GoogleSpeechModel {
134    /// Creates the model.
135    #[must_use]
136    pub fn new(config: SharedConfig, model_id: impl Into<ModelId>) -> Self {
137        Self {
138            provider: config.provider_id(FAMILY),
139            config,
140            model_id: model_id.into(),
141        }
142    }
143
144    /// Builds the request body for `options`.
145    ///
146    /// # Errors
147    ///
148    /// Returns [`ProviderError::InvalidArgument`] for invalid provider options.
149    pub fn prepare_request(
150        &self,
151        options: &SpeechOptions,
152    ) -> Result<PreparedSpeechRequest, ProviderError> {
153        let google = parse_merged::<GoogleSpeechOptions>(
154            &self.config,
155            &options.provider_options,
156            |canonical, custom| {
157                if custom.multi_speaker_voice_config.is_some() {
158                    custom
159                } else {
160                    canonical
161                }
162            },
163        )?;
164        if let Some(multi) = &google.multi_speaker_voice_config {
165            let valid = multi
166                .get("speakerVoiceConfigs")
167                .and_then(JsonValue::as_array)
168                .is_some_and(|speakers| {
169                    speakers.iter().all(|speaker| {
170                        speaker.get("speaker").is_some_and(JsonValue::is_string)
171                            && speaker
172                                .pointer("/voiceConfig/prebuiltVoiceConfig/voiceName")
173                                .is_some_and(JsonValue::is_string)
174                    })
175                });
176            if !valid {
177                return Err(InvalidArgumentError::new("multiSpeakerVoiceConfig", "speaker entries require speaker and voiceConfig.prebuiltVoiceConfig.voiceName strings").into());
178            }
179        }
180        let mut warnings = Vec::new();
181        let speech_config = match &google.multi_speaker_voice_config {
182            Some(multi) => json!({"multiSpeakerVoiceConfig": multi}),
183            None => json!({"voiceConfig": {"prebuiltVoiceConfig": {
184                "voiceName": options.voice.as_deref().unwrap_or(DEFAULT_VOICE)
185            }}}),
186        };
187        let mut prompt = options.text.clone();
188        if let Some(instructions) = &options.instructions {
189            if google.multi_speaker_voice_config.is_some() {
190                warnings.push(Warning::unsupported_with_details(
191                    "instructions",
192                    "Google Gemini TTS ignores `instructions` when `multiSpeakerVoiceConfig` is set, because prepending them would break multi-speaker transcript parsing.",
193                ));
194            } else {
195                prompt = format!("{instructions}: {}", options.text);
196            }
197        }
198        if options.speed.is_some() {
199            warnings.push(Warning::unsupported_with_details(
200                "speed",
201                "Google Gemini TTS models do not support the `speed` option. It was ignored.",
202            ));
203        }
204        if options.language.is_some() {
205            warnings.push(Warning::unsupported_with_details(
206                "language",
207                "Google Gemini TTS models do not support the `language` option. Language is detected automatically from the input text.",
208            ));
209        }
210        let raw_pcm = match options.output_format.as_deref() {
211            Some("pcm") => true,
212            None | Some("wav") => false,
213            Some(other) => {
214                warnings.push(Warning::unsupported_with_details(
215                    "outputFormat",
216                    format!("Unsupported output format: {other}. Using wav instead."),
217                ));
218                false
219            }
220        };
221        let body = json!({
222            "contents": [{"role": "user", "parts": [{"text": prompt}]}],
223            "generationConfig": {
224                "responseModalities": ["AUDIO"],
225                "speechConfig": speech_config,
226            },
227        });
228        Ok(PreparedSpeechRequest {
229            body,
230            warnings,
231            raw_pcm,
232        })
233    }
234}
235
236impl SpeechModel for GoogleSpeechModel {
237    fn provider(&self) -> &ProviderId {
238        &self.provider
239    }
240
241    fn model_id(&self) -> &ModelId {
242        &self.model_id
243    }
244
245    #[tracing::instrument(skip_all, fields(model = %self.model_id))]
246    async fn do_generate(&self, options: SpeechOptions) -> Result<SpeechResult, ProviderError> {
247        let prepared = self.prepare_request(&options)?;
248        let mut warnings = prepared.warnings;
249        let handlers = ResponseHandlers::new(
250            json_response_handler::<SpeechResponse>(),
251            failed_response_handler(),
252        );
253        let response = post_json(
254            self.config.transport.as_ref(),
255            self.config
256                .model_url(self.model_id.as_str(), "generateContent"),
257            self.config.headers(&options.headers)?,
258            &prepared.body,
259            &handlers,
260            options.cancellation.clone(),
261        )
262        .await?;
263        let inline = response
264            .value
265            .candidates
266            .iter()
267            .flatten()
268            .filter_map(|candidate| candidate.content.as_ref())
269            .flat_map(|content| content.parts.iter().flatten())
270            .find_map(|part| {
271                part.inline_data
272                    .as_ref()
273                    .filter(|data| data.data.as_ref().is_some_and(|data| !data.is_empty()))
274            });
275        let mime_type = inline.and_then(|data| data.mime_type.clone());
276        let sample_rate = mime_type
277            .as_deref()
278            .and_then(parse_sample_rate)
279            .unwrap_or(DEFAULT_SAMPLE_RATE);
280        let pcm = match inline {
281            Some(data) => base64::engine::general_purpose::STANDARD
282                .decode(data.data.as_deref().unwrap_or_default())
283                .map_err(|error| {
284                    ProviderError::InvalidResponseData(Box::new(
285                        ferrin_spec::error::InvalidResponseDataError::new(
286                            format!("invalid base64 audio data: {error}"),
287                            JsonValue::Null,
288                        ),
289                    ))
290                })?,
291            None => Vec::new(),
292        };
293        let (audio, media_type) = if prepared.raw_pcm || pcm.is_empty() {
294            if prepared.raw_pcm && !pcm.is_empty() {
295                warnings.push(Warning::unsupported_with_details(
296                    "outputFormat",
297                    format!(
298                        "Returning raw PCM audio (signed 16-bit little-endian, mono, {sample_rate} Hz). These bytes have no container header and are not directly playable; see providerMetadata.google for the sample rate and mime type."
299                    ),
300                ));
301            }
302            let media_type = mime_type.clone().map(MediaType::new);
303            (Bytes::from(pcm), media_type)
304        } else {
305            (
306                add_wav_header(&pcm, sample_rate),
307                Some(MediaType::new("audio/wav")),
308            )
309        };
310        let mapper = OutputMapper::new(self.config.clone(), Default::default());
311        let mut metadata = JsonObject::new();
312        metadata.insert("sampleRate".to_owned(), JsonValue::from(sample_rate));
313        metadata.insert(
314            "mimeType".to_owned(),
315            mime_type.map_or(JsonValue::Null, JsonValue::from),
316        );
317        Ok(SpeechResult {
318            audio,
319            media_type,
320            warnings,
321            request: RequestMetadata::with_body(prepared.body),
322            response: ResponseMetadata {
323                id: response.value.response_id.clone(),
324                timestamp: Some(chrono::Utc::now()),
325                model_id: Some(self.model_id.clone()),
326                headers: Some(response.response_headers),
327                body: response.raw,
328            },
329            provider_metadata: Some(mapper.metadata(metadata)),
330        })
331    }
332}