diff options
| author | sigoden <sigoden@gmail.com> | 2024-04-25 10:15:54 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2024-04-25 10:15:54 +0800 |
| commit | 4db9b309803796bc5f996d0b3713344eb44207ec (patch) | |
| tree | a0ba8f3a2da32e5392e22272c641e3939de49be6 /src/client/common.rs | |
| parent | 8dacea4deb0fe67c0cdbe15bd47e6e50c1dd12a1 (diff) | |
| download | aichat-4db9b309803796bc5f996d0b3713344eb44207ec.tar.gz | |
refactor: rewrite list models of all clients (#436)
Diffstat (limited to 'src/client/common.rs')
| -rw-r--r-- | src/client/common.rs | 34 |
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, |
