summaryrefslogtreecommitdiffstats
path: root/src
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2024-04-30 12:52:58 +0800
committerGitHub <noreply@github.com>2024-04-30 12:52:58 +0800
commit8dba46becfbcc4669867db4487aa9b91c51e4afa (patch)
tree7b67d556751d9f5fed380b50b818ddb8680c8f01 /src
parent8a65337d590729f96a5f3c0b35dc5a08fae5bf94 (diff)
downloadaichat-8dba46becfbcc4669867db4487aa9b91c51e4afa.tar.gz
feat: openai-compatible platforms share the same client (#469)
Diffstat (limited to 'src')
-rw-r--r--src/client/azure_openai.rs6
-rw-r--r--src/client/bedrock.rs6
-rw-r--r--src/client/claude.rs4
-rw-r--r--src/client/cloudflare.rs4
-rw-r--r--src/client/cohere.rs4
-rw-r--r--src/client/common.rs114
-rw-r--r--src/client/ernie.rs4
-rw-r--r--src/client/gemini.rs4
-rw-r--r--src/client/groq.rs1
-rw-r--r--src/client/mistral.rs1
-rw-r--r--src/client/mod.rs34
-rw-r--r--src/client/model.rs6
-rw-r--r--src/client/moonshot.rs1
-rw-r--r--src/client/ollama.rs4
-rw-r--r--src/client/openai.rs4
-rw-r--r--src/client/openai_compatible.rs25
-rw-r--r--src/client/perplexity.rs5
-rw-r--r--src/client/qianwen.rs4
-rw-r--r--src/client/replicate.rs6
-rw-r--r--src/client/vertexai.rs5
-rw-r--r--src/config/mod.rs70
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)
}