From aed243c3aa0dd6d6c7dcba088304170f9d5cb696 Mon Sep 17 00:00:00 2001 From: sigoden Date: Sun, 10 Mar 2024 08:01:54 +0800 Subject: feat: allow use of temporary role in a session (#348) --- src/config/input.rs | 46 ++++++++++++++++++++++++++++++++++++++++++++-- 1 file changed, 44 insertions(+), 2 deletions(-) (limited to 'src/config/input.rs') 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, data_urls: HashMap, + 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) -> Result { + pub fn new(text: &str, files: Vec, context: InputContext) -> Result { 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 { self.data_urls.clone() } + pub fn role(&self) -> Option<&Role> { + self.context.role.as_ref() + } + + pub fn session<'a>(&self, session: &'a Option) -> Option<&'a Session> { + if self.context.in_session { + session.as_ref() + } else { + None + } + } + + pub fn session_mut<'a>(&self, session: &'a mut Option) -> 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, + in_session: bool, +} + +impl InputContext { + pub fn new(role: Option, in_session: bool) -> Self { + Self { role, in_session } + } +} + pub fn resolve_data_url(data_urls: &HashMap, data_url: String) -> String { if data_url.starts_with("data:") { let hash = sha256sum(&data_url); -- cgit v1.2.3