Initial work on Vision support in text conversations and Image Editing
This is a huge patch which does some major refactoring like: - renaming "Image Generation" to "Image Creation" in most places, to better match its new command (`!bai image create`) - relocating image creation command (`!bai image` -> `!bai image create`), so it wouldn't conflict with the new image editing command (`!bai image edit`) - introducing a new image editing command (`!bai image edit`), which is meant to work only with the OpenAI provider, but doesn't fully work yet due to https://github.com/64bit/async-openai/issues/364, though a next patch will fix it - adding support for reading images off of Matrix conversations and forwarding them to text conversations. Works for OpenAI, but not for Anthropic yet (requires custom patches) and not for OpenAI-Compat (no support for images there) - relocating some utils around (base64, mime)
This commit is contained in:
@@ -69,7 +69,7 @@ pub async fn handle(bot: &Bot, message_context: &MessageContext) -> anyhow::Resu
|
||||
);
|
||||
message.push_str("\n\n");
|
||||
|
||||
// Image Generation
|
||||
// Image Creation
|
||||
message.push_str(
|
||||
&generate_image_generation_section(agent_manager, message_context.room_config_context())
|
||||
.await,
|
||||
|
||||
@@ -39,6 +39,8 @@ pub enum ChatCompletionControllerType {
|
||||
|
||||
Audio,
|
||||
|
||||
Image,
|
||||
|
||||
ThreadMention,
|
||||
ReplyMention,
|
||||
}
|
||||
@@ -416,7 +418,8 @@ async fn handle_stage_text_generation(
|
||||
ChatCompletionControllerType::TextCommand
|
||||
| ChatCompletionControllerType::TextMention
|
||||
| ChatCompletionControllerType::TextDirect
|
||||
| ChatCompletionControllerType::Audio => {
|
||||
| ChatCompletionControllerType::Audio
|
||||
| ChatCompletionControllerType::Image => {
|
||||
Some(message_context.combined_admin_and_user_regexes())
|
||||
}
|
||||
|
||||
@@ -438,6 +441,7 @@ async fn handle_stage_text_generation(
|
||||
// When we're triggered via a reply mention, the context is the whole reply chain upward of the message that triggered us.
|
||||
ChatCompletionControllerType::ReplyMention => {
|
||||
create_llm_conversation_for_matrix_reply_chain(
|
||||
&matrix_link,
|
||||
&bot.room_event_fetcher().clone(),
|
||||
message_context.room(),
|
||||
message_context.thread_info().last_event_id.clone(),
|
||||
@@ -449,7 +453,7 @@ async fn handle_stage_text_generation(
|
||||
// Everything else is happening in a thread, so the context is the whole thread.
|
||||
_ => {
|
||||
create_llm_conversation_for_matrix_thread(
|
||||
matrix_link.clone(),
|
||||
&matrix_link,
|
||||
message_context.room(),
|
||||
message_context.thread_info().root_event_id.clone(),
|
||||
¶ms,
|
||||
|
||||
@@ -23,5 +23,6 @@ pub enum ControllerType {
|
||||
ChatCompletion(super::chat_completion::ChatCompletionControllerType),
|
||||
|
||||
ImageGeneration(String),
|
||||
ImageEdit(String),
|
||||
StickerGeneration(String),
|
||||
}
|
||||
|
||||
@@ -36,6 +36,20 @@ pub fn determine_controller(
|
||||
first_thread_message.is_mentioning_bot,
|
||||
)
|
||||
}
|
||||
MessagePayload::Image(_image_message_content) => {
|
||||
let prefix_requirement_type = message_context
|
||||
.room_config_context()
|
||||
.text_generation_prefix_requirement_type();
|
||||
|
||||
match prefix_requirement_type {
|
||||
TextGenerationPrefixRequirementType::CommandPrefix => {
|
||||
ControllerType::Ignore
|
||||
}
|
||||
TextGenerationPrefixRequirementType::No => {
|
||||
ControllerType::ChatCompletion(ChatCompletionControllerType::Image)
|
||||
}
|
||||
}
|
||||
}
|
||||
MessagePayload::Encrypted(thread_info) => {
|
||||
if thread_info.is_thread_root_only() {
|
||||
ControllerType::Error(strings::error::message_is_encrypted().to_owned())
|
||||
@@ -84,7 +98,7 @@ fn determine_text_controller(
|
||||
}
|
||||
|
||||
if let Some(prompt) = text.strip_prefix(&format!("{command_prefix} image")) {
|
||||
return ControllerType::ImageGeneration(prompt.trim().to_owned());
|
||||
return super::image::determine_controller(prompt.trim());
|
||||
}
|
||||
|
||||
if let Some(prompt) = text.strip_prefix(&format!("{command_prefix} sticker")) {
|
||||
|
||||
@@ -84,9 +84,17 @@ fn determine_text_controller() {
|
||||
expected: ControllerType::Config(controller::cfg::ConfigControllerType::Help),
|
||||
},
|
||||
TestCase {
|
||||
name: "Image generation",
|
||||
name: "Generic image command causes usage help",
|
||||
input: "!bai image Draw a cat!",
|
||||
is_mentioning_bot: false,
|
||||
room_text_generation_prefix_requirement_type:
|
||||
super::TextGenerationPrefixRequirementType::No,
|
||||
expected: ControllerType::UsageHelp,
|
||||
},
|
||||
TestCase {
|
||||
name: "Image generation",
|
||||
input: "!bai image create Draw a cat!",
|
||||
is_mentioning_bot: false,
|
||||
room_text_generation_prefix_requirement_type:
|
||||
super::TextGenerationPrefixRequirementType::No,
|
||||
expected: ControllerType::ImageGeneration("Draw a cat!".to_owned()),
|
||||
|
||||
@@ -77,6 +77,15 @@ pub async fn dispatch_controller(
|
||||
)
|
||||
.await
|
||||
}
|
||||
ControllerType::ImageEdit(prompt) => {
|
||||
super::image::edit::handle(
|
||||
bot,
|
||||
bot.matrix_link().clone(),
|
||||
message_context,
|
||||
prompt,
|
||||
)
|
||||
.await
|
||||
}
|
||||
ControllerType::StickerGeneration(prompt) => {
|
||||
super::image::generation::handle_sticker(
|
||||
bot,
|
||||
|
||||
18
src/controller/image/determination/mod.rs
Normal file
18
src/controller/image/determination/mod.rs
Normal file
@@ -0,0 +1,18 @@
|
||||
use crate::controller::ControllerType;
|
||||
mod tests;
|
||||
|
||||
pub fn determine_controller(
|
||||
text: &str,
|
||||
) -> ControllerType {
|
||||
let text = text.trim();
|
||||
|
||||
if let Some(prompt) = text.strip_prefix(&format!("create")) {
|
||||
return ControllerType::ImageGeneration(prompt.trim().to_owned());
|
||||
}
|
||||
|
||||
if let Some(prompt) = text.strip_prefix(&format!("edit")) {
|
||||
return ControllerType::ImageEdit(prompt.trim().to_owned());
|
||||
}
|
||||
|
||||
ControllerType::UsageHelp
|
||||
}
|
||||
37
src/controller/image/determination/tests.rs
Normal file
37
src/controller/image/determination/tests.rs
Normal file
@@ -0,0 +1,37 @@
|
||||
#[test]
|
||||
fn determine_controller() {
|
||||
struct TestCase {
|
||||
name: &'static str,
|
||||
input: &'static str,
|
||||
expected: super::ControllerType,
|
||||
}
|
||||
|
||||
let test_cases = vec![
|
||||
TestCase {
|
||||
name: "Top-level is usage help",
|
||||
input: "",
|
||||
expected: super::ControllerType::UsageHelp,
|
||||
},
|
||||
TestCase {
|
||||
name: "Top-level with some text is usage help",
|
||||
input: "Some text",
|
||||
expected: super::ControllerType::UsageHelp,
|
||||
},
|
||||
TestCase {
|
||||
name: "Image generation triggered by create prefix",
|
||||
input: "create Some prompt",
|
||||
expected: super::ControllerType::ImageGeneration("Some prompt".to_owned()),
|
||||
},
|
||||
|
||||
TestCase {
|
||||
name: "Image edit triggered by edit prefix",
|
||||
input: "edit Turn this into an anime-style image",
|
||||
expected: super::ControllerType::ImageEdit("Turn this into an anime-style image".to_owned()),
|
||||
},
|
||||
];
|
||||
|
||||
for test_case in test_cases {
|
||||
let result = super::determine_controller(test_case.input);
|
||||
assert_eq!(result, test_case.expected, "Test case: {}", test_case.name);
|
||||
}
|
||||
}
|
||||
163
src/controller/image/edit.rs
Normal file
163
src/controller/image/edit.rs
Normal file
@@ -0,0 +1,163 @@
|
||||
use mxlink::{MatrixLink, MessageResponseType};
|
||||
|
||||
use tracing::Instrument;
|
||||
|
||||
use crate::agent::provider::ImageSource;
|
||||
use crate::agent::AgentPurpose;
|
||||
use crate::agent::ControllerTrait;
|
||||
use crate::agent::provider::ImageEditParams;
|
||||
use crate::controller::utils::agent::get_effective_agent_for_purpose_or_complain;
|
||||
use crate::utils::mime::get_file_extension;
|
||||
use crate::conversation::create_llm_conversation_for_matrix_thread;
|
||||
use crate::conversation::matrix::MatrixMessageProcessingParams;
|
||||
use crate::strings;
|
||||
use crate::{Bot, entity::MessageContext};
|
||||
|
||||
pub async fn handle(
|
||||
bot: &Bot,
|
||||
matrix_link: MatrixLink,
|
||||
message_context: &MessageContext,
|
||||
original_prompt: &str,
|
||||
) -> anyhow::Result<()> {
|
||||
let response_type = MessageResponseType::InThread(message_context.thread_info().clone());
|
||||
|
||||
let Some(agent) = get_effective_agent_for_purpose_or_complain(
|
||||
bot,
|
||||
message_context,
|
||||
AgentPurpose::ImageGeneration,
|
||||
response_type.clone(),
|
||||
true,
|
||||
)
|
||||
.await
|
||||
else {
|
||||
return Ok(());
|
||||
};
|
||||
|
||||
if message_context.thread_info().is_thread_root_only() {
|
||||
return send_guide(bot, message_context).await;
|
||||
}
|
||||
|
||||
let _typing_notice_guard = bot.start_typing_notice(message_context.room()).await;
|
||||
|
||||
let params = MatrixMessageProcessingParams::new(
|
||||
bot.user_id().to_owned(),
|
||||
Some(message_context.combined_admin_and_user_regexes()),
|
||||
);
|
||||
|
||||
let conversation = create_llm_conversation_for_matrix_thread(
|
||||
&matrix_link,
|
||||
message_context.room(),
|
||||
message_context.thread_info().root_event_id.clone(),
|
||||
¶ms,
|
||||
)
|
||||
.await?;
|
||||
|
||||
let prompt = if conversation.messages.len() >= 2 {
|
||||
// Skip the first message, which contains the original prompt (which we already have)
|
||||
let other_messages = conversation.messages.iter().skip(1).cloned().collect();
|
||||
|
||||
super::prompt::build(original_prompt, other_messages)
|
||||
} else {
|
||||
original_prompt.to_owned()
|
||||
};
|
||||
|
||||
let got_go_signal = conversation.messages.iter().any(|message| {
|
||||
if let crate::conversation::llm::MessageContent::Text(text) = &message.content {
|
||||
text.to_lowercase() == "go"
|
||||
} else {
|
||||
false
|
||||
}
|
||||
});
|
||||
|
||||
let image_sources: Vec<ImageSource> = conversation.messages.iter().filter_map(|message| {
|
||||
if let crate::conversation::llm::MessageContent::Image(image_content) = &message.content {
|
||||
Some(image_content.into())
|
||||
} else {
|
||||
None
|
||||
}
|
||||
}).collect();
|
||||
|
||||
if !got_go_signal || image_sources.is_empty() {
|
||||
// We don't send the guide again here to avoid being annoying.
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
let span = tracing::debug_span!(
|
||||
"image_edit",
|
||||
agent_id = agent.identifier().as_string()
|
||||
);
|
||||
|
||||
let result = agent
|
||||
.controller()
|
||||
.create_image_edit(&prompt, image_sources, ImageEditParams::default())
|
||||
.instrument(span)
|
||||
.await;
|
||||
|
||||
let response = match result {
|
||||
Ok(response) => response,
|
||||
Err(err) => {
|
||||
tracing::warn!(
|
||||
"Error in room {} while trying to generate image edit via agent {}: {:?}",
|
||||
message_context.room_id(),
|
||||
agent.identifier(),
|
||||
err,
|
||||
);
|
||||
|
||||
bot.messaging()
|
||||
.send_error_markdown_no_fail(
|
||||
message_context.room(),
|
||||
&strings::agent::error_while_serving_purpose(
|
||||
agent.identifier(),
|
||||
&AgentPurpose::ImageGeneration,
|
||||
&err,
|
||||
),
|
||||
response_type,
|
||||
)
|
||||
.await;
|
||||
|
||||
return Ok(());
|
||||
}
|
||||
};
|
||||
|
||||
let attachment_body_text = format!(
|
||||
"generated-image-edit.{}",
|
||||
get_file_extension(&response.mime_type)
|
||||
);
|
||||
|
||||
let mut event_content = matrix_link
|
||||
.media()
|
||||
.upload_and_prepare_event_content(
|
||||
message_context.room(),
|
||||
&response.mime_type,
|
||||
response.bytes,
|
||||
&attachment_body_text,
|
||||
)
|
||||
.await
|
||||
.map_err(|e| anyhow::anyhow!("Failed to upload and prepare event: {}", e))?;
|
||||
|
||||
matrix_link
|
||||
.messaging()
|
||||
.send_event(
|
||||
message_context.room(),
|
||||
&mut event_content,
|
||||
response_type.clone(),
|
||||
)
|
||||
.await?;
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn send_guide(
|
||||
bot: &Bot,
|
||||
message_context: &MessageContext,
|
||||
) -> anyhow::Result<()> {
|
||||
bot.messaging()
|
||||
.send_text_markdown_no_fail(
|
||||
message_context.room(),
|
||||
strings::image_edit::guide_how_to_proceed(),
|
||||
MessageResponseType::InThread(message_context.thread_info().clone()),
|
||||
)
|
||||
.await;
|
||||
|
||||
Ok(())
|
||||
}
|
||||
@@ -6,7 +6,7 @@ use crate::agent::AgentPurpose;
|
||||
use crate::agent::ControllerTrait;
|
||||
use crate::agent::provider::ImageGenerationParams;
|
||||
use crate::controller::utils::agent::get_effective_agent_for_purpose_or_complain;
|
||||
use crate::controller::utils::mime::get_file_extension;
|
||||
use crate::utils::mime::get_file_extension;
|
||||
use crate::conversation::create_llm_conversation_for_matrix_thread;
|
||||
use crate::conversation::matrix::MatrixMessageProcessingParams;
|
||||
use crate::strings;
|
||||
@@ -43,7 +43,7 @@ pub async fn handle_image(
|
||||
);
|
||||
|
||||
let conversation = create_llm_conversation_for_matrix_thread(
|
||||
matrix_link.clone(),
|
||||
&matrix_link,
|
||||
message_context.room(),
|
||||
message_context.thread_info().root_event_id.clone(),
|
||||
¶ms,
|
||||
|
||||
@@ -1,2 +1,6 @@
|
||||
pub mod generation;
|
||||
pub mod edit;
|
||||
mod prompt;
|
||||
mod determination;
|
||||
|
||||
pub use determination::determine_controller;
|
||||
|
||||
@@ -1,11 +1,11 @@
|
||||
use crate::conversation::llm::{Author, Message};
|
||||
use crate::conversation::llm::{Author, Message, MessageContent};
|
||||
|
||||
/// Builds a prompt from the original prompt and other messages in the conversation.
|
||||
///
|
||||
/// Only messages authored by the user are considered.
|
||||
///
|
||||
/// Messages that say "Again" (regardless of casing) are ignored. They are considered special messages
|
||||
/// which trigger re-generation, but do not need to be included in the prompt criteria.
|
||||
/// Messages that say "Again" or "Go" (regardless of casing) are ignored. They are considered special messages
|
||||
/// which trigger re-generation and "start" respectively, and do not need to be included in the prompt criteria.
|
||||
pub fn build(original_prompt: &str, other_messages: Vec<Message>) -> String {
|
||||
let mut prompt = original_prompt.to_owned();
|
||||
|
||||
@@ -14,7 +14,11 @@ pub fn build(original_prompt: &str, other_messages: Vec<Message>) -> String {
|
||||
.into_iter()
|
||||
.filter(|message| {
|
||||
if let Author::User = message.author {
|
||||
message.message_text.to_lowercase() != "again"
|
||||
if let MessageContent::Text(text) = &message.content {
|
||||
text.to_lowercase() != "again" && text.to_lowercase() != "go"
|
||||
} else {
|
||||
false
|
||||
}
|
||||
} else {
|
||||
false
|
||||
}
|
||||
@@ -24,9 +28,11 @@ pub fn build(original_prompt: &str, other_messages: Vec<Message>) -> String {
|
||||
if !other_messages.is_empty() {
|
||||
prompt.push_str("\nOther criteria:");
|
||||
for message in other_messages {
|
||||
prompt.push_str(
|
||||
format!("\n- {}", message.message_text.replace("\n", ". ").as_str()).as_str(),
|
||||
);
|
||||
if let MessageContent::Text(text) = &message.content {
|
||||
prompt.push_str(
|
||||
format!("\n- {}", text.replace("\n", ". ").as_str()).as_str(),
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -36,7 +42,7 @@ pub fn build(original_prompt: &str, other_messages: Vec<Message>) -> String {
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::build;
|
||||
use super::{Author, Message};
|
||||
use super::{Author, Message, MessageContent};
|
||||
|
||||
struct TestCase {
|
||||
original_prompt: &'static str,
|
||||
@@ -60,7 +66,7 @@ mod tests {
|
||||
original_prompt: "Generate a picture of a dog",
|
||||
messages: vec![Message {
|
||||
author: Author::User,
|
||||
message_text: "Must be blue".to_owned(),
|
||||
content: MessageContent::Text("Must be blue".to_owned()),
|
||||
timestamp,
|
||||
}],
|
||||
expected_prompt: "Generate a picture of a dog\nOther criteria:\n- Must be blue",
|
||||
@@ -71,17 +77,17 @@ mod tests {
|
||||
messages: vec![
|
||||
Message {
|
||||
author: Author::User,
|
||||
message_text: "Must be blue".to_owned(),
|
||||
content: MessageContent::Text("Must be blue".to_owned()),
|
||||
timestamp,
|
||||
},
|
||||
Message {
|
||||
author: Author::Assistant,
|
||||
message_text: "Whatever".to_owned(),
|
||||
content: MessageContent::Text("Whatever".to_owned()),
|
||||
timestamp,
|
||||
},
|
||||
Message {
|
||||
author: Author::User,
|
||||
message_text: "Must be 3-legged.\nMust be flying.".to_owned(),
|
||||
content: MessageContent::Text("Must be 3-legged.\nMust be flying.".to_owned()),
|
||||
timestamp,
|
||||
},
|
||||
],
|
||||
@@ -93,22 +99,22 @@ mod tests {
|
||||
messages: vec![
|
||||
Message {
|
||||
author: Author::User,
|
||||
message_text: "Must be blue".to_owned(),
|
||||
content: MessageContent::Text("Must be blue".to_owned()),
|
||||
timestamp,
|
||||
},
|
||||
Message {
|
||||
author: Author::Assistant,
|
||||
message_text: "Whatever".to_owned(),
|
||||
content: MessageContent::Text("Whatever".to_owned()),
|
||||
timestamp,
|
||||
},
|
||||
Message {
|
||||
author: Author::User,
|
||||
message_text: "Again".to_owned(),
|
||||
content: MessageContent::Text("Again".to_owned()),
|
||||
timestamp,
|
||||
},
|
||||
Message {
|
||||
author: Author::User,
|
||||
message_text: "again".to_owned(),
|
||||
content: MessageContent::Text("again".to_owned()),
|
||||
timestamp,
|
||||
},
|
||||
],
|
||||
|
||||
@@ -1,17 +0,0 @@
|
||||
use mxlink::mime;
|
||||
|
||||
pub fn get_file_extension(mime_type: &mime::Mime) -> String {
|
||||
match (mime_type.type_(), mime_type.subtype()) {
|
||||
(mime::AUDIO, mime::BASIC) => "au",
|
||||
(mime::AUDIO, mime::MPEG) => "mp3",
|
||||
(mime::AUDIO, mime::MP4) => "m4a",
|
||||
(mime::AUDIO, mime::OGG) => "ogg",
|
||||
(mime::IMAGE, mime::BMP) => "bmp",
|
||||
(mime::IMAGE, mime::GIF) => "gif",
|
||||
(mime::IMAGE, mime::JPEG) => "jpg",
|
||||
(mime::IMAGE, mime::PNG) => "png",
|
||||
(mime::IMAGE, mime::SVG) => "svg",
|
||||
_ => "bin",
|
||||
}
|
||||
.to_string()
|
||||
}
|
||||
@@ -6,7 +6,6 @@ use crate::{
|
||||
};
|
||||
|
||||
pub mod agent;
|
||||
pub(super) mod mime;
|
||||
pub mod text_to_speech;
|
||||
|
||||
pub async fn get_text_body_or_complain<'a>(
|
||||
|
||||
@@ -3,7 +3,7 @@ use mxlink::{MatrixLink, MessageResponseType};
|
||||
|
||||
use tracing::Instrument;
|
||||
|
||||
use crate::controller::utils::mime::get_file_extension;
|
||||
use crate::utils::mime::get_file_extension;
|
||||
use crate::{
|
||||
Bot,
|
||||
agent::{AgentInstance, AgentPurpose, ControllerTrait, provider::TextToSpeechParams},
|
||||
|
||||
Reference in New Issue
Block a user