summaryrefslogtreecommitdiffstats
path: root/src/client/openai.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/openai.rs
parent557bed14597e72f93836d3ed06c06dabb555898d (diff)
downloadaichat-985f8c094682d35fd67672cc18cb47093977cd9d.tar.gz
chore: improve client-related code quality
Diffstat (limited to 'src/client/openai.rs')
-rw-r--r--src/client/openai.rs38
1 files changed, 24 insertions, 14 deletions
diff --git a/src/client/openai.rs b/src/client/openai.rs
index 885a02d..0141d2d 100644
--- a/src/client/openai.rs
+++ b/src/client/openai.rs
@@ -1,4 +1,4 @@
-use super::{set_proxy, Client, ModelInfo};
+use super::{set_proxy, Client, ClientConfig, ModelInfo};
use crate::config::SharedConfig;
use crate::repl::ReplyStreamHandler;
@@ -56,29 +56,39 @@ impl Client for OpenAIClient {
}
impl OpenAIClient {
- pub fn new(
- global_config: SharedConfig,
- local_config: OpenAIConfig,
- 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 != OpenAIClient::name() {
+ return None;
+ }
+ let local_config = {
+ if let ClientConfig::OpenAI(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 {
"openai"
}
- pub fn list_models(_local_config: &OpenAIConfig) -> Vec<(String, usize)> {
- vec![
- ("gpt-3.5-turbo".into(), 4096),
- ("gpt-3.5-turbo-16k".into(), 16384),
- ("gpt-4".into(), 8192),
- ("gpt-4-32k".into(), 32768),
+ pub fn list_models(_local_config: &OpenAIConfig, index: usize) -> Vec<ModelInfo> {
+ [
+ ("gpt-3.5-turbo", 4096),
+ ("gpt-3.5-turbo-16k", 16384),
+ ("gpt-4", 8192),
+ ("gpt-4-32k", 32768),
]
+ .into_iter()
+ .map(|(name, max_tokens)| ModelInfo::new(Self::name(), name, max_tokens, index))
+ .collect()
}
pub fn create_config() -> Result<String> {