summaryrefslogtreecommitdiffstats
path: root/src/config/role.rs
diff options
context:
space:
mode:
Diffstat (limited to 'src/config/role.rs')
-rw-r--r--src/config/role.rs154
1 files changed, 116 insertions, 38 deletions
diff --git a/src/config/role.rs b/src/config/role.rs
index 1b49bad..135bc50 100644
--- a/src/config/role.rs
+++ b/src/config/role.rs
@@ -1,46 +1,58 @@
-use super::{Config, Input};
+use super::*;
use crate::{
- client::{list_chat_models, Message, MessageContent, MessageRole, Model},
+ client::{Message, MessageContent, MessageRole, Model},
+ function::FUNCTION_ALL_MATCHER,
utils::{detect_os, detect_shell},
};
use anyhow::{Context, Result};
use serde::{Deserialize, Serialize};
-pub const TEMP_ROLE: &str = "%%";
pub const SHELL_ROLE: &str = "%shell%";
pub const EXPLAIN_SHELL_ROLE: &str = "%explain-shell%";
pub const CODE_ROLE: &str = "%code%";
pub const INPUT_PLACEHOLDER: &str = "__INPUT__";
-#[derive(Debug, Clone, Deserialize, Serialize)]
+pub trait RoleLike {
+ fn to_role(&self) -> Role;
+ fn model(&self) -> &Model;
+ fn temperature(&self) -> Option<f64>;
+ fn top_p(&self) -> Option<f64>;
+ fn function_matcher(&self) -> Option<String>;
+ fn set_model(&mut self, model: &Model);
+ fn set_temperature(&mut self, value: Option<f64>);
+ fn set_top_p(&mut self, value: Option<f64>);
+ fn set_function_matcher(&mut self, value: Option<String>);
+}
+
+#[derive(Debug, Clone, Default, Deserialize, Serialize)]
pub struct Role {
- pub name: String,
- pub prompt: String,
+ name: String,
+ prompt: String,
#[serde(
rename(serialize = "model", deserialize = "model"),
skip_serializing_if = "Option::is_none"
)]
- pub model_id: Option<String>,
+ model_id: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
- pub temperature: Option<f64>,
+ temperature: Option<f64>,
#[serde(skip_serializing_if = "Option::is_none")]
- pub top_p: Option<f64>,
+ top_p: Option<f64>,
#[serde(skip_serializing_if = "Option::is_none")]
- pub function_matcher: Option<String>,
+ function_matcher: Option<String>,
+
+ #[serde(skip)]
+ model: Model,
}
impl Role {
- pub fn temp(prompt: &str) -> Self {
+ pub fn new(name: &str, prompt: &str) -> Self {
Self {
- name: TEMP_ROLE.into(),
+ name: name.into(),
prompt: prompt.into(),
- temperature: None,
- model_id: None,
- top_p: None,
- function_matcher: None,
+ ..Default::default()
}
}
@@ -71,16 +83,18 @@ async function timeout(ms) {
.into(),
None,
),
- ("%functions%", String::new(), Some(".*".into())),
+ (
+ "%functions%",
+ String::new(),
+ Some(FUNCTION_ALL_MATCHER.into()),
+ ),
]
.into_iter()
.map(|(name, prompt, function_matcher)| Self {
name: name.into(),
prompt,
- model_id: None,
- temperature: None,
- top_p: None,
function_matcher,
+ ..Default::default()
})
.collect()
}
@@ -91,30 +105,55 @@ async function timeout(ms) {
Ok(output.trim_end().to_string())
}
- pub fn empty_prompt(&self) -> bool {
- self.prompt.is_empty()
+ pub fn sync<T: RoleLike>(&mut self, role_like: &T) {
+ let model = role_like.model();
+ let temperature = role_like.temperature();
+ let top_p = role_like.top_p();
+ let function_matcher = role_like.function_matcher();
+ self.batch_set(model, temperature, top_p, function_matcher);
}
- pub fn embedded_prompt(&self) -> bool {
- self.prompt.contains(INPUT_PLACEHOLDER)
+ pub fn batch_set(
+ &mut self,
+ model: &Model,
+ temperature: Option<f64>,
+ top_p: Option<f64>,
+ function_matcher: Option<String>,
+ ) {
+ self.set_model(model);
+ if temperature.is_some() {
+ self.set_temperature(temperature);
+ }
+ if top_p.is_some() {
+ self.set_top_p(top_p);
+ }
+ if function_matcher.is_some() {
+ self.set_function_matcher(function_matcher);
+ }
}
- pub fn retrieve_model(&self, config: &Config) -> Option<Model> {
- self.model_id
- .as_ref()
- .and_then(|model_id| Model::find(&list_chat_models(config), model_id))
+ pub fn is_derived(&self) -> bool {
+ self.name.is_empty()
}
- pub fn set_model(&mut self, model: &Model) {
- self.model_id = Some(model.id());
+ pub fn name(&self) -> &str {
+ &self.name
}
- pub fn set_temperature(&mut self, value: Option<f64>) {
- self.temperature = value;
+ pub fn model_id(&self) -> Option<&str> {
+ self.model_id.as_deref()
}
- pub fn set_top_p(&mut self, value: Option<f64>) {
- self.top_p = value;
+ pub fn prompt(&self) -> &str {
+ &self.prompt
+ }
+
+ pub fn is_empty_prompt(&self) -> bool {
+ self.prompt.is_empty()
+ }
+
+ pub fn is_embedded_prompt(&self) -> bool {
+ self.prompt.contains(INPUT_PLACEHOLDER)
}
pub fn complete_prompt_args(&mut self, name: &str) {
@@ -134,9 +173,9 @@ async function timeout(ms) {
pub fn echo_messages(&self, input: &Input) -> String {
let input_markdown = input.render();
- if self.empty_prompt() {
+ if self.is_empty_prompt() {
input_markdown
- } else if self.embedded_prompt() {
+ } else if self.is_embedded_prompt() {
self.prompt.replace(INPUT_PLACEHOLDER, &input_markdown)
} else {
format!("{}\n\n{}", self.prompt, input.render())
@@ -145,9 +184,9 @@ async function timeout(ms) {
pub fn build_messages(&self, input: &Input) -> Vec<Message> {
let mut content = input.message_content();
- if self.empty_prompt() {
+ if self.is_empty_prompt() {
vec![Message::new(MessageRole::User, content)]
- } else if self.embedded_prompt() {
+ } else if self.is_embedded_prompt() {
content.merge_prompt(|v: &str| self.prompt.replace(INPUT_PLACEHOLDER, v));
vec![Message::new(MessageRole::User, content)]
} else {
@@ -173,6 +212,45 @@ async function timeout(ms) {
}
}
+impl RoleLike for Role {
+ fn to_role(&self) -> Role {
+ self.clone()
+ }
+
+ fn model(&self) -> &Model {
+ &self.model
+ }
+
+ fn temperature(&self) -> Option<f64> {
+ self.temperature
+ }
+
+ fn top_p(&self) -> Option<f64> {
+ self.top_p
+ }
+
+ fn function_matcher(&self) -> Option<String> {
+ self.function_matcher.clone()
+ }
+
+ fn set_model(&mut self, model: &Model) {
+ self.model_id = Some(model.id());
+ self.model = model.clone();
+ }
+
+ fn set_temperature(&mut self, value: Option<f64>) {
+ self.temperature = value;
+ }
+
+ fn set_top_p(&mut self, value: Option<f64>) {
+ self.top_p = value;
+ }
+
+ fn set_function_matcher(&mut self, matcher: Option<String>) {
+ self.function_matcher = matcher;
+ }
+}
+
fn complete_prompt_args(prompt: &str, name: &str) -> String {
let mut prompt = prompt.trim().to_string();
for (i, arg) in name.split(':').skip(1).enumerate() {