diff options
| author | sigoden <sigoden@gmail.com> | 2025-01-22 20:51:10 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2025-01-22 20:51:10 +0800 |
| commit | e0417e8d5bebe25476aaafa22a9ee23d9bd61457 (patch) | |
| tree | 8edcd3aa9bcd1b012a3a429ad6240e6186caef36 /src/client | |
| parent | df4440a2a049d26c61a540d3254cd885345fbd6c (diff) | |
| download | aichat-e0417e8d5bebe25476aaafa22a9ee23d9bd61457.tar.gz | |
feat: add `--sync-models` cli option (#1114)
Diffstat (limited to 'src/client')
| -rw-r--r-- | src/client/common.rs | 13 | ||||
| -rw-r--r-- | src/client/macros.rs | 8 | ||||
| -rw-r--r-- | src/client/mod.rs | 2 | ||||
| -rw-r--r-- | src/client/model.rs | 23 | ||||
| -rw-r--r-- | src/client/openai_compatible.rs | 2 |
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() { |
