summaryrefslogtreecommitdiffstats
path: root/src/config
diff options
context:
space:
mode:
Diffstat (limited to 'src/config')
-rw-r--r--src/config/input.rs162
-rw-r--r--src/config/mod.rs39
-rw-r--r--src/config/role.rs25
-rw-r--r--src/config/session.rs36
4 files changed, 220 insertions, 42 deletions
diff --git a/src/config/input.rs b/src/config/input.rs
new file mode 100644
index 0000000..3997929
--- /dev/null
+++ b/src/config/input.rs
@@ -0,0 +1,162 @@
+use crate::client::{ImageUrl, MessageContent, MessageContentPart};
+use crate::utils::sha256sum;
+
+use anyhow::{bail, Context, Result};
+use base64::{self, engine::general_purpose::STANDARD, Engine};
+use mime_guess::from_path;
+use std::{
+ collections::HashMap,
+ fs::{self, File},
+ io::Read,
+ path::{Path, PathBuf},
+};
+
+const IMAGE_EXTS: [&str; 5] = ["png", "jpeg", "jpg", "webp", "gif"];
+
+#[derive(Debug, Clone)]
+pub struct Input {
+ text: String,
+ medias: Vec<String>,
+ data_urls: HashMap<String, String>,
+}
+
+impl Input {
+ pub fn from_str(text: &str) -> Self {
+ Self {
+ text: text.to_string(),
+ medias: Default::default(),
+ data_urls: Default::default(),
+ }
+ }
+
+ pub fn new(text: &str, files: Vec<String>) -> Result<Self> {
+ let mut texts = vec![text.to_string()];
+ let mut medias = vec![];
+ let mut data_urls = HashMap::new();
+ for file_item in files.into_iter() {
+ match resolve_path(&file_item) {
+ Some(file_path) => {
+ let file_path = fs::canonicalize(file_path)
+ .with_context(|| format!("Unable to use file '{file_item}"))?;
+ if is_image_ext(&file_path) {
+ let data_url = read_media_to_data_url(&file_path)?;
+ data_urls.insert(sha256sum(&data_url), file_path.display().to_string());
+ medias.push(data_url)
+ } else {
+ let mut text = String::new();
+ let mut file = File::open(&file_path)
+ .with_context(|| format!("Unable to open file '{file_item}'"))?;
+ file.read_to_string(&mut text)
+ .with_context(|| format!("Unable to read file '{file_item}'"))?;
+ texts.push(text);
+ }
+ }
+ None => {
+ if is_image_ext(Path::new(&file_item)) {
+ medias.push(file_item)
+ } else {
+ bail!("Unable to use file '{file_item}");
+ }
+ }
+ }
+ }
+
+ Ok(Self {
+ text: texts.join("\n"),
+ medias,
+ data_urls,
+ })
+ }
+
+ pub fn data_urls(&self) -> HashMap<String, String> {
+ self.data_urls.clone()
+ }
+
+ pub fn render(&self) -> String {
+ if self.medias.is_empty() {
+ return self.text.clone();
+ }
+ let text = if self.text.is_empty() {
+ self.text.to_string()
+ } else {
+ format!(" -- {}", self.text)
+ };
+ let files: Vec<String> = self
+ .medias
+ .iter()
+ .cloned()
+ .map(|url| resolve_data_url(&self.data_urls, url))
+ .collect();
+ format!(".file {}{}", files.join(" "), text)
+ }
+
+ pub fn to_message_content(&self) -> MessageContent {
+ if self.medias.is_empty() {
+ MessageContent::Text(self.text.clone())
+ } else {
+ let mut list: Vec<MessageContentPart> = self
+ .medias
+ .iter()
+ .cloned()
+ .map(|url| MessageContentPart::ImageUrl {
+ image_url: ImageUrl { url },
+ })
+ .collect();
+ if !self.text.is_empty() {
+ list.insert(
+ 0,
+ MessageContentPart::Text {
+ text: self.text.clone(),
+ },
+ );
+ }
+ MessageContent::Array(list)
+ }
+ }
+}
+
+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);
+ if let Some(path) = data_urls.get(&hash) {
+ return path.to_string();
+ }
+ data_url
+ } else {
+ data_url
+ }
+}
+
+fn resolve_path(file: &str) -> Option<PathBuf> {
+ if ["https://", "http://", "data:"]
+ .iter()
+ .any(|v| file.starts_with(v))
+ {
+ return None;
+ }
+ let path = if let (Some(file), Some(home)) = (file.strip_prefix('~'), dirs::home_dir()) {
+ home.join(file)
+ } else {
+ std::env::current_dir().ok()?.join(file)
+ };
+ Some(path)
+}
+
+fn is_image_ext(path: &Path) -> bool {
+ path.extension()
+ .map(|v| IMAGE_EXTS.iter().any(|ext| *ext == v.to_string_lossy()))
+ .unwrap_or_default()
+}
+
+fn read_media_to_data_url<P: AsRef<Path>>(image_path: P) -> Result<String> {
+ let mime_type = from_path(&image_path).first_or_octet_stream().to_string();
+
+ let mut file = File::open(image_path)?;
+ let mut buffer = Vec::new();
+ file.read_to_end(&mut buffer)?;
+
+ let encoded_image = STANDARD.encode(buffer);
+ let data_url = format!("data:{};base64,{}", mime_type, encoded_image);
+
+ Ok(data_url)
+}
diff --git a/src/config/mod.rs b/src/config/mod.rs
index 9c86e4c..08509be 100644
--- a/src/config/mod.rs
+++ b/src/config/mod.rs
@@ -1,6 +1,8 @@
+mod input;
mod role;
mod session;
+pub use self::input::Input;
use self::role::Role;
use self::session::{Session, TEMP_SESSION_NAME};
@@ -78,7 +80,7 @@ pub struct Config {
#[serde(skip)]
pub model: Model,
#[serde(skip)]
- pub last_message: Option<(String, String)>,
+ pub last_message: Option<(Input, String)>,
#[serde(skip)]
pub temperature: Option<f64>,
}
@@ -200,15 +202,15 @@ impl Config {
Ok(path)
}
- pub fn save_message(&mut self, input: &str, output: &str) -> Result<()> {
- self.last_message = Some((input.to_string(), output.to_string()));
+ pub fn save_message(&mut self, input: Input, output: &str) -> Result<()> {
+ self.last_message = Some((input.clone(), output.to_string()));
if self.dry_run {
return Ok(());
}
if let Some(session) = self.session.as_mut() {
- session.add_message(input, output)?;
+ session.add_message(&input, output)?;
return Ok(());
}
@@ -220,13 +222,14 @@ impl Config {
return Ok(());
}
let timestamp = now();
+ let input_markdown = input.render();
let output = match self.role.as_ref() {
None => {
- format!("# CHAT:[{timestamp}]\n{input}\n--------\n{output}\n--------\n\n",)
+ format!("# CHAT:[{timestamp}]\n{input_markdown}\n--------\n{output}\n--------\n\n",)
}
Some(v) => {
format!(
- "# CHAT:[{timestamp}] ({})\n{input}\n--------\n{output}\n--------\n\n",
+ "# CHAT:[{timestamp}] ({})\n{input_markdown}\n--------\n{output}\n--------\n\n",
v.name,
)
}
@@ -292,23 +295,23 @@ impl Config {
Ok(())
}
- pub fn echo_messages(&self, content: &str) -> String {
+ pub fn echo_messages(&self, input: &Input) -> String {
if let Some(session) = self.session.as_ref() {
- session.echo_messages(content)
+ session.echo_messages(input)
} else if let Some(role) = self.role.as_ref() {
- role.echo_messages(content)
+ role.echo_messages(input)
} else {
- content.to_string()
+ input.render()
}
}
- pub fn build_messages(&self, content: &str) -> Result<Vec<Message>> {
+ pub fn build_messages(&self, input: &Input) -> Result<Vec<Message>> {
let messages = if let Some(session) = self.session.as_ref() {
- session.build_emssages(content)
+ session.build_emssages(input)
} else if let Some(role) = self.role.as_ref() {
- role.build_messages(content)
+ role.build_messages(input)
} else {
- let message = Message::new(content);
+ let message = Message::new(input);
vec![message]
};
Ok(messages)
@@ -586,7 +589,7 @@ impl Config {
Ok(dir) => dir,
Err(_) => return vec![],
};
- match read_dir(&sessions_dir) {
+ match read_dir(sessions_dir) {
Ok(rd) => {
let mut names = vec![];
for entry in rd.flatten() {
@@ -643,8 +646,8 @@ impl Config {
}
}
- pub fn prepare_send_data(&self, content: &str, stream: bool) -> Result<SendData> {
- let messages = self.build_messages(content)?;
+ pub fn prepare_send_data(&self, input: &Input, stream: bool) -> Result<SendData> {
+ let messages = self.build_messages(input)?;
self.model.max_tokens_limit(&messages)?;
Ok(SendData {
messages,
@@ -653,7 +656,7 @@ impl Config {
})
}
- pub fn maybe_print_send_tokens(&self, input: &str) {
+ pub fn maybe_print_send_tokens(&self, input: &Input) {
if self.dry_run {
if let Ok(messages) = self.build_messages(input) {
let tokens = self.model.total_tokens(&messages);
diff --git a/src/config/role.rs b/src/config/role.rs
index 2b8fea1..bd7216c 100644
--- a/src/config/role.rs
+++ b/src/config/role.rs
@@ -1,8 +1,10 @@
-use crate::client::{Message, MessageRole};
+use crate::client::{Message, MessageContent, MessageRole};
use anyhow::{Context, Result};
use serde::{Deserialize, Serialize};
+use super::Input;
+
const INPUT_PLACEHOLDER: &str = "__INPUT__";
#[derive(Debug, Clone, Deserialize, Serialize)]
@@ -41,17 +43,20 @@ impl Role {
}
}
- pub fn echo_messages(&self, content: &str) -> String {
+ pub fn echo_messages(&self, input: &Input) -> String {
+ let input_markdown = input.render();
if self.embedded() {
- merge_prompt_content(&self.prompt, content)
+ self.prompt.replace(INPUT_PLACEHOLDER, &input_markdown)
} else {
- format!("{}\n\n{content}", self.prompt)
+ format!("{}\n\n{}", self.prompt, input.render())
}
}
- pub fn build_messages(&self, content: &str) -> Vec<Message> {
+ pub fn build_messages(&self, input: &Input) -> Vec<Message> {
+ let mut content = input.to_message_content();
+
if self.embedded() {
- let content = merge_prompt_content(&self.prompt, content);
+ content.merge_prompt(|v: &str| self.prompt.replace(INPUT_PLACEHOLDER, v));
vec![Message {
role: MessageRole::User,
content,
@@ -60,21 +65,17 @@ impl Role {
vec![
Message {
role: MessageRole::System,
- content: self.prompt.clone(),
+ content: MessageContent::Text(self.prompt.clone()),
},
Message {
role: MessageRole::User,
- content: content.to_string(),
+ content,
},
]
}
}
}
-fn merge_prompt_content(prompt: &str, content: &str) -> String {
- prompt.replace(INPUT_PLACEHOLDER, content)
-}
-
fn complete_prompt_args(prompt: &str, name: &str) -> String {
let mut prompt = prompt.trim().to_string();
for (i, arg) in name.split(':').skip(1).enumerate() {
diff --git a/src/config/session.rs b/src/config/session.rs
index 1aebd64..644255f 100644
--- a/src/config/session.rs
+++ b/src/config/session.rs
@@ -1,12 +1,14 @@
+use super::input::resolve_data_url;
use super::role::Role;
-use super::Model;
+use super::{Input, Model};
-use crate::client::{Message, MessageRole};
+use crate::client::{Message, MessageContent, MessageRole};
use crate::render::MarkdownRender;
use anyhow::{bail, Context, Result};
use serde::{Deserialize, Serialize};
use serde_json::json;
+use std::collections::HashMap;
use std::fs::{self, read_to_string};
use std::path::Path;
@@ -18,6 +20,7 @@ pub struct Session {
model_id: String,
temperature: Option<f64>,
messages: Vec<Message>,
+ data_urls: HashMap<String, String>,
#[serde(skip)]
pub name: String,
#[serde(skip)]
@@ -37,6 +40,7 @@ impl Session {
model_id: model.id(),
temperature,
messages: vec![],
+ data_urls: Default::default(),
name: name.to_string(),
path: None,
dirty: false,
@@ -121,6 +125,7 @@ impl Session {
if !self.is_empty() {
lines.push("".into());
+ let resolve_url_fn = |url: &str| resolve_data_url(&self.data_urls, url.to_string());
for message in &self.messages {
match message.role {
@@ -128,11 +133,17 @@ impl Session {
continue;
}
MessageRole::Assistant => {
- lines.push(render.render(&message.content));
+ if let MessageContent::Text(text) = &message.content {
+ lines.push(render.render(text));
+ }
lines.push("".into());
}
MessageRole::User => {
- lines.push(format!("{}){}", self.name, message.content));
+ lines.push(format!(
+ "{}){}",
+ self.name,
+ message.content.render_input(resolve_url_fn)
+ ));
}
}
}
@@ -218,7 +229,7 @@ impl Session {
self.messages.is_empty()
}
- pub fn add_message(&mut self, input: &str, output: &str) -> Result<()> {
+ 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) = self.role.as_ref() {
@@ -229,35 +240,36 @@ impl Session {
if need_add_msg {
self.messages.push(Message {
role: MessageRole::User,
- content: input.to_string(),
+ content: input.to_message_content(),
});
}
+ self.data_urls.extend(input.data_urls());
self.messages.push(Message {
role: MessageRole::Assistant,
- content: output.to_string(),
+ content: MessageContent::Text(output.to_string()),
});
self.dirty = true;
Ok(())
}
- pub fn echo_messages(&self, content: &str) -> String {
- let messages = self.build_emssages(content);
+ pub fn echo_messages(&self, input: &Input) -> String {
+ let messages = self.build_emssages(input);
serde_yaml::to_string(&messages).unwrap_or_else(|_| "Unable to echo message".into())
}
- pub fn build_emssages(&self, content: &str) -> Vec<Message> {
+ pub fn build_emssages(&self, input: &Input) -> Vec<Message> {
let mut messages = self.messages.clone();
let mut need_add_msg = true;
if messages.is_empty() {
if let Some(role) = self.role.as_ref() {
- messages = role.build_messages(content);
+ messages = role.build_messages(input);
need_add_msg = false;
}
};
if need_add_msg {
messages.push(Message {
role: MessageRole::User,
- content: content.into(),
+ content: input.to_message_content(),
});
}
messages