From 7f2210dbcab017ee30e8f95771e68494b4860c59 Mon Sep 17 00:00:00 2001 From: sigoden Date: Wed, 1 Nov 2023 10:28:54 +0800 Subject: refactor: clients/* and config.rs (#193) - add register_clients macro to make it easier to add a new client - no create_client_config, just add const PROMPTS - move ModelInfo from clients/ to config/ - model's max_tokens are optional - improve code quanity on config/mod.rs - add/use macro config_get_fn --- src/client/openai.rs | 105 ++++++++++++++------------------------------------- 1 file changed, 29 insertions(+), 76 deletions(-) (limited to 'src/client/openai.rs') diff --git a/src/client/openai.rs b/src/client/openai.rs index 57cd5fd..80969e4 100644 --- a/src/client/openai.rs +++ b/src/client/openai.rs @@ -1,4 +1,4 @@ -use super::{prompt_input_api_key, Client, ClientConfig, ExtraConfig, ModelInfo, SendData}; +use super::{Client, ExtraConfig, ModelInfo, OpenAIClient, PromptKind, PromptType, SendData}; use crate::config::SharedConfig; use crate::repl::ReplyStreamHandler; @@ -14,12 +14,12 @@ use std::env; const API_BASE: &str = "https://api.openai.com/v1"; -#[derive(Debug)] -pub struct OpenAIClient { - global_config: SharedConfig, - config: OpenAIConfig, - model_info: ModelInfo, -} +const MODELS: [(&str, usize); 4] = [ + ("gpt-3.5-turbo", 4096), + ("gpt-3.5-turbo-16k", 16384), + ("gpt-4", 8192), + ("gpt-4-32k", 32768), +]; #[derive(Debug, Clone, Deserialize, Default)] pub struct OpenAIConfig { @@ -31,12 +31,8 @@ pub struct OpenAIConfig { #[async_trait] impl Client for OpenAIClient { - fn config(&self) -> &SharedConfig { - &self.global_config - } - - fn extra_config(&self) -> &Option { - &self.config.extra + fn config(&self) -> (&SharedConfig, &Option) { + (&self.global_config, &self.config.extra) } async fn send_message_inner(&self, client: &ReqwestClient, data: SendData) -> Result { @@ -56,61 +52,25 @@ impl Client for OpenAIClient { } impl OpenAIClient { - pub const NAME: &str = "openai"; - - pub fn init(global_config: SharedConfig) -> Option> { - let model_info = global_config.read().model_info.clone(); - let config = { - if let ClientConfig::OpenAI(c) = &global_config.read().clients[model_info.index] { - c.clone() - } else { - return None; - } - }; - Some(Box::new(Self { - global_config, - config, - model_info, - })) - } + config_get_fn!(api_key, get_api_key); - pub fn name(local_config: &OpenAIConfig) -> &str { - local_config.name.as_deref().unwrap_or(Self::NAME) - } + pub const PROMPTS: [PromptType<'static>; 1] = + [("api_key", "API Key:", true, PromptKind::String)]; pub fn list_models(local_config: &OpenAIConfig, index: usize) -> Vec { let client = Self::name(local_config); - - [ - ("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(client, name, max_tokens, index)) - .collect() - } - - pub fn create_config() -> Result { - let mut client_config = format!("clients:\n - type: {}\n", Self::NAME); - - let api_key = prompt_input_api_key()?; - client_config.push_str(&format!(" api_key: {api_key}\n")); - - Ok(client_config) + MODELS + .into_iter() + .map(|(name, max_tokens)| ModelInfo::new(client, name, Some(max_tokens), index)) + .collect() } fn request_builder(&self, client: &ReqwestClient, data: SendData) -> Result { - let env_prefix = Self::name(&self.config).to_uppercase(); - - let api_key = self.config.api_key.clone(); - let api_key = api_key - .or_else(|| env::var(format!("{env_prefix}_API_KEY")).ok()) - .ok_or_else(|| anyhow!("Miss api_key"))?; + let api_key = self.get_api_key()?; let body = openai_build_body(data, self.model_info.name.clone()); + let env_prefix = Self::name(&self.config).to_uppercase(); let api_base = env::var(format!("{env_prefix}_API_BASE")) .ok() .unwrap_or_else(|| API_BASE.to_string()); @@ -127,20 +87,20 @@ impl OpenAIClient { } } -pub(crate) async fn openai_send_message(builder: RequestBuilder) -> Result { +pub async fn openai_send_message(builder: RequestBuilder) -> Result { let data: Value = builder.send().await?.json().await?; if let Some(err_msg) = data["error"]["message"].as_str() { - bail!("Request failed, {err_msg}"); + bail!("{err_msg}"); } let output = data["choices"][0]["message"]["content"] .as_str() - .ok_or_else(|| anyhow!("Unexpected response {data}"))?; + .ok_or_else(|| anyhow!("Invalid response data: {data}"))?; Ok(output.to_string()) } -pub(crate) async fn openai_send_message_streaming( +pub async fn openai_send_message_streaming( builder: RequestBuilder, handler: &mut ReplyStreamHandler, ) -> Result<()> { @@ -148,7 +108,7 @@ pub(crate) async fn openai_send_message_streaming( if !res.status().is_success() { let data: Value = res.json().await?; if let Some(err_msg) = data["error"]["message"].as_str() { - bail!("Request failed, {err_msg}"); + bail!("{err_msg}"); } bail!("Request failed"); } @@ -159,37 +119,30 @@ pub(crate) async fn openai_send_message_streaming( break; } let data: Value = serde_json::from_str(&chunk)?; - let text = data["choices"][0]["delta"]["content"] - .as_str() - .unwrap_or_default(); - if text.is_empty() { - continue; + if let Some(text) = data["choices"][0]["delta"]["content"].as_str() { + handler.text(text)?; } - handler.text(text)?; } Ok(()) } -pub(crate) fn openai_build_body(data: SendData, model: String) -> Value { +pub fn openai_build_body(data: SendData, model: String) -> Value { let SendData { messages, temperature, stream, } = data; + let mut body = json!({ "model": model, "messages": messages, }); - if let Some(v) = temperature { - body.as_object_mut() - .and_then(|m| m.insert("temperature".into(), json!(v))); + body["temperature"] = v.into(); } - if stream { - body.as_object_mut() - .and_then(|m| m.insert("stream".into(), json!(true))); + body["stream"] = true.into(); } body } -- cgit v1.2.3