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 | |
| parent | df4440a2a049d26c61a540d3254cd885345fbd6c (diff) | |
| download | aichat-e0417e8d5bebe25476aaafa22a9ee23d9bd61457.tar.gz | |
feat: add `--sync-models` cli option (#1114)
Diffstat (limited to 'src')
| -rw-r--r-- | src/cli.rs | 3 | ||||
| -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 | ||||
| -rw-r--r-- | src/config/input.rs | 2 | ||||
| -rw-r--r-- | src/config/mod.rs | 118 | ||||
| -rw-r--r-- | src/main.rs | 6 | ||||
| -rw-r--r-- | src/utils/loader.rs | 2 | ||||
| -rw-r--r-- | src/utils/request.rs | 12 |
11 files changed, 132 insertions, 59 deletions
@@ -60,6 +60,9 @@ pub struct Cli { /// Display information #[clap(long)] pub info: bool, + /// Sync models updates + #[clap(long)] + pub sync_models: bool, /// List all available chat models #[clap(long)] pub list_models: bool, 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() { diff --git a/src/config/input.rs b/src/config/input.rs index 58c0065..e468f19 100644 --- a/src/config/input.rs +++ b/src/config/input.rs @@ -442,7 +442,7 @@ async fn load_documents( } for file_url in remote_urls { - let (contents, extension) = fetch(&loaders, &file_url, true) + let (contents, extension) = fetch_with_loaders(&loaders, &file_url, true) .await .with_context(|| format!("Failed to load url '{file_url}'"))?; if extension == MEDIA_URL_EXTENSION { diff --git a/src/config/mod.rs b/src/config/mod.rs index 8813930..341470f 100644 --- a/src/config/mod.rs +++ b/src/config/mod.rs @@ -12,7 +12,7 @@ use self::session::Session; use crate::client::{ create_client_config, list_client_types, list_models, ClientConfig, MessageContentToolCalls, - Model, ModelType, OPENAI_COMPATIBLE_PLATFORMS, + Model, ModelType, ProviderModels, OPENAI_COMPATIBLE_PROVIDERS, }; use crate::function::{FunctionDeclaration, Functions, ToolResult}; use crate::rag::Rag; @@ -24,7 +24,7 @@ use anyhow::{anyhow, bail, Context, Result}; use indexmap::IndexMap; use inquire::{list_option::ListOption, validator::Validation, Confirm, MultiSelect, Select, Text}; use parking_lot::RwLock; -use serde::Deserialize; +use serde::{Deserialize, Serialize}; use serde_json::json; use simplelog::LevelFilter; use std::collections::{HashMap, HashSet}; @@ -64,6 +64,8 @@ const CLIENTS_FIELD: &str = "clients"; const SERVE_ADDR: &str = "127.0.0.1:8000"; +const SYNC_MODELS_URL: &str = "https://cdn.jsdelivr.net/gh/sigoden/aichat/models.yaml"; + const SUMMARIZE_PROMPT: &str = "Summarize the discussion briefly in 200 words or less to use as a prompt for future context."; const SUMMARY_PROMPT: &str = "This is a summary of the chat history as a recap: "; @@ -140,6 +142,7 @@ pub struct Config { pub serve_addr: Option<String>, pub user_agent: Option<String>, pub save_shell_history: bool, + pub sync_models_url: Option<String>, pub clients: Vec<ClientConfig>, @@ -214,6 +217,7 @@ impl Default for Config { serve_addr: None, user_agent: None, save_shell_history: true, + sync_models_url: None, clients: vec![], @@ -240,9 +244,12 @@ impl Config { pub fn init(working_mode: WorkingMode, info_flag: bool) -> Result<Self> { let config_path = Self::config_file(); let mut config = if !config_path.exists() { - match env::var(get_env_name("platform")) { - Ok(v) => Self::load_dynamic(&v)?, - Err(_) => { + match env::var(get_env_name("provider")) + .ok() + .or_else(|| env::var(get_env_name("platform")).ok()) + { + Some(v) => Self::load_dynamic(&v)?, + None => { if *IS_STDOUT_TERMINAL { create_config_file(&config_path)?; } @@ -417,6 +424,10 @@ impl Config { } } + pub fn models_override_file() -> PathBuf { + Self::local_path("models-override.json") + } + pub fn state(&self) -> StateFlags { let mut flags = StateFlags::empty(); if let Some(session) = &self.session { @@ -1362,23 +1373,12 @@ impl Config { let temp_file = temp_file(&format!("-rag-{}", rag.name()), ".txt"); tokio::fs::write(&temp_file, &document_paths.join("\n")) .await - .with_context(|| { - format!( - "Failed to write current document paths to '{}'", - temp_file.display() - ) - })?; + .with_context(|| format!("Failed to write to '{}'", temp_file.display()))?; let editor = config.read().editor()?; edit_file(&editor, &temp_file)?; - let new_document_paths = - tokio::fs::read_to_string(&temp_file) - .await - .with_context(|| { - format!( - "Failed to read new document paths from '{}'", - temp_file.display() - ) - })?; + let new_document_paths = tokio::fs::read_to_string(&temp_file) + .await + .with_context(|| format!("Failed to read '{}'", temp_file.display()))?; let new_document_paths = new_document_paths .split('\n') .filter_map(|v| { @@ -1535,12 +1535,7 @@ impl Config { &agent_config_path, "# see https://github.com/sigoden/aichat/blob/main/config.agent.example.yaml\n", ) - .with_context(|| { - format!( - "Failed to write to agent config file at '{}'", - agent_config_path.display() - ) - })?; + .with_context(|| format!("Failed to write to '{}'", agent_config_path.display()))?; } let editor = self.editor()?; edit_file(&editor, &agent_config_path)?; @@ -1860,6 +1855,50 @@ impl Config { .collect() } + pub fn sync_models_url(&self) -> String { + self.sync_models_url + .clone() + .unwrap_or_else(|| SYNC_MODELS_URL.into()) + } + + pub async fn sync_models(url: &str, abort_signal: AbortSignal) -> Result<()> { + let content = abortable_run_with_spinner(fetch(url), "Fetching models.yaml", abort_signal) + .await + .with_context(|| format!("Failed to fetch '{url}'"))?; + println!("✓ Fetched '{url}'"); + let list = serde_yaml::from_str::<Vec<ProviderModels>>(&content) + .with_context(|| "Failed to parse models.yaml")?; + let models_override = ModelsOverride { + version: env!("CARGO_PKG_VERSION").to_string(), + list, + }; + let models_override_data = + serde_json::to_string_pretty(&models_override).with_context(|| "Failed to serde {}")?; + + let model_override_path = Self::models_override_file(); + ensure_parent_exists(&model_override_path)?; + std::fs::write(&model_override_path, models_override_data) + .with_context(|| format!("Failed to write to '{}'", model_override_path.display()))?; + println!("✓ Updated '{}'", model_override_path.display()); + Ok(()) + } + + pub fn loal_models_override() -> Result<Vec<ProviderModels>> { + let model_override_path = Self::models_override_file(); + let err = || { + format!( + "Failed to load models at '{}'", + model_override_path.display() + ) + }; + let content = read_to_string(&model_override_path).with_context(err)?; + let models_override: ModelsOverride = serde_json::from_str(&content).with_context(err)?; + if models_override.version != env!("CARGO_PKG_VERSION") { + bail!("Incompatible version") + } + Ok(models_override.list) + } + pub fn render_options(&self) -> Result<RenderOptions> { let theme = if self.highlight { let theme_mode = if self.light_theme { "light" } else { "dark" }; @@ -2175,17 +2214,17 @@ impl Config { } fn load_dynamic(model_id: &str) -> Result<Self> { - let platform = match model_id.split_once(':') { + let provider = match model_id.split_once(':') { Some((v, _)) => v, _ => model_id, }; - let is_openai_compatible = OPENAI_COMPATIBLE_PLATFORMS + let is_openai_compatible = OPENAI_COMPATIBLE_PROVIDERS .into_iter() - .any(|(name, _)| platform == name); + .any(|(name, _)| provider == name); let client = if is_openai_compatible { - json!({ "type": "openai-compatible", "name": platform }) + json!({ "type": "openai-compatible", "name": provider }) } else { - json!({ "type": platform }) + json!({ "type": provider }) }; let config = json!({ "model": model_id.to_string(), @@ -2323,6 +2362,9 @@ impl Config { if let Some(Some(v)) = read_env_bool(&get_env_name("save_shell_history")) { self.save_shell_history = v; } + if let Some(v) = read_env_value::<String>(&get_env_name("sync_models_url")) { + self.sync_models_url = v; + } } fn load_functions(&mut self) -> Result<()> { @@ -2502,6 +2544,12 @@ pub struct MacroVariable { pub default: Option<String>, } +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct ModelsOverride { + pub version: String, + pub list: Vec<ProviderModels>, +} + #[derive(Debug, Clone)] pub struct LastMessage { pub input: Input, @@ -2581,12 +2629,8 @@ fn create_config_file(config_path: &Path) -> Result<()> { ); ensure_parent_exists(config_path)?; - std::fs::write(config_path, config_data).with_context(|| { - format!( - "Failed to write to config file at '{}'", - config_path.display() - ) - })?; + std::fs::write(config_path, config_data) + .with_context(|| format!("Failed to write to '{}'", config_path.display()))?; #[cfg(unix)] { use std::os::unix::prelude::PermissionsExt; diff --git a/src/main.rs b/src/main.rs index e30efde..5251bad 100644 --- a/src/main.rs +++ b/src/main.rs @@ -46,6 +46,7 @@ async fn main() -> Result<()> { WorkingMode::Cmd }; let info_flag = cli.info + || cli.sync_models || cli.list_models || cli.list_roles || cli.list_agents @@ -64,6 +65,11 @@ async fn main() -> Result<()> { async fn run(config: GlobalConfig, cli: Cli, text: Option<String>) -> Result<()> { let abort_signal = create_abort_signal(); + if cli.sync_models { + let url = config.read().sync_models_url(); + return Config::sync_models(&url, abort_signal.clone()).await; + } + if cli.list_models { for model in list_models(&config.read(), ModelType::Chat) { println!("{}", model.id()); diff --git a/src/utils/loader.rs b/src/utils/loader.rs index 519563c..2ac671d 100644 --- a/src/utils/loader.rs +++ b/src/utils/loader.rs @@ -61,7 +61,7 @@ pub async fn load_file(loaders: &HashMap<String, String>, path: &str) -> Result< } pub async fn load_url(loaders: &HashMap<String, String>, path: &str) -> Result<LoadedDocument> { - let (contents, extension) = fetch(loaders, path, false).await?; + let (contents, extension) = fetch_with_loaders(loaders, path, false).await?; let mut metadata: DocumentMetadata = Default::default(); metadata.insert(EXTENSION_METADATA.into(), extension); Ok(LoadedDocument::new(path.into(), contents, metadata)) diff --git a/src/utils/request.rs b/src/utils/request.rs index 9f2804b..a479e0e 100644 --- a/src/utils/request.rs +++ b/src/utils/request.rs @@ -50,7 +50,17 @@ lazy_static::lazy_static! { static ref GITHUB_REPO_RE: Regex = Regex::new(r"^https://github\.com/([^/]+)/([^/]+)/tree/([^/]+)").unwrap(); } -pub async fn fetch( +pub async fn fetch(url: &str) -> Result<String> { + let client = match *CLIENT { + Ok(ref client) => client, + Err(ref err) => bail!("{err}"), + }; + let res = client.get(url).send().await?; + let output = res.text().await?; + Ok(output) +} + +pub async fn fetch_with_loaders( loaders: &HashMap<String, String>, path: &str, allow_media: bool, |
