summaryrefslogtreecommitdiffstats
path: root/src/config/role.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/config/role.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/config/role.rs')
-rw-r--r--src/config/role.rs25
1 files changed, 13 insertions, 12 deletions
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() {