Skip to main content

rig_core/providers/gemini/
embedding.rs

1//! Batch text embeddings through the [Gemini API](https://ai.google.dev/api/embeddings).
2//!
3//! ```no_run
4//! use rig_core::providers::gemini::{Gemini, embedding::EMBEDDING_001};
5//!
6//! # fn main() -> Result<(), Box<dyn std::error::Error>> {
7//! let model = Gemini::from_env()?.embedding(EMBEDDING_001, None);
8//! # Ok(())
9//! # }
10//! ```
11
12use crate::error::ProviderError;
13use crate::wire::Flow;
14use serde_json::json;
15
16use crate::embeddings;
17use crate::error::EncodeError;
18use crate::providers::internal::wire::classify_marker_keyed_frame;
19use crate::wire::{
20    Body, Capabilities, Decoder, Descriptor, Encoded, Framing, Mode, Out, Wire, WireEvent,
21    WireFrame,
22};
23
24/// `gemini-embedding-001` embedding model (3072 dimensions by default)
25pub const EMBEDDING_001: &str = "gemini-embedding-001";
26/// `text-embedding-004` embedding model (768 dimensions by default)
27pub const EMBEDDING_004: &str = "text-embedding-004";
28
29/// Returns the default output dimensionality for known Gemini embedding models.
30///
31/// See <https://ai.google.dev/gemini-api/docs/models#gemini-embedding>
32fn model_default_ndims(model: &str) -> Option<usize> {
33    match model {
34        EMBEDDING_001 => Some(3072),
35        EMBEDDING_004 => Some(768),
36        _ => None,
37    }
38}
39
40/// Gemini's batch embedding endpoint.
41///
42/// `POST /v1beta/models/{model}:batchEmbedContents`, authenticated by the
43/// `key` query parameter the GenerateContent family uses.
44#[derive(Clone, Debug, PartialEq, serde::Serialize, serde::Deserialize)]
45pub struct Embeddings {
46    /// The provider this wire speaks to.
47    pub provider: super::GeminiConfig,
48    /// The embedding model, as the path names it.
49    pub model: String,
50    /// The `output_dimensionality` every document in the batch asks for.
51    pub ndims: usize,
52}
53
54impl Embeddings {
55    /// Build a wire for `model` with the requested output dimensions.
56    /// If `ndims` is absent, use the model default or 768 for unknown models.
57    pub fn new(
58        provider: super::GeminiConfig,
59        model: impl Into<String>,
60        ndims: Option<usize>,
61    ) -> Self {
62        let model = model.into();
63        let ndims = ndims.or_else(|| model_default_ndims(&model)).unwrap_or(768);
64        Self {
65            provider,
66            model,
67            ndims,
68        }
69    }
70}
71
72impl Wire for Embeddings {
73    type Op = crate::operation::Embedding;
74    type Payload = crate::wire::Encoded;
75    type Frame = crate::wire::WireFrame;
76    type Decoder<'id> = EmbeddingsDecoder;
77    type Reassembler = crate::wire::document::Unreassembled;
78
79    fn describe(&self) -> Descriptor<'_> {
80        Descriptor::new(super::PROVIDER_NAME)
81            .model(self.model.as_str())
82            .capabilities(Capabilities::embedding(1024, self.ndims))
83    }
84
85    fn encode(&self, request: Vec<String>, _mode: Mode) -> Result<Encoded, EncodeError> {
86        let requests: Vec<_> = request
87            .iter()
88            .map(|doc| {
89                json!({
90                    "model": format!("models/{}", self.model),
91                    "content": json!({
92                        "parts": [json!({
93                            "text": doc
94                        })]
95                    }),
96                    "output_dimensionality": self.ndims,
97                })
98            })
99            .collect();
100
101        let body = json!({ "requests": requests  });
102
103        // Avoid allocating formatted batch JSON unless tracing consumes it.
104        if tracing::enabled!(target: "rig::embedding", tracing::Level::TRACE)
105            && let Ok(pretty_body) = serde_json::to_string_pretty(&body)
106        {
107            tracing::trace!(
108                target: "rig::embedding",
109                "Sending embedding request to Gemini API {pretty_body}"
110            );
111        }
112
113        let request = http::Request::post(format!(
114            "{}/v1beta/models/{}:batchEmbedContents?key={}",
115            self.provider.base_url,
116            self.model,
117            self.provider.api_key.expose()
118        ))
119        .header(http::header::CONTENT_TYPE, "application/json")
120        .body(Body::Bytes(serde_json::to_vec(&body)?))?;
121        // `batchEmbedContents` has no streaming variant, so a streamed call
122        // sends the same bytes and reads the same whole reply.
123        Ok(Encoded::new(request, Framing::Whole))
124    }
125
126    fn decoder<'id>(&self) -> Self::Decoder<'id> {
127        EmbeddingsDecoder
128    }
129}
130
131/// Decode a `batchEmbedContents` reply containing vectors without usage or identity.
132#[derive(Default)]
133pub struct EmbeddingsDecoder;
134
135impl<'id> Decoder<'id, crate::operation::Embedding> for EmbeddingsDecoder {
136    type Event = gemini_api_types::EmbeddingResponse;
137
138    fn classify(&self, frame: WireFrame) -> WireEvent<Self::Event> {
139        classify_marker_keyed_frame(&frame.as_str(), &["embeddings"])
140    }
141
142    fn decode(
143        &mut self,
144        event: Self::Event,
145        out: Out<'id, crate::operation::Embedding>,
146    ) -> Result<Flow, ProviderError> {
147        let vectors = event.embeddings.into_iter().map(|embedding| {
148            embedding
149                .values
150                .into_iter()
151                .filter_map(|value| value.as_f64())
152                .collect()
153        });
154        // Gemini supplies no usage or response id; the driver attaches the raw body.
155        Ok(out.end(embeddings::EmbeddingResponse::from_vectors(vectors)))
156    }
157}
158
159/// Request and response types for [Gemini embeddings](https://ai.google.dev/api/embeddings).
160///
161/// ```
162/// use rig_core::providers::gemini::embedding::gemini_api_types::EmbeddingValues;
163///
164/// let vector = EmbeddingValues { values: vec![1.into(), 2.into()] };
165/// ```
166pub mod gemini_api_types {
167    use serde::{Deserialize, Serialize};
168
169    #[derive(Debug, Clone, Serialize, Deserialize)]
170    pub struct EmbeddingResponse {
171        pub embeddings: Vec<EmbeddingValues>,
172    }
173
174    #[derive(Debug, Clone, Serialize, Deserialize)]
175    pub struct EmbeddingValues {
176        #[serde(default)]
177        pub values: Vec<serde_json::Number>,
178    }
179}
180
181impl super::GeminiConfig {
182    /// The `batchEmbedContents` embedding wire. `ndims` defaults from the
183    /// model identifier.
184    pub(crate) fn embedding(&self, model: impl Into<String>, ndims: Option<usize>) -> Embeddings {
185        Embeddings::new(self.clone(), model, ndims)
186    }
187}
188
189#[cfg(test)]
190mod tests;