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:
Slavi Pantaleev
2025-05-10 09:18:01 +03:00
parent e0dcc39a72
commit 8f86289373
57 changed files with 1074 additions and 335 deletions

View File

@@ -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,

View File

@@ -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(),
&params,

View File

@@ -23,5 +23,6 @@ pub enum ControllerType {
ChatCompletion(super::chat_completion::ChatCompletionControllerType),
ImageGeneration(String),
ImageEdit(String),
StickerGeneration(String),
}

View File

@@ -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")) {

View File

@@ -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()),

View File

@@ -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,

View 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
}

View 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);
}
}

View 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(),
&params,
)
.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(())
}

View File

@@ -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(),
&params,

View File

@@ -1,2 +1,6 @@
pub mod generation;
pub mod edit;
mod prompt;
mod determination;
pub use determination::determine_controller;

View File

@@ -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,
},
],

View File

@@ -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()
}

View File

@@ -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>(

View File

@@ -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},