summaryrefslogtreecommitdiffstats
path: root/src/client
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2024-11-14 06:03:06 +0800
committerGitHub <noreply@github.com>2024-11-14 06:03:06 +0800
commitcfa9217422dfa6cf10bed6e6e3fab0722f58588d (patch)
treef004b226caae309c860914e7393e516464eecee8 /src/client
parentff0ea19b48a18e9a1849bd2a653173a3ed4e4560 (diff)
downloadaichat-cfa9217422dfa6cf10bed6e6e3fab0722f58588d.tar.gz
feat: save function calls in the session (#994)
Diffstat (limited to 'src/client')
-rw-r--r--src/client/message.rs31
-rw-r--r--src/client/model.rs26
2 files changed, 52 insertions, 5 deletions
diff --git a/src/client/message.rs b/src/client/message.rs
index 77adf47..061c0e2 100644
--- a/src/client/message.rs
+++ b/src/client/message.rs
@@ -1,5 +1,7 @@
use super::ToolResults;
+use crate::utils::dimmed_text;
+
use serde::{Deserialize, Serialize};
#[derive(Debug, Clone, Deserialize, Serialize)]
@@ -53,6 +55,7 @@ pub enum MessageRole {
System,
Assistant,
User,
+ Tool,
}
#[allow(dead_code)]
@@ -76,7 +79,11 @@ pub enum MessageContent {
}
impl MessageContent {
- pub fn render_input(&self, resolve_url_fn: impl Fn(&str) -> String) -> String {
+ pub fn render_input(
+ &self,
+ resolve_url_fn: impl Fn(&str) -> String,
+ agent_info: &Option<(String, Vec<String>)>,
+ ) -> String {
match self {
MessageContent::Text(text) => text.to_string(),
MessageContent::Array(list) => {
@@ -96,7 +103,27 @@ impl MessageContent {
}
format!(".file {}{}", files.join(" "), concated_text)
}
- MessageContent::ToolResults(_) => String::new(),
+ MessageContent::ToolResults(results) => {
+ let ToolResults {
+ tool_results, text, ..
+ } = results;
+ let mut lines = vec![];
+ if !text.is_empty() {
+ lines.push(text.clone())
+ }
+ for tool_result in tool_results {
+ let mut parts = vec!["Call".to_string()];
+ if let Some((agent_name, functions)) = agent_info {
+ if functions.contains(&tool_result.call.name) {
+ parts.push(agent_name.clone())
+ }
+ }
+ parts.push(tool_result.call.name.clone());
+ parts.push(tool_result.call.arguments.to_string());
+ lines.push(dimmed_text(&parts.join(" ")));
+ }
+ lines.join("\n")
+ }
}
}
diff --git a/src/client/model.rs b/src/client/model.rs
index f8a23ff..af864e2 100644
--- a/src/client/model.rs
+++ b/src/client/model.rs
@@ -1,6 +1,7 @@
use super::{
list_chat_models, list_embedding_models, list_reranker_models,
- message::{Message, MessageContent},
+ message::{Message, MessageContent, MessageContentPart},
+ ToolResults,
};
use crate::config::Config;
@@ -229,8 +230,27 @@ impl Model {
.iter()
.map(|v| match &v.content {
MessageContent::Text(text) => estimate_token_length(text),
- MessageContent::Array(_) => 0,
- MessageContent::ToolResults(_) => 0,
+ MessageContent::Array(list) => list
+ .iter()
+ .map(|v| match v {
+ MessageContentPart::Text { text } => estimate_token_length(text),
+ MessageContentPart::ImageUrl { .. } => 0,
+ })
+ .sum(),
+ MessageContent::ToolResults(results) => {
+ let ToolResults {
+ tool_results, text, ..
+ } = results;
+ estimate_token_length(text)
+ + tool_results
+ .iter()
+ .map(|v| {
+ serde_json::to_string(v)
+ .map(|v| estimate_token_length(&v))
+ .unwrap_or_default()
+ })
+ .sum::<usize>()
+ }
})
.sum()
}