summaryrefslogtreecommitdiffstats
path: root/src
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2024-04-28 10:14:12 +0800
committerGitHub <noreply@github.com>2024-04-28 10:14:12 +0800
commitd6df1e84a7d4fece6c9107de867a1c4d7d2fee56 (patch)
tree5557d35d080e665370a7663d1317df14700d0ea4 /src
parenteac01fb129ce5a06e00cc2cc24e31b40e81eb66b (diff)
downloadaichat-d6df1e84a7d4fece6c9107de867a1c4d7d2fee56.tar.gz
refactor: extract prelude models to models.yaml (#451)
Diffstat (limited to 'src')
-rw-r--r--src/client/azure_openai.rs1
-rw-r--r--src/client/claude.rs10
-rw-r--r--src/client/cohere.rs8
-rw-r--r--src/client/common.rs85
-rw-r--r--src/client/ernie.rs24
-rw-r--r--src/client/gemini.rs9
-rw-r--r--src/client/groq.rs8
-rw-r--r--src/client/mistral.rs9
-rw-r--r--src/client/mod.rs26
-rw-r--r--src/client/model.rs69
-rw-r--r--src/client/moonshot.rs6
-rw-r--r--src/client/ollama.rs1
-rw-r--r--src/client/openai.rs14
-rw-r--r--src/client/openai_compatible.rs1
-rw-r--r--src/client/perplexity.rs14
-rw-r--r--src/client/qianwen.rs13
-rw-r--r--src/client/vertexai.rs9
17 files changed, 82 insertions, 225 deletions
diff --git a/src/client/azure_openai.rs b/src/client/azure_openai.rs
index ab8c995..0ec0c14 100644
--- a/src/client/azure_openai.rs
+++ b/src/client/azure_openai.rs
@@ -17,7 +17,6 @@ pub struct AzureOpenAIConfig {
}
impl AzureOpenAIClient {
- list_models_fn!(AzureOpenAIConfig);
config_get_fn!(api_base, get_api_base);
config_get_fn!(api_key, get_api_key);
diff --git a/src/client/claude.rs b/src/client/claude.rs
index d3fc68d..e4081e7 100644
--- a/src/client/claude.rs
+++ b/src/client/claude.rs
@@ -24,16 +24,6 @@ pub struct ClaudeConfig {
}
impl ClaudeClient {
- list_models_fn!(
- ClaudeConfig,
- [
- // https://docs.anthropic.com/claude/docs/models-overview
- ("claude-3-opus-20240229", "text,vision", 200000, 4096),
- ("claude-3-sonnet-20240229", "text,vision", 200000, 4096),
- ("claude-3-haiku-20240307", "text,vision", 200000, 4096),
- ]
- );
-
config_get_fn!(api_key, get_api_key);
pub const PROMPTS: [PromptType<'static>; 1] =
diff --git a/src/client/cohere.rs b/src/client/cohere.rs
index 3f97ece..d99276c 100644
--- a/src/client/cohere.rs
+++ b/src/client/cohere.rs
@@ -22,14 +22,6 @@ pub struct CohereConfig {
}
impl CohereClient {
- list_models_fn!(
- CohereConfig,
- [
- // https://docs.cohere.com/docs/command-r
- ("command-r", "text", 128000),
- ("command-r-plus", "text", 128000),
- ]
- );
config_get_fn!(api_key, get_api_key);
pub const PROMPTS: [PromptType<'static>; 1] =
diff --git a/src/client/common.rs b/src/client/common.rs
index c543373..695245a 100644
--- a/src/client/common.rs
+++ b/src/client/common.rs
@@ -1,4 +1,4 @@
-use super::{openai::OpenAIConfig, ClientConfig, Message, Model, ReplyHandler};
+use super::{openai::OpenAIConfig, ClientConfig, ClientModel, Message, Model, ReplyHandler};
use crate::{
config::{GlobalConfig, Input},
@@ -9,12 +9,19 @@ use crate::{
use anyhow::{bail, Context, Result};
use async_trait::async_trait;
use futures_util::{Stream, StreamExt};
+use lazy_static::lazy_static;
use reqwest::{Client as ReqwestClient, ClientBuilder, Proxy, RequestBuilder};
use serde::Deserialize;
use serde_json::{json, Value};
use std::{env, future::Future, time::Duration};
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();
+}
+
#[macro_export]
macro_rules! register_client {
(
@@ -38,6 +45,17 @@ macro_rules! register_client {
Unknown,
}
+ #[derive(Debug, Clone, serde::Deserialize)]
+ #[serde(tag = "type")]
+ pub enum ClientModel {
+ $(
+ #[serde(rename = $name)]
+ $config { models: Vec<ModelConfig> },
+ )+
+ #[serde(other)]
+ Unknown,
+ }
+
$(
#[derive(Debug)]
@@ -68,6 +86,23 @@ 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);
+ }
+ _ => {}
+ }
+ }
+ vec![]
+ } else {
+ Model::from_config(client_name, &local_config.models)
+ }
+ }
+
pub fn name(config: &$config) -> &str {
config.name.as_deref().unwrap_or(Self::NAME)
}
@@ -131,10 +166,9 @@ macro_rules! openai_compatible_client {
$config:ident,
$client:ident,
$api_base:literal,
- [$(($name:literal, $capabilities:literal, $max_input_tokens:literal $(, $max_output_tokens:literal)? )),+$(,)?]
) => {
use $crate::client::openai::openai_build_body;
- use $crate::client::{ExtraConfig, $client, Model, ModelConfig, PromptType, SendData};
+ use $crate::client::{$client, ExtraConfig, Model, ModelConfig, PromptType, SendData};
use $crate::utils::PromptKind;
@@ -159,22 +193,17 @@ macro_rules! openai_compatible_client {
$crate::client::openai::openai_send_message_streaming
);
-
impl $client {
- list_models_fn!(
- $config,
- [
- $(
- ($name, $capabilities, $max_input_tokens $(, $max_output_tokens)?),
- )+
- ]
- );
config_get_fn!(api_key, get_api_key);
pub const PROMPTS: [PromptType<'static>; 1] =
[("api_key", "API Key:", false, PromptKind::String)];
- fn request_builder(&self, client: &ReqwestClient, data: SendData) -> Result<RequestBuilder> {
+ 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);
@@ -191,8 +220,7 @@ macro_rules! openai_compatible_client {
Ok(builder)
}
}
-
- }
+ };
}
#[macro_export]
@@ -270,33 +298,6 @@ macro_rules! config_get_fn {
}
#[macro_export]
-macro_rules! list_models_fn {
- ($config:ident) => {
- pub fn list_models(local_config: &$config) -> Vec<Model> {
- let client_name = Self::name(local_config);
- Model::from_config(client_name, &local_config.models)
- }
- };
- ($config:ident, [$(($name:literal, $capabilities:literal, $max_input_tokens:literal $(, $max_output_tokens:literal)? )),+$(,)?]) => {
- pub fn list_models(local_config: &$config) -> Vec<Model> {
- let client_name = Self::name(local_config);
- if local_config.models.is_empty() {
- vec![
- $(
- Model::new(client_name, $name)
- .set_capabilities($capabilities.into())
- .set_max_input_tokens(Some($max_input_tokens))
- $(.set_max_output_tokens(Some($max_output_tokens)))?
- ),+
- ]
- } else {
- Model::from_config(client_name, &local_config.models)
- }
- }
- };
-}
-
-#[macro_export]
macro_rules! unsupported_model {
($name:expr) => {
anyhow::bail!("Unsupported model '{}'", $name)
diff --git a/src/client/ernie.rs b/src/client/ernie.rs
index 5cc546c..7695061 100644
--- a/src/client/ernie.rs
+++ b/src/client/ernie.rs
@@ -31,19 +31,6 @@ pub struct ErnieConfig {
}
impl ErnieClient {
- list_models_fn!(
- ErnieConfig,
- [
- // https://cloud.baidu.com/doc/WENXINWORKSHOP/s/clntwmv7t
- ("ernie-4.0-8k", "text", 5120, 2048),
- ("ernie-3.5-8k", "text", 5120, 2048),
- ("ernie-3.5-4k", "text", 2048, 2048),
- ("ernie-speed-8k", "text", 7168, 2048),
- ("ernie-speed-128k", "text", 124000, 4096),
- ("ernie-lite-8k", "text", 7168, 2048),
- ("ernie-tiny-8k", "text", 7168, 2048),
- ]
- );
pub const PROMPTS: [PromptType<'static>; 2] = [
("api_key", "API Key:", true, PromptKind::String),
@@ -53,16 +40,9 @@ impl ErnieClient {
fn request_builder(&self, client: &ReqwestClient, data: SendData) -> Result<RequestBuilder> {
let body = build_body(data, &self.model);
- let endpoint = match self.model.name.as_str() {
- "ernie-4.0-8k" => "completions_pro",
- "ernie-3.5-8k" => "ernie-3.5-8k-0205",
- "ernie-3.5-4k" => "ernie-3.5-4k-0205",
- "ernie-speed-8k" => "ernie_speed",
- _ => &self.model.name,
- };
-
let url = format!(
- "{API_BASE}/wenxinworkshop/chat/{endpoint}?access_token={}",
+ "{API_BASE}/wenxinworkshop/chat/{}?access_token={}",
+ &self.model.name,
unsafe { &ACCESS_TOKEN.0 }
);
diff --git a/src/client/gemini.rs b/src/client/gemini.rs
index 4b6f053..7f51739 100644
--- a/src/client/gemini.rs
+++ b/src/client/gemini.rs
@@ -20,15 +20,6 @@ pub struct GeminiConfig {
}
impl GeminiClient {
- list_models_fn!(
- GeminiConfig,
- [
- // https://ai.google.dev/models/gemini
- ("gemini-1.0-pro-latest", "text", 30720),
- ("gemini-1.0-pro-vision-latest", "text,vision", 12288),
- ("gemini-1.5-pro-latest", "text,vision", 1048576),
- ]
- );
config_get_fn!(api_key, get_api_key);
pub const PROMPTS: [PromptType<'static>; 1] =
diff --git a/src/client/groq.rs b/src/client/groq.rs
index ba76b68..24abc6e 100644
--- a/src/client/groq.rs
+++ b/src/client/groq.rs
@@ -2,12 +2,4 @@ openai_compatible_client!(
GroqConfig,
GroqClient,
"https://api.groq.com/openai/v1",
- [
- // https://console.groq.com/docs/models
- ("llama3-8b-8192", "text", 8192),
- ("llama3-70b-8192", "text", 8192),
- ("llama2-70b-4096", "text", 4096),
- ("mixtral-8x7b-32768", "text", 32768),
- ("gemma-7b-it", "text", 8192),
- ]
);
diff --git a/src/client/mistral.rs b/src/client/mistral.rs
index 481ca1f..7279d2b 100644
--- a/src/client/mistral.rs
+++ b/src/client/mistral.rs
@@ -2,13 +2,4 @@ openai_compatible_client!(
MistralConfig,
MistralClient,
"https://api.mistral.ai/v1",
- [
- // https://docs.mistral.ai/platform/endpoints/
- ("open-mistral-7b", "text", 32000),
- ("open-mixtral-8x7b", "text", 32000),
- ("open-mixtral-8x22b", "text", 64000),
- ("mistral-small-latest", "text", 32000),
- ("mistral-medium-latest", "text", 32000),
- ("mistral-large-latest", "text", 32000),
- ]
);
diff --git a/src/client/mod.rs b/src/client/mod.rs
index f7d448f..07520c2 100644
--- a/src/client/mod.rs
+++ b/src/client/mod.rs
@@ -11,26 +11,26 @@ pub use reply_handler::*;
register_client!(
(openai, "openai", OpenAIConfig, OpenAIClient),
- (
- azure_openai,
- "azure-openai",
- AzureOpenAIConfig,
- AzureOpenAIClient
- ),
- (
- openai_compatible,
- "openai-compatible",
- OpenAICompatibleConfig,
- OpenAICompatibleClient
- ),
(gemini, "gemini", GeminiConfig, GeminiClient),
- (vertexai, "vertexai", VertexAIConfig, VertexAIClient),
(claude, "claude", ClaudeConfig, ClaudeClient),
(mistral, "mistral", MistralConfig, MistralClient),
(cohere, "cohere", CohereConfig, CohereClient),
(perplexity, "perplexity", PerplexityConfig, PerplexityClient),
(groq, "groq", GroqConfig, GroqClient),
+ (
+ openai_compatible,
+ "openai-compatible",
+ OpenAICompatibleConfig,
+ OpenAICompatibleClient
+ ),
(ollama, "ollama", OllamaConfig, OllamaClient),
+ (
+ azure_openai,
+ "azure-openai",
+ AzureOpenAIConfig,
+ AzureOpenAIClient
+ ),
+ (vertexai, "vertexai", VertexAIConfig, VertexAIClient),
(ernie, "ernie", ErnieConfig, ErnieClient),
(qianwen, "qianwen", QianwenConfig, QianwenClient),
(moonshot, "moonshot", MoonshotConfig, MoonshotClient),
diff --git a/src/client/model.rs b/src/client/model.rs
index 3f2cbdd..459d94e 100644
--- a/src/client/model.rs
+++ b/src/client/model.rs
@@ -3,7 +3,7 @@ use super::message::{Message, MessageContent};
use crate::utils::count_tokens;
use anyhow::{bail, Result};
-use serde::{Deserialize, Deserializer};
+use serde::Deserialize;
const PER_MESSAGES_TOKENS: usize = 5;
const BASIS_TOKENS: usize = 2;
@@ -41,10 +41,10 @@ impl Model {
.iter()
.map(|v| {
Model::new(client_name, &v.name)
- .set_capabilities(v.capabilities)
.set_max_input_tokens(v.max_input_tokens)
.set_max_output_tokens(v.max_output_tokens)
- .set_extra_fields(v.extra_fields.clone())
+ .set_supports_vision(v.supports_vision)
+ .set_extra_fields(&v.extra_fields)
})
.collect()
}
@@ -84,19 +84,6 @@ impl Model {
format!("{}:{}", self.client_name, self.name)
}
- pub fn set_capabilities(mut self, capabilities: ModelCapabilities) -> Self {
- self.capabilities = capabilities;
- self
- }
-
- pub fn set_extra_fields(
- mut self,
- extra_fields: Option<serde_json::Map<String, serde_json::Value>>,
- ) -> Self {
- self.extra_fields = extra_fields;
- self
- }
-
pub fn set_max_input_tokens(mut self, max_input_tokens: Option<usize>) -> Self {
match max_input_tokens {
None | Some(0) => self.max_input_tokens = None,
@@ -113,6 +100,23 @@ impl Model {
self
}
+ pub fn set_supports_vision(mut self, supports_vision: bool) -> Self {
+ if supports_vision {
+ self.capabilities |= ModelCapabilities::Vision;
+ } else {
+ self.capabilities &= !ModelCapabilities::Vision;
+ }
+ self
+ }
+
+ pub fn set_extra_fields(
+ mut self,
+ extra_fields: &Option<serde_json::Map<String, serde_json::Value>>,
+ ) -> Self {
+ self.extra_fields = extra_fields.clone();
+ self
+ }
+
pub fn messages_tokens(&self, messages: &[Message]) -> usize {
messages
.iter()
@@ -174,10 +178,11 @@ pub struct ModelConfig {
pub name: String,
pub max_input_tokens: Option<usize>,
pub max_output_tokens: Option<isize>,
+ pub input_price: Option<f64>,
+ pub output_price: Option<f64>,
+ #[serde(default)]
+ pub supports_vision: bool,
pub extra_fields: Option<serde_json::Map<String, serde_json::Value>>,
- #[serde(deserialize_with = "deserialize_capabilities")]
- #[serde(default = "default_capabilities")]
- pub capabilities: ModelCapabilities,
}
bitflags::bitflags! {
@@ -187,29 +192,3 @@ bitflags::bitflags! {
const Vision = 0b00000010;
}
}
-
-impl From<&str> for ModelCapabilities {
- fn from(value: &str) -> Self {
- let value = if value.is_empty() { "text" } else { value };
- let mut output = ModelCapabilities::empty();
- if value.contains("text") {
- output |= ModelCapabilities::Text;
- }
- if value.contains("vision") {
- output |= ModelCapabilities::Vision;
- }
- output
- }
-}
-
-fn deserialize_capabilities<'de, D>(deserializer: D) -> Result<ModelCapabilities, D::Error>
-where
- D: Deserializer<'de>,
-{
- let value: String = Deserialize::deserialize(deserializer)?;
- Ok(value.as_str().into())
-}
-
-fn default_capabilities() -> ModelCapabilities {
- ModelCapabilities::Text
-}
diff --git a/src/client/moonshot.rs b/src/client/moonshot.rs
index 6a8303e..471b2b1 100644
--- a/src/client/moonshot.rs
+++ b/src/client/moonshot.rs
@@ -2,10 +2,4 @@ openai_compatible_client!(
MoonshotConfig,
MoonshotClient,
"https://api.moonshot.cn/v1",
- [
- // https://platform.moonshot.cn/docs/intro#%E6%A8%A1%E5%9E%8B%E5%88%97%E8%A1%A8
- ("moonshot-v1-8k", "text", 8000),
- ("moonshot-v1-32k", "text", 32000),
- ("moonshot-v1-128k", "text", 128000),
- ]
);
diff --git a/src/client/ollama.rs b/src/client/ollama.rs
index a801242..e869e87 100644
--- a/src/client/ollama.rs
+++ b/src/client/ollama.rs
@@ -22,7 +22,6 @@ pub struct OllamaConfig {
}
impl OllamaClient {
- list_models_fn!(OllamaConfig);
config_get_fn!(api_key, get_api_key);
pub const PROMPTS: [PromptType<'static>; 4] = [
diff --git a/src/client/openai.rs b/src/client/openai.rs
index a0abc25..54197ac 100644
--- a/src/client/openai.rs
+++ b/src/client/openai.rs
@@ -25,20 +25,6 @@ pub struct OpenAIConfig {
}
impl OpenAIClient {
- list_models_fn!(
- OpenAIConfig,
- [
- // https://platform.openai.com/docs/models
- ("gpt-3.5-turbo", "text", 16385),
- ("gpt-3.5-turbo-1106", "text", 16385),
- ("gpt-4-turbo", "text,vision", 128000),
- ("gpt-4-turbo-preview", "text", 128000),
- ("gpt-4-1106-preview", "text", 128000),
- ("gpt-4-vision-preview", "text,vision", 128000, 4096),
- ("gpt-4", "text", 8192),
- ("gpt-4-32k", "text", 32768),
- ]
- );
config_get_fn!(api_key, get_api_key);
config_get_fn!(api_base, get_api_base);
diff --git a/src/client/openai_compatible.rs b/src/client/openai_compatible.rs
index 1626b17..304e748 100644
--- a/src/client/openai_compatible.rs
+++ b/src/client/openai_compatible.rs
@@ -18,7 +18,6 @@ pub struct OpenAICompatibleConfig {
}
impl OpenAICompatibleClient {
- list_models_fn!(OpenAICompatibleConfig);
config_get_fn!(api_key, get_api_key);
pub const PROMPTS: [PromptType<'static>; 5] = [
diff --git a/src/client/perplexity.rs b/src/client/perplexity.rs
index c8bb6f0..df3d462 100644
--- a/src/client/perplexity.rs
+++ b/src/client/perplexity.rs
@@ -2,18 +2,4 @@ openai_compatible_client!(
PerplexityConfig,
PerplexityClient,
"https://api.perplexity.ai",
- [
- // https://docs.perplexity.ai/docs/model-cards
- ("sonar-small-chat", "text", 16384),
- ("sonar-small-online", "text", 12000),
- ("sonar-medium-chat", "text", 16384),
- ("sonar-medium-online", "text", 12000),
-
- ("llama-3-8b-instruct", "text", 8192),
- ("llama-3-70b-instruct", "text", 8192),
- ("codellama-70b-instruct", "text", 16384),
- ("mistral-7b-instruct", "text", 16384),
- ("mixtral-8x7b-instruct", "text", 16384),
- ("mixtral-8x22b-instruct", "text", 16384),
- ]
);
diff --git a/src/client/qianwen.rs b/src/client/qianwen.rs
index 386716f..64d27d4 100644
--- a/src/client/qianwen.rs
+++ b/src/client/qianwen.rs
@@ -34,19 +34,6 @@ pub struct QianwenConfig {
}
impl QianwenClient {
- list_models_fn!(
- QianwenConfig,
- [
- // https://help.aliyun.com/zh/dashscope/developer-reference/api-details
- ("qwen-turbo", "text", 6000),
- ("qwen-plus", "text", 30000),
- ("qwen-max", "text", 6000),
- ("qwen-max-longcontext", "text", 28000),
- // https://help.aliyun.com/zh/dashscope/developer-reference/tongyi-qianwen-vl-plus-api
- ("qwen-vl-plus", "text,vision", 0),
- ("qwen-vl-max", "text,vision", 0),
- ]
- );
config_get_fn!(api_key, get_api_key);
pub const PROMPTS: [PromptType<'static>; 1] =
diff --git a/src/client/vertexai.rs b/src/client/vertexai.rs
index 566067f..afb1136 100644
--- a/src/client/vertexai.rs
+++ b/src/client/vertexai.rs
@@ -27,15 +27,6 @@ pub struct VertexAIConfig {
}
impl VertexAIClient {
- list_models_fn!(
- VertexAIConfig,
- [
- // https://cloud.google.com/vertex-ai/generative-ai/docs/learn/models
- ("gemini-1.0-pro", "text", 24568),
- ("gemini-1.0-pro-vision", "text,vision", 14336),
- ("gemini-1.5-pro-preview-0409", "text,vision", 1000000),
- ]
- );
config_get_fn!(api_base, get_api_base);
pub const PROMPTS: [PromptType<'static>; 1] =