Skip to main content

rig_core/providers/voyageai/
wire.rs

1//! Voyage AI configuration and embedding and reranking endpoint wires.
2//!
3//! ```no_run
4//! use rig_core::providers::voyageai::{VoyageAi, VOYAGE_3_5};
5//! let wire = VoyageAi::from_env()?.embedding(VOYAGE_3_5, None);
6//! # Ok::<(), Box<dyn std::error::Error>>(())
7//! ```
8
9use 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
26/// The provider descriptor name, as records and telemetry spell it.
27const PROVIDER_NAME: &str = "voyageai";
28
29/// The environment variable carrying the API key.
30const API_KEY_ENV: &str = "VOYAGE_API_KEY";
31
32/// The most texts `POST /embeddings` accepts in one call.
33const MAX_DOCUMENTS: usize = 1024;
34
35/// The most documents `POST /rerank` orders in one call.
36const MAX_RERANK_DOCUMENTS: usize = 1000;
37
38/// The settings of a Voyage AI provider: serializable, and the credential
39/// is never serialized. [`connect`](Self::connect) puts it on a transport as
40/// a [`VoyageAi`](super::VoyageAi) client.
41#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
42pub struct VoyageAiConfig {
43    /// The API key, sent as `Authorization: Bearer`.
44    pub api_key: Secret,
45    /// The API root, without a trailing slash.
46    pub base_url: String,
47}
48
49impl VoyageAiConfig {
50    /// Voyage AI with default settings.
51    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    /// Voyage AI from `VOYAGE_API_KEY`.
59    pub fn from_env() -> Result<Self, EnvError> {
60        Ok(Self::new(env::required(API_KEY_ENV)?))
61    }
62
63    /// Point the wires at another API root.
64    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    /// Build an embedding wire reporting `ndims`, or the known model width,
70    /// or zero if unknown. This does not send an output-dimension override;
71    /// use [`Embeddings::with_output_dimension`] for that.
72    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    /// The rerank wire for `model`.
88    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    /// One request, authenticated and typed as JSON.
99    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/// The embedding wire: `POST /embeddings`.
110///
111/// Every option defaults to `None`, which is Voyage's own server default:
112/// no retrieval prompt, truncation enabled, and the model's default output
113/// dimension.
114#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
115pub struct Embeddings {
116    /// The provider this wire speaks to.
117    pub provider: VoyageAiConfig,
118    /// The model to address.
119    pub model: String,
120    /// The width this wire reports, from the caller or the model's published
121    /// dimensions. `0` means neither named one.
122    pub ndims: usize,
123    /// Prepends a retrieval prompt to the input text: `"document"` when
124    /// embedding stored chunks, `"query"` when embedding search queries.
125    /// Embeddings produced with and without it are compatible.
126    pub input_type: Option<String>,
127    /// Whether to truncate inputs longer than the model's context. Voyage's
128    /// server default is `true`.
129    pub truncation: Option<bool>,
130    /// Dimensionality of the returned vectors, when overriding the model's
131    /// default output dimension.
132    pub output_dimension: Option<usize>,
133}
134
135impl Embeddings {
136    /// Prepend Voyage's retrieval prompt for this input's role.
137    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    /// Set whether overlong inputs are truncated (`true`) or rejected (`false`).
143    pub fn with_truncation(mut self, truncation: bool) -> Self {
144        self.truncation = Some(truncation);
145        self
146    }
147
148    /// Ask Voyage for vectors of this width.
149    ///
150    /// This is the width the request carries; [`Self::ndims`] is the width
151    /// the wire reports to a vector store, so they are set together.
152    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
200/// Decodes one `/embeddings` reply.
201pub 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        // Voyage reports one count; every token of an embedding is input.
216        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/// The rerank wire: `POST /rerank`.
231#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
232pub struct Rerank {
233    /// The provider this wire speaks to.
234    pub provider: VoyageAiConfig,
235    /// The model to address.
236    pub model: String,
237    /// Return only the `top_k` most relevant documents. `None` returns them
238    /// all.
239    pub top_k: Option<usize>,
240    /// Whether the reply echoes each document's text back.
241    pub return_documents: bool,
242    /// Whether to truncate documents longer than the model's context.
243    /// Voyage's server default is `true`.
244    pub truncation: Option<bool>,
245}
246
247impl Rerank {
248    /// Order only the `top_k` most relevant documents.
249    pub fn with_top_k(mut self, top_k: usize) -> Self {
250        self.top_k = Some(top_k);
251        self
252    }
253
254    /// Ask Voyage to echo each ordered document's text back.
255    pub fn with_return_documents(mut self, return_documents: bool) -> Self {
256        self.return_documents = return_documents;
257        self
258    }
259
260    /// Set whether overlong documents are truncated (`true`) or rejected (`false`).
261    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
307/// Decodes one `/rerank` reply.
308pub struct RerankDecoder;
309
310impl<'id> Decoder<'id, RerankOp> for RerankDecoder {
311    /// The ordering, or the error envelope Voyage can answer a **200** with.
312    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        // Voyage reports one count; every token of a rerank is input.
328        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;