summaryrefslogtreecommitdiffstats
path: root/src/client/cohere.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/cohere.rs
parent2eab71a641827e503b14952373aec82661192ba2 (diff)
downloadaichat-abc588daac6053ec2edbdcde3f5a2dc5eb7d50b8.tar.gz
feat: support rerank (#620)
Diffstat (limited to 'src/client/cohere.rs')
-rw-r--r--src/client/cohere.rs51
1 files changed, 49 insertions, 2 deletions
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();