diff options
| author | sigoden <sigoden@gmail.com> | 2024-06-25 07:39:35 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2024-06-25 07:39:35 +0800 |
| commit | 2fbb5271af27eee7ba3ae6d2b3c1aa0818648c8f (patch) | |
| tree | 5e72e72d592025c00a46cebd0eb7012ccd765702 | |
| parent | ed71901611247d41daed8112a5106b42eb12395b (diff) | |
| download | aichat-2fbb5271af27eee7ba3ae6d2b3c1aa0818648c8f.tar.gz | |
feat: support rag-dedicated clients (jina and voyageai) (#645)
| -rw-r--r-- | config.example.yaml | 12 | ||||
| -rw-r--r-- | models.yaml | 109 | ||||
| -rw-r--r-- | src/client/cohere.rs | 36 | ||||
| -rw-r--r-- | src/client/common.rs | 10 | ||||
| -rw-r--r-- | src/client/ernie.rs | 9 | ||||
| -rw-r--r-- | src/client/mod.rs | 11 | ||||
| -rw-r--r-- | src/client/openai_compatible.rs | 6 | ||||
| -rw-r--r-- | src/client/rag_dedicated.rs | 160 | ||||
| -rw-r--r-- | src/rag/mod.rs | 1 |
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 |
