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 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#[derive(Clone, Default, Deserialize, Serialize, Debug)]
107pub struct Embedding {
108 pub document: String,
111 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
126pub 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
144pub 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}