diff options
| author | sigoden <sigoden@gmail.com> | 2024-06-25 07:39:35 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2024-06-25 07:39:35 +0800 |
| commit | 2fbb5271af27eee7ba3ae6d2b3c1aa0818648c8f (patch) | |
| tree | 5e72e72d592025c00a46cebd0eb7012ccd765702 /src/client/cohere.rs | |
| parent | ed71901611247d41daed8112a5106b42eb12395b (diff) | |
| download | aichat-2fbb5271af27eee7ba3ae6d2b3c1aa0818648c8f.tar.gz | |
feat: support rag-dedicated clients (jina and voyageai) (#645)
Diffstat (limited to 'src/client/cohere.rs')
| -rw-r--r-- | src/client/cohere.rs | 36 |
1 files changed, 3 insertions, 33 deletions
diff --git a/src/client/cohere.rs b/src/client/cohere.rs index 3e1cb2b..b2857e7 100644 --- a/src/client/cohere.rs +++ b/src/client/cohere.rs @@ -1,3 +1,4 @@ +use super::rag_dedicated::*; use super::*; use anyhow::{bail, Context, Result}; @@ -74,7 +75,7 @@ impl CohereClient { 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 body = rag_dedicated_build_rerank_body(data, &self.model); let url = RERANK_API_URL; @@ -91,7 +92,7 @@ impl_client_trait!( chat_completions, chat_completions_streaming, embeddings, - cohere_rerank + rag_dedicated_rerank ); async fn chat_completions(builder: RequestBuilder) -> Result<ChatCompletionsOutput> { @@ -162,22 +163,6 @@ 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, @@ -309,21 +294,6 @@ 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(); |
