diff options
Diffstat (limited to 'src/client/cohere.rs')
| -rw-r--r-- | src/client/cohere.rs | 69 |
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, |
