diff options
| author | sigoden <sigoden@gmail.com> | 2024-03-10 08:01:54 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2024-03-10 08:01:54 +0800 |
| commit | aed243c3aa0dd6d6c7dcba088304170f9d5cb696 (patch) | |
| tree | be1ceff108f766b33d9e185ca42acab312180101 /src/config | |
| parent | 8f144989695e089fe3b3e7f4e97ac2b862574bd3 (diff) | |
| download | aichat-aed243c3aa0dd6d6c7dcba088304170f9d5cb696.tar.gz | |
feat: allow use of temporary role in a session (#348)
Diffstat (limited to 'src/config')
| -rw-r--r-- | src/config/input.rs | 46 | ||||
| -rw-r--r-- | src/config/mod.rs | 18 |
2 files changed, 55 insertions, 9 deletions
diff --git a/src/config/input.rs b/src/config/input.rs index cb3cdf2..c5ba1a2 100644 --- a/src/config/input.rs +++ b/src/config/input.rs @@ -1,3 +1,6 @@ +use super::role::Role; +use super::session::Session; + use crate::client::{ImageUrl, MessageContent, MessageContentPart, ModelCapabilities}; use crate::utils::sha256sum; @@ -25,18 +28,20 @@ pub struct Input { text: String, medias: Vec<String>, data_urls: HashMap<String, String>, + context: InputContext, } impl Input { - pub fn from_str(text: &str) -> Self { + pub fn from_str(text: &str, context: InputContext) -> Self { Self { text: text.to_string(), medias: Default::default(), data_urls: Default::default(), + context, } } - pub fn new(text: &str, files: Vec<String>) -> Result<Self> { + pub fn new(text: &str, files: Vec<String>, context: InputContext) -> Result<Self> { let mut texts = vec![text.to_string()]; let mut medias = vec![]; let mut data_urls = HashMap::new(); @@ -72,13 +77,38 @@ impl Input { text: texts.join("\n"), medias, data_urls, + context, }) } + pub fn is_empty(&self) -> bool { + self.text.is_empty() && self.medias.is_empty() + } + pub fn data_urls(&self) -> HashMap<String, String> { self.data_urls.clone() } + pub fn role(&self) -> Option<&Role> { + self.context.role.as_ref() + } + + pub fn session<'a>(&self, session: &'a Option<Session>) -> Option<&'a Session> { + if self.context.in_session { + session.as_ref() + } else { + None + } + } + + pub fn session_mut<'a>(&self, session: &'a mut Option<Session>) -> Option<&'a mut Session> { + if self.context.in_session { + session.as_mut() + } else { + None + } + } + pub fn summary(&self) -> String { let text: String = self .text @@ -154,6 +184,18 @@ impl Input { } } +#[derive(Debug, Clone, Default)] +pub struct InputContext { + role: Option<Role>, + in_session: bool, +} + +impl InputContext { + pub fn new(role: Option<Role>, in_session: bool) -> Self { + Self { role, in_session } + } +} + pub fn resolve_data_url(data_urls: &HashMap<String, String>, data_url: String) -> String { if data_url.starts_with("data:") { let hash = sha256sum(&data_url); diff --git a/src/config/mod.rs b/src/config/mod.rs index 65b573a..4ee5319 100644 --- a/src/config/mod.rs +++ b/src/config/mod.rs @@ -2,7 +2,7 @@ mod input; mod role; mod session; -pub use self::input::Input; +pub use self::input::{Input, InputContext}; use self::role::Role; use self::session::{Session, TEMP_SESSION_NAME}; @@ -226,7 +226,7 @@ impl Config { return Ok(()); } - if let Some(session) = self.session.as_mut() { + if let Some(session) = input.session_mut(&mut self.session) { session.add_message(&input, output)?; return Ok(()); } @@ -241,7 +241,7 @@ impl Config { let timestamp = now(); let summary = input.summary(); let input_markdown = input.render(); - let output = match self.role.as_ref() { + let output = match input.role() { None => { format!("# CHAT: {summary} [{timestamp}]\n{input_markdown}\n--------\n{output}\n--------\n\n",) } @@ -369,9 +369,9 @@ impl Config { } pub fn echo_messages(&self, input: &Input) -> String { - if let Some(session) = self.session.as_ref() { + if let Some(session) = input.session(&self.session) { session.echo_messages(input) - } else if let Some(role) = self.role.as_ref() { + } else if let Some(role) = input.role() { role.echo_messages(input) } else { input.render() @@ -379,9 +379,9 @@ impl Config { } pub fn build_messages(&self, input: &Input) -> Result<Vec<Message>> { - let messages = if let Some(session) = self.session.as_ref() { + let messages = if let Some(session) = input.session(&self.session) { session.build_emssages(input) - } else if let Some(role) = self.role.as_ref() { + } else if let Some(role) = input.role() { role.build_messages(input) } else { let message = Message::new(input); @@ -762,6 +762,10 @@ impl Config { }) } + pub fn input_context(&self) -> InputContext { + InputContext::new(self.role.clone(), self.has_session()) + } + pub fn maybe_print_send_tokens(&self, input: &Input) { if self.dry_run { if let Ok(messages) = self.build_messages(input) { |
