rig_core/providers/gemini/
transcription.rs1use std::path::Path;
2
3use base64::{Engine, prelude::BASE64_STANDARD};
4use serde_json::{Map, Value, json};
5
6use super::completion::usage_of;
7use crate::error::{EncodeError, ProviderError};
8use crate::json_utils::Lenient;
9use crate::operation::Transcription;
10use crate::providers::internal::wire::classify_marker_keyed_frame;
11use crate::transcription;
12use crate::wire::{
13 Body, Decoder, Descriptor, Encoded, Flow, Framing, Mode, Out, Wire, WireEvent, WireFrame,
14};
15
16const TRANSCRIPTION_PREAMBLE: &str =
17 "Translate the provided audio exactly. Do not add additional information.";
18
19fn transcription_body(
22 request: transcription::TranscriptionRequest,
23) -> Result<Vec<u8>, EncodeError> {
24 let mut generation_config = match request.additional_params {
25 None | Some(Value::Null) => Map::new(),
26 Some(Value::Object(config)) => config,
27 Some(other) => {
28 return Err(EncodeError::request(format!(
29 "Gemini transcription `additional_params` should be an object, got {other}"
30 )));
31 }
32 };
33 if let Some(temp) = request.temperature {
36 generation_config.insert("temperature".to_owned(), Value::from(temp));
37 }
38 let mime_type = mime_guess::from_path(Path::new(&request.filename))
40 .first()
41 .map_or_else(|| "audio/mpeg".to_string(), |mime| mime.to_string());
42 let data = BASE64_STANDARD.encode(request.data);
43 let body = json!({
44 "contents": [{
45 "parts": [{ "inlineData": { "mimeType": mime_type, "data": data }, "thought": false }],
46 "role": "user",
47 }],
48 "generationConfig": generation_config,
49 "safetySettings": null,
50 "toolConfig": null,
51 "systemInstruction": {
52 "parts": [{ "text": TRANSCRIPTION_PREAMBLE, "thought": false }],
53 "role": "model",
54 },
55 });
56 tracing::trace!(
57 target: "rig::transcription",
58 "Sending completion request to Gemini API {}",
59 serde_json::to_string_pretty(&body)?
60 );
61 Ok(serde_json::to_vec(&body)?)
62}
63
64#[derive(Clone, Debug, PartialEq, serde::Serialize, serde::Deserialize)]
68pub struct Transcriptions {
69 pub provider: super::GeminiConfig,
71 pub model: String,
74}
75
76impl Transcriptions {
77 pub fn new(provider: super::GeminiConfig, model: impl Into<String>) -> Self {
79 Self {
80 provider,
81 model: model.into(),
82 }
83 }
84}
85
86impl Wire for Transcriptions {
87 type Op = Transcription;
88 type Payload = crate::wire::Encoded;
89 type Frame = crate::wire::WireFrame;
90 type Decoder<'id> = TranscriptionsDecoder;
91 type Reassembler = crate::wire::document::Unreassembled;
92
93 fn describe(&self) -> Descriptor<'_> {
94 Descriptor::new(super::PROVIDER_NAME).model(self.model.as_str())
95 }
96
97 fn encode(
98 &self,
99 request: transcription::TranscriptionRequest,
100 _mode: Mode,
101 ) -> Result<Encoded, EncodeError> {
102 let body = transcription_body(request)?;
103 let request = http::Request::post(format!(
104 "{}/v1beta/models/{}:generateContent?key={}",
105 self.provider.base_url,
106 self.model,
107 self.provider.api_key.expose()
108 ))
109 .header(http::header::CONTENT_TYPE, "application/json")
110 .body(Body::Bytes(body))?;
111 Ok(Encoded::new(request, Framing::Whole))
113 }
114
115 fn decoder<'id>(&self) -> Self::Decoder<'id> {
116 TranscriptionsDecoder
117 }
118}
119
120#[derive(Default)]
123pub struct TranscriptionsDecoder;
124
125impl<'id> Decoder<'id, Transcription> for TranscriptionsDecoder {
126 type Event = Value;
127
128 fn classify(&self, frame: WireFrame) -> WireEvent<Self::Event> {
129 classify_marker_keyed_frame(
130 &frame.as_str(),
131 &["candidates", "promptFeedback", "usageMetadata"],
132 )
133 }
134
135 fn decode(
136 &mut self,
137 event: Self::Event,
138 out: Out<'id, Transcription>,
139 ) -> Result<Flow, ProviderError> {
140 Ok(out.end(transcript_of(&event)?))
141 }
142}
143
144pub fn transcript_of(reply: &Value) -> Result<transcription::TranscriptionResponse, ProviderError> {
148 let candidate = reply
149 .arr("candidates")
150 .first()
151 .ok_or_else(|| ProviderError::Response("No response candidates in response".into()))?;
152 let parts: Vec<&str> = candidate
153 .get("content")
154 .map(|content| content.arr("parts"))
155 .unwrap_or_default()
156 .iter()
157 .filter(|part| part.bool("thought") != Some(true))
158 .filter_map(|part| part.str("text"))
159 .collect();
160 if parts.is_empty() {
161 return Err(ProviderError::Response(
162 "Response content contains no text".to_string(),
163 ));
164 }
165 Ok(transcription::TranscriptionResponse {
166 model: reply.str("modelVersion").map(str::to_owned),
167 response_id: Some(reply.str("responseId").unwrap_or_default().to_owned()),
168 usage: reply.get("usageMetadata").map(usage_of).unwrap_or_default(),
169 ..transcription::TranscriptionResponse::new(parts.concat())
170 })
171}
172
173impl super::GeminiConfig {
174 pub(crate) fn transcription(&self, model: impl Into<String>) -> Transcriptions {
176 Transcriptions::new(self.clone(), model)
177 }
178}
179
180#[cfg(test)]
181mod tests;