diff options
| author | sigoden <sigoden@gmail.com> | 2023-11-27 14:04:50 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2023-11-27 14:04:50 +0800 |
| commit | 35c75506e2fd94e2285d8a3cb66208d518d5f992 (patch) | |
| tree | 97fb643bb5e231022442c1cc617f2daeaa152a66 /src/client/message.rs | |
| parent | 5bfe95d31110e75e84626598f033805b0ae4326c (diff) | |
| download | aichat-35c75506e2fd94e2285d8a3cb66208d518d5f992.tar.gz | |
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
Diffstat (limited to 'src/client/message.rs')
| -rw-r--r-- | src/client/message.rs | 69 |
1 files changed, 65 insertions, 4 deletions
diff --git a/src/client/message.rs b/src/client/message.rs index 55b2663..dc8c3e1 100644 --- a/src/client/message.rs +++ b/src/client/message.rs @@ -1,16 +1,18 @@ +use crate::config::Input; + use serde::{Deserialize, Serialize}; #[derive(Debug, Clone, Deserialize, Serialize)] pub struct Message { pub role: MessageRole, - pub content: String, + pub content: MessageContent, } impl Message { - pub fn new(content: &str) -> Self { + pub fn new(input: &Input) -> Self { Self { role: MessageRole::User, - content: content.to_string(), + content: input.to_message_content(), } } } @@ -38,6 +40,65 @@ impl MessageRole { } } +#[derive(Debug, Clone, Deserialize, Serialize)] +#[serde(untagged)] +pub enum MessageContent { + Text(String), + Array(Vec<MessageContentPart>), +} + +impl MessageContent { + pub fn render_input(&self, resolve_url_fn: impl Fn(&str) -> String) -> String { + match self { + MessageContent::Text(text) => text.to_string(), + MessageContent::Array(list) => { + let (mut concated_text, mut files) = (String::new(), vec![]); + for item in list { + match item { + MessageContentPart::Text { text } => { + concated_text = format!("{concated_text} {text}") + } + MessageContentPart::ImageUrl { image_url } => { + files.push(resolve_url_fn(&image_url.url)) + } + } + } + if !concated_text.is_empty() { + concated_text = format!(" -- {concated_text}") + } + format!(".file {}{}", files.join(" "), concated_text) + } + } + } + + pub fn merge_prompt(&mut self, replace_fn: impl Fn(&str) -> String) { + match self { + MessageContent::Text(text) => *text = replace_fn(text), + MessageContent::Array(list) => { + if list.is_empty() { + list.push(MessageContentPart::Text { + text: replace_fn(""), + }) + } else if let Some(MessageContentPart::Text { text }) = list.get_mut(0) { + *text = replace_fn(text) + } + } + } + } +} + +#[derive(Debug, Clone, Deserialize, Serialize)] +#[serde(tag = "type", rename_all = "snake_case")] +pub enum MessageContentPart { + Text { text: String }, + ImageUrl { image_url: ImageUrl }, +} + +#[derive(Debug, Clone, Deserialize, Serialize)] +pub struct ImageUrl { + pub url: String, +} + #[cfg(test)] mod tests { use super::*; @@ -45,7 +106,7 @@ mod tests { #[test] fn test_serde() { assert_eq!( - serde_json::to_string(&Message::new("Hello World")).unwrap(), + serde_json::to_string(&Message::new(&Input::from_str("Hello World"))).unwrap(), "{\"role\":\"user\",\"content\":\"Hello World\"}" ); } |
