diff options
| author | sigoden <sigoden@gmail.com> | 2024-05-14 12:43:16 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2024-05-14 12:43:16 +0800 |
| commit | 79d0bba640d954cd3e6acd7f4e83900eb9d56a1c (patch) | |
| tree | e266ca283735b91bc2964ec4040b088cf7b865ed /src/config/role.rs | |
| parent | 5284a18248bb8e48eaa4a1e6ddcf73d944d23783 (diff) | |
| download | aichat-79d0bba640d954cd3e6acd7f4e83900eb9d56a1c.tar.gz | |
feat: allow binding model to the role (#505)
Diffstat (limited to 'src/config/role.rs')
| -rw-r--r-- | src/config/role.rs | 10 |
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; } |
