summaryrefslogtreecommitdiffstats
path: root/src/client/localai.rs
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2023-10-29 10:09:33 +0800
committersigoden <sigoden@gmail.com>2023-10-29 10:09:33 +0800
commit985f8c094682d35fd67672cc18cb47093977cd9d (patch)
tree6eff3300feda69047036904f6b66ffd0acb070b2 /src/client/localai.rs
parent557bed14597e72f93836d3ed06c06dabb555898d (diff)
downloadaichat-985f8c094682d35fd67672cc18cb47093977cd9d.tar.gz
chore: improve client-related code quality
Diffstat (limited to 'src/client/localai.rs')
-rw-r--r--src/client/localai.rs27
1 files changed, 17 insertions, 10 deletions
diff --git a/src/client/localai.rs b/src/client/localai.rs
index 5fe5f01..db33480 100644
--- a/src/client/localai.rs
+++ b/src/client/localai.rs
@@ -1,5 +1,5 @@
use super::openai::{openai_send_message, openai_send_message_streaming};
-use super::{set_proxy, Client, ModelInfo};
+use super::{set_proxy, Client, ClientConfig, ModelInfo};
use crate::config::SharedConfig;
use crate::repl::ReplyStreamHandler;
@@ -59,27 +59,34 @@ impl Client for LocalAIClient {
}
impl LocalAIClient {
- pub fn new(
- global_config: SharedConfig,
- local_config: LocalAIConfig,
- model_info: ModelInfo,
- ) -> Self {
- Self {
+ pub fn init(global_config: SharedConfig) -> Option<Box<dyn Client>> {
+ let model_info = global_config.read().model_info.clone();
+ if model_info.client != LocalAIClient::name() {
+ return None;
+ }
+ let local_config = {
+ if let ClientConfig::LocalAI(c) = &global_config.read().clients[model_info.index] {
+ c.clone()
+ } else {
+ return None;
+ }
+ };
+ Some(Box::new(Self {
global_config,
local_config,
model_info,
- }
+ }))
}
pub fn name() -> &'static str {
"localai"
}
- pub fn list_models(local_config: &LocalAIConfig) -> Vec<(String, usize)> {
+ pub fn list_models(local_config: &LocalAIConfig, index: usize) -> Vec<ModelInfo> {
local_config
.models
.iter()
- .map(|v| (v.name.to_string(), v.max_tokens))
+ .map(|v| ModelInfo::new(Self::name(), &v.name, v.max_tokens, index))
.collect()
}