diff options
| author | sigoden <sigoden@gmail.com> | 2023-10-26 19:19:22 +0800 |
|---|---|---|
| committer | sigoden <sigoden@gmail.com> | 2023-10-26 19:19:22 +0800 |
| commit | 66fd547c0fac30d98565a92d93824237a4813669 (patch) | |
| tree | 332ca0d0927739de69ff4e7b8ac589f4f001f9f1 /src/client/mod.rs | |
| parent | 7d8564cafb45afc4bccb666a310615b164800a9a (diff) | |
| download | aichat-66fd547c0fac30d98565a92d93824237a4813669.tar.gz | |
refactor: improve code quanity
remove tokio::runtime::Runtime from client
Diffstat (limited to 'src/client/mod.rs')
| -rw-r--r-- | src/client/mod.rs | 17 |
1 files changed, 10 insertions, 7 deletions
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<String> { - 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<Box<dyn Client>> { +pub fn init_client(config: SharedConfig) -> Result<Box<dyn Client>> { 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<Box<dyn Cli config, local_config, model_info, - runtime, ))) } else if model_info.client == LocalAIClient::name() { let local_config = { @@ -158,7 +155,6 @@ pub fn init_client(config: SharedConfig, runtime: Runtime) -> Result<Box<dyn Cli config, local_config, model_info, - runtime, ))) } else { bail!("Unknown client {}", &model_info.client) @@ -196,3 +192,10 @@ pub fn list_models(config: &Config) -> Vec<ModelInfo> { }) .collect() } + +pub fn init_runtime() -> Result<Runtime> { + tokio::runtime::Builder::new_current_thread() + .enable_all() + .build() + .with_context(|| "Failed to init tokio") +} |
