summaryrefslogtreecommitdiffstats
path: root/src/client/openai_compatible.rs
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2024-07-28 06:04:36 +0800
committerGitHub <noreply@github.com>2024-07-28 06:04:36 +0800
commit0e740d81e94505bd57036755abaaecb12c3b26e3 (patch)
tree49000370fb12e4e5e1f1bd4d145104f2f502aa16 /src/client/openai_compatible.rs
parentf5e8e872b08d2aa9b3c063b0a98c59f842f94ecc (diff)
downloadaichat-0e740d81e94505bd57036755abaaecb12c3b26e3.tar.gz
feat: abandon rag_dedicated client and improve (#757)
Diffstat (limited to 'src/client/openai_compatible.rs')
-rw-r--r--src/client/openai_compatible.rs176
1 files changed, 114 insertions, 62 deletions
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<RequestData> {
- 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<RequestData> {
+ 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<RequestData> {
- 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<RequestData> {
+ 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<RequestData> {
- 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<RequestData> {
+ 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<String> {
- 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<String> {
+ 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<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: 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