diff options
| author | sigoden <sigoden@gmail.com> | 2024-01-13 19:52:07 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2024-01-13 19:52:07 +0800 |
| commit | fe35cfd9419302f01baf9672493c0b0a4b41d889 (patch) | |
| tree | 94e763745fb7989c97af39cc1dfb44250440a5eb /src/client/common.rs | |
| parent | 4e99df4c1bd4028a77251bdb00ff23c664372b5f (diff) | |
| download | aichat-fe35cfd9419302f01baf9672493c0b0a4b41d889.tar.gz | |
feat: supports model capabilities (#297)
1. automatically switch to the model that has the necessary capabilities.
2. throw an error if the client does not have a model with the necessary capabilities
Diffstat (limited to 'src/client/common.rs')
| -rw-r--r-- | src/client/common.rs | 69 |
1 files changed, 54 insertions, 15 deletions
diff --git a/src/client/common.rs b/src/client/common.rs index 9ff02ba..c35ba1b 100644 --- a/src/client/common.rs +++ b/src/client/common.rs @@ -1,4 +1,4 @@ -use super::{openai::OpenAIConfig, ClientConfig, Message, MessageContent}; +use super::{openai::OpenAIConfig, ClientConfig, Message, MessageContent, Model}; use crate::{ config::{GlobalConfig, Input}, @@ -78,12 +78,26 @@ macro_rules! register_client { )+ pub fn init_client(config: &$crate::config::GlobalConfig) -> anyhow::Result<Box<dyn Client>> { - None - $(.or_else(|| $client::init(config)))+ - .ok_or_else(|| { - let model = config.read().model.clone(); - anyhow::anyhow!("Unknown client '{}'", &model.client_name) - }) + None + $(.or_else(|| $client::init(config)))+ + .ok_or_else(|| { + let model = config.read().model.clone(); + anyhow::anyhow!("Unknown client '{}'", &model.client_name) + }) + } + + pub fn ensure_model_capabilities(client: &mut dyn Client, capabilities: $crate::client::ModelCapabilities) -> anyhow::Result<()> { + if !client.model().capabilities.contains(capabilities) { + let models = client.models(); + if let Some(model) = models.into_iter().find(|v| v.capabilities.contains(capabilities)) { + client.set_model(model); + } else { + anyhow::bail!( + "The current model lacks the corresponding capability." + ); + } + } + Ok(()) } pub fn list_client_types() -> Vec<&'static str> { @@ -114,18 +128,37 @@ macro_rules! register_client { } #[macro_export] +macro_rules! client_common_fns { + () => { + fn config( + &self, + ) -> ( + &$crate::config::GlobalConfig, + &Option<$crate::client::ExtraConfig>, + ) { + (&self.global_config, &self.config.extra) + } + + fn models(&self) -> Vec<Model> { + Self::list_models(&self.config) + } + + fn model(&self) -> &Model { + &self.model + } + + fn set_model(&mut self, model: Model) { + self.model = model; + } + }; +} + +#[macro_export] macro_rules! openai_compatible_client { ($client:ident) => { #[async_trait] impl $crate::client::Client for $crate::client::$client { - fn config( - &self, - ) -> ( - &$crate::config::GlobalConfig, - &Option<$crate::client::ExtraConfig>, - ) { - (&self.global_config, &self.config.extra) - } + client_common_fns!(); async fn send_message_inner( &self, @@ -170,6 +203,12 @@ macro_rules! config_get_fn { pub trait Client { fn config(&self) -> (&GlobalConfig, &Option<ExtraConfig>); + fn models(&self) -> Vec<Model>; + + fn model(&self) -> &Model; + + fn set_model(&mut self, model: Model); + fn build_client(&self) -> Result<ReqwestClient> { let mut builder = ReqwestClient::builder(); let options = self.config().1; |
