summaryrefslogtreecommitdiffstats
path: root/src
diff options
context:
space:
mode:
Diffstat (limited to 'src')
-rw-r--r--src/cli.rs3
-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
-rw-r--r--src/config/input.rs2
-rw-r--r--src/config/mod.rs118
-rw-r--r--src/main.rs6
-rw-r--r--src/utils/loader.rs2
-rw-r--r--src/utils/request.rs12
11 files changed, 132 insertions, 59 deletions
diff --git a/src/cli.rs b/src/cli.rs
index de88776..3204c58 100644
--- a/src/cli.rs
+++ b/src/cli.rs
@@ -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,