From 66fd547c0fac30d98565a92d93824237a4813669 Mon Sep 17 00:00:00 2001 From: sigoden Date: Thu, 26 Oct 2023 19:19:22 +0800 Subject: refactor: improve code quanity remove tokio::runtime::Runtime from client --- src/client/localai.rs | 8 -------- src/client/mod.rs | 17 ++++++++++------- src/client/openai.rs | 8 -------- 3 files changed, 10 insertions(+), 23 deletions(-) (limited to 'src/client') diff --git a/src/client/localai.rs b/src/client/localai.rs index 4375c1d..52d7ab3 100644 --- a/src/client/localai.rs +++ b/src/client/localai.rs @@ -11,7 +11,6 @@ use reqwest::{Client as ReqwestClient, Proxy, RequestBuilder}; use serde::Deserialize; use serde_json::json; use std::time::Duration; -use tokio::runtime::Runtime; #[allow(clippy::module_name_repetitions)] #[derive(Debug)] @@ -19,7 +18,6 @@ pub struct LocalAIClient { global_config: SharedConfig, local_config: LocalAIConfig, model_info: ModelInfo, - runtime: Runtime, } #[derive(Debug, Clone, Deserialize)] @@ -44,10 +42,6 @@ impl Client for LocalAIClient { &self.global_config } - fn get_runtime(&self) -> &Runtime { - &self.runtime - } - async fn send_message_inner(&self, content: &str) -> Result { let builder = self.request_builder(content, false)?; openai_send_message(builder).await @@ -68,13 +62,11 @@ impl LocalAIClient { global_config: SharedConfig, local_config: LocalAIConfig, model_info: ModelInfo, - runtime: Runtime, ) -> Self { Self { global_config, local_config, model_info, - runtime, } } diff --git a/src/client/mod.rs b/src/client/mod.rs index b541cff..bf83c60 100644 --- a/src/client/mod.rs +++ b/src/client/mod.rs @@ -63,10 +63,8 @@ impl ModelInfo { pub trait Client { fn get_config(&self) -> &SharedConfig; - fn get_runtime(&self) -> &Runtime; - fn send_message(&self, content: &str) -> Result { - self.get_runtime().block_on(async { + init_runtime()?.block_on(async { if self.get_config().read().dry_run { return Ok(self.get_config().read().echo_messages(content)); } @@ -90,7 +88,7 @@ pub trait Client { } } let abort = handler.get_abort(); - self.get_runtime().block_on(async { + init_runtime()?.block_on(async { tokio::select! { ret = async { if self.get_config().read().dry_run { @@ -123,7 +121,7 @@ pub trait Client { ) -> Result<()>; } -pub fn init_client(config: SharedConfig, runtime: Runtime) -> Result> { +pub fn init_client(config: SharedConfig) -> Result> { let model_info = config.read().model_info.clone(); let model_info_err = |model_info: &ModelInfo| { bail!( @@ -144,7 +142,6 @@ pub fn init_client(config: SharedConfig, runtime: Runtime) -> Result Result Vec { }) .collect() } + +pub fn init_runtime() -> Result { + tokio::runtime::Builder::new_current_thread() + .enable_all() + .build() + .with_context(|| "Failed to init tokio") +} diff --git a/src/client/openai.rs b/src/client/openai.rs index 932dc79..8ad3bab 100644 --- a/src/client/openai.rs +++ b/src/client/openai.rs @@ -13,7 +13,6 @@ use serde::Deserialize; use serde_json::{json, Value}; use std::env; use std::time::Duration; -use tokio::runtime::Runtime; const API_URL: &str = "https://api.openai.com/v1/chat/completions"; @@ -23,7 +22,6 @@ pub struct OpenAIClient { global_config: SharedConfig, local_config: OpenAIConfig, model_info: ModelInfo, - runtime: Runtime, } #[allow(clippy::struct_excessive_bools)] @@ -42,10 +40,6 @@ impl Client for OpenAIClient { &self.global_config } - fn get_runtime(&self) -> &Runtime { - &self.runtime - } - async fn send_message_inner(&self, content: &str) -> Result { let builder = self.request_builder(content, false)?; openai_send_message(builder).await @@ -66,13 +60,11 @@ impl OpenAIClient { global_config: SharedConfig, local_config: OpenAIConfig, model_info: ModelInfo, - runtime: Runtime, ) -> Self { Self { global_config, local_config, model_info, - runtime, } } -- cgit v1.2.3