From 0e740d81e94505bd57036755abaaecb12c3b26e3 Mon Sep 17 00:00:00 2001 From: sigoden Date: Sun, 28 Jul 2024 06:04:36 +0800 Subject: feat: abandon rag_dedicated client and improve (#757) --- src/client/rag_dedicated.rs | 150 -------------------------------------------- 1 file changed, 150 deletions(-) delete mode 100644 src/client/rag_dedicated.rs (limited to 'src/client/rag_dedicated.rs') diff --git a/src/client/rag_dedicated.rs b/src/client/rag_dedicated.rs deleted file mode 100644 index 7d2b846..0000000 --- a/src/client/rag_dedicated.rs +++ /dev/null @@ -1,150 +0,0 @@ -use super::openai::*; -use super::*; - -use anyhow::bail; -use anyhow::Context; -use anyhow::Result; -use reqwest::RequestBuilder; -use serde::Deserialize; -use serde_json::json; -use serde_json::Value; - -#[derive(Debug, Clone, Deserialize)] -pub struct RagDedicatedConfig { - pub name: Option, - pub api_base: Option, - pub api_key: Option, - #[serde(default)] - pub models: Vec, - pub patch: Option, - pub extra: Option, -} - -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 prepare_chat_completions(&self, _data: ChatCompletionsData) -> Result { - bail!("The client doesn't support chat-completions api"); - } - - fn prepare_embeddings(&self, data: EmbeddingsData) -> Result { - let api_key = self.get_api_key().ok(); - let api_base = self.get_api_base_ext()?; - - let url = format!("{api_base}/embeddings"); - - let body = openai_build_embeddings_body(data, &self.model); - - let mut request_data = RequestData::new(url, body); - - if let Some(api_key) = api_key { - request_data.bearer_auth(api_key); - } - - Ok(request_data) - } - - fn prepare_rerank(&self, data: RerankData) -> Result { - let api_key = self.get_api_key().ok(); - let api_base = self.get_api_base_ext()?; - - let url = format!("{api_base}/rerank"); - - let body = rag_dedicated_build_rerank_body(data, &self.model); - - let mut request_data = RequestData::new(url, body); - - if let Some(api_key) = api_key { - request_data.bearer_auth(api_key); - } - - Ok(request_data) - } - - fn get_api_base_ext(&self) -> Result { - 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 { - 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 { - 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 -} -- cgit v1.2.3