diff options
Diffstat (limited to 'src')
| -rw-r--r-- | src/client/localai.rs | 8 | ||||
| -rw-r--r-- | src/client/mod.rs | 17 | ||||
| -rw-r--r-- | src/client/openai.rs | 8 | ||||
| -rw-r--r-- | src/main.rs | 19 | ||||
| -rw-r--r-- | src/repl/handler.rs | 13 | ||||
| -rw-r--r-- | src/repl/mod.rs | 5 |
6 files changed, 21 insertions, 49 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, } } diff --git a/src/main.rs b/src/main.rs index ec8abee..573755b 100644 --- a/src/main.rs +++ b/src/main.rs @@ -11,7 +11,7 @@ use crate::cli::Cli; use crate::client::Client; use crate::config::{Config, SharedConfig}; -use anyhow::{anyhow, Context, Result}; +use anyhow::{anyhow, Result}; use clap::Parser; use client::{init_client, list_models}; use crossbeam::sync::WaitGroup; @@ -22,7 +22,6 @@ use repl::{AbortSignal, Repl}; use std::io::{stdin, Read}; use std::sync::Arc; use std::{io::stdout, process::exit}; -use tokio::runtime::Runtime; use utils::cl100k_base_singleton; fn main() -> Result<()> { @@ -71,8 +70,7 @@ fn main() -> Result<()> { exit(0); } let no_stream = cli.no_stream; - let runtime = init_runtime()?; - let client = init_client(config.clone(), runtime)?; + let client = init_client(config.clone())?; if atty::isnt(atty::Stream::Stdin) { let mut input = String::new(); stdin().read_to_string(&mut input)?; @@ -83,7 +81,7 @@ fn main() -> Result<()> { } else { match text { Some(text) => start_directive(client.as_ref(), &config, &text, no_stream), - None => start_interactive(client, config), + None => start_interactive(config), } } } @@ -123,16 +121,9 @@ fn start_directive( config.read().save_message(input, &output) } -fn start_interactive(client: Box<dyn Client>, config: SharedConfig) -> Result<()> { +fn start_interactive(config: SharedConfig) -> Result<()> { cl100k_base_singleton(); config.write().on_repl()?; let mut repl = Repl::init(config.clone())?; - repl.run(client, config) -} - -fn init_runtime() -> Result<Runtime> { - tokio::runtime::Builder::new_current_thread() - .enable_all() - .build() - .with_context(|| "Failed to init tokio") + repl.run(config) } diff --git a/src/repl/handler.rs b/src/repl/handler.rs index 53cdd98..e512cb5 100644 --- a/src/repl/handler.rs +++ b/src/repl/handler.rs @@ -1,4 +1,4 @@ -use crate::client::Client; +use crate::client::init_client; use crate::config::SharedConfig; use crate::print_now; use crate::render::render_stream; @@ -26,7 +26,6 @@ pub enum ReplCmd { #[allow(clippy::module_name_repetitions)] pub struct ReplCmdHandler { - client: Box<dyn Client>, config: SharedConfig, reply: RefCell<String>, abort: SharedAbortSignal, @@ -34,14 +33,9 @@ pub struct ReplCmdHandler { impl ReplCmdHandler { #[allow(clippy::unnecessary_wraps)] - pub fn init( - client: Box<dyn Client>, - config: SharedConfig, - abort: SharedAbortSignal, - ) -> Result<Self> { + pub fn init(config: SharedConfig, abort: SharedAbortSignal) -> Result<Self> { let reply = RefCell::new(String::new()); Ok(Self { - client, config, reply, abort, @@ -57,9 +51,10 @@ impl ReplCmdHandler { } self.config.read().maybe_print_send_tokens(&input); let wg = WaitGroup::new(); + let client = init_client(self.config.clone())?; let ret = render_stream( &input, - self.client.as_ref(), + client.as_ref(), &self.config, true, self.abort.clone(), diff --git a/src/repl/mod.rs b/src/repl/mod.rs index ce640f8..311ca72 100644 --- a/src/repl/mod.rs +++ b/src/repl/mod.rs @@ -9,7 +9,6 @@ pub use self::abort::*; pub use self::handler::*; pub use self::init::Repl; -use crate::client::Client; use crate::config::SharedConfig; use crate::print_now; use crate::term; @@ -35,9 +34,9 @@ pub const REPL_COMMANDS: [(&str, &str); 13] = [ ]; impl Repl { - pub fn run(&mut self, client: Box<dyn Client>, config: SharedConfig) -> Result<()> { + pub fn run(&mut self, config: SharedConfig) -> Result<()> { let abort = AbortSignal::new(); - let handler = ReplCmdHandler::init(client, config, abort.clone())?; + let handler = ReplCmdHandler::init(config, abort.clone())?; print_now!("Welcome to aichat {}\n", env!("CARGO_PKG_VERSION")); print_now!("Type \".help\" for more information.\n"); let mut already_ctrlc = false; |
