summaryrefslogtreecommitdiffstats
path: root/src/client/common.rs
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2024-04-30 12:52:58 +0800
committerGitHub <noreply@github.com>2024-04-30 12:52:58 +0800
commit8dba46becfbcc4669867db4487aa9b91c51e4afa (patch)
tree7b67d556751d9f5fed380b50b818ddb8680c8f01 /src/client/common.rs
parent8a65337d590729f96a5f3c0b35dc5a08fae5bf94 (diff)
downloadaichat-8dba46becfbcc4669867db4487aa9b91c51e4afa.tar.gz
feat: openai-compatible platforms share the same client (#469)
Diffstat (limited to 'src/client/common.rs')
-rw-r--r--src/client/common.rs114
1 files changed, 29 insertions, 85 deletions
diff --git a/src/client/common.rs b/src/client/common.rs
index ee190d5..91d6bca 100644
--- a/src/client/common.rs
+++ b/src/client/common.rs
@@ -1,4 +1,4 @@
-use super::{openai::OpenAIConfig, ClientConfig, ClientModel, Message, Model, SseHandler};
+use super::{openai::OpenAIConfig, BuiltinModels, ClientConfig, Message, Model, SseHandler};
use crate::{
config::{GlobalConfig, Input},
@@ -20,7 +20,8 @@ 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();
+ pub static ref ALL_CLIENT_MODELS: Vec<BuiltinModels> =
+ serde_yaml::from_str(MODELS_YAML).unwrap();
}
#[macro_export]
@@ -90,13 +91,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() {
- for model in $crate::client::CLIENT_MODELS.iter() {
- match model {
- $crate::client::ClientModel::$config { models } => {
- return Model::from_config(client_name, models);
- }
- _ => {}
- }
+ if let Some(client_models) = $crate::client::ALL_CLIENT_MODELS.iter().find(|v| {
+ v.platform == $name || ($name == "openai-compatible" && local_config.name.as_deref() == Some(&v.platform))
+ }) {
+ return Model::from_config(client_name, &client_models.models);
}
vec![]
} else {
@@ -135,7 +133,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::KNOWN_OPENAI_COMPATIBLE_PLATFORMS.iter().map(|(name, _)| *name));
+ client_types.extend($crate::client::OPENAI_COMPATIBLE_PLATFORMS.iter().map(|(name, _)| *name));
client_types
}
@@ -171,69 +169,6 @@ macro_rules! register_client {
}
#[macro_export]
-macro_rules! openai_compatible_client {
- (
- $config:ident,
- $client:ident,
- $api_base:literal,
- ) => {
- use $crate::client::openai::openai_build_body;
- use $crate::client::{$client, ExtraConfig, Model, ModelConfig, PromptType, SendData};
-
- use $crate::utils::PromptKind;
-
- use anyhow::Result;
- use reqwest::{Client as ReqwestClient, RequestBuilder};
- use serde::Deserialize;
-
- const API_BASE: &str = $api_base;
-
- #[derive(Debug, Clone, Deserialize)]
- pub struct $config {
- pub name: Option<String>,
- pub api_key: Option<String>,
- #[serde(default)]
- pub models: Vec<ModelConfig>,
- pub extra: Option<ExtraConfig>,
- }
-
- impl_client_trait!(
- $client,
- $crate::client::openai::openai_send_message,
- $crate::client::openai::openai_send_message_streaming
- );
-
- impl $client {
- config_get_fn!(api_key, get_api_key);
-
- pub const PROMPTS: [PromptType<'static>; 1] =
- [("api_key", "API Key:", true, PromptKind::String)];
-
- 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);
-
- let url = format!("{API_BASE}/chat/completions");
-
- debug!("Request: {url} {body}");
-
- let mut builder = client.post(url).json(&body);
- if let Some(api_key) = api_key {
- builder = builder.bearer_auth(api_key);
- }
-
- Ok(builder)
- }
- }
- };
-}
-
-#[macro_export]
macro_rules! client_common_fns {
() => {
fn config(
@@ -437,36 +372,45 @@ pub struct CompletionDetails {
pub output_tokens: Option<u64>,
}
-pub type PromptType<'a> = (&'a str, &'a str, bool, PromptKind);
+pub type PromptAction<'a> = (&'a str, &'a str, bool, PromptKind);
-pub fn create_config(list: &[PromptType], client: &str) -> Result<(String, Value)> {
+pub fn create_config(prompts: &[PromptAction], client: &str) -> Result<(String, Value)> {
let mut config = json!({
"type": client,
});
let mut model = client.to_string();
- set_client_config_values(list, &mut model, &mut config)?;
+ set_client_config_values(prompts, &mut model, &mut config)?;
let clients = json!(vec![config]);
Ok((model, clients))
}
pub fn create_openai_compatible_client_config(client: &str) -> Result<Option<(String, Value)>> {
- match super::KNOWN_OPENAI_COMPATIBLE_PLATFORMS
+ match super::OPENAI_COMPATIBLE_PLATFORMS
.iter()
.find(|(name, _)| client == *name)
{
None => Ok(None),
- Some((name, api_base)) => {
+ Some((name, _)) => {
let mut config = json!({
"type": "openai-compatible",
"name": name,
- "api_base": api_base,
});
+ let prompts = if ALL_CLIENT_MODELS.iter().any(|v| &v.platform == name) {
+ vec![("api_key", "API Key:", false, PromptKind::String)]
+ } else {
+ vec![
+ ("api_key", "API Key:", false, PromptKind::String),
+ ("models[].name", "Model Name:", true, PromptKind::String),
+ (
+ "models[].max_input_tokens",
+ "Max Input Tokens:",
+ false,
+ PromptKind::Integer,
+ ),
+ ]
+ };
let mut model = client.to_string();
- set_client_config_values(
- &super::KNOWN_OPENAI_COMPATIBLE_PROMPTS,
- &mut model,
- &mut config,
- )?;
+ set_client_config_values(&prompts, &mut model, &mut config)?;
let clients = json!(vec![config]);
Ok(Some((model, clients)))
}
@@ -683,7 +627,7 @@ where
}
fn set_client_config_values(
- list: &[PromptType],
+ list: &[PromptAction],
model: &mut String,
client_config: &mut Value,
) -> Result<()> {