use crate::client::env::{self, EnvError};
use crate::error::EncodeError;
use crate::error::ProviderError;
use crate::operation::{Embedding, Rerank as RerankOp, RerankRequest};
use crate::rerank::{RerankResponse, RerankResult};
use crate::wire::Flow;
use crate::wire::{
Body, Capabilities, Decoder, Descriptor, Encoded, Framing, Mode, Out, Secret, Wire, WireEvent,
WireFrame,
};
use serde::{Deserialize, Serialize};
use super::{
EmbeddingResponse as VoyageEmbeddingResponse, RerankApiResponse, VOYAGEAI_API_BASE_URL,
model_dimensions_from_identifier,
};
const PROVIDER_NAME: &str = "voyageai";
const API_KEY_ENV: &str = "VOYAGE_API_KEY";
const MAX_DOCUMENTS: usize = 1024;
const MAX_RERANK_DOCUMENTS: usize = 1000;
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct VoyageAiConfig {
pub api_key: Secret,
pub base_url: String,
}
impl VoyageAiConfig {
pub fn new(api_key: impl Into<Secret>) -> Self {
Self {
api_key: api_key.into(),
base_url: VOYAGEAI_API_BASE_URL.to_owned(),
}
}
pub fn from_env() -> Result<Self, EnvError> {
Ok(Self::new(env::required(API_KEY_ENV)?))
}
pub fn with_base_url(mut self, base_url: impl AsRef<str>) -> Self {
self.base_url = base_url.as_ref().trim_end_matches('/').to_owned();
self
}
pub(crate) fn embedding(&self, model: impl Into<String>, ndims: Option<usize>) -> Embeddings {
let model = model.into();
let ndims = ndims
.or_else(|| model_dimensions_from_identifier(&model))
.unwrap_or_default();
Embeddings {
provider: self.clone(),
model,
ndims,
input_type: None,
truncation: None,
output_dimension: None,
}
}
pub(crate) fn rerank(&self, model: impl Into<String>) -> Rerank {
Rerank {
provider: self.clone(),
model: model.into(),
top_k: None,
return_documents: false,
truncation: None,
}
}
fn post(&self, path: &str) -> http::request::Builder {
http::Request::post(format!("{}{path}", self.base_url))
.header(http::header::CONTENT_TYPE, "application/json")
.header(
http::header::AUTHORIZATION,
format!("Bearer {}", self.api_key.expose()),
)
}
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct Embeddings {
pub provider: VoyageAiConfig,
pub model: String,
pub ndims: usize,
pub input_type: Option<String>,
pub truncation: Option<bool>,
pub output_dimension: Option<usize>,
}
impl Embeddings {
pub fn with_input_type(mut self, input_type: impl Into<String>) -> Self {
self.input_type = Some(input_type.into());
self
}
pub fn with_truncation(mut self, truncation: bool) -> Self {
self.truncation = Some(truncation);
self
}
pub fn with_output_dimension(mut self, output_dimension: usize) -> Self {
self.output_dimension = Some(output_dimension);
self.ndims = output_dimension;
self
}
}
impl Wire for Embeddings {
type Op = Embedding;
type Payload = crate::wire::Encoded;
type Frame = crate::wire::WireFrame;
type Decoder<'id> = EmbeddingsDecoder;
type Reassembler = crate::wire::document::Unreassembled;
fn describe(&self) -> Descriptor<'_> {
Descriptor::new(PROVIDER_NAME)
.model(self.model.as_str())
.capabilities(Capabilities::embedding(MAX_DOCUMENTS, self.ndims))
}
fn encode(&self, texts: Vec<String>, _mode: Mode) -> Result<Encoded, EncodeError> {
let mut body = serde_json::Map::new();
body.insert("model".to_owned(), serde_json::json!(self.model));
body.insert("input".to_owned(), serde_json::json!(texts));
if let Some(input_type) = &self.input_type {
body.insert("input_type".to_owned(), serde_json::json!(input_type));
}
if let Some(truncation) = self.truncation {
body.insert("truncation".to_owned(), serde_json::json!(truncation));
}
if let Some(output_dimension) = self.output_dimension {
body.insert(
"output_dimension".to_owned(),
serde_json::json!(output_dimension),
);
}
let request = self
.provider
.post("/embeddings")
.body(Body::Bytes(serde_json::to_vec(&body)?))?;
Ok(Encoded::new(request, Framing::Whole))
}
fn decoder<'id>(&self) -> Self::Decoder<'id> {
EmbeddingsDecoder
}
}
pub struct EmbeddingsDecoder;
impl<'id> Decoder<'id, Embedding> for EmbeddingsDecoder {
type Event = VoyageEmbeddingResponse;
fn classify(&self, frame: WireFrame) -> WireEvent<Self::Event> {
crate::providers::internal::wire::classify_marker_keyed_frame(&frame.as_str(), &["data"])
}
fn decode(
&mut self,
reply: Self::Event,
out: Out<'id, Embedding>,
) -> Result<Flow, ProviderError> {
let usage = crate::completion::Usage {
input_tokens: Some(reply.usage.total_tokens as u64),
total_tokens: Some(reply.usage.total_tokens as u64),
..Default::default()
};
let vectors = reply.data.into_iter().map(|embedding| embedding.embedding);
Ok(out.end(crate::embeddings::EmbeddingResponse {
model: Some(reply.model),
usage,
..crate::embeddings::EmbeddingResponse::from_vectors(vectors)
}))
}
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct Rerank {
pub provider: VoyageAiConfig,
pub model: String,
pub top_k: Option<usize>,
pub return_documents: bool,
pub truncation: Option<bool>,
}
impl Rerank {
pub fn with_top_k(mut self, top_k: usize) -> Self {
self.top_k = Some(top_k);
self
}
pub fn with_return_documents(mut self, return_documents: bool) -> Self {
self.return_documents = return_documents;
self
}
pub fn with_truncation(mut self, truncation: bool) -> Self {
self.truncation = Some(truncation);
self
}
}
impl Wire for Rerank {
type Op = RerankOp;
type Payload = crate::wire::Encoded;
type Frame = crate::wire::WireFrame;
type Decoder<'id> = RerankDecoder;
type Reassembler = crate::wire::document::Unreassembled;
fn describe(&self) -> Descriptor<'_> {
Descriptor::new(PROVIDER_NAME)
.model(self.model.as_str())
.capabilities(Capabilities::rerank(MAX_RERANK_DOCUMENTS))
}
fn encode(&self, request: RerankRequest, _mode: Mode) -> Result<Encoded, EncodeError> {
let mut body = serde_json::Map::new();
body.insert("query".to_owned(), serde_json::json!(request.query));
body.insert("documents".to_owned(), serde_json::json!(request.documents));
body.insert("model".to_owned(), serde_json::json!(self.model));
if let Some(top_k) = self.top_k {
body.insert("top_k".to_owned(), serde_json::json!(top_k));
}
body.insert(
"return_documents".to_owned(),
serde_json::json!(self.return_documents),
);
if let Some(truncation) = self.truncation {
body.insert("truncation".to_owned(), serde_json::json!(truncation));
}
let request = self
.provider
.post("/rerank")
.body(Body::Bytes(serde_json::to_vec(&body)?))?;
Ok(Encoded::new(request, Framing::Whole))
}
fn decoder<'id>(&self) -> Self::Decoder<'id> {
RerankDecoder
}
}
pub struct RerankDecoder;
impl<'id> Decoder<'id, RerankOp> for RerankDecoder {
type Event = Result<RerankApiResponse, String>;
fn classify(&self, frame: WireFrame) -> WireEvent<Self::Event> {
crate::providers::internal::wire::classify_reply_or_message_envelope(
&frame.as_str(),
"data",
)
}
fn decode(
&mut self,
reply: Self::Event,
out: Out<'id, RerankOp>,
) -> Result<Flow, ProviderError> {
let reply = reply.map_err(ProviderError::from_provider_body)?;
let usage = crate::completion::Usage {
input_tokens: Some(reply.usage.total_tokens as u64),
total_tokens: Some(reply.usage.total_tokens as u64),
..Default::default()
};
let results = reply
.data
.into_iter()
.map(|result| RerankResult {
index: result.index,
document: result.document,
relevance_score: result.relevance_score,
})
.collect();
Ok(out.end(RerankResponse {
model: Some(reply.model),
usage,
..RerankResponse::new(results)
}))
}
}
#[cfg(test)]
mod tests;