diff options
Diffstat (limited to 'src/client')
| -rw-r--r-- | src/client/localai.rs | 8 | ||||
| -rw-r--r-- | src/client/mod.rs | 17 | ||||
| -rw-r--r-- | src/client/openai.rs | 8 |
3 files changed, 10 insertions, 23 deletions
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<String> { 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<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") +} 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<String> { 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, } } |
