diff options
| author | sigoden <sigoden@gmail.com> | 2024-04-26 07:43:35 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2024-04-26 07:43:35 +0800 |
| commit | 740ca2413a1992d151c0a9ea62c12e69321c0973 (patch) | |
| tree | ccacd2fd5789064c091b212f26f3c46561f71f69 /src/client/cohere.rs | |
| parent | a21e1213ccdd11d76f6338436c431e459ab8e574 (diff) | |
| download | aichat-740ca2413a1992d151c0a9ea62c12e69321c0973.tar.gz | |
refactor: simplify impl client trait (#445)
Diffstat (limited to 'src/client/cohere.rs')
| -rw-r--r-- | src/client/cohere.rs | 32 |
1 files changed, 5 insertions, 27 deletions
diff --git a/src/client/cohere.rs b/src/client/cohere.rs index 27acb6c..3f97ece 100644 --- a/src/client/cohere.rs +++ b/src/client/cohere.rs @@ -1,12 +1,11 @@ use super::{ - catch_error, extract_system_message, json_stream, message::*, Client, CohereClient, + catch_error, extract_system_message, json_stream, message::*, CohereClient, ExtraConfig, Model, ModelConfig, PromptType, ReplyHandler, SendData, }; use crate::utils::PromptKind; use anyhow::{bail, Result}; -use async_trait::async_trait; use reqwest::{Client as ReqwestClient, RequestBuilder}; use serde::Deserialize; use serde_json::{json, Value}; @@ -22,26 +21,6 @@ pub struct CohereConfig { pub extra: Option<ExtraConfig>, } -#[async_trait] -impl Client for CohereClient { - client_common_fns!(); - - async fn send_message_inner(&self, client: &ReqwestClient, data: SendData) -> Result<String> { - let builder = self.request_builder(client, data)?; - send_message(builder).await - } - - async fn send_message_streaming_inner( - &self, - client: &ReqwestClient, - handler: &mut ReplyHandler, - data: SendData, - ) -> Result<()> { - let builder = self.request_builder(client, data)?; - send_message_streaming(builder, handler).await - } -} - impl CohereClient { list_models_fn!( CohereConfig, @@ -74,7 +53,9 @@ impl CohereClient { } } -pub(crate) async fn send_message(builder: RequestBuilder) -> Result<String> { +impl_client_trait!(CohereClient, send_message, send_message_streaming); + +async fn send_message(builder: RequestBuilder) -> Result<String> { let res = builder.send().await?; let status = res.status(); let data: Value = res.json().await?; @@ -85,10 +66,7 @@ pub(crate) async fn send_message(builder: RequestBuilder) -> Result<String> { Ok(output.to_string()) } -pub(crate) async fn send_message_streaming( - builder: RequestBuilder, - handler: &mut ReplyHandler, -) -> Result<()> { +async fn send_message_streaming(builder: RequestBuilder, handler: &mut ReplyHandler) -> Result<()> { let res = builder.send().await?; let status = res.status(); if status != 200 { |
