summaryrefslogtreecommitdiffstats
path: root/src/config
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2024-03-10 08:01:54 +0800
committerGitHub <noreply@github.com>2024-03-10 08:01:54 +0800
commitaed243c3aa0dd6d6c7dcba088304170f9d5cb696 (patch)
treebe1ceff108f766b33d9e185ca42acab312180101 /src/config
parent8f144989695e089fe3b3e7f4e97ac2b862574bd3 (diff)
downloadaichat-aed243c3aa0dd6d6c7dcba088304170f9d5cb696.tar.gz
feat: allow use of temporary role in a session (#348)
Diffstat (limited to 'src/config')
-rw-r--r--src/config/input.rs46
-rw-r--r--src/config/mod.rs18
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) {