summaryrefslogtreecommitdiffstats
diff options
context:
space:
mode:
-rw-r--r--config.example.yaml12
-rw-r--r--models.yaml109
-rw-r--r--src/client/cohere.rs36
-rw-r--r--src/client/common.rs10
-rw-r--r--src/client/ernie.rs9
-rw-r--r--src/client/mod.rs11
-rw-r--r--src/client/openai_compatible.rs6
-rw-r--r--src/client/rag_dedicated.rs160
-rw-r--r--src/rag/mod.rs1
9 files changed, 307 insertions, 47 deletions
diff --git a/config.example.yaml b/config.example.yaml
index 0c3d985..7d16730 100644
--- a/config.example.yaml
+++ b/config.example.yaml
@@ -286,3 +286,15 @@ clients:
name: together
api_base: https://api.together.xyz/v1
api_key: xxx # ENV: {client}_API_KEY
+
+ # See https://jina.ai
+ - type: rag-dedicated
+ name: jina
+ api_base: https://api.jina.ai/v1
+ api_key: xxx # ENV: {client}_API_KEY
+
+ # See https://docs.voyageai.com/docs/introduction
+ - type: rag-dedicated
+ name: voyageai
+ api_base: https://api.voyageai.ai/v1
+ api_key: xxx # ENV: {client}_API_KEY \ No newline at end of file
diff --git a/models.yaml b/models.yaml
index 51cbb76..0b24d78 100644
--- a/models.yaml
+++ b/models.yaml
@@ -1145,4 +1145,111 @@
input_price: 0.008
output_vector_size: 768
default_chunk_size: 1000
- max_batch_size: 100 \ No newline at end of file
+ max_batch_size: 100
+
+- platform: jina
+ # docs:
+ # - https://jina.ai/
+ # - https://api.jina.ai/redoc
+ models:
+ - name: jina-embeddings-v2-base-en
+ type: embedding
+ max_input_tokens: 8192
+ input_price: 0.02
+ output_vector_size: 768
+ default_chunk_size: 1500
+ max_batch_size: 100
+ - name: jina-embeddings-v2-small-en
+ type: embedding
+ max_input_tokens: 8192
+ input_price: 0.02
+ output_vector_size: 512
+ default_chunk_size: 1000
+ max_batch_size: 100
+ - name: jina-embeddings-v2-base-zsh
+ type: embedding
+ max_input_tokens: 8192
+ input_price: 0.02
+ output_vector_size: 768
+ default_chunk_size: 1500
+ max_batch_size: 100
+ - name: jina-embeddings-v2-base-code
+ type: embedding
+ max_input_tokens: 8192
+ input_price: 0.02
+ output_vector_size: 768
+ default_chunk_size: 1500
+ max_batch_size: 100
+ - name: jina-colbert-v1-en
+ type: embedding
+ max_input_tokens: 8192
+ input_price: 0.02
+ output_vector_size: 768
+ default_chunk_size: 1500
+ max_batch_size: 100
+ - name: jina-reranker-v1-base-en
+ type: rerank
+ max_input_tokens: 8192
+ input_price: 0.02
+ - name: jina-reranker-v1-turbo-en
+ type: rerank
+ max_input_tokens: 8192
+ input_price: 0.02
+ - name: jina-colbert-v1-en
+ type: rerank
+ max_input_tokens: 8192
+ input_price: 0.02
+ - name: jina-reranker-v1-base-multilingual
+ type: rerank
+ max_input_tokens: 8192
+ input_price: 0.02
+
+- platform: voyageai
+ # docs:
+ # - https://docs.voyageai.com/docs/embeddings
+ # - https://docs.voyageai.com/docs/pricing
+ # - https://docs.voyageai.com/reference/embeddings-api
+ models:
+ - name: voyage-large-2-instruct
+ type: embedding
+ max_input_tokens: 16000
+ input_price: 0.12
+ output_vector_size: 1024
+ default_chunk_size: 2000
+ max_batch_size: 128
+ - name: voyage-large-2
+ type: embedding
+ max_input_tokens: 16000
+ input_price: 0.12
+ output_vector_size: 1536
+ default_chunk_size: 3000
+ max_batch_size: 128
+ - name: voyage-multilingual-2
+ type: embedding
+ max_input_tokens: 32000
+ input_price: 0.12
+ output_vector_size: 1024
+ default_chunk_size: 2000
+ max_batch_size: 128
+ - name: voyage-code-2
+ type: embedding
+ max_input_tokens: 16000
+ input_price: 0.12
+ output_vector_size: 1536
+ default_chunk_size: 3000
+ max_batch_size: 128
+ - name: voyage-2
+ type: embedding
+ max_input_tokens: 4000
+ input_price: 0.1
+ output_vector_size: 1024
+ default_chunk_size: 2000
+ max_batch_size: 128
+ - name: rerank-1
+ type: rerank
+ max_input_tokens: 8000
+ input_price: 0.05
+ - name: rerank-lite-1
+ type: rerank
+ max_input_tokens: 4000
+ input_price: 0.02
diff --git a/src/client/cohere.rs b/src/client/cohere.rs
index 3e1cb2b..b2857e7 100644
--- a/src/client/cohere.rs
+++ b/src/client/cohere.rs
@@ -1,3 +1,4 @@
+use super::rag_dedicated::*;
use super::*;
use anyhow::{bail, Context, Result};
@@ -74,7 +75,7 @@ impl CohereClient {
fn rerank_builder(&self, client: &ReqwestClient, data: RerankData) -> Result<RequestBuilder> {
let api_key = self.get_api_key()?;
- let body = cohere_build_rerank_body(data, &self.model);
+ let body = rag_dedicated_build_rerank_body(data, &self.model);
let url = RERANK_API_URL;
@@ -91,7 +92,7 @@ impl_client_trait!(
chat_completions,
chat_completions_streaming,
embeddings,
- cohere_rerank
+ rag_dedicated_rerank
);
async fn chat_completions(builder: RequestBuilder) -> Result<ChatCompletionsOutput> {
@@ -162,22 +163,6 @@ struct EmbeddingsResBody {
embeddings: Vec<Vec<f32>>,
}
-pub async fn cohere_rerank(builder: RequestBuilder) -> Result<RerankOutput> {
- let res = builder.send().await?;
- let status = res.status();
- let data: Value = res.json().await?;
- if !status.is_success() {
- catch_error(&data, status.as_u16())?;
- }
- let res_body: RerankResBody = serde_json::from_value(data).context("Invalid rerank data")?;
- Ok(res_body.results)
-}
-
-#[derive(Deserialize)]
-struct RerankResBody {
- results: RerankOutput,
-}
-
fn build_chat_completions_body(data: ChatCompletionsData, model: &Model) -> Result<Value> {
let ChatCompletionsData {
mut messages,
@@ -309,21 +294,6 @@ fn build_chat_completions_body(data: ChatCompletionsData, model: &Model) -> Resu
Ok(body)
}
-pub fn cohere_build_rerank_body(data: RerankData, model: &Model) -> Value {
- let RerankData {
- query,
- documents,
- top_n,
- } = data;
-
- json!({
- "model": model.name(),
- "query": query,
- "documents": documents,
- "top_n": top_n
- })
-}
-
fn extract_chat_completions(data: &Value) -> Result<ChatCompletionsOutput> {
let text = data["text"].as_str().unwrap_or_default();
diff --git a/src/client/common.rs b/src/client/common.rs
index 1b1e723..bf2f336 100644
--- a/src/client/common.rs
+++ b/src/client/common.rs
@@ -83,7 +83,9 @@ macro_rules! register_client {
let client_name = Self::name(local_config);
if local_config.models.is_empty() {
if let Some(models) = $crate::client::ALL_MODELS.iter().find(|v| {
- v.platform == $name || ($name == "openai-compatible" && local_config.name.as_deref() == Some(&v.platform))
+ v.platform == $name ||
+ ($name == OpenAICompatibleClient::NAME && local_config.name.as_deref() == Some(&v.platform)) ||
+ ($name == RagDedicatedClient::NAME && local_config.name.as_deref() == Some(&v.platform))
}) {
return Model::from_config(client_name, &models.models);
}
@@ -432,7 +434,7 @@ pub trait Client: Sync + Send {
_client: &ReqwestClient,
_data: EmbeddingsData,
) -> Result<EmbeddingsOutput> {
- bail!("No embeddings api")
+ bail!("The client doesn't support embeddings api")
}
async fn rerank_inner(
@@ -440,7 +442,7 @@ pub trait Client: Sync + Send {
_client: &ReqwestClient,
_data: RerankData,
) -> Result<RerankOutput> {
- bail!("No rerank api")
+ bail!("The client doesn't support rerank api")
}
}
@@ -566,7 +568,7 @@ pub fn create_openai_compatible_client_config(client: &str) -> Result<Option<(St
None => Ok(None),
Some((name, api_base)) => {
let mut config = json!({
- "type": "openai-compatible",
+ "type": OpenAICompatibleClient::NAME,
"name": name,
"api_base": api_base,
});
diff --git a/src/client/ernie.rs b/src/client/ernie.rs
index 428d263..64c1820 100644
--- a/src/client/ernie.rs
+++ b/src/client/ernie.rs
@@ -1,4 +1,5 @@
use super::access_token::*;
+use super::rag_dedicated::*;
use super::*;
use anyhow::{anyhow, bail, Context, Result};
@@ -220,15 +221,11 @@ struct EmbeddingsResBodyEmbedding {
async fn rerank(builder: RequestBuilder) -> Result<RerankOutput> {
let data: Value = builder.send().await?.json().await?;
maybe_catch_error(&data)?;
- let res_body: RerankResBody = serde_json::from_value(data).context("Invalid rerank data")?;
+ let res_body: RagDedicatedRerankResBody =
+ serde_json::from_value(data).context("Invalid rerank data")?;
Ok(res_body.results)
}
-#[derive(Deserialize)]
-struct RerankResBody {
- results: RerankOutput,
-}
-
fn build_chat_completions_body(data: ChatCompletionsData, model: &Model) -> Value {
let ChatCompletionsData {
mut messages,
diff --git a/src/client/mod.rs b/src/client/mod.rs
index 579c2a3..6878541 100644
--- a/src/client/mod.rs
+++ b/src/client/mod.rs
@@ -21,6 +21,12 @@ register_client!(
OpenAICompatibleConfig,
OpenAICompatibleClient
),
+ (
+ rag_dedicated,
+ "rag-dedicated",
+ RagDedicatedConfig,
+ RagDedicatedClient
+ ),
(gemini, "gemini", GeminiConfig, GeminiClient),
(claude, "claude", ClaudeConfig, ClaudeClient),
(cohere, "cohere", CohereConfig, CohereClient),
@@ -60,3 +66,8 @@ pub const OPENAI_COMPATIBLE_PLATFORMS: [(&str, &str); 13] = [
("zhipuai", "https://open.bigmodel.cn/api/paas/v4"),
("lingyiwanwu", "https://api.lingyiwanwu.com/v1"),
];
+
+pub const RAG_DEDICATED_PLATFORMS: [(&str, &str); 2] = [
+ ("jina", "https://api.jina.ai/v1"),
+ ("voyageai", "https://api.voyageai.com/v1"),
+];
diff --git a/src/client/openai_compatible.rs b/src/client/openai_compatible.rs
index 789604f..132b40a 100644
--- a/src/client/openai_compatible.rs
+++ b/src/client/openai_compatible.rs
@@ -1,5 +1,5 @@
-use super::cohere::*;
use super::openai::*;
+use super::rag_dedicated::*;
use super::*;
use anyhow::Result;
@@ -90,7 +90,7 @@ impl OpenAICompatibleClient {
let api_key = self.get_api_key().ok();
let api_base = self.get_api_base_ext()?;
- let body = cohere_build_rerank_body(data, &self.model);
+ let body = rag_dedicated_build_rerank_body(data, &self.model);
let url = format!("{api_base}/rerank");
@@ -131,5 +131,5 @@ impl_client_trait!(
openai_chat_completions,
openai_chat_completions_streaming,
openai_embeddings,
- cohere_rerank
+ rag_dedicated_rerank
);
diff --git a/src/client/rag_dedicated.rs b/src/client/rag_dedicated.rs
new file mode 100644
index 0000000..a9a9f18
--- /dev/null
+++ b/src/client/rag_dedicated.rs
@@ -0,0 +1,160 @@
+use super::openai::*;
+use super::*;
+
+use anyhow::bail;
+use anyhow::Context;
+use anyhow::Result;
+use reqwest::{Client as ReqwestClient, RequestBuilder};
+use serde::Deserialize;
+use serde_json::json;
+use serde_json::Value;
+
+#[derive(Debug, Clone, Deserialize)]
+pub struct RagDedicatedConfig {
+ pub name: Option<String>,
+ pub api_base: Option<String>,
+ pub api_key: Option<String>,
+ #[serde(default)]
+ pub models: Vec<ModelData>,
+ pub patches: Option<ModelPatches>,
+ pub extra: Option<ExtraConfig>,
+}
+
+impl RagDedicatedClient {
+ config_get_fn!(api_base, get_api_base);
+ config_get_fn!(api_key, get_api_key);
+
+ pub const PROMPTS: [PromptAction<'static>; 0] = [];
+
+ fn chat_completions_builder(
+ &self,
+ _client: &ReqwestClient,
+ _data: ChatCompletionsData,
+ ) -> Result<RequestBuilder> {
+ bail!("The client doesn't support chat-completions api");
+ }
+
+ fn embeddings_builder(
+ &self,
+ client: &ReqwestClient,
+ data: EmbeddingsData,
+ ) -> Result<RequestBuilder> {
+ let api_key = self.get_api_key().ok();
+ let api_base = self.get_api_base_ext()?;
+
+ let body = openai_build_embeddings_body(data, &self.model);
+
+ let url = format!("{api_base}/embeddings");
+
+ debug!("RagDedicated Embeddings Request: {url} {body}");
+
+ let mut builder = client.post(url).json(&body);
+ if let Some(api_key) = api_key {
+ builder = builder.bearer_auth(api_key);
+ }
+
+ Ok(builder)
+ }
+
+ fn rerank_builder(&self, client: &ReqwestClient, data: RerankData) -> Result<RequestBuilder> {
+ let api_key = self.get_api_key().ok();
+ let api_base = self.get_api_base_ext()?;
+
+ let body = rag_dedicated_build_rerank_body(data, &self.model);
+
+ let url = format!("{api_base}/rerank");
+
+ debug!("RagDedicated Rerank Request: {url} {body}");
+
+ let mut builder = client.post(url).json(&body);
+ if let Some(api_key) = api_key {
+ builder = builder.bearer_auth(api_key);
+ }
+
+ Ok(builder)
+ }
+
+ fn get_api_base_ext(&self) -> Result<String> {
+ let api_base = match self.get_api_base() {
+ Ok(v) => v,
+ Err(err) => {
+ match RAG_DEDICATED_PLATFORMS
+ .into_iter()
+ .find_map(|(name, api_base)| {
+ if name == self.model.client_name() {
+ Some(api_base.to_string())
+ } else {
+ None
+ }
+ }) {
+ Some(v) => v,
+ None => return Err(err),
+ }
+ }
+ };
+ Ok(api_base)
+ }
+}
+
+impl_client_trait!(
+ RagDedicatedClient,
+ no_chat_completions,
+ no_chat_completions_streaming,
+ openai_embeddings,
+ rag_dedicated_rerank
+);
+
+pub async fn no_chat_completions(_builder: RequestBuilder) -> Result<ChatCompletionsOutput> {
+ bail!("The client doesn't support chat-completions api");
+}
+
+pub async fn no_chat_completions_streaming(
+ _builder: RequestBuilder,
+ _handler: &mut SseHandler,
+) -> Result<()> {
+ bail!("The client doesn't support chat-completions api")
+}
+
+pub async fn rag_dedicated_rerank(builder: RequestBuilder) -> Result<RerankOutput> {
+ let res = builder.send().await?;
+ let status = res.status();
+ let mut data: Value = res.json().await?;
+ if !status.is_success() {
+ catch_error(&data, status.as_u16())?;
+ }
+ if data.get("results").is_none() && data.get("data").is_some() {
+ if let Some(data_obj) = data.as_object_mut() {
+ if let Some(value) = data_obj.remove("data") {
+ data_obj.insert("results".to_string(), value);
+ }
+ }
+ }
+ let res_body: RagDedicatedRerankResBody =
+ serde_json::from_value(data).context("Invalid rerank data")?;
+ Ok(res_body.results)
+}
+
+#[derive(Deserialize)]
+pub struct RagDedicatedRerankResBody {
+ pub results: RerankOutput,
+}
+
+pub fn rag_dedicated_build_rerank_body(data: RerankData, model: &Model) -> Value {
+ let RerankData {
+ query,
+ documents,
+ top_n,
+ } = data;
+
+ let mut body = json!({
+ "model": model.name(),
+ "query": query,
+ "documents": documents,
+ });
+ if model.client_name() == "voyageai" {
+ body["top_k"] = top_n.into()
+ } else {
+ body["top_n"] = top_n.into()
+ }
+ body
+}
diff --git a/src/rag/mod.rs b/src/rag/mod.rs
index 30ca92f..3fdad02 100644
--- a/src/rag/mod.rs
+++ b/src/rag/mod.rs
@@ -334,6 +334,7 @@ impl Rag {
let list = client.rerank(data).await?;
let ids = list
.into_iter()
+ .take(top_k)
.filter_map(|item| {
if item.relevance_score < min_score {
None