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/openai_compatible.rs | 176 ++++++++++++++++++++++++++-------------- 1 file changed, 114 insertions(+), 62 deletions(-) (limited to 'src/client/openai_compatible.rs') diff --git a/src/client/openai_compatible.rs b/src/client/openai_compatible.rs index a2302de..8acac58 100644 --- a/src/client/openai_compatible.rs +++ b/src/client/openai_compatible.rs @@ -1,9 +1,10 @@ use super::openai::*; -use super::rag_dedicated::*; use super::*; -use anyhow::Result; +use anyhow::{Context, Result}; +use reqwest::RequestBuilder; use serde::Deserialize; +use serde_json::{json, Value}; #[derive(Debug, Clone, Deserialize)] pub struct OpenAICompatibleConfig { @@ -33,90 +34,141 @@ impl OpenAICompatibleClient { PromptKind::Integer, ), ]; +} + - fn prepare_chat_completions(&self, data: ChatCompletionsData) -> Result { - let api_key = self.get_api_key().ok(); - let api_base = self.get_api_base_ext()?; +impl_client_trait!( + OpenAICompatibleClient, + ( + prepare_chat_completions, + openai_chat_completions, + openai_chat_completions_streaming + ), + (prepare_embeddings, openai_embeddings), + (prepare_rerank, generic_rerank), +); - let chat_endpoint = self - .config - .chat_endpoint - .as_deref() - .unwrap_or("/chat/completions"); +fn prepare_chat_completions( + self_: &OpenAICompatibleClient, + data: ChatCompletionsData, +) -> Result { + let api_key = self_.get_api_key().ok(); + let api_base = get_api_base_ext(self_)?; - let url = format!("{api_base}{chat_endpoint}"); + let chat_endpoint = self_ + .config + .chat_endpoint + .as_deref() + .unwrap_or("/chat/completions"); - let body = openai_build_chat_completions_body(data, &self.model); + let url = format!("{api_base}{chat_endpoint}"); - let mut request_data = RequestData::new(url, body); + let body = openai_build_chat_completions_body(data, &self_.model); - if let Some(api_key) = api_key { - request_data.bearer_auth(api_key); - } + let mut request_data = RequestData::new(url, body); - Ok(request_data) + if let Some(api_key) = api_key { + request_data.bearer_auth(api_key); } - fn prepare_embeddings(&self, data: EmbeddingsData) -> Result { - let api_key = self.get_api_key().ok(); - let api_base = self.get_api_base_ext()?; + Ok(request_data) +} - let url = format!("{api_base}/embeddings"); +fn prepare_embeddings(self_: &OpenAICompatibleClient, data: EmbeddingsData) -> Result { + let api_key = self_.get_api_key().ok(); + let api_base = get_api_base_ext(self_)?; - let body = openai_build_embeddings_body(data, &self.model); + let url = format!("{api_base}/embeddings"); - let mut request_data = RequestData::new(url, body); + let body = openai_build_embeddings_body(data, &self_.model); - if let Some(api_key) = api_key { - request_data.bearer_auth(api_key); - } + let mut request_data = RequestData::new(url, body); - Ok(request_data) + if let Some(api_key) = api_key { + request_data.bearer_auth(api_key); } - fn prepare_rerank(&self, data: RerankData) -> Result { - let api_key = self.get_api_key().ok(); - let api_base = self.get_api_base_ext()?; + Ok(request_data) +} - let url = format!("{api_base}/rerank"); +fn prepare_rerank(self_: &OpenAICompatibleClient, data: RerankData) -> Result { + let api_key = self_.get_api_key().ok(); + let api_base = get_api_base_ext(self_)?; - let body = rag_dedicated_build_rerank_body(data, &self.model); + let url = format!("{api_base}/rerank"); - let mut request_data = RequestData::new(url, body); + let body = generic_build_rerank_body(data, &self_.model); - if let Some(api_key) = api_key { - request_data.bearer_auth(api_key); - } + let mut request_data = RequestData::new(url, body); - Ok(request_data) + if let Some(api_key) = api_key { + request_data.bearer_auth(api_key); } - fn get_api_base_ext(&self) -> Result { - let api_base = match self.get_api_base() { - Ok(v) => v, - Err(err) => { - match OPENAI_COMPATIBLE_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(request_data) +} + +fn get_api_base_ext(self_: &OpenAICompatibleClient) -> Result { + let api_base = match self_.get_api_base() { + Ok(v) => v, + Err(err) => { + match OPENAI_COMPATIBLE_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) + } + }; + Ok(api_base) +} + +pub async fn generic_rerank(builder: RequestBuilder, _model: &Model) -> 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: GenericRerankResBody = + serde_json::from_value(data).context("Invalid rerank data")?; + Ok(res_body.results) } -impl_client_trait!( - OpenAICompatibleClient, - openai_chat_completions, - openai_chat_completions_streaming, - openai_embeddings, - rag_dedicated_rerank -); +#[derive(Deserialize)] +pub struct GenericRerankResBody { + pub results: RerankOutput, +} + +pub fn generic_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 +} \ No newline at end of file -- cgit v1.2.3