From 338b0438dc34280f31cd7642a3da4f565490c749 Mon Sep 17 00:00:00 2001 From: sigoden Date: Sun, 28 Apr 2024 13:28:24 +0800 Subject: feat: run without config file by set `AICHAT_CLIENT_TYPE` (#452) --- config.example.yaml | 73 ++++++++++++++++++-------------------- src/config/mod.rs | 100 +++++++++++++++++++--------------------------------- 2 files changed, 71 insertions(+), 102 deletions(-) diff --git a/config.example.yaml b/config.example.yaml index cd08273..a0c6da8 100644 --- a/config.example.yaml +++ b/config.example.yaml @@ -4,7 +4,7 @@ top_p: null # Set default top-p parameter save: true # Indicates whether to persist the message save_session: null # Controls the persistence of the session, if null, asking the user highlight: true # Controls syntax highlighting -light_theme: false # Activates a light color theme when true +light_theme: false # Activates a light color theme when true. ENV: AICHAT_LIGHT_THEME wrap: no # Controls text wrapping (no, auto, ) wrap_code: false # Enables or disables wrapping of code blocks auto_copy: false # Enables or disables automatic copying the last LLM response to the clipboard @@ -32,83 +32,78 @@ clients: # name: xxxx # Only use it to distinguish clients with the same client type. Optional # models: # - name: xxxx # The model name - # max_input_tokens: 100000 # Optional field - # max_output_tokens: 4096 # Optional field - # capabilities: text,vision # Optional field, supported capabilities: text, vision - # extra_fields: # Optional field, set custom parameters, will merge with the body json + # max_input_tokens: 100000 + # max_output_tokens: 4096 + # supports_vision: true + # extra_fields: # Set custom parameters, will merge with the body json # key: value # extra: - # proxy: socks5://127.0.0.1:1080 # Specify https/socks5 proxy server. Note HTTPS_PROXY/ALL_PROXY also works. + # proxy: socks5://127.0.0.1:1080 # Specify https/socks5 proxy server. ENV: HTTPS_PROXY/ALL_PROXY # connect_timeout: 10 # Set a timeout in seconds for connect to server # See https://platform.openai.com/docs/quickstart - type: openai - api_key: sk-xxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxx - api_base: https://api.openai.com/v1 # Optional field - organization_id: org-xxxxxxxxxxxxxxxxxxxxxxxx # Optional field + api_key: sk-xxx # ENV: {client_name}_API_KEY + api_base: https://api.openai.com/v1 # ENV: {client_name}_API_BASE + organization_id: org-xxx # Optional # See https://ai.google.dev/docs - type: gemini - api_key: xxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxx - # Optional field, possible values: BLOCK_NONE, BLOCK_ONLY_HIGH, BLOCK_MEDIUM_AND_ABOVE, BLOCK_LOW_AND_ABOVE - block_threshold: BLOCK_NONE + api_key: xxx # ENV: {client_name}_API_KEY + # possible values: BLOCK_NONE, BLOCK_ONLY_HIGH, BLOCK_MEDIUM_AND_ABOVE, BLOCK_LOW_AND_ABOVE + block_threshold: BLOCK_NONE # Optional # See https://docs.anthropic.com/claude/reference/getting-started-with-the-api - type: claude - api_key: sk-ant-xxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxx + api_key: sk-ant-xxx # ENV: {client_name}_API_KEY # See https://docs.mistral.ai/ - type: mistral - api_key: xxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxx + api_key: xxx # ENV: {client_name}_API_KEY # See https://docs.cohere.com/docs/the-cohere-platform - type: cohere - api_key: xxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxx + api_key: xxx # ENV: {client_name}_API_KEY # See https://docs.perplexity.ai/docs/getting-started - type: perplexity - api_key: pplx-xxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxx + api_key: pplx-xxx # ENV: {client_name}_API_KEY # See https://console.groq.com/docs/quickstart - type: groq - api_key: gsk_xxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxx + api_key: gsk_xxx # ENV: {client_name}_API_KEY # Any openai-compatible API providers - type: openai-compatible name: localai api_base: http://localhost:8080/v1 - api_key: sk-xxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxx - chat_endpoint: /chat/completions # Optional field + api_key: sk-xxx # ENV: {client_name}_API_BASE + chat_endpoint: /chat/completions models: - - name: llama2 - max_input_tokens: 8192 - extra_fields: # Optional field, set custom parameters - key: value - - name: llava + - name: llama3 max_input_tokens: 8192 - capabilities: text,vision # Optional field, choices: text, vision # See https://github.com/jmorganca/ollama - type: ollama api_base: http://localhost:11434 - api_key: Basic xxx # Set authorization header - chat_endpoint: /api/chat # Optional field + api_key: Basic xxx # Set authorization header, ENV: {client_name}_API_BASE + chat_endpoint: /api/chat # Optional models: - - name: llama2 + - name: llama3 max_input_tokens: 8192 # See https://learn.microsoft.com/en-us/azure/ai-services/openai/chatgpt-quickstart - type: azure-openai api_base: https://{RESOURCE}.openai.azure.com - api_key: xxx + api_key: xxx # ENV: {client_name}_API_BASE models: - - name: MyGPT4 # Model deployment name + - name: gpt-35-turbo # Model deployment name max_input_tokens: 8192 # See https://cloud.google.com/vertex-ai - type: vertexai - project_id: xxx - location: xxx + project_id: xxx # ENV: {client_name}_PROJECT_ID + location: xxx # ENV: {client_name}_LOCATION # Specifies a application-default-credentials (adc) file, Optional field # Run `gcloud auth application-default login` to init the adc file # see https://cloud.google.com/docs/authentication/external/set-up-adc @@ -118,19 +113,19 @@ clients: # See https://docs.aws.amazon.com/bedrock/latest/userguide/ - type: bedrock - access_key_id: xxxxxxxxxxxxxxxxxxxx - secret_access_key: xxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxx - region: xxx + access_key_id: xxx # ENV: {client_name}_ACCESS_KEY_ID + secret_access_key: xxx # ENV: {client_name}_SECRET_ACCESS_KEY + region: xxx # ENV: {client_name}_REGION # See https://cloud.baidu.com/doc/WENXINWORKSHOP/index.html - type: ernie - api_key: xxxxxxxxxxxxxxxxxxxxxxxx - secret_key: xxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxx + api_key: xxx # ENV: {client_name}_API_KEY + secret_key: xxxx # ENV: {client_name}_SECRET_KEY # See https://help.aliyun.com/zh/dashscope/ - type: qianwen - api_key: sk-xxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxx + api_key: sk-xxx # ENV: {client_name}_API_KEY # See https://platform.moonshot.cn/docs/intro - type: moonshot - api_key: sk-xxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxx + api_key: sk-xxx # ENV: {client_name}_API_KEY diff --git a/src/config/mod.rs b/src/config/mod.rs index 87be519..7807204 100644 --- a/src/config/mod.rs +++ b/src/config/mod.rs @@ -8,8 +8,7 @@ pub use self::role::{CODE_ROLE, EXPLAIN_ROLE, SHELL_ROLE}; use self::session::{Session, TEMP_SESSION_NAME}; use crate::client::{ - create_client_config, list_client_types, list_models, ClientConfig, ExtraConfig, Message, - Model, OpenAIClient, SendData, + create_client_config, list_client_types, list_models, ClientConfig, Message, Model, SendData, }; use crate::render::{MarkdownRender, RenderOptions}; use crate::utils::{get_env_name, light_theme_from_colorfgbg, now, render_prompt, set_text}; @@ -91,7 +90,7 @@ impl Default for Config { model_id: None, temperature: None, top_p: None, - save: true, + save: false, save_session: None, highlight: true, dry_run: false, @@ -124,23 +123,16 @@ impl Config { pub fn init(working_mode: WorkingMode) -> Result { let config_path = Self::config_file()?; - let api_key = env::var("OPENAI_API_KEY").ok(); - - let exist_config_path = config_path.exists(); - if working_mode != WorkingMode::Command && api_key.is_none() && !exist_config_path { + let client_type = env::var(get_env_name("client_type")).ok(); + if working_mode != WorkingMode::Command && client_type.is_none() && !config_path.exists() { create_config_file(&config_path)?; } - let mut config = if api_key.is_some() && !exist_config_path { - Self::default() + let mut config = if client_type.is_some() { + Self::load_config_env(&client_type.unwrap())? } else { - Self::load_config(&config_path)? + Self::load_config_file(&config_path)? }; - // Compatible with old configuration files - if exist_config_path { - config.compat_old_config(&config_path)?; - } - if let Some(wrap) = config.wrap.clone() { config.set_wrap(&wrap)?; } @@ -898,20 +890,39 @@ impl Config { Ok(()) } - fn load_config(config_path: &Path) -> Result { + fn load_config_file(config_path: &Path) -> Result { let ctx = || format!("Failed to load config at {}", config_path.display()); let content = read_to_string(config_path).with_context(ctx)?; + let config = Self::load_config(&content).with_context(ctx)?; + Ok(config) + } - let config: Self = serde_yaml::from_str(&content) - .map_err(|err| { - let err_msg = err.to_string(); - if err_msg.starts_with(&format!("{}: ", CLIENTS_FIELD)) { - anyhow!("clients: invalid value") - } else { - anyhow!("{err_msg}") - } - }) - .with_context(ctx)?; + fn load_config_env(client_type: &str) -> Result { + let model_id = match env::var(get_env_name("model_name")) { + Ok(model_name) => format!("{client_type}:{model_name}"), + Err(_) => client_type.to_string(), + }; + let content = format!( + r#" +model: {model_id} +save: false +clients: + - type: {client_type} +"# + ); + let config = Self::load_config(&content).with_context(|| "Failed to load config")?; + Ok(config) + } + + fn load_config(content: &str) -> Result { + let config: Self = serde_yaml::from_str(content).map_err(|err| { + let err_msg = err.to_string(); + if err_msg.starts_with(&format!("{}: ", CLIENTS_FIELD)) { + anyhow!("clients: invalid value") + } else { + anyhow!("{err_msg}") + } + })?; Ok(config) } @@ -969,43 +980,6 @@ impl Config { }; Ok(()) } - - fn compat_old_config(&mut self, config_path: &PathBuf) -> Result<()> { - let content = read_to_string(config_path)?; - let value: serde_json::Value = serde_yaml::from_str(&content)?; - if value.get(CLIENTS_FIELD).is_some() { - return Ok(()); - } - - if let Some(model_name) = value.get("model").and_then(|v| v.as_str()) { - if model_name.starts_with("gpt") { - self.model_id = Some(format!("{}:{}", OpenAIClient::NAME, model_name)); - } - } - - if let Some(ClientConfig::OpenAIConfig(client_config)) = self.clients.first_mut() { - if let Some(api_key) = value.get("api_key").and_then(|v| v.as_str()) { - client_config.api_key = Some(api_key.to_string()) - } - - if let Some(organization_id) = value.get("organization_id").and_then(|v| v.as_str()) { - client_config.organization_id = Some(organization_id.to_string()) - } - - let mut extra_config = ExtraConfig::default(); - - if let Some(proxy) = value.get("proxy").and_then(|v| v.as_str()) { - extra_config.proxy = Some(proxy.to_string()) - } - - if let Some(connect_timeout) = value.get("connect_timeout").and_then(|v| v.as_i64()) { - extra_config.connect_timeout = Some(connect_timeout as _) - } - - client_config.extra = Some(extra_config); - } - Ok(()) - } } #[derive(Debug, Clone, Deserialize, Default)] -- cgit v1.2.3