diff options
| author | sigoden <sigoden@gmail.com> | 2024-04-30 12:52:58 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2024-04-30 12:52:58 +0800 |
| commit | 8dba46becfbcc4669867db4487aa9b91c51e4afa (patch) | |
| tree | 7b67d556751d9f5fed380b50b818ddb8680c8f01 /src | |
| parent | 8a65337d590729f96a5f3c0b35dc5a08fae5bf94 (diff) | |
| download | aichat-8dba46becfbcc4669867db4487aa9b91c51e4afa.tar.gz | |
feat: openai-compatible platforms share the same client (#469)
Diffstat (limited to 'src')
| -rw-r--r-- | src/client/azure_openai.rs | 6 | ||||
| -rw-r--r-- | src/client/bedrock.rs | 6 | ||||
| -rw-r--r-- | src/client/claude.rs | 4 | ||||
| -rw-r--r-- | src/client/cloudflare.rs | 4 | ||||
| -rw-r--r-- | src/client/cohere.rs | 4 | ||||
| -rw-r--r-- | src/client/common.rs | 114 | ||||
| -rw-r--r-- | src/client/ernie.rs | 4 | ||||
| -rw-r--r-- | src/client/gemini.rs | 4 | ||||
| -rw-r--r-- | src/client/groq.rs | 1 | ||||
| -rw-r--r-- | src/client/mistral.rs | 1 | ||||
| -rw-r--r-- | src/client/mod.rs | 34 | ||||
| -rw-r--r-- | src/client/model.rs | 6 | ||||
| -rw-r--r-- | src/client/moonshot.rs | 1 | ||||
| -rw-r--r-- | src/client/ollama.rs | 4 | ||||
| -rw-r--r-- | src/client/openai.rs | 4 | ||||
| -rw-r--r-- | src/client/openai_compatible.rs | 25 | ||||
| -rw-r--r-- | src/client/perplexity.rs | 5 | ||||
| -rw-r--r-- | src/client/qianwen.rs | 4 | ||||
| -rw-r--r-- | src/client/replicate.rs | 6 | ||||
| -rw-r--r-- | src/client/vertexai.rs | 5 | ||||
| -rw-r--r-- | src/config/mod.rs | 70 |
21 files changed, 138 insertions, 174 deletions
diff --git a/src/client/azure_openai.rs b/src/client/azure_openai.rs index 005351f..315a4ce 100644 --- a/src/client/azure_openai.rs +++ b/src/client/azure_openai.rs @@ -1,5 +1,7 @@ use super::openai::openai_build_body; -use super::{AzureOpenAIClient, ExtraConfig, Model, ModelConfig, PromptKind, PromptType, SendData}; +use super::{ + AzureOpenAIClient, ExtraConfig, Model, ModelConfig, PromptAction, PromptKind, SendData, +}; use anyhow::Result; use reqwest::{Client as ReqwestClient, RequestBuilder}; @@ -18,7 +20,7 @@ impl AzureOpenAIClient { config_get_fn!(api_base, get_api_base); config_get_fn!(api_key, get_api_key); - pub const PROMPTS: [PromptType<'static>; 4] = [ + pub const PROMPTS: [PromptAction<'static>; 4] = [ ("api_base", "API Base:", true, PromptKind::String), ("api_key", "API Key:", true, PromptKind::String), ("models[].name", "Model Name:", true, PromptKind::String), diff --git a/src/client/bedrock.rs b/src/client/bedrock.rs index a889162..fb4dad8 100644 --- a/src/client/bedrock.rs +++ b/src/client/bedrock.rs @@ -1,8 +1,8 @@ use super::claude::{claude_build_body, claude_extract_completion}; use super::{ catch_error, generate_prompt, BedrockClient, Client, CompletionDetails, ExtraConfig, Model, - ModelConfig, PromptFormat, PromptKind, PromptType, SendData, SseHandler, LLAMA2_PROMPT_FORMAT, - LLAMA3_PROMPT_FORMAT, + ModelConfig, PromptAction, PromptFormat, PromptKind, SendData, SseHandler, + LLAMA2_PROMPT_FORMAT, LLAMA3_PROMPT_FORMAT, }; use crate::utils::{base64_decode, encode_uri, hex_encode, hmac_sha256, sha256}; @@ -65,7 +65,7 @@ impl BedrockClient { config_get_fn!(secret_access_key, get_secret_access_key); config_get_fn!(region, get_region); - pub const PROMPTS: [PromptType<'static>; 3] = [ + pub const PROMPTS: [PromptAction<'static>; 3] = [ ( "access_key_id", "AWS Access Key ID", diff --git a/src/client/claude.rs b/src/client/claude.rs index 84b4a6c..0a230e9 100644 --- a/src/client/claude.rs +++ b/src/client/claude.rs @@ -1,6 +1,6 @@ use super::{ catch_error, extract_system_message, sse_stream, ClaudeClient, CompletionDetails, ExtraConfig, - ImageUrl, MessageContent, MessageContentPart, Model, ModelConfig, PromptKind, PromptType, + ImageUrl, MessageContent, MessageContentPart, Model, ModelConfig, PromptAction, PromptKind, SendData, SsMmessage, SseHandler, }; @@ -23,7 +23,7 @@ pub struct ClaudeConfig { impl ClaudeClient { config_get_fn!(api_key, get_api_key); - pub const PROMPTS: [PromptType<'static>; 1] = + pub const PROMPTS: [PromptAction<'static>; 1] = [("api_key", "API Key:", true, PromptKind::String)]; fn request_builder(&self, client: &ReqwestClient, data: SendData) -> Result<RequestBuilder> { diff --git a/src/client/cloudflare.rs b/src/client/cloudflare.rs index 14a5828..9758032 100644 --- a/src/client/cloudflare.rs +++ b/src/client/cloudflare.rs @@ -1,6 +1,6 @@ use super::{ catch_error, sse_stream, CloudflareClient, CompletionDetails, ExtraConfig, Model, ModelConfig, - PromptKind, PromptType, SendData, SsMmessage, SseHandler, + PromptAction, PromptKind, SendData, SsMmessage, SseHandler, }; use anyhow::{anyhow, Result}; @@ -24,7 +24,7 @@ impl CloudflareClient { config_get_fn!(account_id, get_account_id); config_get_fn!(api_key, get_api_key); - pub const PROMPTS: [PromptType<'static>; 2] = [ + pub const PROMPTS: [PromptAction<'static>; 2] = [ ("account_id", "Account ID:", true, PromptKind::String), ("api_key", "API Key:", true, PromptKind::String), ]; diff --git a/src/client/cohere.rs b/src/client/cohere.rs index 0069c2c..e0ef6f0 100644 --- a/src/client/cohere.rs +++ b/src/client/cohere.rs @@ -1,6 +1,6 @@ use super::{ catch_error, extract_system_message, json_stream, message::*, CohereClient, CompletionDetails, - ExtraConfig, Model, ModelConfig, PromptKind, PromptType, SendData, SseHandler, + ExtraConfig, Model, ModelConfig, PromptAction, PromptKind, SendData, SseHandler, }; use anyhow::{anyhow, bail, Result}; @@ -22,7 +22,7 @@ pub struct CohereConfig { impl CohereClient { config_get_fn!(api_key, get_api_key); - pub const PROMPTS: [PromptType<'static>; 1] = + pub const PROMPTS: [PromptAction<'static>; 1] = [("api_key", "API Key:", true, PromptKind::String)]; fn request_builder(&self, client: &ReqwestClient, data: SendData) -> Result<RequestBuilder> { diff --git a/src/client/common.rs b/src/client/common.rs index ee190d5..91d6bca 100644 --- a/src/client/common.rs +++ b/src/client/common.rs @@ -1,4 +1,4 @@ -use super::{openai::OpenAIConfig, ClientConfig, ClientModel, Message, Model, SseHandler}; +use super::{openai::OpenAIConfig, BuiltinModels, ClientConfig, Message, Model, SseHandler}; use crate::{ config::{GlobalConfig, Input}, @@ -20,7 +20,8 @@ use tokio::{sync::mpsc::unbounded_channel, time::sleep}; const MODELS_YAML: &str = include_str!("../../models.yaml"); lazy_static! { - pub static ref CLIENT_MODELS: Vec<ClientModel> = serde_yaml::from_str(MODELS_YAML).unwrap(); + pub static ref ALL_CLIENT_MODELS: Vec<BuiltinModels> = + serde_yaml::from_str(MODELS_YAML).unwrap(); } #[macro_export] @@ -90,13 +91,10 @@ macro_rules! register_client { pub fn list_models(local_config: &$config) -> Vec<Model> { let client_name = Self::name(local_config); if local_config.models.is_empty() { - for model in $crate::client::CLIENT_MODELS.iter() { - match model { - $crate::client::ClientModel::$config { models } => { - return Model::from_config(client_name, models); - } - _ => {} - } + if let Some(client_models) = $crate::client::ALL_CLIENT_MODELS.iter().find(|v| { + v.platform == $name || ($name == "openai-compatible" && local_config.name.as_deref() == Some(&v.platform)) + }) { + return Model::from_config(client_name, &client_models.models); } vec![] } else { @@ -135,7 +133,7 @@ macro_rules! register_client { pub fn list_client_types() -> Vec<&'static str> { let mut client_types: Vec<_> = vec![$($client::NAME,)+]; - client_types.extend($crate::client::KNOWN_OPENAI_COMPATIBLE_PLATFORMS.iter().map(|(name, _)| *name)); + client_types.extend($crate::client::OPENAI_COMPATIBLE_PLATFORMS.iter().map(|(name, _)| *name)); client_types } @@ -171,69 +169,6 @@ macro_rules! register_client { } #[macro_export] -macro_rules! openai_compatible_client { - ( - $config:ident, - $client:ident, - $api_base:literal, - ) => { - use $crate::client::openai::openai_build_body; - use $crate::client::{$client, ExtraConfig, Model, ModelConfig, PromptType, SendData}; - - use $crate::utils::PromptKind; - - use anyhow::Result; - use reqwest::{Client as ReqwestClient, RequestBuilder}; - use serde::Deserialize; - - const API_BASE: &str = $api_base; - - #[derive(Debug, Clone, Deserialize)] - pub struct $config { - pub name: Option<String>, - pub api_key: Option<String>, - #[serde(default)] - pub models: Vec<ModelConfig>, - pub extra: Option<ExtraConfig>, - } - - impl_client_trait!( - $client, - $crate::client::openai::openai_send_message, - $crate::client::openai::openai_send_message_streaming - ); - - impl $client { - config_get_fn!(api_key, get_api_key); - - pub const PROMPTS: [PromptType<'static>; 1] = - [("api_key", "API Key:", true, PromptKind::String)]; - - fn request_builder( - &self, - client: &ReqwestClient, - data: SendData, - ) -> Result<RequestBuilder> { - let api_key = self.get_api_key().ok(); - - let body = openai_build_body(data, &self.model); - - let url = format!("{API_BASE}/chat/completions"); - - debug!("Request: {url} {body}"); - - let mut builder = client.post(url).json(&body); - if let Some(api_key) = api_key { - builder = builder.bearer_auth(api_key); - } - - Ok(builder) - } - } - }; -} - -#[macro_export] macro_rules! client_common_fns { () => { fn config( @@ -437,36 +372,45 @@ pub struct CompletionDetails { pub output_tokens: Option<u64>, } -pub type PromptType<'a> = (&'a str, &'a str, bool, PromptKind); +pub type PromptAction<'a> = (&'a str, &'a str, bool, PromptKind); -pub fn create_config(list: &[PromptType], client: &str) -> Result<(String, Value)> { +pub fn create_config(prompts: &[PromptAction], client: &str) -> Result<(String, Value)> { let mut config = json!({ "type": client, }); let mut model = client.to_string(); - set_client_config_values(list, &mut model, &mut config)?; + set_client_config_values(prompts, &mut model, &mut config)?; let clients = json!(vec![config]); Ok((model, clients)) } pub fn create_openai_compatible_client_config(client: &str) -> Result<Option<(String, Value)>> { - match super::KNOWN_OPENAI_COMPATIBLE_PLATFORMS + match super::OPENAI_COMPATIBLE_PLATFORMS .iter() .find(|(name, _)| client == *name) { None => Ok(None), - Some((name, api_base)) => { + Some((name, _)) => { let mut config = json!({ "type": "openai-compatible", "name": name, - "api_base": api_base, }); + let prompts = if ALL_CLIENT_MODELS.iter().any(|v| &v.platform == name) { + vec![("api_key", "API Key:", false, PromptKind::String)] + } else { + vec![ + ("api_key", "API Key:", false, PromptKind::String), + ("models[].name", "Model Name:", true, PromptKind::String), + ( + "models[].max_input_tokens", + "Max Input Tokens:", + false, + PromptKind::Integer, + ), + ] + }; let mut model = client.to_string(); - set_client_config_values( - &super::KNOWN_OPENAI_COMPATIBLE_PROMPTS, - &mut model, - &mut config, - )?; + set_client_config_values(&prompts, &mut model, &mut config)?; let clients = json!(vec![config]); Ok(Some((model, clients))) } @@ -683,7 +627,7 @@ where } fn set_client_config_values( - list: &[PromptType], + list: &[PromptAction], model: &mut String, client_config: &mut Value, ) -> Result<()> { diff --git a/src/client/ernie.rs b/src/client/ernie.rs index a0187a4..982edae 100644 --- a/src/client/ernie.rs +++ b/src/client/ernie.rs @@ -1,6 +1,6 @@ use super::{ maybe_catch_error, patch_system_message, sse_stream, Client, CompletionDetails, ErnieClient, - ExtraConfig, Model, ModelConfig, PromptKind, PromptType, SendData, SsMmessage, SseHandler, + ExtraConfig, Model, ModelConfig, PromptAction, PromptKind, SendData, SsMmessage, SseHandler, }; use anyhow::{anyhow, Context, Result}; @@ -27,7 +27,7 @@ pub struct ErnieConfig { } impl ErnieClient { - pub const PROMPTS: [PromptType<'static>; 2] = [ + pub const PROMPTS: [PromptAction<'static>; 2] = [ ("api_key", "API Key:", true, PromptKind::String), ("secret_key", "Secret Key:", true, PromptKind::String), ]; diff --git a/src/client/gemini.rs b/src/client/gemini.rs index 783b674..8f6a76d 100644 --- a/src/client/gemini.rs +++ b/src/client/gemini.rs @@ -1,5 +1,5 @@ use super::vertexai::gemini_build_body; -use super::{ExtraConfig, GeminiClient, Model, ModelConfig, PromptKind, PromptType, SendData}; +use super::{ExtraConfig, GeminiClient, Model, ModelConfig, PromptAction, PromptKind, SendData}; use anyhow::Result; use reqwest::{Client as ReqwestClient, RequestBuilder}; @@ -20,7 +20,7 @@ pub struct GeminiConfig { impl GeminiClient { config_get_fn!(api_key, get_api_key); - pub const PROMPTS: [PromptType<'static>; 1] = + pub const PROMPTS: [PromptAction<'static>; 1] = [("api_key", "API Key:", true, PromptKind::String)]; fn request_builder(&self, client: &ReqwestClient, data: SendData) -> Result<RequestBuilder> { diff --git a/src/client/groq.rs b/src/client/groq.rs deleted file mode 100644 index 23ca33d..0000000 --- a/src/client/groq.rs +++ /dev/null @@ -1 +0,0 @@ -openai_compatible_client!(GroqConfig, GroqClient, "https://api.groq.com/openai/v1",); diff --git a/src/client/mistral.rs b/src/client/mistral.rs deleted file mode 100644 index 351502d..0000000 --- a/src/client/mistral.rs +++ /dev/null @@ -1 +0,0 @@ -openai_compatible_client!(MistralConfig, MistralClient, "https://api.mistral.ai/v1",); diff --git a/src/client/mod.rs b/src/client/mod.rs index a8519b3..4ea5533 100644 --- a/src/client/mod.rs +++ b/src/client/mod.rs @@ -14,12 +14,15 @@ pub use sse_handler::*; register_client!( (openai, "openai", OpenAIConfig, OpenAIClient), + ( + openai_compatible, + "openai-compatible", + OpenAICompatibleConfig, + OpenAICompatibleClient + ), (gemini, "gemini", GeminiConfig, GeminiClient), (claude, "claude", ClaudeConfig, ClaudeClient), - (mistral, "mistral", MistralConfig, MistralClient), (cohere, "cohere", CohereConfig, CohereClient), - (perplexity, "perplexity", PerplexityConfig, PerplexityClient), - (groq, "groq", GroqConfig, GroqClient), (ollama, "ollama", OllamaConfig, OllamaClient), ( azure_openai, @@ -33,30 +36,17 @@ register_client!( (replicate, "replicate", ReplicateConfig, ReplicateClient), (ernie, "ernie", ErnieConfig, ErnieClient), (qianwen, "qianwen", QianwenConfig, QianwenClient), - (moonshot, "moonshot", MoonshotConfig, MoonshotClient), - ( - openai_compatible, - "openai-compatible", - OpenAICompatibleConfig, - OpenAICompatibleClient - ), ); -pub const KNOWN_OPENAI_COMPATIBLE_PLATFORMS: [(&str, &str); 5] = [ +pub const OPENAI_COMPATIBLE_PLATFORMS: [(&str, &str); 10] = [ ("anyscale", "https://api.endpoints.anyscale.com/v1"), ("deepinfra", "https://api.deepinfra.com/v1/openai"), ("fireworks", "https://api.fireworks.ai/inference/v1"), + ("groq", "https://api.groq.com/openai/v1"), + ("mistral", "https://api.mistral.ai/v1"), + ("moonshot", "https://api.moonshot.cn/v1"), + ("openrouter", "https://openrouter.ai/api/v1"), ("octoai", "https://text.octoai.run/v1"), + ("perplexity", "https://api.perplexity.ai"), ("together", "https://api.together.xyz/v1"), ]; - -pub const KNOWN_OPENAI_COMPATIBLE_PROMPTS: [PromptType<'static>; 3] = [ - ("api_key", "API Key:", false, PromptKind::String), - ("models[].name", "Model Name:", true, PromptKind::String), - ( - "models[].max_input_tokens", - "Max Input Tokens:", - false, - PromptKind::Integer, - ), -]; diff --git a/src/client/model.rs b/src/client/model.rs index 42040fb..4213556 100644 --- a/src/client/model.rs +++ b/src/client/model.rs @@ -242,6 +242,12 @@ pub struct ModelConfig { pub extra_fields: Option<serde_json::Map<String, serde_json::Value>>, } +#[derive(Debug, Clone, Deserialize)] +pub struct BuiltinModels { + pub platform: String, + pub models: Vec<ModelConfig>, +} + bitflags::bitflags! { #[derive(Debug, Clone, Copy, PartialEq)] pub struct ModelCapabilities: u32 { diff --git a/src/client/moonshot.rs b/src/client/moonshot.rs deleted file mode 100644 index 903d60f..0000000 --- a/src/client/moonshot.rs +++ /dev/null @@ -1 +0,0 @@ -openai_compatible_client!(MoonshotConfig, MoonshotClient, "https://api.moonshot.cn/v1",); diff --git a/src/client/ollama.rs b/src/client/ollama.rs index a2688a5..b61417a 100644 --- a/src/client/ollama.rs +++ b/src/client/ollama.rs @@ -1,6 +1,6 @@ use super::{ catch_error, message::*, CompletionDetails, ExtraConfig, Model, ModelConfig, OllamaClient, - PromptKind, PromptType, SendData, SseHandler, + PromptAction, PromptKind, SendData, SseHandler, }; use anyhow::{anyhow, bail, Result}; @@ -23,7 +23,7 @@ impl OllamaClient { config_get_fn!(api_base, get_api_base); config_get_fn!(api_auth, get_api_auth); - pub const PROMPTS: [PromptType<'static>; 4] = [ + pub const PROMPTS: [PromptAction<'static>; 4] = [ ("api_base", "API Base:", true, PromptKind::String), ("api_auth", "API Auth:", false, PromptKind::String), ("models[].name", "Model Name:", true, PromptKind::String), diff --git a/src/client/openai.rs b/src/client/openai.rs index 7e8fb87..08bb94d 100644 --- a/src/client/openai.rs +++ b/src/client/openai.rs @@ -1,6 +1,6 @@ use super::{ catch_error, sse_stream, CompletionDetails, ExtraConfig, Model, ModelConfig, OpenAIClient, - PromptKind, PromptType, SendData, SsMmessage, SseHandler, + PromptAction, PromptKind, SendData, SsMmessage, SseHandler, }; use anyhow::{anyhow, Result}; @@ -25,7 +25,7 @@ impl OpenAIClient { config_get_fn!(api_key, get_api_key); config_get_fn!(api_base, get_api_base); - pub const PROMPTS: [PromptType<'static>; 1] = + pub const PROMPTS: [PromptAction<'static>; 1] = [("api_key", "API Key:", true, PromptKind::String)]; fn request_builder(&self, client: &ReqwestClient, data: SendData) -> Result<RequestBuilder> { diff --git a/src/client/openai_compatible.rs b/src/client/openai_compatible.rs index d7aff2b..6eae77b 100644 --- a/src/client/openai_compatible.rs +++ b/src/client/openai_compatible.rs @@ -1,6 +1,8 @@ +use crate::client::OPENAI_COMPATIBLE_PLATFORMS; + use super::openai::openai_build_body; use super::{ - ExtraConfig, Model, ModelConfig, OpenAICompatibleClient, PromptKind, PromptType, SendData, + ExtraConfig, Model, ModelConfig, OpenAICompatibleClient, PromptAction, PromptKind, SendData, }; use anyhow::Result; @@ -13,6 +15,7 @@ pub struct OpenAICompatibleConfig { pub api_base: Option<String>, pub api_key: Option<String>, pub chat_endpoint: Option<String>, + #[serde(default)] pub models: Vec<ModelConfig>, pub extra: Option<ExtraConfig>, } @@ -21,7 +24,7 @@ impl OpenAICompatibleClient { config_get_fn!(api_base, get_api_base); config_get_fn!(api_key, get_api_key); - pub const PROMPTS: [PromptType<'static>; 5] = [ + pub const PROMPTS: [PromptAction<'static>; 5] = [ ("name", "Platform Name:", true, PromptKind::String), ("api_base", "API Base:", true, PromptKind::String), ("api_key", "API Key:", false, PromptKind::String), @@ -35,7 +38,23 @@ impl OpenAICompatibleClient { ]; fn request_builder(&self, client: &ReqwestClient, data: SendData) -> Result<RequestBuilder> { - let api_base = self.get_api_base()?; + let api_base = match self.get_api_base() { + Ok(v) => v, + Err(err) => { + match OPENAI_COMPATIBLE_PLATFORMS + .into_iter() + .find_map(|(name, api_base)| { + if name == self.model.client_name { + Some(api_base.to_string()) + } else { + None + } + }) { + Some(v) => v, + None => return Err(err), + } + } + }; let api_key = self.get_api_key().ok(); let mut body = openai_build_body(data, &self.model); diff --git a/src/client/perplexity.rs b/src/client/perplexity.rs deleted file mode 100644 index df3d462..0000000 --- a/src/client/perplexity.rs +++ /dev/null @@ -1,5 +0,0 @@ -openai_compatible_client!( - PerplexityConfig, - PerplexityClient, - "https://api.perplexity.ai", -); diff --git a/src/client/qianwen.rs b/src/client/qianwen.rs index b1ba093..76d7436 100644 --- a/src/client/qianwen.rs +++ b/src/client/qianwen.rs @@ -1,6 +1,6 @@ use super::{ maybe_catch_error, message::*, sse_stream, Client, CompletionDetails, ExtraConfig, Model, - ModelConfig, PromptKind, PromptType, QianwenClient, SendData, SsMmessage, SseHandler, + ModelConfig, PromptAction, PromptKind, QianwenClient, SendData, SsMmessage, SseHandler, }; use crate::utils::{base64_decode, sha256}; @@ -33,7 +33,7 @@ pub struct QianwenConfig { impl QianwenClient { config_get_fn!(api_key, get_api_key); - pub const PROMPTS: [PromptType<'static>; 1] = + pub const PROMPTS: [PromptAction<'static>; 1] = [("api_key", "API Key:", true, PromptKind::String)]; fn request_builder(&self, client: &ReqwestClient, data: SendData) -> Result<RequestBuilder> { diff --git a/src/client/replicate.rs b/src/client/replicate.rs index aef992d..a20ce71 100644 --- a/src/client/replicate.rs +++ b/src/client/replicate.rs @@ -2,8 +2,8 @@ use std::time::Duration; use super::{ catch_error, generate_prompt, smart_prompt_format, sse_stream, Client, CompletionDetails, - ExtraConfig, Model, ModelConfig, PromptKind, PromptType, ReplicateClient, SendData, SsMmessage, - SseHandler, + ExtraConfig, Model, ModelConfig, PromptAction, PromptKind, ReplicateClient, SendData, + SsMmessage, SseHandler, }; use anyhow::{anyhow, Result}; @@ -26,7 +26,7 @@ pub struct ReplicateConfig { impl ReplicateClient { config_get_fn!(api_key, get_api_key); - pub const PROMPTS: [PromptType<'static>; 1] = + pub const PROMPTS: [PromptAction<'static>; 1] = [("api_key", "API Key:", true, PromptKind::String)]; fn request_builder( diff --git a/src/client/vertexai.rs b/src/client/vertexai.rs index e801025..2c1edd1 100644 --- a/src/client/vertexai.rs +++ b/src/client/vertexai.rs @@ -1,7 +1,8 @@ use super::claude::{claude_build_body, claude_send_message, claude_send_message_streaming}; use super::{ catch_error, json_stream, message::*, patch_system_message, Client, CompletionDetails, - ExtraConfig, Model, ModelConfig, PromptKind, PromptType, SendData, SseHandler, VertexAIClient, + ExtraConfig, Model, ModelConfig, PromptAction, PromptKind, SendData, SseHandler, + VertexAIClient, }; use anyhow::{anyhow, bail, Context, Result}; @@ -30,7 +31,7 @@ impl VertexAIClient { config_get_fn!(project_id, get_project_id); config_get_fn!(location, get_location); - pub const PROMPTS: [PromptType<'static>; 2] = [ + pub const PROMPTS: [PromptAction<'static>; 2] = [ ("project_id", "Project ID", true, PromptKind::String), ("location", "Location", true, PromptKind::String), ]; diff --git a/src/config/mod.rs b/src/config/mod.rs index 0c89324..3dbd8ac 100644 --- a/src/config/mod.rs +++ b/src/config/mod.rs @@ -9,6 +9,7 @@ use self::session::{Session, TEMP_SESSION_NAME}; use crate::client::{ create_client_config, list_client_types, list_models, ClientConfig, Message, Model, SendData, + OPENAI_COMPATIBLE_PLATFORMS, }; use crate::render::{MarkdownRender, RenderOptions}; use crate::utils::{ @@ -21,6 +22,7 @@ use inquire::{Confirm, Select, Text}; use is_terminal::IsTerminal; use parking_lot::RwLock; use serde::Deserialize; +use serde_json::json; use std::collections::{HashMap, HashSet}; use std::{ env, @@ -126,12 +128,12 @@ impl Config { pub fn init(working_mode: WorkingMode) -> Result<Self> { let config_path = Self::config_file()?; - let client_type = env::var(get_env_name("client_type")).ok(); - if working_mode != WorkingMode::Command && client_type.is_none() && !config_path.exists() { + let platform = env::var(get_env_name("platform")).ok(); + if working_mode != WorkingMode::Command && platform.is_none() && !config_path.exists() { create_config_file(&config_path)?; } - let mut config = if client_type.is_some() { - Self::load_config_env(&client_type.unwrap())? + let mut config = if platform.is_some() { + Self::load_config_env(&platform.unwrap())? } else { Self::load_config_file(&config_path)? }; @@ -926,37 +928,45 @@ impl Config { fn load_config_file(config_path: &Path) -> Result<Self> { 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)?; + let config: Self = serde_yaml::from_str(&content).map_err(|err| { + let err_msg = err.to_string(); + let err_msg = if err_msg.starts_with(&format!("{}: ", CLIENTS_FIELD)) { + // location is incorrect, get rid of it + err_msg + .split_once(" at line") + .map(|(v, _)| { + format!("{v} (Sorry for being unable to provide an exact location)") + }) + .unwrap_or_else(|| "clients: invalid value".into()) + } else { + err_msg + }; + anyhow!("{err_msg}") + })?; + Ok(config) } - fn load_config_env(client_type: &str) -> Result<Self> { + fn load_config_env(platform: &str) -> Result<Self> { let model_id = match env::var(get_env_name("model_name")) { - Ok(model_name) => format!("{client_type}:{model_name}"), - Err(_) => client_type.to_string(), + Ok(model_name) => format!("{platform}:{model_name}"), + Err(_) => platform.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<Self> { - 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}") - } - })?; - + let is_openai_compatible = OPENAI_COMPATIBLE_PLATFORMS + .into_iter() + .any(|(name, _)| platform == name); + let client = if is_openai_compatible { + json!({ "type": "openai-compatible", "name": platform }) + } else { + json!({ "type": platform }) + }; + let config = json!({ + "model": model_id, + "save": false, + "clients": vec![client], + }); + let config = + serde_json::from_value(config).with_context(|| "Failed to load config from env")?; Ok(config) } |
