rig_core/providers/gemini/
embedding.rs1use 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
24pub const EMBEDDING_001: &str = "gemini-embedding-001";
26pub const EMBEDDING_004: &str = "text-embedding-004";
28
29fn 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#[derive(Clone, Debug, PartialEq, serde::Serialize, serde::Deserialize)]
45pub struct Embeddings {
46 pub provider: super::GeminiConfig,
48 pub model: String,
50 pub ndims: usize,
52}
53
54impl Embeddings {
55 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 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 Ok(Encoded::new(request, Framing::Whole))
124 }
125
126 fn decoder<'id>(&self) -> Self::Decoder<'id> {
127 EmbeddingsDecoder
128 }
129}
130
131#[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 Ok(out.end(embeddings::EmbeddingResponse::from_vectors(vectors)))
156 }
157}
158
159pub 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 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;