summaryrefslogtreecommitdiffstats
path: root/src/config/role.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/role.rs
parent5284a18248bb8e48eaa4a1e6ddcf73d944d23783 (diff)
downloadaichat-79d0bba640d954cd3e6acd7f4e83900eb9d56a1c.tar.gz
feat: allow binding model to the role (#505)
Diffstat (limited to 'src/config/role.rs')
-rw-r--r--src/config/role.rs10
1 files changed, 9 insertions, 1 deletions
diff --git a/src/config/role.rs b/src/config/role.rs
index 4fc34c9..2a4b30c 100644
--- a/src/config/role.rs
+++ b/src/config/role.rs
@@ -1,6 +1,6 @@
use super::Input;
use crate::{
- client::{Message, MessageContent, MessageRole},
+ client::{Message, MessageContent, MessageRole, Model},
utils::{detect_os, detect_shell},
};
@@ -18,6 +18,8 @@ pub const INPUT_PLACEHOLDER: &str = "__INPUT__";
pub struct Role {
pub name: String,
pub prompt: String,
+ #[serde(rename(serialize = "model", deserialize = "model"))]
+ pub model_id: Option<String>,
pub temperature: Option<f64>,
pub top_p: Option<f64>,
}
@@ -28,6 +30,7 @@ impl Role {
name: TEMP_ROLE.into(),
prompt: prompt.into(),
temperature: None,
+ model_id: None,
top_p: None,
}
}
@@ -62,6 +65,7 @@ async function timeout(ms) {
.map(|(name, prompt)| Self {
name: name.into(),
prompt,
+ model_id: None,
temperature: None,
top_p: None,
})
@@ -78,6 +82,10 @@ async function timeout(ms) {
self.prompt.contains(INPUT_PLACEHOLDER)
}
+ pub fn set_model(&mut self, model: &Model) {
+ self.model_id = Some(model.id());
+ }
+
pub fn set_temperature(&mut self, value: Option<f64>) {
self.temperature = value;
}