summaryrefslogtreecommitdiffstats
path: root/src/client/cohere.rs
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2024-06-05 09:02:23 +0800
committerGitHub <noreply@github.com>2024-06-05 09:02:23 +0800
commit1ec6abfaee2fdc189b348b7e3a8145bd9a84da74 (patch)
tree6f68a860a39b0fbe87784de925b5a31e00e74e33 /src/client/cohere.rs
parent71f2e94579511d7524f5534377001ab3f02a9597 (diff)
downloadaichat-1ec6abfaee2fdc189b348b7e3a8145bd9a84da74.tar.gz
feat: support RAG (#560)
* feat: support RAG * support more embeddings models and implement concurrent embedding api * show the progress of addings paths * ignore embedding context when saving message * embedding model max_chunk_size => default_chunk_size * support pdf and pandoc formats (docx, epub, ipynb)
Diffstat (limited to 'src/client/cohere.rs')
-rw-r--r--src/client/cohere.rs69
1 files changed, 58 insertions, 11 deletions
diff --git a/src/client/cohere.rs b/src/client/cohere.rs
index e0a5eec..69c343b 100644
--- a/src/client/cohere.rs
+++ b/src/client/cohere.rs
@@ -1,15 +1,12 @@
-use super::{
- catch_error, extract_system_message, json_stream, message::*, ChatCompletionsData,
- ChatCompletionsOutput, Client, CohereClient, ExtraConfig, Model, ModelData, ModelPatches,
- PromptAction, PromptKind, SseHandler, ToolCall,
-};
+use super::*;
-use anyhow::{bail, Result};
+use anyhow::{bail, Context, Result};
use reqwest::{Client as ReqwestClient, RequestBuilder};
use serde::Deserialize;
use serde_json::{json, Value};
-const API_URL: &str = "https://api.cohere.ai/v1/chat";
+const CHAT_COMPLETIONS_API_URL: &str = "https://api.cohere.ai/v1/chat";
+const EMBEDDINGS_API_URL: &str = "https://api.cohere.ai/v1/embed";
#[derive(Debug, Clone, Deserialize, Default)]
pub struct CohereConfig {
@@ -35,11 +32,38 @@ impl CohereClient {
let api_key = self.get_api_key()?;
let mut body = build_chat_completions_body(data, &self.model)?;
- self.patch_request_body(&mut body);
+ self.patch_chat_completions_body(&mut body);
- let url = API_URL;
+ let url = CHAT_COMPLETIONS_API_URL;
- debug!("Cohere Request: {url} {body}");
+ debug!("Cohere Chat Completions Request: {url} {body}");
+
+ let builder = client.post(url).bearer_auth(api_key).json(&body);
+
+ Ok(builder)
+ }
+
+ fn embeddings_builder(
+ &self,
+ client: &ReqwestClient,
+ data: EmbeddingsData,
+ ) -> Result<RequestBuilder> {
+ let api_key = self.get_api_key()?;
+
+ let input_type = match data.query {
+ true => "search_query",
+ false => "search_document",
+ };
+
+ let body = json!({
+ "model": self.model.name(),
+ "texts": data.texts,
+ "input_type": input_type,
+ });
+
+ let url = EMBEDDINGS_API_URL;
+
+ debug!("Cohere Embeddings Request: {url} {body}");
let builder = client.post(url).bearer_auth(api_key).json(&body);
@@ -47,7 +71,12 @@ impl CohereClient {
}
}
-impl_client_trait!(CohereClient, chat_completions, chat_completions_streaming);
+impl_client_trait!(
+ CohereClient,
+ chat_completions,
+ chat_completions_streaming,
+ embeddings
+);
async fn chat_completions(builder: RequestBuilder) -> Result<ChatCompletionsOutput> {
let res = builder.send().await?;
@@ -100,6 +129,24 @@ async fn chat_completions_streaming(
Ok(())
}
+async fn embeddings(
+ builder: RequestBuilder,
+) -> Result<EmbeddingsOutput> {
+ 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: EmbeddingsResBody = serde_json::from_value(data).context("Invalid request data")?;
+ Ok(res_body.embeddings)
+}
+
+#[derive(Deserialize)]
+struct EmbeddingsResBody {
+ embeddings: Vec<Vec<f32>>,
+}
+
fn build_chat_completions_body(data: ChatCompletionsData, model: &Model) -> Result<Value> {
let ChatCompletionsData {
mut messages,