summaryrefslogtreecommitdiffstats
path: root/src/client
diff options
context:
space:
mode:
Diffstat (limited to 'src/client')
-rw-r--r--src/client/localai.rs8
-rw-r--r--src/client/mod.rs17
-rw-r--r--src/client/openai.rs8
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,
}
}