summaryrefslogtreecommitdiffstats
path: root/src/client
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
parent2eab71a641827e503b14952373aec82661192ba2 (diff)
downloadaichat-abc588daac6053ec2edbdcde3f5a2dc5eb7d50b8.tar.gz
feat: support rerank (#620)
Diffstat (limited to 'src/client')
-rw-r--r--src/client/claude.rs3
-rw-r--r--src/client/cohere.rs51
-rw-r--r--src/client/common.rs101
-rw-r--r--src/client/gemini.rs2
-rw-r--r--src/client/model.rs13
-rw-r--r--src/client/ollama.rs2
-rw-r--r--src/client/openai.rs2
-rw-r--r--src/client/openai_compatible.rs29
-rw-r--r--src/client/qianwen.rs2
-rw-r--r--src/client/reka.rs10
-rw-r--r--src/client/vertexai.rs2
11 files changed, 189 insertions, 28 deletions
diff --git a/src/client/claude.rs b/src/client/claude.rs
index f95fafc..2b836ee 100644
--- a/src/client/claude.rs
+++ b/src/client/claude.rs
@@ -38,8 +38,7 @@ impl ClaudeClient {
debug!("Claude Request: {url} {body}");
let mut builder = client.post(url).json(&body);
- builder = builder
- .header("anthropic-version", "2023-06-01");
+ builder = builder.header("anthropic-version", "2023-06-01");
if let Some(api_key) = api_key {
builder = builder.header("x-api-key", api_key)
}
diff --git a/src/client/cohere.rs b/src/client/cohere.rs
index 8745347..698c2a6 100644
--- a/src/client/cohere.rs
+++ b/src/client/cohere.rs
@@ -7,6 +7,7 @@ use serde_json::{json, Value};
const CHAT_COMPLETIONS_API_URL: &str = "https://api.cohere.ai/v1/chat";
const EMBEDDINGS_API_URL: &str = "https://api.cohere.ai/v1/embed";
+const RERANK_API_URL: &str = "https://api.cohere.ai/v1/rerank";
#[derive(Debug, Clone, Deserialize, Default)]
pub struct CohereConfig {
@@ -69,13 +70,28 @@ impl CohereClient {
Ok(builder)
}
+
+ fn rerank_builder(&self, client: &ReqwestClient, data: RerankData) -> Result<RequestBuilder> {
+ let api_key = self.get_api_key()?;
+
+ let body = cohere_build_rerank_body(data, &self.model);
+
+ let url = RERANK_API_URL;
+
+ debug!("Cohere Rerank Request: {url} {body}");
+
+ let builder = client.post(url).bearer_auth(api_key).json(&body);
+
+ Ok(builder)
+ }
}
impl_client_trait!(
CohereClient,
chat_completions,
chat_completions_streaming,
- embeddings
+ embeddings,
+ cohere_rerank
);
async fn chat_completions(builder: RequestBuilder) -> Result<ChatCompletionsOutput> {
@@ -137,7 +153,7 @@ async fn embeddings(builder: RequestBuilder) -> Result<EmbeddingsOutput> {
catch_error(&data, status.as_u16())?;
}
let res_body: EmbeddingsResBody =
- serde_json::from_value(data).context("Invalid request data")?;
+ serde_json::from_value(data).context("Invalid embeddings data")?;
Ok(res_body.embeddings)
}
@@ -146,6 +162,22 @@ struct EmbeddingsResBody {
embeddings: Vec<Vec<f32>>,
}
+pub async fn cohere_rerank(builder: RequestBuilder) -> Result<RerankOutput> {
+ let res = builder.send().await?;
+ let status = res.status();
+ let data: Value = res.json().await?;
+ if !status.is_success() {
+ catch_error(&data, status.as_u16())?;
+ }
+ let res_body: RerankResBody = serde_json::from_value(data).context("Invalid rerank data")?;
+ Ok(res_body.results)
+}
+
+#[derive(Deserialize)]
+struct RerankResBody {
+ results: RerankOutput,
+}
+
fn build_chat_completions_body(data: ChatCompletionsData, model: &Model) -> Result<Value> {
let ChatCompletionsData {
mut messages,
@@ -277,6 +309,21 @@ fn build_chat_completions_body(data: ChatCompletionsData, model: &Model) -> Resu
Ok(body)
}
+pub fn cohere_build_rerank_body(data: RerankData, model: &Model) -> Value {
+ let RerankData {
+ query,
+ documents,
+ top_n,
+ } = data;
+
+ json!({
+ "model": model.name(),
+ "query": query,
+ "documents": documents,
+ "top_n": top_n
+ })
+}
+
fn extract_chat_completions(data: &Value) -> Result<ChatCompletionsOutput> {
let text = data["text"].as_str().unwrap_or_default();
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)> {
diff --git a/src/client/gemini.rs b/src/client/gemini.rs
index 03eef7a..6382fb8 100644
--- a/src/client/gemini.rs
+++ b/src/client/gemini.rs
@@ -94,7 +94,7 @@ async fn gemini_embeddings(builder: RequestBuilder) -> Result<EmbeddingsOutput>
catch_error(&data, status.as_u16())?;
}
let res_body: EmbeddingsResBody =
- serde_json::from_value(data).context("Invalid request data")?;
+ serde_json::from_value(data).context("Invalid embeddings data")?;
let output = vec![res_body.embedding.values];
Ok(output)
}
diff --git a/src/client/model.rs b/src/client/model.rs
index 56421bf..ebf1264 100644
--- a/src/client/model.rs
+++ b/src/client/model.rs
@@ -1,5 +1,5 @@
use super::{
- list_chat_models, list_embedding_models,
+ list_chat_models, list_embedding_models, list_rerank_models,
message::{Message, MessageContent},
EmbeddingsData,
};
@@ -46,14 +46,21 @@ impl Model {
pub fn retrieve_chat(config: &Config, model_id: &str) -> Result<Self> {
match Self::find(&list_chat_models(config), model_id) {
Some(v) => Ok(v),
- None => bail!("Invalid model '{model_id}'"),
+ None => bail!("Invalid chat model '{model_id}'"),
}
}
pub fn retrieve_embedding(config: &Config, model_id: &str) -> Result<Self> {
match Self::find(&list_embedding_models(config), model_id) {
Some(v) => Ok(v),
- None => bail!("Invalid model '{model_id}'"),
+ None => bail!("Invalid embedding model '{model_id}'"),
+ }
+ }
+
+ pub fn retrieve_rerank(config: &Config, model_id: &str) -> Result<Self> {
+ match Self::find(&list_rerank_models(config), model_id) {
+ Some(v) => Ok(v),
+ None => bail!("Invalid rerank model '{model_id}'"),
}
}
diff --git a/src/client/ollama.rs b/src/client/ollama.rs
index 9bc8978..be065c1 100644
--- a/src/client/ollama.rs
+++ b/src/client/ollama.rs
@@ -141,7 +141,7 @@ async fn embeddings(builder: RequestBuilder) -> Result<EmbeddingsOutput> {
catch_error(&data, status.as_u16())?;
}
let res_body: EmbeddingsResBody =
- serde_json::from_value(data).context("Invalid request data")?;
+ serde_json::from_value(data).context("Invalid embeddings data")?;
let output = vec![res_body.embedding];
Ok(output)
}
diff --git a/src/client/openai.rs b/src/client/openai.rs
index 0c51b33..40058e6 100644
--- a/src/client/openai.rs
+++ b/src/client/openai.rs
@@ -148,7 +148,7 @@ pub async fn openai_embeddings(builder: RequestBuilder) -> Result<EmbeddingsOutp
catch_error(&data, status.as_u16())?;
}
let res_body: EmbeddingsResBody =
- serde_json::from_value(data).context("Invalid request data")?;
+ serde_json::from_value(data).context("Invalid embeddings data")?;
let output = res_body.data.into_iter().map(|v| v.embedding).collect();
Ok(output)
}
diff --git a/src/client/openai_compatible.rs b/src/client/openai_compatible.rs
index f5f446a..789604f 100644
--- a/src/client/openai_compatible.rs
+++ b/src/client/openai_compatible.rs
@@ -1,3 +1,4 @@
+use super::cohere::*;
use super::openai::*;
use super::*;
@@ -68,7 +69,7 @@ impl OpenAICompatibleClient {
client: &ReqwestClient,
data: EmbeddingsData,
) -> Result<RequestBuilder> {
- let api_key = self.get_api_key()?;
+ let api_key = self.get_api_key().ok();
let api_base = self.get_api_base_ext()?;
let body = openai_build_embeddings_body(data, &self.model);
@@ -77,7 +78,28 @@ impl OpenAICompatibleClient {
debug!("OpenAICompatible Embeddings Request: {url} {body}");
- let builder = client.post(url).bearer_auth(api_key).json(&body);
+ let mut builder = client.post(url).json(&body);
+ if let Some(api_key) = api_key {
+ builder = builder.bearer_auth(api_key);
+ }
+
+ Ok(builder)
+ }
+
+ fn rerank_builder(&self, client: &ReqwestClient, data: RerankData) -> Result<RequestBuilder> {
+ let api_key = self.get_api_key().ok();
+ let api_base = self.get_api_base_ext()?;
+
+ let body = cohere_build_rerank_body(data, &self.model);
+
+ let url = format!("{api_base}/rerank");
+
+ debug!("OpenAICompatible Rerank Request: {url} {body}");
+
+ let mut builder = client.post(url).json(&body);
+ if let Some(api_key) = api_key {
+ builder = builder.bearer_auth(api_key);
+ }
Ok(builder)
}
@@ -108,5 +130,6 @@ impl_client_trait!(
OpenAICompatibleClient,
openai_chat_completions,
openai_chat_completions_streaming,
- openai_embeddings
+ openai_embeddings,
+ cohere_rerank
);
diff --git a/src/client/qianwen.rs b/src/client/qianwen.rs
index b0d0b58..ec84011 100644
--- a/src/client/qianwen.rs
+++ b/src/client/qianwen.rs
@@ -326,7 +326,7 @@ async fn embeddings(builder: RequestBuilder) -> Result<EmbeddingsOutput> {
let data: Value = builder.send().await?.json().await?;
maybe_catch_error(&data)?;
let res_body: EmbeddingsResBody =
- serde_json::from_value(data).context("Invalid request data")?;
+ serde_json::from_value(data).context("Invalid embeddings data")?;
let output = res_body
.output
.embeddings
diff --git a/src/client/reka.rs b/src/client/reka.rs
index 46f9118..2e9b88f 100644
--- a/src/client/reka.rs
+++ b/src/client/reka.rs
@@ -43,11 +43,7 @@ impl RekaClient {
}
}
-impl_client_trait!(
- RekaClient,
- chat_completions,
- chat_completions_streaming
-);
+impl_client_trait!(RekaClient, chat_completions, chat_completions_streaming);
async fn chat_completions(builder: RequestBuilder) -> Result<ChatCompletionsOutput> {
let res = builder.send().await?;
@@ -113,7 +109,9 @@ fn build_chat_completions_body(data: ChatCompletionsData, model: &Model) -> Valu
}
fn extract_chat_completions(data: &Value) -> Result<ChatCompletionsOutput> {
- let text = data["responses"][0]["message"]["content"].as_str().ok_or_else(|| anyhow!("Invalid response data: {data}"))?;
+ let text = data["responses"][0]["message"]["content"]
+ .as_str()
+ .ok_or_else(|| anyhow!("Invalid response data: {data}"))?;
let output = ChatCompletionsOutput {
text: text.to_string(),
diff --git a/src/client/vertexai.rs b/src/client/vertexai.rs
index dc75c9f..69d1bc4 100644
--- a/src/client/vertexai.rs
+++ b/src/client/vertexai.rs
@@ -185,7 +185,7 @@ async fn embeddings(builder: RequestBuilder) -> Result<EmbeddingsOutput> {
catch_error(&data, status.as_u16())?;
}
let res_body: EmbeddingsResBody =
- serde_json::from_value(data).context("Invalid request data")?;
+ serde_json::from_value(data).context("Invalid embeddings data")?;
let output = res_body
.predictions
.into_iter()