summaryrefslogtreecommitdiffstats
path: root/src/client
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2025-01-22 20:51:10 +0800
committerGitHub <noreply@github.com>2025-01-22 20:51:10 +0800
commite0417e8d5bebe25476aaafa22a9ee23d9bd61457 (patch)
tree8edcd3aa9bcd1b012a3a429ad6240e6186caef36 /src/client
parentdf4440a2a049d26c61a540d3254cd885345fbd6c (diff)
downloadaichat-e0417e8d5bebe25476aaafa22a9ee23d9bd61457.tar.gz
feat: add `--sync-models` cli option (#1114)
Diffstat (limited to 'src/client')
-rw-r--r--src/client/common.rs13
-rw-r--r--src/client/macros.rs8
-rw-r--r--src/client/mod.rs2
-rw-r--r--src/client/model.rs23
-rw-r--r--src/client/openai_compatible.rs2
5 files changed, 29 insertions, 19 deletions
diff --git a/src/client/common.rs b/src/client/common.rs
index b4d01ce..80f585d 100644
--- a/src/client/common.rs
+++ b/src/client/common.rs
@@ -1,7 +1,7 @@
use super::*;
use crate::{
- config::{GlobalConfig, Input},
+ config::{Config, GlobalConfig, Input},
function::{eval_tool_calls, FunctionDeclaration, ToolCall, ToolResult},
render::render_stream,
utils::*,
@@ -20,7 +20,9 @@ use tokio::sync::mpsc::unbounded_channel;
const MODELS_YAML: &str = include_str!("../../models.yaml");
lazy_static::lazy_static! {
- pub static ref ALL_PREDEFINED_MODELS: Vec<PredefinedModels> = serde_yaml::from_str(MODELS_YAML).unwrap();
+ pub static ref ALL_PROVIDER_MODELS: Vec<ProviderModels> = {
+ Config::loal_models_override().ok().unwrap_or_else(|| serde_yaml::from_str(MODELS_YAML).unwrap())
+ };
static ref ESCAPE_SLASH_RE: Regex = Regex::new(r"(?<!\\)/").unwrap();
}
@@ -338,14 +340,15 @@ pub fn create_config(prompts: &[PromptAction], client: &str) -> Result<(String,
}
pub fn create_openai_compatible_client_config(client: &str) -> Result<Option<(String, Value)>> {
- let api_base = super::OPENAI_COMPATIBLE_PLATFORMS
+ let api_base = super::OPENAI_COMPATIBLE_PROVIDERS
.into_iter()
.find(|(name, _)| client == *name)
.map(|(_, api_base)| api_base)
.unwrap_or("http(s)://{API_ADDR}/v1");
let name = if client == OpenAICompatibleClient::NAME {
- prompt_input_string("Provider Name", true, None)?
+ let value = prompt_input_string("Provider Name", true, None)?;
+ value.replace(' ', "-")
} else {
client.to_string()
};
@@ -548,7 +551,7 @@ fn set_client_config(list: &[PromptAction], client_config: &mut Value, client: &
}
fn set_client_models_config(client_config: &mut Value, client: &str) -> Result<()> {
- if ALL_PREDEFINED_MODELS.iter().any(|v| v.platform == client) {
+ if ALL_PROVIDER_MODELS.iter().any(|v| v.provider == client) {
return Ok(());
}
diff --git a/src/client/macros.rs b/src/client/macros.rs
index a76e62b..97171db 100644
--- a/src/client/macros.rs
+++ b/src/client/macros.rs
@@ -52,10 +52,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() {
- if let Some(models) = $crate::client::ALL_PREDEFINED_MODELS.iter().find(|v| {
- v.platform == $name ||
+ if let Some(models) = $crate::client::ALL_PROVIDER_MODELS.iter().find(|v| {
+ v.provider == $name ||
($name == OpenAICompatibleClient::NAME
- && local_config.name.as_ref().map(|name| name.starts_with(&v.platform)).unwrap_or_default())
+ && local_config.name.as_ref().map(|name| name.starts_with(&v.provider)).unwrap_or_default())
}) {
return Model::from_config(client_name, &models.models);
}
@@ -83,7 +83,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::OPENAI_COMPATIBLE_PLATFORMS.iter().map(|(name, _)| *name));
+ client_types.extend($crate::client::OPENAI_COMPATIBLE_PROVIDERS.iter().map(|(name, _)| *name));
client_types
}
diff --git a/src/client/mod.rs b/src/client/mod.rs
index 3d8d4da..bf11107 100644
--- a/src/client/mod.rs
+++ b/src/client/mod.rs
@@ -34,7 +34,7 @@ register_client!(
(ernie, "ernie", ErnieConfig, ErnieClient),
);
-pub const OPENAI_COMPATIBLE_PLATFORMS: [(&str, &str); 22] = [
+pub const OPENAI_COMPATIBLE_PROVIDERS: [(&str, &str); 22] = [
("ai21", "https://api.ai21.com/studio/v1"),
(
"cloudflare",
diff --git a/src/client/model.rs b/src/client/model.rs
index 4b0457f..b562705 100644
--- a/src/client/model.rs
+++ b/src/client/model.rs
@@ -277,26 +277,33 @@ pub struct ModelData {
pub name: String,
#[serde(default = "default_model_type", rename = "type")]
pub model_type: String,
+ #[serde(skip_serializing_if = "Option::is_none")]
pub max_input_tokens: Option<usize>,
+ #[serde(skip_serializing_if = "Option::is_none")]
pub input_price: Option<f64>,
+ #[serde(skip_serializing_if = "Option::is_none")]
pub output_price: Option<f64>,
// chat-only properties
+ #[serde(skip_serializing_if = "Option::is_none")]
pub max_output_tokens: Option<isize>,
- #[serde(default)]
+ #[serde(default, skip_serializing_if = "std::ops::Not::not")]
pub require_max_tokens: bool,
- #[serde(default)]
+ #[serde(default, skip_serializing_if = "std::ops::Not::not")]
pub supports_vision: bool,
- #[serde(default)]
+ #[serde(default, skip_serializing_if = "std::ops::Not::not")]
pub supports_function_calling: bool,
- #[serde(default)]
+ #[serde(default, skip_serializing_if = "std::ops::Not::not")]
no_stream: bool,
- #[serde(default)]
+ #[serde(default, skip_serializing_if = "std::ops::Not::not")]
no_system_message: bool,
// embedding-only properties
+ #[serde(skip_serializing_if = "Option::is_none")]
pub max_tokens_per_chunk: Option<usize>,
+ #[serde(skip_serializing_if = "Option::is_none")]
pub default_chunk_size: Option<usize>,
+ #[serde(skip_serializing_if = "Option::is_none")]
pub max_batch_size: Option<usize>,
}
@@ -310,9 +317,9 @@ impl ModelData {
}
}
-#[derive(Debug, Clone, Deserialize)]
-pub struct PredefinedModels {
- pub platform: String,
+#[derive(Debug, Clone, Serialize, Deserialize)]
+pub struct ProviderModels {
+ pub provider: String,
pub models: Vec<ModelData>,
}
diff --git a/src/client/openai_compatible.rs b/src/client/openai_compatible.rs
index 18acafb..ce1eea3 100644
--- a/src/client/openai_compatible.rs
+++ b/src/client/openai_compatible.rs
@@ -96,7 +96,7 @@ fn get_api_base_ext(self_: &OpenAICompatibleClient) -> Result<String> {
let api_base = match self_.get_api_base() {
Ok(v) => v,
Err(err) => {
- match OPENAI_COMPATIBLE_PLATFORMS
+ match OPENAI_COMPATIBLE_PROVIDERS
.into_iter()
.find_map(|(name, api_base)| {
if name == self_.model.client_name() {