Skip to main content

rig_core/embeddings/
embedding.rs

1//! Text and image embedding models, responses, and input identifiers.
2//!
3//! ```no_run
4//! use rig_core::DynModel;
5//! use rig_core::operation::Embedding;
6//!
7//! # async fn example(model: DynModel<Embedding>) -> Result<(), Box<dyn std::error::Error>> {
8//! let embedding = model.embed_text("A document").await?;
9//! # let _ = embedding;
10//! # Ok(())
11//! # }
12//! ```
13
14use crate::completion::Usage;
15use crate::error::ProviderError;
16use serde::{Deserialize, Serialize};
17
18impl<W, T> crate::driver::Model<W, T>
19where
20    W: crate::wire::Wire<Op = crate::operation::Embedding>,
21    T: crate::driver::Transport<W>,
22{
23    /// Embed one text, returning the last vector or an error if none is returned.
24    pub async fn embed_text(&self, text: &str) -> Result<Embedding, ProviderError> {
25        last_embedding(self.call(vec![text.to_owned()]).await?)
26    }
27}
28
29impl crate::driver::DynModel<crate::operation::Embedding> {
30    /// Embed one text, returning the last vector or an error if none is returned.
31    pub async fn embed_text(&self, text: &str) -> Result<Embedding, ProviderError> {
32        last_embedding(self.call(vec![text.to_owned()]).await?)
33    }
34}
35
36/// The last vector of a one-text batch, or the empty-reply error.
37fn last_embedding(response: EmbeddingResponse) -> Result<Embedding, ProviderError> {
38    let mut embeddings = response.embeddings;
39    embeddings.pop().ok_or_else(|| {
40        ProviderError::Response(
41            "embedding provider returned an empty response for embed_text".to_string(),
42        )
43    })
44}
45
46/// Text or image embeddings and normalized provider metadata.
47#[derive(Debug, Clone, Serialize, Deserialize)]
48pub struct EmbeddingResponse {
49    /// The embeddings returned by the provider, one per input, in input order.
50    pub embeddings: Vec<Embedding>,
51    /// Token usage for this request; every counter is `None` when the
52    /// provider reported none (see [`Usage`]).
53    #[serde(default)]
54    pub usage: Usage,
55    /// Stable descriptor name of the provider that produced this response,
56    /// for example `"openai"`. Always populated.
57    pub provider: String,
58    /// Provider-reported model identifier, when the wire response named one.
59    #[serde(default)]
60    pub model: Option<String>,
61    /// Provider-assigned response-scoped identifier, when reported.
62    #[serde(default, skip_serializing_if = "Option::is_none")]
63    pub response_id: Option<String>,
64    /// Transport request identifier from HTTP response headers, or `None` when absent.
65    #[serde(default, skip_serializing_if = "Option::is_none")]
66    pub provider_request_id: Option<String>,
67    /// Provider response payload, or null when no raw payload was attached.
68    #[serde(default, skip_serializing_if = "serde_json::Value::is_null")]
69    pub raw: serde_json::Value,
70}
71
72impl EmbeddingResponse {
73    /// A response carrying `embeddings`. The driver writes the provider, the
74    /// transport request id and the reply document; decoders set what the
75    /// provider reported.
76    pub fn new(embeddings: Vec<Embedding>) -> Self {
77        Self {
78            embeddings,
79            usage: Usage::default(),
80            provider: String::new(),
81            model: None,
82            response_id: None,
83            provider_request_id: None,
84            raw: serde_json::Value::Null,
85        }
86    }
87
88    /// A response carrying bare `vectors` in input order. The operation's
89    /// fold pairs each with the input it belongs to, since a reply is not
90    /// trusted to echo its inputs back.
91    pub(crate) fn from_vectors(vectors: impl IntoIterator<Item = Vec<f64>>) -> Self {
92        Self::new(
93            vectors
94                .into_iter()
95                .map(|vec| Embedding {
96                    document: String::new(),
97                    vec,
98                })
99                .collect(),
100        )
101    }
102}
103
104/// A document identifier and its vector. Equality compares only the document,
105/// not vector values.
106#[derive(Clone, Default, Deserialize, Serialize, Debug)]
107pub struct Embedding {
108    /// The text that was embedded, or a non-sensitive input identifier for
109    /// non-text embeddings. Used for debugging and equality.
110    pub document: String,
111    /// The embedding vector
112    pub vec: Vec<f64>,
113}
114
115impl PartialEq for Embedding {
116    fn eq(&self, other: &Self) -> bool {
117        self.document == other.document
118    }
119}
120
121impl Eq for Embedding {}
122
123#[cfg(test)]
124mod provider_response_tests;
125
126/// The media type of an encoded image, sniffed from its magic bytes.
127///
128/// The image-embedding wires need it twice: once to reject a format the
129/// provider does not accept, and once to name the vector's input.
130pub fn image_media_type(bytes: &[u8]) -> Option<&'static str> {
131    if bytes.starts_with(b"\x89PNG\r\n\x1a\n") {
132        Some("image/png")
133    } else if bytes.starts_with(b"\xff\xd8\xff") {
134        Some("image/jpeg")
135    } else if bytes.starts_with(b"GIF87a") || bytes.starts_with(b"GIF89a") {
136        Some("image/gif")
137    } else if bytes.starts_with(b"RIFF") && bytes.get(8..12) == Some(b"WEBP".as_slice()) {
138        Some("image/webp")
139    } else {
140        None
141    }
142}
143
144/// Identifies image bytes by media type and a URL-safe, unpadded SHA-256 digest,
145/// without retaining the image or a reversible encoding. Unknown formats use
146/// `application/octet-stream`; this function does not validate provider support.
147pub fn image_document(bytes: &[u8]) -> String {
148    use base64::Engine as _;
149    use sha2::Digest as _;
150    let media_type = image_media_type(bytes).unwrap_or("application/octet-stream");
151    let digest = sha2::Sha256::digest(bytes);
152    format!(
153        "{media_type};sha256={}",
154        base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(digest)
155    )
156}