summaryrefslogtreecommitdiffstats
path: root/src
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2023-10-26 19:19:22 +0800
committersigoden <sigoden@gmail.com>2023-10-26 19:19:22 +0800
commit66fd547c0fac30d98565a92d93824237a4813669 (patch)
tree332ca0d0927739de69ff4e7b8ac589f4f001f9f1 /src
parent7d8564cafb45afc4bccb666a310615b164800a9a (diff)
downloadaichat-66fd547c0fac30d98565a92d93824237a4813669.tar.gz
refactor: improve code quanity
remove tokio::runtime::Runtime from client
Diffstat (limited to 'src')
-rw-r--r--src/client/localai.rs8
-rw-r--r--src/client/mod.rs17
-rw-r--r--src/client/openai.rs8
-rw-r--r--src/main.rs19
-rw-r--r--src/repl/handler.rs13
-rw-r--r--src/repl/mod.rs5
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;