rig_core/embeddings/
embedding.rs1use 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 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 pub async fn embed_text(&self, text: &str) -> Result<Embedding, ProviderError> {
32 last_embedding(self.call(vec![text.to_owned()]).await?)
33 }
34}
35
36fn 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#[derive(Debug, Clone, Serialize, Deserialize)]
48pub struct EmbeddingResponse {
49 pub embeddings: Vec<Embedding>,
51 #[serde(default)]
54 pub usage: Usage,
55 pub provider: String,
58 #[serde(default)]
60 pub model: Option<String>,
61 #[serde(default, skip_serializing_if = "Option::is_none")]
63 pub response_id: Option<String>,
64 #[serde(default, skip_serializing_if = "Option::is_none")]
66 pub provider_request_id: Option<String>,
67 #[serde(default, skip_serializing_if = "serde_json::Value::is_null")]
69 pub raw: serde_json::Value,
70}
71
72impl EmbeddingResponse {
73 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#[derive(Debug, Clone, Serialize, Deserialize)]
91pub struct ImageEmbeddingResponse {
92 pub embeddings: Vec<Embedding>,
94 #[serde(default)]
97 pub usage: Usage,
98 pub provider: String,
101 #[serde(default)]
103 pub model: Option<String>,
104 #[serde(default, skip_serializing_if = "Option::is_none")]
106 pub response_id: Option<String>,
107 #[serde(default, skip_serializing_if = "Option::is_none")]
109 pub provider_request_id: Option<String>,
110 #[serde(default, skip_serializing_if = "serde_json::Value::is_null")]
112 pub raw: serde_json::Value,
113}
114
115impl ImageEmbeddingResponse {
116 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#[derive(Clone, Default, Deserialize, Serialize, Debug)]
135pub struct Embedding {
136 pub document: String,
139 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
154pub 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
172pub 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}