1use 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
29pub const FAMILY: &str = "speech";
31
32pub const DEFAULT_VOICE: &str = "Kore";
34
35pub const DEFAULT_SAMPLE_RATE: u32 = 24_000;
37
38#[derive(Debug, Clone, Default, PartialEq, Deserialize)]
40#[serde(rename_all = "camelCase")]
41pub struct GoogleSpeechOptions {
42 #[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#[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#[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#[derive(Debug, Clone)]
116pub struct GoogleSpeechModel {
117 config: SharedConfig,
118 provider: ProviderId,
119 model_id: ModelId,
120}
121
122#[derive(Debug, Clone)]
124pub struct PreparedSpeechRequest {
125 pub body: JsonValue,
127 pub warnings: Vec<Warning>,
129 pub raw_pcm: bool,
131}
132
133impl GoogleSpeechModel {
134 #[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 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}