summaryrefslogtreecommitdiffstats
path: root/src/client/cohere.rs
diff options
context:
space:
mode:
Diffstat (limited to 'src/client/cohere.rs')
-rw-r--r--src/client/cohere.rs90
1 files changed, 50 insertions, 40 deletions
diff --git a/src/client/cohere.rs b/src/client/cohere.rs
index b6e9755..aff919e 100644
--- a/src/client/cohere.rs
+++ b/src/client/cohere.rs
@@ -1,5 +1,5 @@
-use super::rag_dedicated::*;
use super::*;
+use super::openai_compatible::*;
use anyhow::{bail, Context, Result};
use reqwest::RequestBuilder;
@@ -25,62 +25,71 @@ impl CohereClient {
pub const PROMPTS: [PromptAction<'static>; 1] =
[("api_key", "API Key:", true, PromptKind::String)];
+}
- fn prepare_chat_completions(&self, data: ChatCompletionsData) -> Result<RequestData> {
- let api_key = self.get_api_key()?;
+impl_client_trait!(
+ CohereClient,
+ (
+ prepare_chat_completions,
+ chat_completions,
+ chat_completions_streaming
+ ),
+ (prepare_embeddings, embeddings),
+ (prepare_rerank, generic_rerank),
+);
- let body = build_chat_completions_body(data, &self.model)?;
+fn prepare_chat_completions(
+ self_: &CohereClient,
+ data: ChatCompletionsData,
+) -> Result<RequestData> {
+ let api_key = self_.get_api_key()?;
- let mut request_data = RequestData::new(CHAT_COMPLETIONS_API_URL, body);
+ let body = build_chat_completions_body(data, &self_.model)?;
- request_data.bearer_auth(api_key);
+ let mut request_data = RequestData::new(CHAT_COMPLETIONS_API_URL, body);
- Ok(request_data)
- }
+ request_data.bearer_auth(api_key);
- fn prepare_embeddings(&self, data: EmbeddingsData) -> Result<RequestData> {
- let api_key = self.get_api_key()?;
+ Ok(request_data)
+}
- let input_type = match data.query {
- true => "search_query",
- false => "search_document",
- };
+fn prepare_embeddings(self_: &CohereClient, data: EmbeddingsData) -> Result<RequestData> {
+ 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 body = json!({
+ "model": self_.model.name(),
+ "texts": data.texts,
+ "input_type": input_type,
+ });
- let mut request_data = RequestData::new(EMBEDDINGS_API_URL, body);
+ let mut request_data = RequestData::new(EMBEDDINGS_API_URL, body);
- request_data.bearer_auth(api_key);
+ request_data.bearer_auth(api_key);
- Ok(request_data)
- }
+ Ok(request_data)
+}
- fn prepare_rerank(&self, data: RerankData) -> Result<RequestData> {
- let api_key = self.get_api_key()?;
+fn prepare_rerank(self_: &CohereClient, data: RerankData) -> Result<RequestData> {
+ let api_key = self_.get_api_key()?;
- let body = rag_dedicated_build_rerank_body(data, &self.model);
+ let body = generic_build_rerank_body(data, &self_.model);
- let mut request_data = RequestData::new(RERANK_API_URL, body);
+ let mut request_data = RequestData::new(RERANK_API_URL, body);
- request_data.bearer_auth(api_key);
+ request_data.bearer_auth(api_key);
- Ok(request_data)
- }
+ Ok(request_data)
}
-impl_client_trait!(
- CohereClient,
- chat_completions,
- chat_completions_streaming,
- embeddings,
- rag_dedicated_rerank
-);
-
-async fn chat_completions(builder: RequestBuilder) -> Result<ChatCompletionsOutput> {
+async fn chat_completions(
+ builder: RequestBuilder,
+ _model: &Model,
+) -> Result<ChatCompletionsOutput> {
let res = builder.send().await?;
let status = res.status();
let data: Value = res.json().await?;
@@ -95,6 +104,7 @@ async fn chat_completions(builder: RequestBuilder) -> Result<ChatCompletionsOutp
async fn chat_completions_streaming(
builder: RequestBuilder,
handler: &mut SseHandler,
+ _model: &Model,
) -> Result<()> {
let res = builder.send().await?;
let status = res.status();
@@ -131,7 +141,7 @@ async fn chat_completions_streaming(
Ok(())
}
-async fn embeddings(builder: RequestBuilder) -> Result<EmbeddingsOutput> {
+async fn embeddings(builder: RequestBuilder, _model: &Model) -> Result<EmbeddingsOutput> {
let res = builder.send().await?;
let status = res.status();
let data: Value = res.json().await?;