use crate::error::ProviderError;
use crate::wire::Flow;
use serde_json::json;
use crate::embeddings;
use crate::error::EncodeError;
use crate::providers::internal::wire::classify_marker_keyed_frame;
use crate::wire::{
Body, Capabilities, Decoder, Descriptor, Encoded, Framing, Mode, Out, Wire, WireEvent,
WireFrame,
};
pub const EMBEDDING_001: &str = "gemini-embedding-001";
pub const EMBEDDING_004: &str = "text-embedding-004";
fn model_default_ndims(model: &str) -> Option<usize> {
match model {
EMBEDDING_001 => Some(3072),
EMBEDDING_004 => Some(768),
_ => None,
}
}
#[derive(Clone, Debug, PartialEq, serde::Serialize, serde::Deserialize)]
pub struct Embeddings {
pub provider: super::GeminiConfig,
pub model: String,
pub ndims: usize,
}
impl Embeddings {
pub fn new(
provider: super::GeminiConfig,
model: impl Into<String>,
ndims: Option<usize>,
) -> Self {
let model = model.into();
let ndims = ndims.or_else(|| model_default_ndims(&model)).unwrap_or(768);
Self {
provider,
model,
ndims,
}
}
}
impl Wire for Embeddings {
type Op = crate::operation::Embedding;
type Payload = crate::wire::Encoded;
type Frame = crate::wire::WireFrame;
type Decoder<'id> = EmbeddingsDecoder;
type Reassembler = crate::wire::document::Unreassembled;
fn describe(&self) -> Descriptor<'_> {
Descriptor::new(super::PROVIDER_NAME)
.model(self.model.as_str())
.capabilities(Capabilities::embedding(1024, self.ndims))
}
fn encode(&self, request: Vec<String>, _mode: Mode) -> Result<Encoded, EncodeError> {
let requests: Vec<_> = request
.iter()
.map(|doc| {
json!({
"model": format!("models/{}", self.model),
"content": json!({
"parts": [json!({
"text": doc
})]
}),
"output_dimensionality": self.ndims,
})
})
.collect();
let body = json!({ "requests": requests });
if tracing::enabled!(target: "rig::embedding", tracing::Level::TRACE)
&& let Ok(pretty_body) = serde_json::to_string_pretty(&body)
{
tracing::trace!(
target: "rig::embedding",
"Sending embedding request to Gemini API {pretty_body}"
);
}
let request = http::Request::post(format!(
"{}/v1beta/models/{}:batchEmbedContents?key={}",
self.provider.base_url,
self.model,
self.provider.api_key.expose()
))
.header(http::header::CONTENT_TYPE, "application/json")
.body(Body::Bytes(serde_json::to_vec(&body)?))?;
Ok(Encoded::new(request, Framing::Whole))
}
fn decoder<'id>(&self) -> Self::Decoder<'id> {
EmbeddingsDecoder
}
}
#[derive(Default)]
pub struct EmbeddingsDecoder;
impl<'id> Decoder<'id, crate::operation::Embedding> for EmbeddingsDecoder {
type Event = gemini_api_types::EmbeddingResponse;
fn classify(&self, frame: WireFrame) -> WireEvent<Self::Event> {
classify_marker_keyed_frame(&frame.as_str(), &["embeddings"])
}
fn decode(
&mut self,
event: Self::Event,
out: Out<'id, crate::operation::Embedding>,
) -> Result<Flow, ProviderError> {
let vectors = event.embeddings.into_iter().map(|embedding| {
embedding
.values
.into_iter()
.filter_map(|value| value.as_f64())
.collect()
});
Ok(out.end(embeddings::EmbeddingResponse::from_vectors(vectors)))
}
}
pub mod gemini_api_types {
use serde::{Deserialize, Serialize};
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct EmbeddingResponse {
pub embeddings: Vec<EmbeddingValues>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct EmbeddingValues {
#[serde(default)]
pub values: Vec<serde_json::Number>,
}
}
impl super::GeminiConfig {
pub(crate) fn embedding(&self, model: impl Into<String>, ndims: Option<usize>) -> Embeddings {
Embeddings::new(self.clone(), model, ndims)
}
}
#[cfg(test)]
mod tests;