From 444f4ebe9de8aa68afdf8d39b734a8543b9b5c4b Mon Sep 17 00:00:00 2001 From: sigoden Date: Thu, 2 Nov 2023 09:53:54 +0800 Subject: refactor: improve code quanity (#196) - rewrite Repl, remove ReplHandler - move ReplyStreamHandler to repl/ and rename it to ReplyHandler - deprecate utils::print_now - refactor session info --- src/client/azure_openai.rs | 4 +++- src/client/common.rs | 15 ++++++--------- src/client/localai.rs | 4 +++- src/client/mod.rs | 6 ------ src/client/openai.rs | 11 ++++++++--- 5 files changed, 20 insertions(+), 20 deletions(-) (limited to 'src/client') diff --git a/src/client/azure_openai.rs b/src/client/azure_openai.rs index 5603666..fe3ec0f 100644 --- a/src/client/azure_openai.rs +++ b/src/client/azure_openai.rs @@ -1,5 +1,7 @@ use super::openai::{openai_build_body, OPENAI_TOKENS_COUNT_FACTORS}; -use super::{AzureOpenAIClient, ExtraConfig, ModelInfo, PromptKind, PromptType, SendData}; +use super::{AzureOpenAIClient, ExtraConfig, PromptType, SendData}; + +use crate::{config::ModelInfo, utils::PromptKind}; use anyhow::{anyhow, Result}; use async_trait::async_trait; diff --git a/src/client/common.rs b/src/client/common.rs index 7963c76..0d0c0e2 100644 --- a/src/client/common.rs +++ b/src/client/common.rs @@ -1,6 +1,7 @@ use crate::{ config::{Message, SharedConfig}, - repl::{ReplyStreamHandler, SharedAbortSignal}, + render::ReplyHandler, + repl::AbortSignal, utils::{init_tokio_runtime, prompt_input_integer, prompt_input_string, tokenize, PromptKind}, }; @@ -139,7 +140,7 @@ macro_rules! openai_compatible_client { async fn send_message_streaming_inner( &self, client: &reqwest::Client, - handler: &mut $crate::repl::ReplyStreamHandler, + handler: &mut $crate::render::ReplyHandler, data: $crate::client::SendData, ) -> Result<()> { let builder = self.request_builder(client, data)?; @@ -201,12 +202,8 @@ pub trait Client { }) } - fn send_message_streaming( - &self, - content: &str, - handler: &mut ReplyStreamHandler, - ) -> Result<()> { - async fn watch_abort(abort: SharedAbortSignal) { + fn send_message_streaming(&self, content: &str, handler: &mut ReplyHandler) -> Result<()> { + async fn watch_abort(abort: AbortSignal) { loop { if abort.aborted() { break; @@ -252,7 +249,7 @@ pub trait Client { async fn send_message_streaming_inner( &self, client: &ReqwestClient, - handler: &mut ReplyStreamHandler, + handler: &mut ReplyHandler, data: SendData, ) -> Result<()>; } diff --git a/src/client/localai.rs b/src/client/localai.rs index d438388..796b574 100644 --- a/src/client/localai.rs +++ b/src/client/localai.rs @@ -1,5 +1,7 @@ use super::openai::{openai_build_body, OPENAI_TOKENS_COUNT_FACTORS}; -use super::{ExtraConfig, LocalAIClient, ModelInfo, PromptKind, PromptType, SendData}; +use super::{ExtraConfig, LocalAIClient, PromptType, SendData}; + +use crate::{config::ModelInfo, utils::PromptKind}; use anyhow::Result; use async_trait::async_trait; diff --git a/src/client/mod.rs b/src/client/mod.rs index 5fa6146..e55055d 100644 --- a/src/client/mod.rs +++ b/src/client/mod.rs @@ -3,12 +3,6 @@ mod common; pub use common::*; -use crate::{ - config::{ModelInfo, TokensCountFactors}, - repl::ReplyStreamHandler, - utils::PromptKind, -}; - register_client!( (openai, "openai", OpenAI, OpenAIConfig, OpenAIClient), (localai, "localai", LocalAI, LocalAIConfig, LocalAIClient), diff --git a/src/client/openai.rs b/src/client/openai.rs index a82ee9a..98de5a8 100644 --- a/src/client/openai.rs +++ b/src/client/openai.rs @@ -1,6 +1,11 @@ use super::{ - ExtraConfig, ModelInfo, OpenAIClient, PromptKind, PromptType, ReplyStreamHandler, SendData, - TokensCountFactors, + ExtraConfig, OpenAIClient, PromptType, SendData, +}; + +use crate::{ + config::{ModelInfo, TokensCountFactors}, + render::ReplyHandler, + utils::PromptKind, }; use anyhow::{anyhow, bail, Result}; @@ -88,7 +93,7 @@ pub async fn openai_send_message(builder: RequestBuilder) -> Result { pub async fn openai_send_message_streaming( builder: RequestBuilder, - handler: &mut ReplyStreamHandler, + handler: &mut ReplyHandler, ) -> Result<()> { let mut es = builder.eventsource()?; while let Some(event) = es.next().await { -- cgit v1.2.3