summaryrefslogtreecommitdiffstats
path: root/src/client/message.rs
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2023-11-27 14:04:50 +0800
committerGitHub <noreply@github.com>2023-11-27 14:04:50 +0800
commit35c75506e2fd94e2285d8a3cb66208d518d5f992 (patch)
tree97fb643bb5e231022442c1cc617f2daeaa152a66 /src/client/message.rs
parent5bfe95d31110e75e84626598f033805b0ae4326c (diff)
downloadaichat-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.rs69
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\"}"
);
}