Skip to main content

rig_core/providers/gemini/
transcription.rs

1use 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
19/// Encode audio as inline base64 with a transcription system instruction.
20/// Reject invalid generation parameters or JSON serialization failures.
21fn 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    // A temperature named on the request outranks one carried inside
34    // `additional_params`.
35    if let Some(temp) = request.temperature {
36        generation_config.insert("temperature".to_owned(), Value::from(temp));
37    }
38    // The request supplies no explicit MIME type, so infer it from the filename.
39    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/// The transcription wire: `POST /v1beta/models/{model}:generateContent`.
65///
66/// Both [`Mode`]s send inline audio in JSON and read a whole response.
67#[derive(Clone, Debug, PartialEq, serde::Serialize, serde::Deserialize)]
68pub struct Transcriptions {
69    /// The provider this wire speaks to.
70    pub provider: super::GeminiConfig,
71    /// The model transcribing, for example
72    /// [`GEMINI_2_0_FLASH`](super::completion::GEMINI_2_0_FLASH).
73    pub model: String,
74}
75
76impl Transcriptions {
77    /// The transcription wire for `model`.
78    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        // Gemini reports no transport request-id header.
112        Ok(Encoded::new(request, Framing::Whole))
113    }
114
115    fn decoder<'id>(&self) -> Self::Decoder<'id> {
116        TranscriptionsDecoder
117    }
118}
119
120/// Decode visible text from the first `generateContent` candidate.
121/// Missing candidates or visible text produce response errors.
122#[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
144/// The first candidate's visible text parts, those not marked `thought`, as
145/// a transcript of `reply`. Errors when there is no candidate or no visible
146/// text part.
147pub 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    /// The audio transcription wire.
175    pub(crate) fn transcription(&self, model: impl Into<String>) -> Transcriptions {
176        Transcriptions::new(self.clone(), model)
177    }
178}
179
180#[cfg(test)]
181mod tests;