summaryrefslogtreecommitdiffstats
path: root/src/client/common.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/common.rs
parentf5e8e872b08d2aa9b3c063b0a98c59f842f94ecc (diff)
downloadaichat-0e740d81e94505bd57036755abaaecb12c3b26e3.tar.gz
feat: abandon rag_dedicated client and improve (#757)
Diffstat (limited to 'src/client/common.rs')
-rw-r--r--src/client/common.rs77
1 files changed, 46 insertions, 31 deletions
diff --git a/src/client/common.rs b/src/client/common.rs
index 3321511..2108941 100644
--- a/src/client/common.rs
+++ b/src/client/common.rs
@@ -8,7 +8,6 @@ use crate::{
};
use anyhow::{bail, Context, Result};
-use async_trait::async_trait;
use fancy_regex::Regex;
use indexmap::IndexMap;
use lazy_static::lazy_static;
@@ -25,7 +24,7 @@ lazy_static! {
static ref ESCAPE_SLASH_RE: Regex = Regex::new(r"(?<!\\)/").unwrap();
}
-#[async_trait]
+#[async_trait::async_trait]
pub trait Client: Sync + Send {
fn global_config(&self) -> &GlobalConfig;
@@ -110,6 +109,35 @@ pub trait Client: Sync + Send {
.context("Failed to call rerank api")
}
+ async fn chat_completions_inner(
+ &self,
+ client: &ReqwestClient,
+ data: ChatCompletionsData,
+ ) -> Result<ChatCompletionsOutput>;
+
+ async fn chat_completions_streaming_inner(
+ &self,
+ client: &ReqwestClient,
+ handler: &mut SseHandler,
+ data: ChatCompletionsData,
+ ) -> Result<()>;
+
+ async fn embeddings_inner(
+ &self,
+ _client: &ReqwestClient,
+ _data: EmbeddingsData,
+ ) -> Result<EmbeddingsOutput> {
+ bail!("The client doesn't support embeddings api")
+ }
+
+ async fn rerank_inner(
+ &self,
+ _client: &ReqwestClient,
+ _data: RerankData,
+ ) -> Result<RerankOutput> {
+ bail!("The client doesn't support rerank api")
+ }
+
fn request_builder(
&self,
client: &reqwest::Client,
@@ -147,35 +175,6 @@ pub trait Client: Sync + Send {
}
}
}
-
- async fn chat_completions_inner(
- &self,
- client: &ReqwestClient,
- data: ChatCompletionsData,
- ) -> Result<ChatCompletionsOutput>;
-
- async fn chat_completions_streaming_inner(
- &self,
- client: &ReqwestClient,
- handler: &mut SseHandler,
- data: ChatCompletionsData,
- ) -> Result<()>;
-
- async fn embeddings_inner(
- &self,
- _client: &ReqwestClient,
- _data: EmbeddingsData,
- ) -> Result<EmbeddingsOutput> {
- bail!("The client doesn't support embeddings api")
- }
-
- async fn rerank_inner(
- &self,
- _client: &ReqwestClient,
- _data: RerankData,
- ) -> Result<RerankOutput> {
- bail!("The client doesn't support rerank api")
- }
}
impl Default for ClientConfig {
@@ -448,6 +447,22 @@ where
Ok(())
}
+pub fn noop_prepare_embeddings<T>(_client: &T, _data: EmbeddingsData) -> Result<RequestData> {
+ bail!("The client doesn't support embeddings api")
+}
+
+pub async fn noop_embeddings(_builder: RequestBuilder, _model: &Model) -> Result<EmbeddingsOutput> {
+ bail!("The client doesn't support embeddings api")
+}
+
+pub fn noop_prepare_rerank<T>(_client: &T, _data: RerankData) -> Result<RequestData> {
+ bail!("The client doesn't support rerank api")
+}
+
+pub async fn noop_rerank(_builder: RequestBuilder, _model: &Model) -> Result<RerankOutput> {
+ bail!("The client doesn't support rerank api")
+}
+
pub fn catch_error(data: &Value, status: u16) -> Result<()> {
if (200..300).contains(&status) {
return Ok(());