summaryrefslogtreecommitdiffstats
path: root/src/config
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2024-05-18 19:06:21 +0800
committerGitHub <noreply@github.com>2024-05-18 19:06:21 +0800
commitb4a40e3fedb438570770a224b890ea24f6e660a9 (patch)
tree344b96102da7cbedf1034d023aa82599940388b1 /src/config
parent1348a62e5f8bc140a7218fbfe1b73f990ab16101 (diff)
downloadaichat-b4a40e3fedb438570770a224b890ea24f6e660a9.tar.gz
feat: support function calling (#514)
* feat: support function calling * fix on Windows OS * implement multi-steps function calling * fix on Windows OS * add error for client not support function calling * refactor message data structure and make claude client supporting function calling * support reuse previous call results * improve error handling for function calling * use prefix `may_` as indicator for `execute` type fucntions
Diffstat (limited to 'src/config')
-rw-r--r--src/config/input.rs62
-rw-r--r--src/config/mod.rs52
-rw-r--r--src/config/role.rs64
-rw-r--r--src/config/session.rs80
4 files changed, 176 insertions, 82 deletions
diff --git a/src/config/input.rs b/src/config/input.rs
index 0a3201a..78da9f2 100644
--- a/src/config/input.rs
+++ b/src/config/input.rs
@@ -1,9 +1,10 @@
use super::{role::Role, session::Session, GlobalConfig};
use crate::client::{
- init_client, list_models, Client, ImageUrl, Message, MessageContent, MessageContentPart, Model,
- ModelCapabilities, SendData,
+ init_client, list_models, Client, ImageUrl, Message, MessageContent, MessageContentPart,
+ MessageRole, Model, SendData,
};
+use crate::function::{ToolCallResult, ToolResults};
use crate::utils::{base64_encode, sha256};
use anyhow::{bail, Context, Result};
@@ -30,6 +31,7 @@ pub struct Input {
text: String,
medias: Vec<String>,
data_urls: HashMap<String, String>,
+ tool_call: Option<ToolResults>,
context: InputContext,
}
@@ -40,6 +42,7 @@ impl Input {
text: text.to_string(),
medias: Default::default(),
data_urls: Default::default(),
+ tool_call: None,
context: context.unwrap_or_else(|| InputContext::from_config(config)),
}
}
@@ -91,6 +94,7 @@ impl Input {
text: texts.join("\n"),
medias,
data_urls,
+ tool_call: Default::default(),
context: context.unwrap_or_else(|| InputContext::from_config(config)),
})
}
@@ -111,6 +115,21 @@ impl Input {
self.text = text;
}
+ pub fn merge_tool_call(
+ mut self,
+ output: String,
+ tool_call_results: Vec<ToolCallResult>,
+ ) -> Self {
+ match self.tool_call.as_mut() {
+ Some(exist_tool_call_results) => {
+ exist_tool_call_results.0.extend(tool_call_results);
+ exist_tool_call_results.1 = output;
+ }
+ None => self.tool_call = Some((tool_call_results, output)),
+ }
+ self
+ }
+
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()) {
@@ -130,7 +149,10 @@ impl Input {
init_client(&self.config, Some(self.model()))
}
- pub fn prepare_send_data(&self, stream: bool) -> Result<SendData> {
+ pub fn prepare_send_data(&self, model: &Model, stream: bool) -> Result<SendData> {
+ if !self.medias.is_empty() && !model.supports_vision() {
+ bail!("The current model does not support vision.");
+ }
let messages = self.build_messages()?;
self.config.read().model.max_input_tokens_limit(&messages)?;
let (temperature, top_p) = if let Some(session) = self.session(&self.config.read().session)
@@ -142,23 +164,41 @@ impl Input {
let config = self.config.read();
(config.temperature, config.top_p)
};
+ let mut functions = None;
+ if self.config.read().function_calling && model.supports_function_calling() {
+ let config = self.config.read();
+ let function_filter = if let Some(session) = self.session(&config.session) {
+ session.function_filter()
+ } else if let Some(role) = self.role() {
+ role.function_filter.as_deref()
+ } else {
+ None
+ };
+ functions = config.function.filtered_declarations(function_filter);
+ };
Ok(SendData {
messages,
temperature,
top_p,
+ functions,
stream,
})
}
pub fn build_messages(&self) -> Result<Vec<Message>> {
- let messages = if let Some(session) = self.session(&self.config.read().session) {
+ let mut messages = if let Some(session) = self.session(&self.config.read().session) {
session.build_messages(self)
} else if let Some(role) = self.role() {
role.build_messages(self)
} else {
- let message = Message::new(self);
- vec![message]
+ vec![Message::new(MessageRole::User, self.message_content())]
};
+ if let Some(tool_results) = &self.tool_call {
+ messages.push(Message::new(
+ MessageRole::Assistant,
+ MessageContent::ToolResults(tool_results.clone()),
+ ))
+ }
Ok(messages)
}
@@ -234,7 +274,7 @@ impl Input {
format!(".file {}{}", files.join(" "), text)
}
- pub fn to_message_content(&self) -> MessageContent {
+ pub fn message_content(&self) -> MessageContent {
if self.medias.is_empty() {
MessageContent::Text(self.text.clone())
} else {
@@ -257,14 +297,6 @@ impl Input {
MessageContent::Array(list)
}
}
-
- pub fn required_capabilities(&self) -> ModelCapabilities {
- if !self.medias.is_empty() {
- ModelCapabilities::Vision
- } else {
- ModelCapabilities::Text
- }
- }
}
#[derive(Debug, Clone, Default)]
diff --git a/src/config/mod.rs b/src/config/mod.rs
index 66abac2..d59a81e 100644
--- a/src/config/mod.rs
+++ b/src/config/mod.rs
@@ -10,6 +10,7 @@ use crate::client::{
create_client_config, list_client_types, list_models, ClientConfig, Model,
OPENAI_COMPATIBLE_PLATFORMS,
};
+use crate::function::{Function, ToolCallResult};
use crate::render::{MarkdownRender, RenderOptions};
use crate::utils::{
format_option_value, fuzzy_match, get_env_name, light_theme_from_colorfgbg, now, render_prompt,
@@ -41,6 +42,7 @@ const CONFIG_FILE_NAME: &str = "config.yaml";
const ROLES_FILE_NAME: &str = "roles.yaml";
const MESSAGES_FILE_NAME: &str = "messages.md";
const SESSIONS_DIR_NAME: &str = "sessions";
+const FUNCTIONS_DIR_NAME: &str = "functions";
const CLIENTS_FIELD: &str = "clients";
@@ -69,6 +71,7 @@ pub struct Config {
pub keybindings: Keybindings,
pub prelude: Option<String>,
pub buffer_editor: Option<String>,
+ pub function_calling: bool,
pub compress_threshold: usize,
pub summarize_prompt: Option<String>,
pub summary_prompt: Option<String>,
@@ -84,6 +87,8 @@ pub struct Config {
#[serde(skip)]
pub model: Model,
#[serde(skip)]
+ pub function: Function,
+ #[serde(skip)]
pub working_mode: WorkingMode,
#[serde(skip)]
pub last_message: Option<(Input, String)>,
@@ -106,6 +111,7 @@ impl Default for Config {
keybindings: Default::default(),
prelude: None,
buffer_editor: None,
+ function_calling: false,
compress_threshold: 2000,
summarize_prompt: None,
summary_prompt: None,
@@ -116,6 +122,7 @@ impl Default for Config {
role: None,
session: None,
model: Default::default(),
+ function: Default::default(),
working_mode: WorkingMode::Command,
last_message: None,
}
@@ -142,6 +149,8 @@ impl Config {
config.set_wrap(&wrap)?;
}
+ config.function = Function::init(&Self::functions_dir()?)?;
+
config.working_mode = working_mode;
config.load_roles()?;
@@ -212,15 +221,20 @@ impl Config {
Ok(path)
}
- pub fn save_message(&mut self, input: Input, output: &str) -> Result<()> {
+ pub fn save_message(
+ &mut self,
+ input: &Input,
+ output: &str,
+ tool_call_results: &[ToolCallResult],
+ ) -> Result<()> {
self.last_message = Some((input.clone(), output.to_string()));
- if self.dry_run {
+ if self.dry_run || output.is_empty() || !tool_call_results.is_empty() {
return Ok(());
}
if let Some(session) = input.session_mut(&mut self.session) {
- session.add_message(&input, output)?;
+ session.add_message(input, output)?;
return Ok(());
}
@@ -275,6 +289,10 @@ impl Config {
Self::local_path(SESSIONS_DIR_NAME)
}
+ pub fn functions_dir() -> Result<PathBuf> {
+ Self::local_path(FUNCTIONS_DIR_NAME)
+ }
+
pub fn session_file(name: &str) -> Result<PathBuf> {
let mut path = Self::sessions_dir()?;
path.push(&format!("{name}.yaml"));
@@ -294,8 +312,7 @@ impl Config {
pub fn set_role_obj(&mut self, role: Role) -> Result<()> {
if let Some(session) = self.session.as_mut() {
session.guard_empty()?;
- session.set_temperature(role.temperature);
- session.set_top_p(role.top_p);
+ session.set_role_properties(&role);
}
if let Some(model_id) = &role.model_id {
self.set_model(model_id)?;
@@ -428,6 +445,8 @@ impl Config {
),
("temperature", format_option_value(&temperature)),
("top_p", format_option_value(&top_p)),
+ ("function_calling", self.function_calling.to_string()),
+ ("compress_threshold", self.compress_threshold.to_string()),
("dry_run", self.dry_run.to_string()),
("save", self.save.to_string()),
("save_session", format_option_value(&self.save_session)),
@@ -438,11 +457,11 @@ impl Config {
("auto_copy", self.auto_copy.to_string()),
("keybindings", self.keybindings.stringify().into()),
("prelude", format_option_value(&self.prelude)),
- ("compress_threshold", self.compress_threshold.to_string()),
("config_file", display_path(&Self::config_file()?)),
("roles_file", display_path(&Self::roles_file()?)),
("messages_file", display_path(&Self::messages_file()?)),
("sessions_dir", display_path(&Self::sessions_dir()?)),
+ ("functions_dir", display_path(&Self::functions_dir()?)),
];
let output = items
.iter()
@@ -508,6 +527,7 @@ impl Config {
"max_output_tokens",
"temperature",
"top_p",
+ "function_calling",
"compress_threshold",
"save",
"save_session",
@@ -523,10 +543,11 @@ impl Config {
(values, args[0])
} else if args.len() == 2 {
let values = match args[0] {
- "max_output_tokens" => match self.model.max_output_tokens {
+ "max_output_tokens" => match self.model.max_output_tokens() {
Some(v) => vec![v.to_string()],
None => vec![],
},
+ "function_calling" => complete_bool(self.function_calling),
"save" => complete_bool(self.save),
"save_session" => {
let save_session = if let Some(session) = &self.session {
@@ -574,6 +595,10 @@ impl Config {
let value = parse_value(value)?;
self.set_top_p(value);
}
+ "function_calling" => {
+ let value = value.parse().with_context(|| "Invalid value")?;
+ self.function_calling = value;
+ }
"compress_threshold" => {
let value = parse_value(value)?;
self.set_compress_threshold(value);
@@ -792,11 +817,14 @@ impl Config {
fn generate_prompt_context(&self) -> HashMap<&str, String> {
let mut output = HashMap::new();
output.insert("model", self.model.id());
- output.insert("client_name", self.model.client_name.clone());
- output.insert("model_name", self.model.name.clone());
+ output.insert("client_name", self.model.client_name().to_string());
+ output.insert("model_name", self.model.name().to_string());
output.insert(
"max_input_tokens",
- self.model.max_input_tokens.unwrap_or_default().to_string(),
+ self.model
+ .max_input_tokens()
+ .unwrap_or_default()
+ .to_string(),
);
if let Some(temperature) = self.temperature {
if temperature != 0.0 {
@@ -884,8 +912,8 @@ impl Config {
}
fn load_config_file(config_path: &Path) -> Result<Self> {
- let ctx = || format!("Failed to load config at {}", config_path.display());
- let content = read_to_string(config_path).with_context(ctx)?;
+ let content = read_to_string(config_path)
+ .with_context(|| format!("Failed to load config at {}", config_path.display()))?;
let config: Self = serde_yaml::from_str(&content).map_err(|err| {
let err_msg = err.to_string();
let err_msg = if err_msg.starts_with(&format!("{}: ", CLIENTS_FIELD)) {
diff --git a/src/config/role.rs b/src/config/role.rs
index 2a4b30c..27858a3 100644
--- a/src/config/role.rs
+++ b/src/config/role.rs
@@ -1,4 +1,5 @@
use super::Input;
+
use crate::{
client::{Message, MessageContent, MessageRole, Model},
utils::{detect_os, detect_shell},
@@ -18,10 +19,17 @@ pub const INPUT_PLACEHOLDER: &str = "__INPUT__";
pub struct Role {
pub name: String,
pub prompt: String,
- #[serde(rename(serialize = "model", deserialize = "model"))]
+ #[serde(
+ rename(serialize = "model", deserialize = "model"),
+ skip_serializing_if = "Option::is_none"
+ )]
pub model_id: Option<String>,
+ #[serde(skip_serializing_if = "Option::is_none")]
pub temperature: Option<f64>,
+ #[serde(skip_serializing_if = "Option::is_none")]
pub top_p: Option<f64>,
+ #[serde(skip_serializing_if = "Option::is_none")]
+ pub function_filter: Option<String>,
}
impl Role {
@@ -32,12 +40,13 @@ impl Role {
temperature: None,
model_id: None,
top_p: None,
+ function_filter: None,
}
}
pub fn builtin() -> Vec<Role> {
[
- (SHELL_ROLE, shell_prompt()),
+ (SHELL_ROLE, shell_prompt(), None),
(
EXPLAIN_SHELL_ROLE,
r#"Provide a terse, single sentence description of the given shell command.
@@ -45,6 +54,7 @@ Describe each argument and option of the command.
Provide short responses in about 80 words.
APPLY MARKDOWN formatting when possible."#
.into(),
+ None,
),
(
CODE_ROLE,
@@ -59,15 +69,18 @@ async function timeout(ms) {
```
"#
.into(),
+ None,
),
+ ("%functions%", String::new(), Some(".*".into())),
]
.into_iter()
- .map(|(name, prompt)| Self {
+ .map(|(name, prompt, function_filter)| Self {
name: name.into(),
prompt,
model_id: None,
temperature: None,
top_p: None,
+ function_filter,
})
.collect()
}
@@ -78,7 +91,11 @@ async function timeout(ms) {
Ok(output.trim_end().to_string())
}
- pub fn embedded(&self) -> bool {
+ pub fn empty_prompt(&self) -> bool {
+ self.prompt.is_empty()
+ }
+
+ pub fn embedded_prompt(&self) -> bool {
self.prompt.contains(INPUT_PLACEHOLDER)
}
@@ -111,7 +128,9 @@ async function timeout(ms) {
pub fn echo_messages(&self, input: &Input) -> String {
let input_markdown = input.render();
- if self.embedded() {
+ if self.empty_prompt() {
+ input_markdown
+ } else if self.embedded_prompt() {
self.prompt.replace(INPUT_PLACEHOLDER, &input_markdown)
} else {
format!("{}\n\n{}", self.prompt, input.render())
@@ -119,41 +138,30 @@ async function timeout(ms) {
}
pub fn build_messages(&self, input: &Input) -> Vec<Message> {
- let mut content = input.to_message_content();
-
- if self.embedded() {
+ let mut content = input.message_content();
+ if self.empty_prompt() {
+ vec![Message::new(MessageRole::User, content)]
+ } else if self.embedded_prompt() {
content.merge_prompt(|v: &str| self.prompt.replace(INPUT_PLACEHOLDER, v));
- vec![Message {
- role: MessageRole::User,
- content,
- }]
+ vec![Message::new(MessageRole::User, content)]
} else {
let mut messages = vec![];
let (system, cases) = parse_structure_prompt(&self.prompt);
if !system.is_empty() {
- messages.push(Message {
- role: MessageRole::System,
- content: MessageContent::Text(system.to_string()),
- })
+ messages.push(Message::new(
+ MessageRole::System,
+ MessageContent::Text(system.to_string()),
+ ));
}
if !cases.is_empty() {
messages.extend(cases.into_iter().flat_map(|(i, o)| {
vec![
- Message {
- role: MessageRole::User,
- content: MessageContent::Text(i.to_string()),
- },
- Message {
- role: MessageRole::Assistant,
- content: MessageContent::Text(o.to_string()),
- },
+ Message::new(MessageRole::User, MessageContent::Text(i.to_string())),
+ Message::new(MessageRole::Assistant, MessageContent::Text(o.to_string())),
]
}));
}
- messages.push(Message {
- role: MessageRole::User,
- content,
- });
+ messages.push(Message::new(MessageRole::User, content));
messages
}
}
diff --git a/src/config/session.rs b/src/config/session.rs
index a458d4e..ef1eab1 100644
--- a/src/config/session.rs
+++ b/src/config/session.rs
@@ -1,5 +1,5 @@
use super::input::resolve_data_url;
-use super::{Config, Input, Model};
+use super::{Config, Input, Model, Role};
use crate::client::{Message, MessageContent, MessageRole};
use crate::render::MarkdownRender;
@@ -17,15 +17,20 @@ pub const TEMP_SESSION_NAME: &str = "temp";
pub struct Session {
#[serde(rename(serialize = "model", deserialize = "model"))]
model_id: String,
+ #[serde(skip_serializing_if = "Option::is_none")]
temperature: Option<f64>,
+ #[serde(skip_serializing_if = "Option::is_none")]
top_p: Option<f64>,
- #[serde(default)]
+ #[serde(skip_serializing_if = "Option::is_none")]
+ function_filter: Option<String>,
+ #[serde(skip_serializing_if = "Option::is_none")]
save_session: Option<bool>,
messages: Vec<Message>,
- #[serde(default)]
+ #[serde(default, skip_serializing_if = "HashMap::is_empty")]
data_urls: HashMap<String, String>,
- #[serde(default)]
+ #[serde(default, skip_serializing_if = "Vec::is_empty")]
compressed_messages: Vec<Message>,
+ #[serde(skip_serializing_if = "Option::is_none")]
compress_threshold: Option<usize>,
#[serde(skip)]
pub name: String,
@@ -41,13 +46,14 @@ pub struct Session {
impl Session {
pub fn new(config: &Config, name: &str) -> Self {
- Self {
+ let mut session = Self {
model_id: config.model.id(),
temperature: config.temperature,
top_p: config.top_p,
+ function_filter: None,
save_session: config.save_session,
- messages: vec![],
- compressed_messages: vec![],
+ messages: Default::default(),
+ compressed_messages: Default::default(),
compress_threshold: None,
data_urls: Default::default(),
name: name.to_string(),
@@ -55,7 +61,11 @@ impl Session {
dirty: false,
compressing: false,
model: config.model.clone(),
+ };
+ if let Some(role) = &config.role {
+ session.set_role_properties(role);
}
+ session
}
pub fn load(name: &str, path: &Path) -> Result<Self> {
@@ -86,6 +96,10 @@ impl Session {
self.top_p
}
+ pub fn function_filter(&self) -> Option<&str> {
+ self.function_filter.as_deref()
+ }
+
pub fn save_session(&self) -> Option<bool> {
self.save_session
}
@@ -120,12 +134,15 @@ impl Session {
if let Some(top_p) = self.top_p() {
data["top_p"] = top_p.into();
}
+ if let Some(function_filter) = self.function_filter() {
+ data["function_filter"] = function_filter.into();
+ }
if let Some(save_session) = self.save_session() {
data["save_session"] = save_session.into();
}
data["total_tokens"] = tokens.into();
- if let Some(context_window) = self.model.max_input_tokens {
- data["max_input_tokens"] = context_window.into();
+ if let Some(max_input_tokens) = self.model.max_input_tokens() {
+ data["max_input_tokens"] = max_input_tokens.into();
}
if percent != 0.0 {
data["total/max"] = format!("{}%", percent).into();
@@ -153,6 +170,10 @@ impl Session {
items.push(("top_p", top_p.to_string()));
}
+ if let Some(function_filter) = self.function_filter() {
+ items.push(("function_filter", function_filter.into()));
+ }
+
if let Some(save_session) = self.save_session() {
items.push(("save_session", save_session.to_string()));
}
@@ -161,7 +182,7 @@ impl Session {
items.push(("compress_threshold", compress_threshold.to_string()));
}
- if let Some(max_input_tokens) = self.model.max_input_tokens {
+ if let Some(max_input_tokens) = self.model.max_input_tokens() {
items.push(("max_input_tokens", max_input_tokens.to_string()));
}
@@ -202,7 +223,7 @@ impl Session {
pub fn tokens_and_percent(&self) -> (usize, f32) {
let tokens = self.tokens();
- let max_input_tokens = self.model.max_input_tokens.unwrap_or_default();
+ let max_input_tokens = self.model.max_input_tokens().unwrap_or_default();
let percent = if max_input_tokens == 0 {
0.0
} else {
@@ -226,6 +247,16 @@ impl Session {
}
}
+ pub fn set_functions(&mut self, function_filter: Option<&str>) {
+ self.function_filter = function_filter.map(|v| v.to_string());
+ }
+
+ pub fn set_role_properties(&mut self, role: &Role) {
+ self.set_temperature(role.temperature);
+ self.set_top_p(role.top_p);
+ self.set_functions(role.function_filter.as_deref());
+ }
+
pub fn set_save_session(&mut self, value: Option<bool>) {
if self.save_session != value {
self.save_session = value;
@@ -251,10 +282,10 @@ impl Session {
pub fn compress(&mut self, prompt: String) {
self.compressed_messages.append(&mut self.messages);
- self.messages.push(Message {
- role: MessageRole::System,
- content: MessageContent::Text(prompt),
- });
+ self.messages.push(Message::new(
+ MessageRole::System,
+ MessageContent::Text(prompt),
+ ));
self.dirty = true;
}
@@ -300,16 +331,14 @@ impl Session {
}
}
if need_add_msg {
- self.messages.push(Message {
- role: MessageRole::User,
- content: input.to_message_content(),
- });
+ self.messages
+ .push(Message::new(MessageRole::User, input.message_content()));
}
self.data_urls.extend(input.data_urls());
- self.messages.push(Message {
- role: MessageRole::Assistant,
- content: MessageContent::Text(output.to_string()),
- });
+ self.messages.push(Message::new(
+ MessageRole::Assistant,
+ MessageContent::Text(output.to_string()),
+ ));
self.dirty = true;
Ok(())
}
@@ -340,10 +369,7 @@ impl Session {
.extend(self.compressed_messages[self.compressed_messages.len() - 2..].to_vec());
}
if need_add_msg {
- messages.push(Message {
- role: MessageRole::User,
- content: input.to_message_content(),
- });
+ messages.push(Message::new(MessageRole::User, input.message_content()));
}
messages
}