From 35c75506e2fd94e2285d8a3cb66208d518d5f992 Mon Sep 17 00:00:00 2001 From: sigoden Date: Mon, 27 Nov 2023 14:04:50 +0800 Subject: feat: support vision (#249) * feat: support vision * clippy * implement vision * resolve data url to local file * add model openai:gpt-4-vision-preview * use newline to concate embeded text files * set max_tokens for gpt-4-vision-preview --- src/config/input.rs | 162 ++++++++++++++++++++++++++++++++++++++++++++++++++ src/config/mod.rs | 39 ++++++------ src/config/role.rs | 25 ++++---- src/config/session.rs | 36 +++++++---- 4 files changed, 220 insertions(+), 42 deletions(-) create mode 100644 src/config/input.rs (limited to 'src/config') 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, + data_urls: HashMap, +} + +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) -> Result { + 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 { + 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 = 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 = 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, 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 { + 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>(image_path: P) -> Result { + 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, } @@ -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> { + pub fn build_messages(&self, input: &Input) -> Result> { 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 { - let messages = self.build_messages(content)?; + pub fn prepare_send_data(&self, input: &Input, stream: bool) -> Result { + 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 { + pub fn build_messages(&self, input: &Input) -> Vec { + 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, messages: Vec, + data_urls: HashMap, #[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 { + pub fn build_emssages(&self, input: &Input) -> Vec { 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 -- cgit v1.2.3