summaryrefslogtreecommitdiffstats
path: root/src/client/common.rs
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2024-04-25 10:15:54 +0800
committerGitHub <noreply@github.com>2024-04-25 10:15:54 +0800
commit4db9b309803796bc5f996d0b3713344eb44207ec (patch)
treea0ba8f3a2da32e5392e22272c641e3939de49be6 /src/client/common.rs
parent8dacea4deb0fe67c0cdbe15bd47e6e50c1dd12a1 (diff)
downloadaichat-4db9b309803796bc5f996d0b3713344eb44207ec.tar.gz
refactor: rewrite list models of all clients (#436)
Diffstat (limited to 'src/client/common.rs')
-rw-r--r--src/client/common.rs34
1 files changed, 10 insertions, 24 deletions
diff --git a/src/client/common.rs b/src/client/common.rs
index 593140d..e003e7e 100644
--- a/src/client/common.rs
+++ b/src/client/common.rs
@@ -1,4 +1,4 @@
-use super::{openai::OpenAIConfig, ClientConfig, Message, MessageContent, Model, ReplyHandler};
+use super::{openai::OpenAIConfig, ClientConfig, Message, Model, ReplyHandler};
use crate::{
config::{GlobalConfig, Input},
@@ -205,11 +205,18 @@ macro_rules! list_models_fn {
Model::from_config(client_name, &local_config.models)
}
};
- ($config:ident, $models:expr) => {
+ ($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() {
- Model::from_static(client_name, $models)
+ 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)
}
@@ -402,27 +409,6 @@ where
Ok(())
}
-pub fn patch_system_message(messages: &mut Vec<Message>) {
- if messages[0].role.is_system() {
- let system_message = messages.remove(0);
- if let (Some(message), MessageContent::Text(system_text)) =
- (messages.get_mut(0), system_message.content)
- {
- if let MessageContent::Text(text) = message.content.clone() {
- message.content = MessageContent::Text(format!("{}\n\n{}", system_text, text))
- }
- }
- }
-}
-
-pub fn extract_sytem_message(messages: &mut Vec<Message>) -> Option<String> {
- if messages[0].role.is_system() {
- let system_message = messages.remove(0);
- return Some(system_message.content.to_text());
- }
- None
-}
-
pub async fn json_stream<S, F>(mut stream: S, mut handle: F) -> Result<()>
where
S: Stream<Item = Result<bytes::Bytes, reqwest::Error>> + Unpin,