Initial commit
This commit is contained in:
179
src/controller/image/generation.rs
Normal file
179
src/controller/image/generation.rs
Normal file
@@ -0,0 +1,179 @@
|
||||
use mxlink::{MatrixLink, MessageResponseType};
|
||||
|
||||
use tracing::Instrument;
|
||||
|
||||
use crate::agent::provider::ImageGenerationParams;
|
||||
use crate::agent::AgentPurpose;
|
||||
use crate::agent::ControllerTrait;
|
||||
use crate::controller::utils::agent::get_effective_agent_for_purpose_or_complain;
|
||||
use crate::conversation::create_llm_conversation_for_matrix_thread;
|
||||
use crate::conversation::matrix::MatrixMessageProcessingParams;
|
||||
use crate::strings;
|
||||
use crate::{entity::MessageContext, Bot};
|
||||
|
||||
// We may make this configurable (per room, etc.) in the future, but for now it's hardcoded.
|
||||
const STICKER_SIZE: &str = "256x256";
|
||||
|
||||
pub async fn handle_image(
|
||||
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(());
|
||||
};
|
||||
|
||||
let params = MatrixMessageProcessingParams::new(
|
||||
bot.user_id().as_str().to_owned(),
|
||||
message_context.combined_admin_and_user_regexes(),
|
||||
);
|
||||
|
||||
let conversation = create_llm_conversation_for_matrix_thread(
|
||||
matrix_link.clone(),
|
||||
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()
|
||||
};
|
||||
|
||||
message_context.room().typing_notice(true).await?;
|
||||
|
||||
let span = tracing::debug_span!(
|
||||
"image_generation",
|
||||
agent_id = agent.identifier().as_string()
|
||||
);
|
||||
|
||||
let response = agent
|
||||
.controller()
|
||||
.generate_image(&prompt, ImageGenerationParams::default())
|
||||
.instrument(span)
|
||||
.await?;
|
||||
|
||||
let actual_prompt = response.revised_prompt.as_deref().unwrap_or(&prompt);
|
||||
|
||||
if *actual_prompt.trim() != *prompt.trim() {
|
||||
bot.messaging()
|
||||
.send_notice_markdown_no_fail(
|
||||
message_context.room(),
|
||||
strings::image_generation::revised_prompt(actual_prompt),
|
||||
response_type.clone(),
|
||||
)
|
||||
.await;
|
||||
}
|
||||
|
||||
let attachment_body_text = format!("Generated image based on: {}", actual_prompt);
|
||||
|
||||
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?;
|
||||
|
||||
if conversation.messages.len() == 1 {
|
||||
// If this is the beginning of the thread, send helpful instructions
|
||||
bot.messaging()
|
||||
.send_notice_markdown_no_fail(
|
||||
message_context.room(),
|
||||
strings::image_generation::guide_how_to_proceed(),
|
||||
response_type.clone(),
|
||||
)
|
||||
.await;
|
||||
}
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub async fn handle_sticker(
|
||||
bot: &Bot,
|
||||
matrix_link: MatrixLink,
|
||||
message_context: &MessageContext,
|
||||
original_prompt: &str,
|
||||
) -> anyhow::Result<()> {
|
||||
// Stickers are always sent directly to the room - no threading.
|
||||
let response_type =
|
||||
MessageResponseType::Reply(message_context.thread_info().root_event_id.clone());
|
||||
|
||||
let Some(agent) = get_effective_agent_for_purpose_or_complain(
|
||||
bot,
|
||||
message_context,
|
||||
AgentPurpose::ImageGeneration,
|
||||
response_type.clone(),
|
||||
true,
|
||||
)
|
||||
.await
|
||||
else {
|
||||
return Ok(());
|
||||
};
|
||||
|
||||
message_context.room().typing_notice(true).await?;
|
||||
|
||||
let span = tracing::debug_span!(
|
||||
"sticker_generation",
|
||||
agent_id = agent.identifier().as_string()
|
||||
);
|
||||
|
||||
let params = ImageGenerationParams::default()
|
||||
.with_size_override(Some(STICKER_SIZE.to_owned()))
|
||||
.with_cheaper_model_switching_allowed(true)
|
||||
.with_cheaper_quality_switching_allowed(true);
|
||||
|
||||
let response = agent
|
||||
.controller()
|
||||
.generate_image(original_prompt, params)
|
||||
.instrument(span)
|
||||
.await?;
|
||||
|
||||
let attachment_body_text = format!("Generated sticker image based on: {}", original_prompt);
|
||||
|
||||
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)
|
||||
.await?;
|
||||
|
||||
Ok(())
|
||||
}
|
||||
2
src/controller/image/mod.rs
Normal file
2
src/controller/image/mod.rs
Normal file
@@ -0,0 +1,2 @@
|
||||
pub mod generation;
|
||||
mod prompt;
|
||||
111
src/controller/image/prompt.rs
Normal file
111
src/controller/image/prompt.rs
Normal file
@@ -0,0 +1,111 @@
|
||||
use crate::conversation::llm::{Author, Message};
|
||||
|
||||
/// 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.
|
||||
pub fn build(original_prompt: &str, other_messages: Vec<Message>) -> String {
|
||||
let mut prompt = original_prompt.to_owned();
|
||||
|
||||
// Make a new messages vector that only contains messages we care about
|
||||
let other_messages: Vec<Message> = other_messages
|
||||
.into_iter()
|
||||
.filter(|message| {
|
||||
if let Author::User = message.author {
|
||||
message.message_text.to_lowercase() != "again"
|
||||
} else {
|
||||
false
|
||||
}
|
||||
})
|
||||
.collect();
|
||||
|
||||
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(),
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
prompt
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::build;
|
||||
use super::{Author, Message};
|
||||
|
||||
struct TestCase {
|
||||
original_prompt: &'static str,
|
||||
messages: Vec<Message>,
|
||||
expected_prompt: &'static str,
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_build_prompt() {
|
||||
let test_cases = vec![
|
||||
// Simple case
|
||||
TestCase {
|
||||
original_prompt: "Generate a picture of a cat",
|
||||
messages: vec![],
|
||||
expected_prompt: "Generate a picture of a cat",
|
||||
},
|
||||
// Only a single user message
|
||||
TestCase {
|
||||
original_prompt: "Generate a picture of a dog",
|
||||
messages: vec![Message {
|
||||
author: Author::User,
|
||||
message_text: "Must be blue".to_owned(),
|
||||
}],
|
||||
expected_prompt: "Generate a picture of a dog\nOther criteria:\n- Must be blue",
|
||||
},
|
||||
// Multiple complex user messages dispersed with assistant messages
|
||||
TestCase {
|
||||
original_prompt: "Generate a picture of an elephant",
|
||||
messages: vec![Message {
|
||||
author: Author::User,
|
||||
message_text: "Must be blue".to_owned(),
|
||||
},
|
||||
Message {
|
||||
author: Author::Assistant,
|
||||
message_text: "Whatever".to_owned(),
|
||||
},
|
||||
Message {
|
||||
author: Author::User,
|
||||
message_text: "Must be 3-legged.\nMust be flying.".to_owned(),
|
||||
}],
|
||||
expected_prompt: "Generate a picture of an elephant\nOther criteria:\n- Must be blue\n- Must be 3-legged.. Must be flying.",
|
||||
},
|
||||
// "Again" is ignored.
|
||||
TestCase {
|
||||
original_prompt: "Generate a picture of a grizzly bear",
|
||||
messages: vec![Message {
|
||||
author: Author::User,
|
||||
message_text: "Must be blue".to_owned(),
|
||||
},
|
||||
Message {
|
||||
author: Author::Assistant,
|
||||
message_text: "Whatever".to_owned(),
|
||||
},
|
||||
Message {
|
||||
author: Author::User,
|
||||
message_text: "Again".to_owned(),
|
||||
},
|
||||
Message {
|
||||
author: Author::User,
|
||||
message_text: "again".to_owned(),
|
||||
}],
|
||||
expected_prompt: "Generate a picture of a grizzly bear\nOther criteria:\n- Must be blue",
|
||||
},
|
||||
];
|
||||
|
||||
for test_case in test_cases {
|
||||
let actual_prompt = build(test_case.original_prompt, test_case.messages);
|
||||
|
||||
assert_eq!(actual_prompt, test_case.expected_prompt);
|
||||
}
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user