summaryrefslogtreecommitdiffstats
path: root/src/config/session.rs
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2024-06-11 11:00:12 +0800
committerGitHub <noreply@github.com>2024-06-11 11:00:12 +0800
commitbb867c4fcbc6f42770471f0a0117cd8909f660d7 (patch)
tree6293f7f1108309160d1951f53f6429e9b004870d /src/config/session.rs
parent5635ca6a58fb4a590419335b098b7317285bfb82 (diff)
downloadaichat-bb867c4fcbc6f42770471f0a0117cd8909f660d7.tar.gz
feat: support bot (#579)
* feat: support bots * refactor with RoleLike * improve exiting session * make bot works with rag * refactor repl assert state * add bot banner * repl complete bots according bots.txt * fix on windows * remove threadpool executing function callings * adjust repl left_prompt * move bot config to global config.yaml * `.bot` throw err if funciton callings is not configured
Diffstat (limited to 'src/config/session.rs')
-rw-r--r--src/config/session.rs217
1 files changed, 126 insertions, 91 deletions
diff --git a/src/config/session.rs b/src/config/session.rs
index 908cc5b..97843bd 100644
--- a/src/config/session.rs
+++ b/src/config/session.rs
@@ -1,19 +1,17 @@
-use super::input::resolve_data_url;
-use super::{Config, Input, Model, Role};
+use super::input::*;
+use super::*;
use crate::client::{Message, MessageContent, MessageRole};
use crate::render::MarkdownRender;
use anyhow::{bail, Context, Result};
-use inquire::{Confirm, Text};
+use inquire::{required, Confirm, Text};
use serde::{Deserialize, Serialize};
use serde_json::json;
use std::collections::HashMap;
use std::fs::{self, create_dir_all, read_to_string};
use std::path::Path;
-pub const TEMP_SESSION_NAME: &str = "temp";
-
#[derive(Debug, Clone, Default, Deserialize, Serialize)]
pub struct Session {
#[serde(rename(serialize = "model", deserialize = "model"))]
@@ -26,68 +24,65 @@ pub struct Session {
function_matcher: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
save_session: Option<bool>,
+ #[serde(skip_serializing_if = "Option::is_none")]
+ compress_threshold: Option<usize>,
+
messages: Vec<Message>,
#[serde(default, skip_serializing_if = "HashMap::is_empty")]
data_urls: HashMap<String, String>,
#[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)]
+ model: Model,
#[serde(skip)]
- pub name: String,
+ role_prompt: String,
#[serde(skip)]
- pub path: Option<String>,
+ role_name: String,
#[serde(skip)]
- pub dirty: bool,
+ name: String,
#[serde(skip)]
- pub compressing: bool,
+ path: Option<String>,
#[serde(skip)]
- pub model: Model,
+ dirty: bool,
+ #[serde(skip)]
+ compressing: bool,
}
impl Session {
pub fn new(config: &Config, name: &str) -> Self {
- let name = if name.is_empty() {
- TEMP_SESSION_NAME
- } else {
- name
- };
let save_session = if name == TEMP_SESSION_NAME {
None
} else {
config.save_session
};
+ let role = config.extract_role();
let mut session = Self {
- model_id: config.model.id(),
- temperature: config.temperature,
- top_p: config.top_p,
- function_matcher: None,
- save_session,
- messages: Default::default(),
- compressed_messages: Default::default(),
- compress_threshold: None,
- data_urls: Default::default(),
name: name.to_string(),
- path: None,
- dirty: false,
- compressing: false,
- model: config.model.clone(),
+ save_session,
+ ..Default::default()
};
- if let Some(role) = &config.role {
- session.set_role_properties(role);
- }
+ session.set_role(role);
+ session.dirty = false;
session
}
- pub fn load(name: &str, path: &Path) -> Result<Self> {
+ pub fn load(config: &Config, name: &str, path: &Path) -> Result<Self> {
let content = read_to_string(path)
.with_context(|| format!("Failed to load session {} at {}", name, path.display()))?;
let mut session: Self =
serde_yaml::from_str(&content).with_context(|| format!("Invalid session {}", name))?;
+ session.model = Model::retrieve(config, &session.model_id)?;
session.name = name.to_string();
session.path = Some(path.display().to_string());
+ if let Some(bot) = &config.bot {
+ session
+ .role_prompt
+ .clone_from(&bot.definition().instructions);
+ }
+
Ok(session)
}
@@ -95,20 +90,12 @@ impl Session {
&self.name
}
- pub fn model_id(&self) -> &str {
- &self.model_id
+ pub fn dirty(&self) -> bool {
+ self.dirty
}
- pub fn temperature(&self) -> Option<f64> {
- self.temperature
- }
-
- pub fn top_p(&self) -> Option<f64> {
- self.top_p
- }
-
- pub fn function_matcher(&self) -> Option<&str> {
- self.function_matcher.as_deref()
+ pub fn compressing(&self) -> bool {
+ self.compressing
}
pub fn save_session(&self) -> Option<bool> {
@@ -123,7 +110,7 @@ impl Session {
}
pub fn tokens(&self) -> usize {
- self.model.total_tokens(&self.messages)
+ self.model().total_tokens(&self.messages)
}
pub fn user_messages_len(&self) -> usize {
@@ -134,10 +121,9 @@ impl Session {
if self.path.is_none() {
bail!("Not found session '{}'", self.name)
}
- let (tokens, percent) = self.tokens_and_percent();
let mut data = json!({
"path": self.path,
- "model": self.model_id(),
+ "model": self.model().id(),
});
if let Some(temperature) = self.temperature() {
data["temperature"] = temperature.into();
@@ -151,8 +137,9 @@ impl Session {
if let Some(save_session) = self.save_session() {
data["save_session"] = save_session.into();
}
+ let (tokens, percent) = self.tokens_usage();
data["total_tokens"] = tokens.into();
- if let Some(max_input_tokens) = self.model.max_input_tokens() {
+ if let Some(max_input_tokens) = self.model().max_input_tokens() {
data["max_input_tokens"] = max_input_tokens.into();
}
if percent != 0.0 {
@@ -165,14 +152,14 @@ impl Session {
Ok(output)
}
- pub fn info(&self, render: &mut MarkdownRender) -> Result<String> {
+ pub fn render(&self, render: &mut MarkdownRender) -> Result<String> {
let mut items = vec![];
if let Some(path) = &self.path {
items.push(("path", path.to_string()));
}
- items.push(("model", self.model.id()));
+ items.push(("model", self.model().id()));
if let Some(temperature) = self.temperature() {
items.push(("temperature", temperature.to_string()));
@@ -182,7 +169,7 @@ impl Session {
}
if let Some(function_matcher) = self.function_matcher() {
- items.push(("function_matcher", function_matcher.into()));
+ items.push(("function_matcher", function_matcher));
}
if let Some(save_session) = self.save_session() {
@@ -193,7 +180,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()));
}
@@ -228,13 +215,17 @@ impl Session {
}
}
+ if lines.last() == Some(&String::new()) {
+ lines.pop();
+ }
+
let output = lines.join("\n");
Ok(output)
}
- pub fn tokens_and_percent(&self) -> (usize, f32) {
+ pub fn tokens_usage(&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 {
@@ -244,28 +235,24 @@ impl Session {
(tokens, percent)
}
- pub fn set_temperature(&mut self, value: Option<f64>) {
- if self.temperature != value {
- self.temperature = value;
- self.dirty = true;
- }
- }
-
- pub fn set_top_p(&mut self, value: Option<f64>) {
- if self.top_p != value {
- self.top_p = value;
- self.dirty = true;
- }
+ pub fn set_name(&mut self, name: &str) {
+ self.name = name.to_string();
}
- pub fn set_function_matcher(&mut self, function_matcher: Option<&str>) {
- self.function_matcher = function_matcher.map(|v| v.to_string());
+ pub fn set_role(&mut self, role: Role) {
+ 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.model = role.model().clone();
+ self.role_name = role.name().to_string();
+ self.role_prompt = role.prompt().to_string();
+ self.dirty = true;
}
- pub fn set_role_properties(&mut self, role: &Role) {
- self.set_temperature(role.temperature);
- self.set_top_p(role.top_p);
- self.set_function_matcher(role.function_matcher.as_deref());
+ pub fn clear_role(&mut self) {
+ self.role_name.clear();
+ self.role_prompt.clear();
}
pub fn set_save_session(&mut self, value: Option<bool>) {
@@ -285,13 +272,8 @@ impl Session {
}
}
- pub fn set_model(&mut self, model: &Model) {
- let model_id = model.id();
- if self.model_id != model_id {
- self.model_id = model_id;
- self.dirty = true;
- }
- self.model = model.clone();
+ pub fn set_compressing(&mut self, compressing: bool) {
+ self.compressing = compressing;
}
pub fn compress(&mut self, prompt: String) {
@@ -314,8 +296,10 @@ impl Session {
if !ans {
return Ok(());
}
- while self.is_temp() {
- self.name = Text::new("Session name:").prompt()?;
+ if self.is_temp() {
+ self.name = Text::new("Session name:")
+ .with_validator(required!("This field is required"))
+ .prompt()?;
}
}
self.save(sessions_dir)?;
@@ -344,6 +328,8 @@ impl Session {
)
})?;
+ println!("✨ Saved session to '{}'", session_path.display());
+
self.dirty = false;
Ok(())
@@ -367,10 +353,8 @@ impl Session {
pub fn add_message(&mut self, input: &Input, output: &str) -> Result<()> {
let mut need_add_msg = true;
if self.messages.is_empty() {
- if let Some(role) = input.role() {
- self.messages.extend(role.build_messages(input));
- need_add_msg = false;
- }
+ self.messages.extend(input.role().build_messages(input));
+ need_add_msg = false;
}
if need_add_msg {
self.messages
@@ -402,10 +386,8 @@ impl Session {
let mut need_add_msg = true;
let len = messages.len();
if len == 0 {
- if let Some(role) = input.role() {
- messages = role.build_messages(input);
- need_add_msg = false;
- }
+ messages = input.role().build_messages(input);
+ need_add_msg = false;
} else if len == 1 && self.compressed_messages.len() >= 2 {
messages
.extend(self.compressed_messages[self.compressed_messages.len() - 2..].to_vec());
@@ -416,3 +398,56 @@ impl Session {
messages
}
}
+
+impl RoleLike for Session {
+ fn to_role(&self) -> Role {
+ let mut role = Role::new(&self.role_name, &self.role_prompt);
+ role.sync(self);
+ role
+ }
+
+ 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) {
+ if self.model().id() != model.id() {
+ self.model_id = model.id();
+ self.model = model.clone();
+ self.dirty = true;
+ }
+ }
+
+ fn set_temperature(&mut self, value: Option<f64>) {
+ if self.temperature != value {
+ self.temperature = value;
+ self.dirty = true;
+ }
+ }
+
+ fn set_top_p(&mut self, value: Option<f64>) {
+ if self.top_p != value {
+ self.top_p = value;
+ self.dirty = true;
+ }
+ }
+
+ fn set_function_matcher(&mut self, value: Option<String>) {
+ if self.function_matcher != value {
+ self.function_matcher = value;
+ self.dirty = true;
+ }
+ }
+}