summaryrefslogtreecommitdiffstats
path: root/src/config/input.rs
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2024-05-14 12:43:16 +0800
committerGitHub <noreply@github.com>2024-05-14 12:43:16 +0800
commit79d0bba640d954cd3e6acd7f4e83900eb9d56a1c (patch)
treee266ca283735b91bc2964ec4040b088cf7b865ed /src/config/input.rs
parent5284a18248bb8e48eaa4a1e6ddcf73d944d23783 (diff)
downloadaichat-79d0bba640d954cd3e6acd7f4e83900eb9d56a1c.tar.gz
feat: allow binding model to the role (#505)
Diffstat (limited to 'src/config/input.rs')
-rw-r--r--src/config/input.rs21
1 files changed, 18 insertions, 3 deletions
diff --git a/src/config/input.rs b/src/config/input.rs
index 20aa755..7210c45 100644
--- a/src/config/input.rs
+++ b/src/config/input.rs
@@ -1,8 +1,8 @@
use super::{role::Role, session::Session, GlobalConfig};
use crate::client::{
- init_client, Client, ImageUrl, Message, MessageContent, MessageContentPart, ModelCapabilities,
- SendData,
+ init_client, list_models, Client, ImageUrl, Message, MessageContent, MessageContentPart, Model,
+ ModelCapabilities, SendData,
};
use crate::utils::{base64_encode, sha256};
@@ -111,8 +111,23 @@ impl Input {
self.text = text;
}
+ pub fn model(&self) -> Model {
+ let model = self.config.read().model.clone();
+ if let Some(model_id) = self.role().and_then(|v| v.model_id.clone()) {
+ if model.id() != model_id {
+ if let Some(model) = list_models(&self.config.read())
+ .into_iter()
+ .find(|v| v.id() == model_id)
+ {
+ return model.clone();
+ }
+ }
+ };
+ model
+ }
+
pub fn create_client(&self) -> Result<Box<dyn Client>> {
- init_client(&self.config)
+ init_client(&self.config, Some(self.model()))
}
pub fn prepare_send_data(&self, stream: bool) -> Result<SendData> {