rig_core/providers/voyageai/
wire.rs1use crate::client::env::{self, EnvError};
10use crate::error::EncodeError;
11use crate::error::ProviderError;
12use crate::operation::{Embedding, Rerank as RerankOp, RerankRequest};
13use crate::rerank::{RerankResponse, RerankResult};
14use crate::wire::Flow;
15use crate::wire::{
16 Body, Capabilities, Decoder, Descriptor, Encoded, Framing, Mode, Out, Secret, Wire, WireEvent,
17 WireFrame,
18};
19use serde::{Deserialize, Serialize};
20
21use super::{
22 EmbeddingResponse as VoyageEmbeddingResponse, RerankApiResponse, VOYAGEAI_API_BASE_URL,
23 model_dimensions_from_identifier,
24};
25
26const PROVIDER_NAME: &str = "voyageai";
28
29const API_KEY_ENV: &str = "VOYAGE_API_KEY";
31
32const MAX_DOCUMENTS: usize = 1024;
34
35const MAX_RERANK_DOCUMENTS: usize = 1000;
37
38#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
42pub struct VoyageAiConfig {
43 pub api_key: Secret,
45 pub base_url: String,
47}
48
49impl VoyageAiConfig {
50 pub fn new(api_key: impl Into<Secret>) -> Self {
52 Self {
53 api_key: api_key.into(),
54 base_url: VOYAGEAI_API_BASE_URL.to_owned(),
55 }
56 }
57
58 pub fn from_env() -> Result<Self, EnvError> {
60 Ok(Self::new(env::required(API_KEY_ENV)?))
61 }
62
63 pub fn with_base_url(mut self, base_url: impl AsRef<str>) -> Self {
65 self.base_url = base_url.as_ref().trim_end_matches('/').to_owned();
66 self
67 }
68
69 pub(crate) fn embedding(&self, model: impl Into<String>, ndims: Option<usize>) -> Embeddings {
73 let model = model.into();
74 let ndims = ndims
75 .or_else(|| model_dimensions_from_identifier(&model))
76 .unwrap_or_default();
77 Embeddings {
78 provider: self.clone(),
79 model,
80 ndims,
81 input_type: None,
82 truncation: None,
83 output_dimension: None,
84 }
85 }
86
87 pub(crate) fn rerank(&self, model: impl Into<String>) -> Rerank {
89 Rerank {
90 provider: self.clone(),
91 model: model.into(),
92 top_k: None,
93 return_documents: false,
94 truncation: None,
95 }
96 }
97
98 fn post(&self, path: &str) -> http::request::Builder {
100 http::Request::post(format!("{}{path}", self.base_url))
101 .header(http::header::CONTENT_TYPE, "application/json")
102 .header(
103 http::header::AUTHORIZATION,
104 format!("Bearer {}", self.api_key.expose()),
105 )
106 }
107}
108
109#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
115pub struct Embeddings {
116 pub provider: VoyageAiConfig,
118 pub model: String,
120 pub ndims: usize,
123 pub input_type: Option<String>,
127 pub truncation: Option<bool>,
130 pub output_dimension: Option<usize>,
133}
134
135impl Embeddings {
136 pub fn with_input_type(mut self, input_type: impl Into<String>) -> Self {
138 self.input_type = Some(input_type.into());
139 self
140 }
141
142 pub fn with_truncation(mut self, truncation: bool) -> Self {
144 self.truncation = Some(truncation);
145 self
146 }
147
148 pub fn with_output_dimension(mut self, output_dimension: usize) -> Self {
153 self.output_dimension = Some(output_dimension);
154 self.ndims = output_dimension;
155 self
156 }
157}
158
159impl Wire for Embeddings {
160 type Op = Embedding;
161 type Payload = crate::wire::Encoded;
162 type Frame = crate::wire::WireFrame;
163 type Decoder<'id> = EmbeddingsDecoder;
164 type Reassembler = crate::wire::document::Unreassembled;
165
166 fn describe(&self) -> Descriptor<'_> {
167 Descriptor::new(PROVIDER_NAME)
168 .model(self.model.as_str())
169 .capabilities(Capabilities::embedding(MAX_DOCUMENTS, self.ndims))
170 }
171
172 fn encode(&self, texts: Vec<String>, _mode: Mode) -> Result<Encoded, EncodeError> {
173 let mut body = serde_json::Map::new();
174 body.insert("model".to_owned(), serde_json::json!(self.model));
175 body.insert("input".to_owned(), serde_json::json!(texts));
176 if let Some(input_type) = &self.input_type {
177 body.insert("input_type".to_owned(), serde_json::json!(input_type));
178 }
179 if let Some(truncation) = self.truncation {
180 body.insert("truncation".to_owned(), serde_json::json!(truncation));
181 }
182 if let Some(output_dimension) = self.output_dimension {
183 body.insert(
184 "output_dimension".to_owned(),
185 serde_json::json!(output_dimension),
186 );
187 }
188 let request = self
189 .provider
190 .post("/embeddings")
191 .body(Body::Bytes(serde_json::to_vec(&body)?))?;
192 Ok(Encoded::new(request, Framing::Whole))
193 }
194
195 fn decoder<'id>(&self) -> Self::Decoder<'id> {
196 EmbeddingsDecoder
197 }
198}
199
200pub struct EmbeddingsDecoder;
202
203impl<'id> Decoder<'id, Embedding> for EmbeddingsDecoder {
204 type Event = VoyageEmbeddingResponse;
205
206 fn classify(&self, frame: WireFrame) -> WireEvent<Self::Event> {
207 crate::providers::internal::wire::classify_marker_keyed_frame(&frame.as_str(), &["data"])
208 }
209
210 fn decode(
211 &mut self,
212 reply: Self::Event,
213 out: Out<'id, Embedding>,
214 ) -> Result<Flow, ProviderError> {
215 let usage = crate::completion::Usage {
217 input_tokens: Some(reply.usage.total_tokens as u64),
218 total_tokens: Some(reply.usage.total_tokens as u64),
219 ..Default::default()
220 };
221 let vectors = reply.data.into_iter().map(|embedding| embedding.embedding);
222 Ok(out.end(crate::embeddings::EmbeddingResponse {
223 model: Some(reply.model),
224 usage,
225 ..crate::embeddings::EmbeddingResponse::from_vectors(vectors)
226 }))
227 }
228}
229
230#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
232pub struct Rerank {
233 pub provider: VoyageAiConfig,
235 pub model: String,
237 pub top_k: Option<usize>,
240 pub return_documents: bool,
242 pub truncation: Option<bool>,
245}
246
247impl Rerank {
248 pub fn with_top_k(mut self, top_k: usize) -> Self {
250 self.top_k = Some(top_k);
251 self
252 }
253
254 pub fn with_return_documents(mut self, return_documents: bool) -> Self {
256 self.return_documents = return_documents;
257 self
258 }
259
260 pub fn with_truncation(mut self, truncation: bool) -> Self {
262 self.truncation = Some(truncation);
263 self
264 }
265}
266
267impl Wire for Rerank {
268 type Op = RerankOp;
269 type Payload = crate::wire::Encoded;
270 type Frame = crate::wire::WireFrame;
271 type Decoder<'id> = RerankDecoder;
272 type Reassembler = crate::wire::document::Unreassembled;
273
274 fn describe(&self) -> Descriptor<'_> {
275 Descriptor::new(PROVIDER_NAME)
276 .model(self.model.as_str())
277 .capabilities(Capabilities::rerank(MAX_RERANK_DOCUMENTS))
278 }
279
280 fn encode(&self, request: RerankRequest, _mode: Mode) -> Result<Encoded, EncodeError> {
281 let mut body = serde_json::Map::new();
282 body.insert("query".to_owned(), serde_json::json!(request.query));
283 body.insert("documents".to_owned(), serde_json::json!(request.documents));
284 body.insert("model".to_owned(), serde_json::json!(self.model));
285 if let Some(top_k) = self.top_k {
286 body.insert("top_k".to_owned(), serde_json::json!(top_k));
287 }
288 body.insert(
289 "return_documents".to_owned(),
290 serde_json::json!(self.return_documents),
291 );
292 if let Some(truncation) = self.truncation {
293 body.insert("truncation".to_owned(), serde_json::json!(truncation));
294 }
295 let request = self
296 .provider
297 .post("/rerank")
298 .body(Body::Bytes(serde_json::to_vec(&body)?))?;
299 Ok(Encoded::new(request, Framing::Whole))
300 }
301
302 fn decoder<'id>(&self) -> Self::Decoder<'id> {
303 RerankDecoder
304 }
305}
306
307pub struct RerankDecoder;
309
310impl<'id> Decoder<'id, RerankOp> for RerankDecoder {
311 type Event = Result<RerankApiResponse, String>;
313
314 fn classify(&self, frame: WireFrame) -> WireEvent<Self::Event> {
315 crate::providers::internal::wire::classify_reply_or_message_envelope(
316 &frame.as_str(),
317 "data",
318 )
319 }
320
321 fn decode(
322 &mut self,
323 reply: Self::Event,
324 out: Out<'id, RerankOp>,
325 ) -> Result<Flow, ProviderError> {
326 let reply = reply.map_err(ProviderError::from_provider_body)?;
327 let usage = crate::completion::Usage {
329 input_tokens: Some(reply.usage.total_tokens as u64),
330 total_tokens: Some(reply.usage.total_tokens as u64),
331 ..Default::default()
332 };
333 let results = reply
334 .data
335 .into_iter()
336 .map(|result| RerankResult {
337 index: result.index,
338 document: result.document,
339 relevance_score: result.relevance_score,
340 })
341 .collect();
342 Ok(out.end(RerankResponse {
343 model: Some(reply.model),
344 usage,
345 ..RerankResponse::new(results)
346 }))
347 }
348}
349
350#[cfg(test)]
351mod tests;