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 embeddings and normalized provider metadata.
47#[derive(Debug, Clone, Serialize, Deserialize)]
48pub struct EmbeddingResponse {
49    /// The embeddings returned by the provider, one per input text, 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
89/// Image embeddings and normalized provider metadata.
90#[derive(Debug, Clone, Serialize, Deserialize)]
91pub struct ImageEmbeddingResponse {
92    /// The embeddings returned by the provider, one per input image, in input order.
93    pub embeddings: Vec<Embedding>,
94    /// Token usage for this request; every counter is `None` when the
95    /// provider reported none (see [`Usage`]).
96    #[serde(default)]
97    pub usage: Usage,
98    /// Stable descriptor name of the provider that produced this response,
99    /// for example `"openai"`. Always populated.
100    pub provider: String,
101    /// Provider-reported model identifier, when the wire response named one.
102    #[serde(default)]
103    pub model: Option<String>,
104    /// Provider-assigned response-scoped identifier, when reported.
105    #[serde(default, skip_serializing_if = "Option::is_none")]
106    pub response_id: Option<String>,
107    /// Transport request identifier from HTTP response headers, or `None` when absent.
108    #[serde(default, skip_serializing_if = "Option::is_none")]
109    pub provider_request_id: Option<String>,
110    /// Provider response payload, or null when no raw payload was attached.
111    #[serde(default, skip_serializing_if = "serde_json::Value::is_null")]
112    pub raw: serde_json::Value,
113}
114
115impl ImageEmbeddingResponse {
116    /// A response carrying `embeddings`. The driver writes the provider, the
117    /// transport request id and the reply document; decoders set what the
118    /// provider reported.
119    pub fn new(embeddings: Vec<Embedding>) -> Self {
120        Self {
121            embeddings,
122            usage: Usage::default(),
123            provider: String::new(),
124            model: None,
125            response_id: None,
126            provider_request_id: None,
127            raw: serde_json::Value::Null,
128        }
129    }
130}
131
132/// A document identifier and its vector. Equality compares only the document,
133/// not vector values.
134#[derive(Clone, Default, Deserialize, Serialize, Debug)]
135pub struct Embedding {
136    /// The text that was embedded, or a non-sensitive input identifier for
137    /// non-text embeddings. Used for debugging and equality.
138    pub document: String,
139    /// The embedding vector
140    pub vec: Vec<f64>,
141}
142
143impl PartialEq for Embedding {
144    fn eq(&self, other: &Self) -> bool {
145        self.document == other.document
146    }
147}
148
149impl Eq for Embedding {}
150
151#[cfg(test)]
152mod provider_response_tests;
153
154/// The media type of an encoded image, sniffed from its magic bytes.
155///
156/// The image-embedding wires need it twice: once to reject a format the
157/// provider does not accept, and once to name the vector's input.
158pub fn image_media_type(bytes: &[u8]) -> Option<&'static str> {
159    if bytes.starts_with(b"\x89PNG\r\n\x1a\n") {
160        Some("image/png")
161    } else if bytes.starts_with(b"\xff\xd8\xff") {
162        Some("image/jpeg")
163    } else if bytes.starts_with(b"GIF87a") || bytes.starts_with(b"GIF89a") {
164        Some("image/gif")
165    } else if bytes.starts_with(b"RIFF") && bytes.get(8..12) == Some(b"WEBP".as_slice()) {
166        Some("image/webp")
167    } else {
168        None
169    }
170}
171
172/// Identifies image bytes by media type and a URL-safe, unpadded SHA-256 digest,
173/// without retaining the image or a reversible encoding. Unknown formats use
174/// `application/octet-stream`; this function does not validate provider support.
175pub fn image_document(bytes: &[u8]) -> String {
176    use base64::Engine as _;
177    use sha2::Digest as _;
178    let media_type = image_media_type(bytes).unwrap_or("application/octet-stream");
179    let digest = sha2::Sha256::digest(bytes);
180    format!(
181        "{media_type};sha256={}",
182        base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(digest)
183    )
184}