diff options
Diffstat (limited to 'src')
| -rw-r--r-- | src/config/bot.rs | 8 | ||||
| -rw-r--r-- | src/config/mod.rs | 8 | ||||
| -rw-r--r-- | src/config/role.rs | 33 | ||||
| -rw-r--r-- | src/config/session.rs | 22 | ||||
| -rw-r--r-- | src/function.rs | 6 |
5 files changed, 39 insertions, 38 deletions
diff --git a/src/config/bot.rs b/src/config/bot.rs index 80d3a5d..a0ea58a 100644 --- a/src/config/bot.rs +++ b/src/config/bot.rs @@ -2,7 +2,7 @@ use super::*; use crate::{ client::Model, - function::{Functions, FUNCTION_ALL_MATCHER}, + function::{Functions, FunctionsFilter, SELECTED_ALL_FUNCTIONS}, }; use anyhow::{Context, Result}; @@ -141,11 +141,11 @@ impl RoleLike for Bot { self.config.top_p } - fn function_matcher(&self) -> Option<String> { + fn selected_functions(&self) -> Option<FunctionsFilter> { if self.functions.is_empty() { None } else { - Some(FUNCTION_ALL_MATCHER.into()) + Some(SELECTED_ALL_FUNCTIONS.into()) } } @@ -162,7 +162,7 @@ impl RoleLike for Bot { self.config.top_p = value; } - fn set_function_matcher(&mut self, _value: Option<String>) {} + fn set_selected_functions(&mut self, _value: Option<FunctionsFilter>) {} } #[derive(Debug, Clone, Default, Deserialize, Serialize)] diff --git a/src/config/mod.rs b/src/config/mod.rs index 3af4467..b088282 100644 --- a/src/config/mod.rs +++ b/src/config/mod.rs @@ -948,11 +948,11 @@ impl Config { pub fn select_functions(&self, model: &Model, role: &Role) -> Option<Vec<FunctionDeclaration>> { let mut functions = None; if self.function_calling { - let function_matcher = role.function_matcher(); - if let Some(matcher) = function_matcher { + let filter = role.selected_functions(); + if let Some(filter) = filter { functions = match &self.bot { - Some(bot) => bot.functions().select(&matcher), - None => self.functions.select(&matcher), + Some(bot) => bot.functions().select(&filter), + None => self.functions.select(&filter), }; if !model.supports_function_calling() { functions = None; diff --git a/src/config/role.rs b/src/config/role.rs index 135bc50..b6b7fc8 100644 --- a/src/config/role.rs +++ b/src/config/role.rs @@ -2,7 +2,7 @@ use super::*; use crate::{ client::{Message, MessageContent, MessageRole, Model}, - function::FUNCTION_ALL_MATCHER, + function::{FunctionsFilter, SELECTED_ALL_FUNCTIONS}, utils::{detect_os, detect_shell}, }; @@ -20,16 +20,17 @@ pub trait RoleLike { fn model(&self) -> &Model; fn temperature(&self) -> Option<f64>; fn top_p(&self) -> Option<f64>; - fn function_matcher(&self) -> Option<String>; + fn selected_functions(&self) -> Option<FunctionsFilter>; 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>); + fn set_selected_functions(&mut self, value: Option<FunctionsFilter>); } #[derive(Debug, Clone, Default, Deserialize, Serialize)] pub struct Role { name: String, + #[serde(default)] prompt: String, #[serde( rename(serialize = "model", deserialize = "model"), @@ -41,7 +42,7 @@ pub struct Role { #[serde(skip_serializing_if = "Option::is_none")] top_p: Option<f64>, #[serde(skip_serializing_if = "Option::is_none")] - function_matcher: Option<String>, + selected_functions: Option<FunctionsFilter>, #[serde(skip)] model: Model, @@ -86,14 +87,14 @@ async function timeout(ms) { ( "%functions%", String::new(), - Some(FUNCTION_ALL_MATCHER.into()), + Some(SELECTED_ALL_FUNCTIONS.into()), ), ] .into_iter() - .map(|(name, prompt, function_matcher)| Self { + .map(|(name, prompt, selected_functions)| Self { name: name.into(), prompt, - function_matcher, + selected_functions, ..Default::default() }) .collect() @@ -109,8 +110,8 @@ async function timeout(ms) { 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); + let selected_functions = role_like.selected_functions(); + self.batch_set(model, temperature, top_p, selected_functions); } pub fn batch_set( @@ -118,7 +119,7 @@ async function timeout(ms) { model: &Model, temperature: Option<f64>, top_p: Option<f64>, - function_matcher: Option<String>, + selected_functions: Option<FunctionsFilter>, ) { self.set_model(model); if temperature.is_some() { @@ -127,8 +128,8 @@ async function timeout(ms) { if top_p.is_some() { self.set_top_p(top_p); } - if function_matcher.is_some() { - self.set_function_matcher(function_matcher); + if selected_functions.is_some() { + self.set_selected_functions(selected_functions); } } @@ -229,8 +230,8 @@ impl RoleLike for Role { self.top_p } - fn function_matcher(&self) -> Option<String> { - self.function_matcher.clone() + fn selected_functions(&self) -> Option<FunctionsFilter> { + self.selected_functions.clone() } fn set_model(&mut self, model: &Model) { @@ -246,8 +247,8 @@ impl RoleLike for Role { self.top_p = value; } - fn set_function_matcher(&mut self, matcher: Option<String>) { - self.function_matcher = matcher; + fn set_selected_functions(&mut self, value: Option<FunctionsFilter>) { + self.selected_functions = value; } } diff --git a/src/config/session.rs b/src/config/session.rs index a525af3..62ba93c 100644 --- a/src/config/session.rs +++ b/src/config/session.rs @@ -21,7 +21,7 @@ pub struct Session { #[serde(skip_serializing_if = "Option::is_none")] top_p: Option<f64>, #[serde(skip_serializing_if = "Option::is_none")] - function_matcher: Option<String>, + selected_functions: Option<FunctionsFilter>, #[serde(skip_serializing_if = "Option::is_none")] save_session: Option<bool>, #[serde(skip_serializing_if = "Option::is_none")] @@ -139,8 +139,8 @@ impl Session { if let Some(top_p) = self.top_p() { data["top_p"] = top_p.into(); } - if let Some(function_matcher) = self.function_matcher() { - data["function_matcher"] = function_matcher.into(); + if let Some(selected_functions) = self.selected_functions() { + data["selected_functions"] = selected_functions.into(); } if let Some(save_session) = self.save_session() { data["save_session"] = save_session.into(); @@ -176,8 +176,8 @@ impl Session { items.push(("top_p", top_p.to_string())); } - if let Some(function_matcher) = self.function_matcher() { - items.push(("function_matcher", function_matcher)); + if let Some(selected_functions) = self.selected_functions() { + items.push(("selected_functions", selected_functions)); } if let Some(save_session) = self.save_session() { @@ -251,7 +251,7 @@ impl Session { self.model_id = role.model().id(); self.temperature = role.temperature(); self.top_p = role.top_p(); - self.function_matcher = role.function_matcher().map(|v| v.to_string()); + self.selected_functions = role.selected_functions(); self.model = role.model().clone(); self.role_name = role.name().to_string(); self.role_prompt = role.prompt().to_string(); @@ -418,8 +418,8 @@ impl RoleLike for Session { self.top_p } - fn function_matcher(&self) -> Option<String> { - self.function_matcher.clone() + fn selected_functions(&self) -> Option<FunctionsFilter> { + self.selected_functions.clone() } fn set_model(&mut self, model: &Model) { @@ -444,9 +444,9 @@ impl RoleLike for Session { } } - fn set_function_matcher(&mut self, value: Option<String>) { - if self.function_matcher != value { - self.function_matcher = value; + fn set_selected_functions(&mut self, value: Option<FunctionsFilter>) { + if self.selected_functions != value { + self.selected_functions = value; self.dirty = true; } } diff --git a/src/function.rs b/src/function.rs index f16fe3f..068cee0 100644 --- a/src/function.rs +++ b/src/function.rs @@ -15,7 +15,7 @@ use std::{ path::Path, }; -pub const FUNCTION_ALL_MATCHER: &str = ".*"; +pub const SELECTED_ALL_FUNCTIONS: &str = ".*"; pub type ToolResults = (Vec<ToolCallResult>, String); pub type FunctionsFilter = String; @@ -80,8 +80,8 @@ impl Functions { }) } - pub fn select(&self, matcher: &str) -> Option<Vec<FunctionDeclaration>> { - let regex = Regex::new(&format!("^({matcher})$")).ok()?; + pub fn select(&self, filter: &FunctionsFilter) -> Option<Vec<FunctionDeclaration>> { + let regex = Regex::new(&format!("^({filter})$")).ok()?; let output: Vec<FunctionDeclaration> = self .declarations .iter() |
