summaryrefslogtreecommitdiffstats
path: root/src/client/common.rs
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2024-06-21 06:00:26 +0800
committerGitHub <noreply@github.com>2024-06-21 06:00:26 +0800
commitabc588daac6053ec2edbdcde3f5a2dc5eb7d50b8 (patch)
tree9f1cd8a40bd959420dfdf4ac3ae8b52f823aa1e9 /src/client/common.rs
parent2eab71a641827e503b14952373aec82661192ba2 (diff)
downloadaichat-abc588daac6053ec2edbdcde3f5a2dc5eb7d50b8.tar.gz
feat: support rerank (#620)
Diffstat (limited to 'src/client/common.rs')
-rw-r--r--src/client/common.rs101
1 files changed, 94 insertions, 7 deletions
diff --git a/src/client/common.rs b/src/client/common.rs
index bc84fe8..6c86acd 100644
--- a/src/client/common.rs
+++ b/src/client/common.rs
@@ -151,6 +151,10 @@ macro_rules! register_client {
pub fn list_embedding_models(config: &$crate::config::Config) -> Vec<&'static $crate::client::Model> {
list_models(config).into_iter().filter(|v| v.mode() == "embedding").collect()
}
+
+ pub fn list_rerank_models(config: &$crate::config::Config) -> Vec<&'static $crate::client::Model> {
+ list_models(config).into_iter().filter(|v| v.mode() == "rerank").collect()
+ }
};
}
@@ -236,12 +240,55 @@ macro_rules! impl_client_trait {
async fn embeddings_inner(
&self,
- client: &ReqwestClient,
- data: EmbeddingsData,
- ) -> Result<Vec<Vec<f32>>> {
+ client: &reqwest::Client,
+ data: $crate::client::EmbeddingsData,
+ ) -> Result<$crate::client::EmbeddingsOutput> {
+ let builder = self.embeddings_builder(client, data)?;
+ $embeddings(builder).await
+ }
+ }
+ };
+ ($client:ident, $chat_completions:path, $chat_completions_streaming:path, $embeddings:path, $rerank:path) => {
+ #[async_trait::async_trait]
+ impl $crate::client::Client for $crate::client::$client {
+ client_common_fns!();
+
+ async fn chat_completions_inner(
+ &self,
+ client: &reqwest::Client,
+ data: $crate::client::ChatCompletionsData,
+ ) -> anyhow::Result<$crate::client::ChatCompletionsOutput> {
+ let builder = self.chat_completions_builder(client, data)?;
+ $chat_completions(builder).await
+ }
+
+ async fn chat_completions_streaming_inner(
+ &self,
+ client: &reqwest::Client,
+ handler: &mut $crate::client::SseHandler,
+ data: $crate::client::ChatCompletionsData,
+ ) -> Result<()> {
+ let builder = self.chat_completions_builder(client, data)?;
+ $chat_completions_streaming(builder, handler).await
+ }
+
+ async fn embeddings_inner(
+ &self,
+ client: &reqwest::Client,
+ data: $crate::client::EmbeddingsData,
+ ) -> Result<$crate::client::EmbeddingsOutput> {
let builder = self.embeddings_builder(client, data)?;
$embeddings(builder).await
}
+
+ async fn rerank_inner(
+ &self,
+ client: &ReqwestClient,
+ data: RerankData,
+ ) -> Result<RerankOutput> {
+ let builder = self.rerank_builder(client, data)?;
+ $rerank(builder).await
+ }
}
};
}
@@ -308,7 +355,7 @@ pub trait Client: Sync + Send {
let data = input.prepare_completion_data(self.model(), false)?;
self.chat_completions_inner(&client, data)
.await
- .with_context(|| "Failed to get chat completions")
+ .with_context(|| "Failed to fetch chat completions")
}
async fn chat_completions_streaming(
@@ -334,7 +381,7 @@ pub trait Client: Sync + Send {
self.chat_completions_streaming_inner(&client, handler, data).await
} => {
handler.done()?;
- ret.with_context(|| "Failed to get chat completions")
+ ret.with_context(|| "Failed to fetch chat completions")
}
_ = watch_abort_signal(abort_signal) => {
handler.done()?;
@@ -348,7 +395,14 @@ pub trait Client: Sync + Send {
self.model().guard_max_concurrent_chunks(&data)?;
self.embeddings_inner(&client, data)
.await
- .with_context(|| "Failed to get embeddings")
+ .context("Failed to fetch embeddings")
+ }
+
+ async fn rerank(&self, data: RerankData) -> Result<RerankOutput> {
+ let client = self.build_client()?;
+ self.rerank_inner(&client, data)
+ .await
+ .context("Failed to fetch rerank")
}
fn patch_chat_completions_body(&self, body: &mut Value) {
@@ -377,9 +431,17 @@ pub trait Client: Sync + Send {
&self,
_client: &ReqwestClient,
_data: EmbeddingsData,
- ) -> Result<Vec<Vec<f32>>> {
+ ) -> Result<EmbeddingsOutput> {
bail!("No embeddings api")
}
+
+ async fn rerank_inner(
+ &self,
+ _client: &ReqwestClient,
+ _data: RerankData,
+ ) -> Result<RerankOutput> {
+ bail!("No rerank api")
+ }
}
impl Default for ClientConfig {
@@ -459,6 +521,31 @@ impl EmbeddingsData {
pub type EmbeddingsOutput = Vec<Vec<f32>>;
+#[derive(Debug)]
+pub struct RerankData {
+ pub query: String,
+ pub documents: Vec<String>,
+ pub top_n: usize,
+}
+
+impl RerankData {
+ pub fn new(query: String, documents: Vec<String>, top_n: usize) -> Self {
+ Self {
+ query,
+ documents,
+ top_n,
+ }
+ }
+}
+
+pub type RerankOutput = Vec<RerankResult>;
+
+#[derive(Debug, Deserialize)]
+pub struct RerankResult {
+ pub index: usize,
+ pub relevance_score: f64,
+}
+
pub type PromptAction<'a> = (&'a str, &'a str, bool, PromptKind);
pub fn create_config(prompts: &[PromptAction], client: &str) -> Result<(String, Value)> {