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

@@ -7,12 +7,14 @@ use anthropic::types::ContentBlock;
use super::super::ControllerTrait;
use crate::agent::AgentPurpose;
use crate::agent::provider::entity::{
ImageGenerationResult, PingResult, TextGenerationParams, TextGenerationResult,
TextToSpeechParams, TextToSpeechResult,
ImageGenerationResult, ImageEditResult, ImageSource, PingResult, TextGenerationParams,
TextGenerationResult, TextToSpeechParams, TextToSpeechResult,
};
use crate::agent::provider::{
ImageGenerationParams, ImageEditParams, SpeechToTextParams, SpeechToTextResult,
};
use crate::agent::provider::{ImageGenerationParams, SpeechToTextParams, SpeechToTextResult};
use crate::conversation::llm::{
Author as LLMAuthor, Conversation as LLMConversation, Message as LLMMessage,
Author as LLMAuthor, Conversation as LLMConversation, Message as LLMMessage, MessageContent as LLMMessageContent,
shorten_messages_list_to_context_size,
};
use crate::strings;
@@ -69,7 +71,7 @@ impl ControllerTrait for Controller {
let messages = vec![LLMMessage {
author: LLMAuthor::User,
message_text: "Hello!".to_string(),
content: LLMMessageContent::Text("Hello!".to_string()),
timestamp: chrono::Utc::now(),
}];
@@ -106,7 +108,7 @@ impl ControllerTrait for Controller {
} else {
Some(LLMMessage {
author: LLMAuthor::Prompt,
message_text: prompt_text,
content: LLMMessageContent::Text(prompt_text),
timestamp: chrono::Utc::now(),
})
};
@@ -145,7 +147,9 @@ impl ControllerTrait for Controller {
.unwrap_or(text_generation_config.temperature);
if let Some(prompt_message) = prompt_message {
request.system = prompt_message.message_text;
if let LLMMessageContent::Text(text) = &prompt_message.content {
request.system = text.clone();
}
}
request.model = text_generation_config.model_id.clone();
@@ -213,6 +217,15 @@ impl ControllerTrait for Controller {
Err(anyhow::anyhow!("Image generation not supported"))
}
async fn create_image_edit(
&self,
_prompt: &str,
_images: Vec<ImageSource>,
_params: ImageEditParams,
) -> anyhow::Result<ImageEditResult> {
Err(anyhow::anyhow!("Image editing is not supported"))
}
async fn text_to_speech(
&self,
_input: &str,

View File

@@ -1,6 +1,6 @@
use anthropic::types::{ContentBlock, Message, MessagesRequest, MessagesRequestBuilder, Role};
use crate::conversation::llm::{Author as LLMAuthor, Message as LLMMessage};
use crate::conversation::llm::{Author as LLMAuthor, Message as LLMMessage, MessageContent as LLMMessageContent};
pub(super) fn create_anthropic_message_request(llm_messages: Vec<LLMMessage>) -> MessagesRequest {
let mut messages = vec![];
@@ -14,13 +14,28 @@ pub(super) fn create_anthropic_message_request(llm_messages: Vec<LLMMessage>) ->
}
};
let content = vec![ContentBlock::Text {
text: message.message_text,
}];
let content = match &message.content {
LLMMessageContent::Text(text) => Some(vec![ContentBlock::Text { text: text.clone() }]),
LLMMessageContent::Image(_image_details) => {
// This cannot be implemented yet, because the Anthropic library does not support it.
// The code below requires the library to be forked and changed a bit.
// vec![ContentBlock::Image {
// source: ImageSource {
// r#type: "base64".to_string(),
// media_type: mime_type.to_string(),
// data: crate::utils::base64::base64_encode(image_details.data),
// },
// }]
tracing::warn!("Image content is not supported by the Anthropic library yet. Skipping it.");
None
}
};
let message = Message { role, content };
if let Some(content) = content {
let message = Message { role, content };
messages.push(message);
messages.push(message);
}
}
MessagesRequestBuilder::default()

View File

@@ -1,10 +1,10 @@
use crate::{agent::AgentPurpose, conversation::llm::Conversation};
use super::{
ImageGenerationParams, SpeechToTextParams, SpeechToTextResult,
ImageGenerationParams, ImageEditParams, SpeechToTextParams, SpeechToTextResult,
entity::{
ImageGenerationResult, PingResult, TextGenerationParams, TextGenerationResult,
TextToSpeechParams, TextToSpeechResult,
ImageGenerationResult, ImageEditResult, ImageSource, PingResult, TextGenerationParams,
TextGenerationResult, TextToSpeechParams, TextToSpeechResult,
},
};
@@ -42,6 +42,13 @@ pub trait ControllerTrait {
params: ImageGenerationParams,
) -> impl std::future::Future<Output = anyhow::Result<ImageGenerationResult>> + Send;
fn create_image_edit(
&self,
prompt: &str,
images: Vec<ImageSource>,
params: ImageEditParams,
) -> impl std::future::Future<Output = anyhow::Result<ImageEditResult>> + Send;
fn text_to_speech(
&self,
text: &str,
@@ -166,6 +173,23 @@ impl ControllerTrait for ControllerType {
}
}
async fn create_image_edit(
&self,
prompt: &str,
images: Vec<ImageSource>,
params: ImageEditParams,
) -> anyhow::Result<ImageEditResult> {
match &self {
ControllerType::OpenAI(controller) => controller.create_image_edit(prompt, images, params).await,
ControllerType::OpenAICompat(controller) => {
controller.create_image_edit(prompt, images, params).await
}
ControllerType::Anthropic(controller) => {
controller.create_image_edit(prompt, images, params).await
}
}
}
async fn text_to_speech(
&self,
text: &str,

View File

@@ -0,0 +1,65 @@
use mxlink::mime;
#[derive(Default)]
pub struct ImageGenerationParams {
pub size_override: Option<String>,
pub cheaper_model_switching_allowed: bool,
pub cheaper_quality_switching_allowed: bool,
}
impl ImageGenerationParams {
pub fn with_size_override(mut self, value: Option<String>) -> Self {
self.size_override = value;
self
}
pub fn with_cheaper_model_switching_allowed(mut self, value: bool) -> Self {
self.cheaper_model_switching_allowed = value;
self
}
pub fn with_cheaper_quality_switching_allowed(mut self, value: bool) -> Self {
self.cheaper_quality_switching_allowed = value;
self
}
}
pub struct ImageGenerationResult {
pub bytes: Vec<u8>,
pub mime_type: mime::Mime,
pub revised_prompt: Option<String>,
}
#[derive(Default)]
pub struct ImageEditParams {
}
pub struct ImageEditResult {
pub bytes: Vec<u8>,
pub mime_type: mime::Mime,
}
pub struct ImageSource {
pub filename: String,
pub bytes: Vec<u8>,
pub mime_type: mime::Mime,
}
impl ImageSource {
pub fn new(filename: String, bytes: Vec<u8>, mime_type: mime::Mime) -> Self {
Self { filename, bytes, mime_type }
}
}
impl Into<async_openai::types::ImageInput> for ImageSource {
fn into(self) -> async_openai::types::ImageInput {
async_openai::types::ImageInput{
source: async_openai::types::InputSource::VecU8 {
filename: self.filename,
vec: self.bytes,
},
}
}
}

View File

@@ -1,31 +0,0 @@
#[derive(Default)]
pub struct ImageGenerationParams {
pub size_override: Option<String>,
pub cheaper_model_switching_allowed: bool,
pub cheaper_quality_switching_allowed: bool,
}
impl ImageGenerationParams {
pub fn with_size_override(mut self, value: Option<String>) -> Self {
self.size_override = value;
self
}
pub fn with_cheaper_model_switching_allowed(mut self, value: bool) -> Self {
self.cheaper_model_switching_allowed = value;
self
}
pub fn with_cheaper_quality_switching_allowed(mut self, value: bool) -> Self {
self.cheaper_quality_switching_allowed = value;
self
}
}
pub struct ImageGenerationResult {
pub bytes: Vec<u8>,
pub mime_type: mxlink::mime::Mime,
pub revised_prompt: Option<String>,
}

View File

@@ -1,12 +1,12 @@
mod agent_provider;
mod image_generation;
mod image;
mod ping;
mod speech_to_text;
mod text_generation;
mod text_to_speech;
pub use agent_provider::{AgentProvider, AgentProviderInfo};
pub use image_generation::{ImageGenerationParams, ImageGenerationResult};
pub use image::{ImageGenerationParams, ImageGenerationResult, ImageEditParams, ImageEditResult, ImageSource};
pub use ping::PingResult;
pub use speech_to_text::{SpeechToTextParams, SpeechToTextResult};
pub use text_generation::{

View File

@@ -20,6 +20,6 @@ pub use controller::{ControllerTrait, ControllerType};
pub use config::ConfigTrait;
pub use entity::{
AgentProvider, AgentProviderInfo, ImageGenerationParams, PingResult, SpeechToTextParams,
AgentProvider, AgentProviderInfo, ImageGenerationParams, ImageEditParams, ImageSource, PingResult, SpeechToTextParams,
SpeechToTextResult, TextGenerationParams, TextGenerationPromptVariables, TextToSpeechParams,
};

View File

@@ -5,7 +5,8 @@ use async_openai::{
config::OpenAIConfig,
types::{
ChatCompletionRequestMessage, CreateChatCompletionRequestArgs, CreateImageRequestArgs,
CreateSpeechRequestArgs, CreateTranscriptionRequestArgs,
CreateSpeechRequestArgs, CreateTranscriptionRequestArgs, CreateImageEditRequestArgs,
ImageInput, ImageModel, DallE2ImageSize, ImageResponseFormat, Image,
},
};
@@ -14,24 +15,22 @@ use crate::{
agent::{
AgentPurpose,
provider::{
entity::{ImageGenerationResult, PingResult, TextToSpeechParams, TextToSpeechResult},
entity::{ImageGenerationResult, ImageEditResult, ImageSource, PingResult, TextToSpeechParams, TextToSpeechResult},
openai::utils::convert_string_to_enum,
},
},
strings,
};
use crate::{
agent::{
provider::{
ImageGenerationParams, SpeechToTextParams, SpeechToTextResult,
entity::{TextGenerationParams, TextGenerationResult},
},
utils::base64_decode,
agent::provider::{
ImageGenerationParams, ImageEditParams, SpeechToTextParams, SpeechToTextResult,
entity::{TextGenerationParams, TextGenerationResult},
},
conversation::llm::{
Author as LLMAuthor, Conversation as LLMConversation, Message as LLMMessage,
shorten_messages_list_to_context_size,
MessageContent as LLMMessageContent, shorten_messages_list_to_context_size,
},
utils::base64::base64_decode,
};
use super::config::Config;
@@ -64,7 +63,7 @@ impl ControllerTrait for Controller {
let messages = vec![LLMMessage {
author: LLMAuthor::User,
message_text: "Hello!".to_string(),
content: LLMMessageContent::Text("Hello!".to_string()),
timestamp: chrono::Utc::now(),
}];
@@ -101,7 +100,7 @@ impl ControllerTrait for Controller {
} else {
Some(LLMMessage {
author: LLMAuthor::Prompt,
message_text: prompt_text,
content: LLMMessageContent::Text(prompt_text),
timestamp: chrono::Utc::now(),
})
};
@@ -290,11 +289,13 @@ impl ControllerTrait for Controller {
.unwrap_or(image_generation_config.size);
let response_format = match model.clone() {
async_openai::types::ImageModel::DallE2 => Some(async_openai::types::ImageResponseFormat::B64Json),
async_openai::types::ImageModel::DallE3 => Some(async_openai::types::ImageResponseFormat::B64Json),
async_openai::types::ImageModel::Other(model_str) => match model_str.as_str() {
ImageModel::DallE2 => Some(ImageResponseFormat::B64Json),
ImageModel::DallE3 => Some(ImageResponseFormat::B64Json),
ImageModel::Other(model_str) => match model_str.as_str() {
// gpt-image-1 only outputs base64 and we don't need to specify the response format.
// In fact, specifying the response format results in an error.
OPENAI_IMAGE_MODEL_GPT_IMAGE_1 => None,
_ => Some(async_openai::types::ImageResponseFormat::B64Json),
_ => Some(ImageResponseFormat::B64Json),
},
};
@@ -355,6 +356,96 @@ impl ControllerTrait for Controller {
))
}
async fn create_image_edit(
&self,
prompt: &str,
images: Vec<ImageSource>,
_params: ImageEditParams,
) -> anyhow::Result<ImageEditResult> {
let Some(image_generation_config) = &self.config.image_generation else {
return Err(anyhow::anyhow!(
strings::agent::no_configuration_for_purpose_so_cannot_be_used(
&AgentPurpose::ImageGeneration
),
));
};
let Some(first_image) = images.into_iter().next() else {
return Err(anyhow::anyhow!("No image sources provided"));
};
let image_input: ImageInput = first_image.into();
let dalle2_size = match image_generation_config.size {
async_openai::types::ImageSize::S256x256 => Some(DallE2ImageSize::S256x256),
async_openai::types::ImageSize::S512x512 => Some(DallE2ImageSize::S512x512),
async_openai::types::ImageSize::S1024x1024 => Some(DallE2ImageSize::S1024x1024),
_ => None,
};
let model = image_generation_config
.model_id_as_openai_image_model()
.map_err(|err| anyhow::anyhow!(err))?;
let response_format = match model.clone() {
async_openai::types::ImageModel::DallE2 => Some(async_openai::types::ImageResponseFormat::B64Json),
async_openai::types::ImageModel::DallE3 => Some(async_openai::types::ImageResponseFormat::B64Json),
async_openai::types::ImageModel::Other(model_str) => match model_str.as_str() {
OPENAI_IMAGE_MODEL_GPT_IMAGE_1 => None,
_ => Some(async_openai::types::ImageResponseFormat::B64Json),
},
};
let mut request_builder = CreateImageEditRequestArgs::default();
request_builder
.image(image_input)
.prompt(prompt.to_owned())
.model(model);
if let Some(size) = dalle2_size {
request_builder.size(size);
}
if let Some(response_format) = response_format {
request_builder.response_format(response_format);
}
let request = request_builder.build()
.map_err(|e| anyhow::anyhow!("Failed to build CreateImageEditRequest: {}", e))?;
tracing::trace!(
model = format!("{:?}", request.model),
size = format!("{:?}", request.size),
response_format = format!("{:?}", request.response_format),
"Sending OpenAI image edit API request"
);
let response = self.client.images().create_edit(request).await?;
if let Some(image_data) = response.data.into_iter().next() {
match image_data.deref() {
Image::B64Json { b64_json, .. } => {
let bytes = base64_decode(b64_json)?;
return Ok(ImageEditResult {
bytes,
mime_type: mxlink::mime::IMAGE_PNG,
});
}
Image::Url { url, .. } => {
tracing::warn!(?url, "Received URL instead of B64Json for image edit");
return Err(anyhow::anyhow!(
"Unexpected image type (URL) when B64Json was requested"
));
}
}
}
Err(anyhow::anyhow!(
"The OpenAI image edit API returned no images"
))
}
async fn text_to_speech(
&self,
input: &str,

View File

@@ -1,9 +1,9 @@
use async_openai::types::{
ChatCompletionRequestAssistantMessageArgs, ChatCompletionRequestMessage,
ChatCompletionRequestSystemMessageArgs, ChatCompletionRequestUserMessageArgs,
ChatCompletionRequestAssistantMessageArgs, ChatCompletionRequestMessage, ChatCompletionRequestMessageContentPartImage, ChatCompletionRequestSystemMessageArgs, ChatCompletionRequestUserMessageArgs, ChatCompletionRequestUserMessageContent, ChatCompletionRequestUserMessageContentPart, ImageUrlArgs
};
use crate::conversation::llm::{Author as LLMAuthor, Message as LLMMessage};
use crate::conversation::llm::{Author as LLMAuthor, Message as LLMMessage, MessageContent as LLMMessageContent};
use crate::utils::base64::base64_encode;
pub fn convert_llm_messages_to_openai_messages(
conversation_messages: Vec<LLMMessage>,
@@ -12,29 +12,65 @@ pub fn convert_llm_messages_to_openai_messages(
Vec::with_capacity(conversation_messages.len());
for message in conversation_messages {
openai_conversation_messages.push(convert_llm_message_to_openai_message(message));
let openai_message = convert_llm_message_to_openai_message(message);
if let Some(openai_message) = openai_message {
openai_conversation_messages.push(openai_message);
}
}
openai_conversation_messages
}
fn convert_llm_message_to_openai_message(llm_message: LLMMessage) -> ChatCompletionRequestMessage {
match llm_message.author {
LLMAuthor::Prompt => ChatCompletionRequestSystemMessageArgs::default()
.content(llm_message.message_text)
.build()
.expect("Failed building OpenAI system message")
.into(),
LLMAuthor::Assistant => ChatCompletionRequestAssistantMessageArgs::default()
.content(llm_message.message_text)
.build()
.expect("Failed building OpenAI assistant message")
.into(),
LLMAuthor::User => ChatCompletionRequestUserMessageArgs::default()
.content(llm_message.message_text)
.build()
.expect("Failed building OpenAI user message")
.into(),
fn convert_llm_message_to_openai_message(llm_message: LLMMessage) -> Option<ChatCompletionRequestMessage> {
match &llm_message.content {
LLMMessageContent::Text(text) => {
Some(match llm_message.author {
LLMAuthor::Prompt => ChatCompletionRequestSystemMessageArgs::default()
.content(text.clone())
.build()
.expect("Failed building OpenAI system message")
.into(),
LLMAuthor::Assistant => ChatCompletionRequestAssistantMessageArgs::default()
.content(text.clone())
.build()
.expect("Failed building OpenAI assistant message")
.into(),
LLMAuthor::User => ChatCompletionRequestUserMessageArgs::default()
.content(text.clone())
.build()
.expect("Failed building OpenAI user message")
.into(),
})
}
LLMMessageContent::Image(image_details) => {
let image_url = format!("data:{};base64,{}", image_details.mime, base64_encode(&image_details.data));
let part = ChatCompletionRequestUserMessageContentPart::ImageUrl(
ChatCompletionRequestMessageContentPartImage{
image_url: ImageUrlArgs::default()
.url(image_url)
.build()
.expect("Failed building OpenAI image url")
}
);
let message_content = ChatCompletionRequestUserMessageContent::Array(vec![part]);
match llm_message.author {
LLMAuthor::User => Some(ChatCompletionRequestUserMessageArgs::default()
.content(message_content)
.build()
.expect("Failed building OpenAI user message")
.into()),
_ => {
tracing::warn!(
"OpenAI API does not support image content for messages authored by {:?}. This message part will be skipped.",
llm_message.author
);
None
}
}
}
}
}

View File

@@ -4,22 +4,23 @@ use etke_openai_api_rust::images::{ImagesApi, ImagesBody};
use etke_openai_api_rust::{Auth, Message, OpenAI};
use super::super::ControllerTrait;
use crate::agent::utils::base64_decode;
use crate::utils::base64::base64_decode;
use crate::{
agent::provider::{
ImageGenerationParams, SpeechToTextParams, SpeechToTextResult,
ImageGenerationParams, ImageEditParams, ImageSource, SpeechToTextParams, SpeechToTextResult,
entity::{TextGenerationParams, TextGenerationResult},
},
conversation::llm::{
Author as LLMAuthor, Conversation as LLMConversation, Message as LLMMessage,
shorten_messages_list_to_context_size,
MessageContent as LLMMessageContent, shorten_messages_list_to_context_size,
},
};
use crate::{
agent::{
AgentPurpose,
provider::entity::{
ImageGenerationResult, PingResult, TextToSpeechParams, TextToSpeechResult,
ImageGenerationResult, ImageEditResult, PingResult, TextToSpeechParams,
TextToSpeechResult,
},
},
strings,
@@ -60,7 +61,7 @@ impl ControllerTrait for Controller {
let messages = vec![LLMMessage {
author: LLMAuthor::User,
message_text: "Hello!".to_string(),
content: LLMMessageContent::Text("Hello!".to_string()),
timestamp: chrono::Utc::now(),
}];
@@ -97,7 +98,7 @@ impl ControllerTrait for Controller {
} else {
Some(LLMMessage {
author: LLMAuthor::Prompt,
message_text: prompt_text,
content: LLMMessageContent::Text(prompt_text),
timestamp: chrono::Utc::now(),
})
};
@@ -366,6 +367,17 @@ impl ControllerTrait for Controller {
))
}
async fn create_image_edit(
&self,
_prompt: &str,
_images: Vec<ImageSource>,
_params: ImageEditParams,
) -> anyhow::Result<ImageEditResult> {
Err(anyhow::anyhow!(
"The OpenAI image edit API is not supported by the OpenAI-compat provider"
))
}
async fn text_to_speech(
&self,
input: &str,

View File

@@ -2,7 +2,7 @@ use etke_openai_api_rust::{Message, Role};
use crate::agent::provider::openai::Config as OpenAIConfig;
use crate::conversation::llm::{Author as LLMAuthor, Message as LLMMessage};
use crate::conversation::llm::{Author as LLMAuthor, Message as LLMMessage, MessageContent as LLMMessageContent};
pub fn convert_llm_messages_to_openai_messages(
conversation_messages: Vec<LLMMessage>,
@@ -11,22 +11,33 @@ pub fn convert_llm_messages_to_openai_messages(
Vec::with_capacity(conversation_messages.len());
for message in conversation_messages {
openai_conversation_messages.push(convert_llm_message_to_openai_message(message));
let openai_message = convert_llm_message_to_openai_message(message);
if let Some(openai_message) = openai_message {
openai_conversation_messages.push(openai_message);
}
}
openai_conversation_messages
}
fn convert_llm_message_to_openai_message(llm_message: LLMMessage) -> Message {
fn convert_llm_message_to_openai_message(llm_message: LLMMessage) -> Option<Message> {
let role = match llm_message.author {
LLMAuthor::Prompt => Role::System,
LLMAuthor::Assistant => Role::Assistant,
LLMAuthor::User => Role::User,
};
Message {
role,
content: llm_message.message_text,
match &llm_message.content {
LLMMessageContent::Text(text) => {
Some(Message {
role,
content: text.clone(),
})
},
LLMMessageContent::Image(_image_details) => {
tracing::warn!("The OpenAI-compat provider's library does not support image content. This image message will be skipped.");
None
},
}
}

View File

@@ -1,5 +1,3 @@
use base64::{Engine as _, engine::general_purpose::STANDARD};
use crate::{
agent::{
AgentInstance, AgentPurpose, ControllerTrait, Manager as AgentManager, PublicIdentifier,
@@ -140,7 +138,3 @@ async fn get_global_agent_id_for_purpose(
.handler
.get_by_purpose_with_catch_all_fallback(purpose)
}
pub(crate) fn base64_decode(base64_string: &str) -> Result<Vec<u8>, base64::DecodeError> {
STANDARD.decode(base64_string)
}

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

View File

@@ -1,4 +1,8 @@
use chrono::{DateTime, Utc};
use mxlink::matrix_sdk::ruma::events::room::message::ImageMessageEventContent;
use mxlink::mime::Mime;
use crate::agent::provider::ImageSource;
#[derive(Debug, Clone, PartialEq)]
pub enum Author {
@@ -10,10 +14,51 @@ pub enum Author {
#[derive(Debug, Clone)]
pub struct Message {
pub author: Author,
pub message_text: String,
pub timestamp: DateTime<Utc>,
pub content: MessageContent,
}
#[derive(Debug, Clone)]
pub struct ImageDetails {
pub event_content: ImageMessageEventContent,
pub mime: Mime,
pub data: Vec<u8>,
}
impl ImageDetails {
pub fn new(event_content: ImageMessageEventContent, mime: Mime, data: Vec<u8>) -> Self {
Self { event_content, mime, data }
}
pub fn filename(&self) -> String {
self.event_content.filename.clone().unwrap_or(self.event_content.body.clone())
}
}
impl Into<ImageSource> for &ImageDetails {
fn into(self) -> ImageSource {
ImageSource::new(self.filename(), self.data.clone(), self.mime.clone())
}
}
#[derive(Debug, Clone)]
pub enum MessageContent {
Text(String),
Image(ImageDetails),
}
impl PartialEq for MessageContent {
fn eq(&self, other: &Self) -> bool {
match (self, other) {
(MessageContent::Text(a), MessageContent::Text(b)) => a == b,
(MessageContent::Image(a), MessageContent::Image(b)) => {
// We can probably do better than this by inspecting `.event_conten1t.source`, but for now this is good enough.
a.filename() == b.filename()
},
_ => false,
}
}
}
#[derive(Debug)]
pub struct Conversation {
pub messages: Vec<Message>,
@@ -28,27 +73,32 @@ impl Conversation {
pub fn combine_consecutive_messages(&self) -> Conversation {
// We'll likely get fewer messages, but let's reserve the maximum we expect.
let mut new_messages = Vec::with_capacity(self.messages.len());
let mut last_seen_author: Option<Author> = None;
let mut last_seen_text_from_author: Option<Author> = None;
for message in &self.messages {
let Some(last_seen_author_clone) = last_seen_author.clone() else {
last_seen_author = Some(message.author.clone());
let MessageContent::Text(message_text_content) = &message.content else {
last_seen_text_from_author = None;
new_messages.push(message.clone());
continue;
};
let Some(last_seen_author_clone) = last_seen_text_from_author.clone() else {
last_seen_text_from_author = Some(message.author.clone());
new_messages.push(message.clone());
continue;
};
if message.author != last_seen_author_clone {
last_seen_author = Some(message.author.clone());
last_seen_text_from_author = Some(message.author.clone());
new_messages.push(message.clone());
continue;
}
new_messages.last_mut().unwrap().message_text.push('\n');
new_messages
.last_mut()
.unwrap()
.message_text
.push_str(&message.message_text);
let last_message = new_messages.last_mut().unwrap();
if let MessageContent::Text(ref mut text) = last_message.content {
text.push('\n');
text.push_str(message_text_content);
}
}
Conversation {
@@ -65,48 +115,76 @@ impl Conversation {
mod tests {
use super::*;
use chrono::{TimeZone, Utc};
use mxlink::matrix_sdk::ruma::OwnedMxcUri;
use mxlink::mime;
#[test]
fn combine_consecutive_messages() {
let timestamp_1 = Utc.with_ymd_and_hms(2024, 9, 20, 18, 34, 15).unwrap();
let timestamp_2 = Utc.with_ymd_and_hms(2024, 9, 21, 18, 34, 15).unwrap();
let timestamp_2 = Utc.with_ymd_and_hms(2024, 9, 21, 18, 34, 16).unwrap();
let timestamp_3 = Utc.with_ymd_and_hms(2024, 9, 22, 18, 34, 15).unwrap();
let timestamp_3 = Utc.with_ymd_and_hms(2024, 9, 22, 18, 34, 17).unwrap();
let timestamp_4 = Utc.with_ymd_and_hms(2024, 9, 23, 18, 34, 18).unwrap();
let image_event_content = ImageMessageEventContent::plain(
"image.png".to_string(),
OwnedMxcUri::from("mxc://example.com/1234567890"),
);
let conversation = Conversation {
messages: vec![
// User's turn
Message {
author: Author::User,
message_text: "Hello".to_string(),
content: MessageContent::Text("Hello".to_string()),
timestamp: timestamp_1,
},
Message {
author: Author::User,
message_text: "How are you?".to_string(),
content: MessageContent::Text("How are you?".to_string()),
timestamp: timestamp_2,
},
Message {
author: Author::User,
message_text: "I'm OK, btw.".to_string(),
content: MessageContent::Text("I'm OK, btw.".to_string()),
timestamp: timestamp_3,
},
Message {
author: Author::User,
content: MessageContent::Image(ImageDetails::new(
image_event_content.clone(),
mime::IMAGE_PNG,
vec![],
)),
timestamp: timestamp_4,
},
Message {
author: Author::User,
content: MessageContent::Text("Above is an image.".to_string()),
timestamp: timestamp_4,
},
Message {
author: Author::User,
content: MessageContent::Text("Would you take a look at it?".to_string()),
timestamp: timestamp_4,
},
// Assistant's turn
Message {
author: Author::Assistant,
message_text: "Hi there!".to_string(),
content: MessageContent::Text("Hi there!".to_string()),
timestamp: timestamp_2,
},
Message {
author: Author::Assistant,
message_text: "I'm doing well, thank you.".to_string(),
content: MessageContent::Text("I'm doing well, thank you.".to_string()),
timestamp: timestamp_3,
},
// User's turn
Message {
author: Author::User,
message_text: "That's great!".to_string(),
content: MessageContent::Text("That's great!".to_string()),
timestamp: timestamp_3,
},
],
@@ -114,23 +192,44 @@ mod tests {
let conversation = conversation.combine_consecutive_messages();
assert_eq!(conversation.messages.len(), 3);
assert_eq!(conversation.messages.len(), 5);
assert_eq!(conversation.messages[0].author, Author::User);
assert_eq!(
conversation.messages[0].message_text,
"Hello\nHow are you?\nI'm OK, btw."
conversation.messages[0].content,
MessageContent::Text("Hello\nHow are you?\nI'm OK, btw.".to_string())
);
assert_eq!(conversation.messages[0].timestamp, timestamp_1);
assert_eq!(conversation.messages[1].author, Author::Assistant);
assert_eq!(conversation.messages[1].author, Author::User);
assert_eq!(
conversation.messages[1].message_text,
"Hi there!\nI'm doing well, thank you."
conversation.messages[1].content,
MessageContent::Image(ImageDetails::new(
image_event_content.clone(),
mime::IMAGE_PNG,
vec![],
))
);
assert_eq!(conversation.messages[1].timestamp, timestamp_2);
assert_eq!(conversation.messages[2].author, Author::User);
assert_eq!(conversation.messages[2].message_text, "That's great!");
assert_eq!(conversation.messages[2].timestamp, timestamp_3);
assert_eq!(
conversation.messages[2].content,
MessageContent::Text("Above is an image.\nWould you take a look at it?".to_string())
);
assert_eq!(conversation.messages[2].timestamp, timestamp_4);
assert_eq!(conversation.messages[3].author, Author::Assistant);
assert_eq!(
conversation.messages[3].content,
MessageContent::Text("Hi there!\nI'm doing well, thank you.".to_string())
);
assert_eq!(conversation.messages[3].timestamp, timestamp_2);
assert_eq!(conversation.messages[4].author, Author::User);
assert_eq!(
conversation.messages[4].content,
MessageContent::Text("That's great!".to_string())
);
assert_eq!(conversation.messages[4].timestamp, timestamp_3);
}
}

View File

@@ -12,8 +12,7 @@ fn test_messages_by_the_bot_are_identified_correctly() {
let matrix_message = super::super::matrix::MatrixMessage {
sender_id: bot_user_id.to_owned(),
message_type: super::super::matrix::MatrixMessageType::Text,
message_text: "Hello!".to_owned(),
content: super::super::matrix::MatrixMessageContent::Text("Hello!".to_owned()),
mentioned_users: vec![],
timestamp: chrono::Utc::now(),
};
@@ -21,12 +20,11 @@ fn test_messages_by_the_bot_are_identified_correctly() {
let llm_message = convert_matrix_message_to_llm_message(&matrix_message, &bot_user_id).unwrap();
assert_eq!(llm_message.author, Author::Assistant);
assert_eq!(llm_message.message_text, "Hello!");
assert_eq!(llm_message.content, MessageContent::Text("Hello!".to_string()));
}
#[test]
fn test_notice_messages_by_bot_with_speech_to_text_prefix_are_cleaned_up_and_considered_sent_by_user()
{
fn test_notice_messages_by_bot_with_speech_to_text_prefix_are_cleaned_up_and_considered_sent_by_user() {
let bot_user_id =
OwnedUserId::try_from("@bot:example.com").expect("Failed to parse bot user ID");
@@ -37,8 +35,7 @@ fn test_notice_messages_by_bot_with_speech_to_text_prefix_are_cleaned_up_and_con
let matrix_message = super::super::matrix::MatrixMessage {
sender_id: bot_user_id.to_owned(),
message_type: super::super::matrix::MatrixMessageType::Notice,
message_text,
content: super::super::matrix::MatrixMessageContent::Notice(message_text),
mentioned_users: vec![],
timestamp: chrono::Utc::now(),
};
@@ -46,7 +43,7 @@ fn test_notice_messages_by_bot_with_speech_to_text_prefix_are_cleaned_up_and_con
let llm_message = convert_matrix_message_to_llm_message(&matrix_message, &bot_user_id).unwrap();
assert_eq!(llm_message.author, Author::User);
assert_eq!(llm_message.message_text, source_message_text);
assert_eq!(llm_message.content, MessageContent::Text(source_message_text.to_string()));
}
#[test]
@@ -61,8 +58,7 @@ fn test_notice_error_messages_by_bot_are_ignored() {
let matrix_message = super::super::matrix::MatrixMessage {
sender_id: bot_user_id.to_owned(),
message_type: super::super::matrix::MatrixMessageType::Notice,
message_text,
content: super::super::matrix::MatrixMessageContent::Notice(message_text),
mentioned_users: vec![],
timestamp: chrono::Utc::now(),
};
@@ -86,8 +82,7 @@ fn test_other_notice_messages_by_the_bot_are_ignored() {
let matrix_message = super::super::matrix::MatrixMessage {
sender_id: bot_user_id.to_owned(),
message_type: super::super::matrix::MatrixMessageType::Notice,
message_text: message_text.to_owned(),
content: super::super::matrix::MatrixMessageContent::Notice(message_text.to_owned()),
mentioned_users: vec![],
timestamp: chrono::Utc::now(),
};

View File

@@ -2,7 +2,7 @@ use tiktoken_rs::CoreBPE;
use tiktoken_rs::get_bpe_from_tokenizer;
use tiktoken_rs::tokenizer;
use super::{Author, Message};
use super::{Author, Message, MessageContent};
fn get_bpe_for_model(model: &str) -> CoreBPE {
let tokenizer = tokenizer::get_tokenizer(model)
@@ -71,7 +71,10 @@ fn calculate_token_size_for_message(bpe: &CoreBPE, model: &str, message: &Messag
Author::Prompt => bpe.encode_with_special_tokens("system").len() as i32,
};
let text_length = bpe.encode_with_special_tokens(&message.message_text).len() as i32;
let text_length = match &message.content {
MessageContent::Text(text) => bpe.encode_with_special_tokens(text).len() as i32,
MessageContent::Image(..) => 0,
};
(text_length + role_length + tokens_per_message + tokens_per_name) as u32
}
@@ -85,7 +88,7 @@ pub mod test {
let message = super::Message {
author: super::Author::User,
message_text: "Hello there!".to_owned(),
content: super::MessageContent::Text("Hello there!".to_string()),
timestamp: chrono::Utc::now(),
};
@@ -104,7 +107,7 @@ pub mod test {
let prompt = super::Message {
author: super::Author::Prompt,
message_text: "You are a bot!".to_owned(),
content: super::MessageContent::Text("You are a bot!".to_string()),
timestamp: chrono::Utc::now(),
};
let prompt_length = 10;
@@ -118,7 +121,7 @@ pub mod test {
let first = super::Message {
author: super::Author::User,
message_text: "Hello there!".to_owned(),
content: super::MessageContent::Text("Hello there!".to_string()),
timestamp: chrono::Utc::now(),
};
let first_length = 8;
@@ -132,7 +135,7 @@ pub mod test {
let second = super::Message {
author: super::Author::Assistant,
message_text: "Hello!".to_owned(),
content: super::MessageContent::Text("Hello!".to_string()),
timestamp: chrono::Utc::now(),
};
let second_length = 7;
@@ -146,8 +149,10 @@ pub mod test {
let third = super::Message {
author: super::Author::User,
message_text: "This is the 3rd message in this conversation. It shall be preserved."
.to_owned(),
content: super::MessageContent::Text(
"This is the 3rd message in this conversation. It shall be preserved."
.to_owned(),
),
timestamp: chrono::Utc::now(),
};
let third_length = 21;
@@ -161,7 +166,9 @@ pub mod test {
let forth = super::Message {
author: super::Author::Assistant,
message_text: "This is yet another message that shall be preserved.".to_owned(),
content: super::MessageContent::Text(
"This is yet another message that shall be preserved.".to_owned(),
),
timestamp: chrono::Utc::now(),
};
let forth_length = 15;
@@ -186,13 +193,13 @@ pub mod test {
assert_eq!(2, new_conversation_messages.len());
assert_eq!(
new_conversation_messages.first().unwrap().message_text,
third.message_text
new_conversation_messages.first().unwrap().content,
third.content
);
assert_eq!(
new_conversation_messages.last().unwrap().message_text,
forth.message_text
new_conversation_messages.last().unwrap().content,
forth.content
);
}
@@ -206,7 +213,7 @@ pub mod test {
let prompt = super::Message {
author: super::Author::User,
message_text: "あなたはボットです。".to_owned(),
content: super::MessageContent::Text("あなたはボットです。".to_string()),
timestamp: chrono::Utc::now(),
};
let prompt_length = 14;
@@ -220,7 +227,7 @@ pub mod test {
let first = super::Message {
author: super::Author::User,
message_text: "こんにちは!".to_owned(),
content: super::MessageContent::Text("こんにちは!".to_string()),
timestamp: chrono::Utc::now(),
};
let first_length = 7;
@@ -234,7 +241,7 @@ pub mod test {
let second = super::Message {
author: super::Author::Assistant,
message_text: "こんにちは。今日は元気ですか。".to_owned(),
content: super::MessageContent::Text("こんにちは。今日は元気ですか。".to_string()),
timestamp: chrono::Utc::now(),
};
let second_length = 15;
@@ -248,7 +255,9 @@ pub mod test {
let third = super::Message {
author: super::Author::User,
message_text: "これは第3のメッセージなので、保存されます。".to_owned(),
content: super::MessageContent::Text(
"これは第3のメッセージなので、保存されます。".to_string(),
),
timestamp: chrono::Utc::now(),
};
let third_length = 22;
@@ -262,7 +271,9 @@ pub mod test {
let forth = super::Message {
author: super::Author::Assistant,
message_text: "これはもう一つの保存されますメッセージです。".to_owned(),
content: super::MessageContent::Text(
"これはもう一つの保存されますメッセージです。".to_string(),
),
timestamp: chrono::Utc::now(),
};
let forth_length = 21;
@@ -287,13 +298,13 @@ pub mod test {
assert_eq!(2, new_conversation_messages.len());
assert_eq!(
new_conversation_messages.first().unwrap().message_text,
third.message_text
new_conversation_messages.first().unwrap().content,
third.content
);
assert_eq!(
new_conversation_messages.last().unwrap().message_text,
forth.message_text
new_conversation_messages.last().unwrap().content,
forth.content
);
}
}

View File

@@ -1,7 +1,7 @@
use matrix_sdk::ruma::OwnedUserId;
use mxlink::matrix_sdk::ruma::OwnedUserId;
use super::{Author, Message};
use crate::conversation::matrix::{MatrixMessage, MatrixMessageType};
use super::entity::{Author, ImageDetails, Message, MessageContent};
use crate::conversation::matrix::{MatrixMessage, MatrixMessageContent};
use crate::utils::text_to_speech as text_to_speech_utils;
pub fn convert_matrix_message_to_llm_message(
@@ -16,12 +16,23 @@ pub fn convert_matrix_message_to_llm_message(
}
fn convert_bot_message(matrix_message: &MatrixMessage) -> Option<Message> {
match matrix_message.message_type {
MatrixMessageType::Text => {
convert_bot_text_message(&matrix_message.message_text, &matrix_message.timestamp)
match &matrix_message.content {
MatrixMessageContent::Text(text) => {
convert_bot_text_message(text, &matrix_message.timestamp)
}
MatrixMessageType::Notice => {
convert_bot_notice_message(&matrix_message.message_text, &matrix_message.timestamp)
MatrixMessageContent::Notice(text) => {
convert_bot_notice_message(text, &matrix_message.timestamp)
}
MatrixMessageContent::Image(image_content, mime_type, media_bytes) => {
Some(Message {
author: Author::Assistant,
content: MessageContent::Image(ImageDetails::new(
image_content.clone(),
mime_type.clone(),
media_bytes.clone()
)),
timestamp: matrix_message.timestamp.to_owned(),
})
}
}
}
@@ -32,7 +43,7 @@ fn convert_bot_text_message(
) -> Option<Message> {
Some(Message {
author: Author::Assistant,
message_text: text.to_owned(),
content: MessageContent::Text(text.to_owned()),
timestamp: timestamp.to_owned(),
})
}
@@ -52,7 +63,7 @@ fn convert_bot_notice_message(
// This is a transcription message. We remove the prefix and consider it as a message sent by the user.
return Some(Message {
author: Author::User,
message_text: text.to_owned(),
content: MessageContent::Text(text.to_owned()),
timestamp: timestamp.to_owned(),
});
}
@@ -61,9 +72,31 @@ fn convert_bot_notice_message(
}
fn convert_user_message(matrix_message: &MatrixMessage) -> Option<Message> {
Some(Message {
author: Author::User,
message_text: matrix_message.message_text.clone(),
timestamp: matrix_message.timestamp.to_owned(),
})
match &matrix_message.content {
MatrixMessageContent::Text(text) => {
Some(Message {
author: Author::User,
content: MessageContent::Text(text.clone()),
timestamp: matrix_message.timestamp.to_owned(),
})
}
MatrixMessageContent::Notice(text) => {
Some(Message {
author: Author::User,
content: MessageContent::Text(text.clone()),
timestamp: matrix_message.timestamp.to_owned(),
})
}
MatrixMessageContent::Image(image_content, mime_type, media_bytes) => {
Some(Message {
author: Author::User,
content: MessageContent::Image(ImageDetails::new(
image_content.clone(),
mime_type.clone(),
media_bytes.clone(),
)),
timestamp: matrix_message.timestamp.to_owned(),
})
}
}
}

View File

@@ -2,20 +2,22 @@ use chrono::{DateTime, Utc};
use regex::Regex;
use mxlink::matrix_sdk::ruma::OwnedUserId;
use mxlink::matrix_sdk::ruma::events::room::message::ImageMessageEventContent;
use mxlink::mime::Mime;
#[derive(Clone)]
pub struct MatrixMessage {
pub sender_id: OwnedUserId,
pub message_type: MatrixMessageType,
pub message_text: String,
pub content: MatrixMessageContent,
pub mentioned_users: Vec<OwnedUserId>,
pub timestamp: DateTime<Utc>,
}
#[derive(Clone)]
pub enum MatrixMessageType {
Text,
Notice,
pub enum MatrixMessageContent {
Text(String),
Notice(String),
Image(ImageMessageEventContent, Mime, Vec<u8>),
}
#[derive(Clone)]

View File

@@ -6,6 +6,6 @@ mod utils;
pub(crate) use room_display_name_fetcher::RoomDisplayNameFetcher;
pub(crate) use room_event_fetcher::RoomEventFetcher;
pub(crate) use entity::{MatrixMessage, MatrixMessageProcessingParams, MatrixMessageType};
pub(crate) use entity::{MatrixMessage, MatrixMessageProcessingParams, MatrixMessageContent};
pub(crate) use utils::*;

View File

@@ -18,9 +18,11 @@ use mxlink::matrix_sdk::{
},
};
use mxlink::{MatrixLink, ThreadGetMessagesParams, ThreadInfo};
use tracing::Instrument;
use super::{MatrixMessage, MatrixMessageProcessingParams, MatrixMessageType, RoomEventFetcher};
use super::{MatrixMessage, MatrixMessageContent, MatrixMessageProcessingParams, RoomEventFetcher};
use crate::entity::{InteractionContext, InteractionTrigger, MessagePayload};
use crate::utils::mime::get_mime_type_from_file_name;
struct DetailedMessagePayload {
is_mentioning_bot: bool,
@@ -28,7 +30,7 @@ struct DetailedMessagePayload {
}
pub async fn get_matrix_messages_in_thread(
matrix_link: MatrixLink,
matrix_link: &MatrixLink,
room: &Room,
thread_id: OwnedEventId,
) -> Result<Vec<MatrixMessage>, mxlink::matrix_sdk::Error> {
@@ -40,18 +42,18 @@ pub async fn get_matrix_messages_in_thread(
let mut messages: Vec<MatrixMessage> = Vec::new();
for matrix_native_message in messages_native {
let Some(message) = convert_matrix_native_event_to_matrix_message(&matrix_native_message)
else {
continue;
};
let message_result = convert_matrix_native_event_to_matrix_message(matrix_link, &matrix_native_message).await?;
messages.push(message);
if let Some(message) = message_result {
messages.push(message);
}
}
Ok(messages)
}
pub async fn get_matrix_messages_in_reply_chain(
matrix_link: &MatrixLink,
event_fetcher: &Arc<RoomEventFetcher>,
room: &Room,
event_id: OwnedEventId,
@@ -62,12 +64,11 @@ pub async fn get_matrix_messages_in_reply_chain(
let mut messages: Vec<MatrixMessage> = Vec::new();
for matrix_native_message in messages_native {
let Some(message) = convert_matrix_native_event_to_matrix_message(&matrix_native_message)
else {
continue;
};
let message_result = convert_matrix_native_event_to_matrix_message(matrix_link, &matrix_native_message).await?;
messages.push(message);
if let Some(message) = message_result {
messages.push(message);
}
}
Ok(messages)
@@ -150,30 +151,34 @@ pub async fn process_matrix_messages(
let mut message = message.clone();
if i == 0 && !params.first_message_prefixes_to_strip.is_empty() {
let mut message_text = message.message_text.clone();
if let MatrixMessageContent::Text(message_text) = &message.content {
let mut message_text = message_text.clone();
for prefix in &params.first_message_prefixes_to_strip {
if let Some(message_text_stripped) = message_text.strip_prefix(prefix) {
message_text = message_text_stripped.to_owned();
for prefix in &params.first_message_prefixes_to_strip {
if let Some(message_text_stripped) = message_text.strip_prefix(prefix) {
message_text = message_text_stripped.to_owned();
}
}
}
message.message_text = message_text.trim().to_owned();
message.content = MatrixMessageContent::Text(message_text.trim().to_owned());
}
}
// We only strip `bot_user_prefixes_to_strip`-defined prefixes from messages that mention the bot user.
if !params.bot_user_prefixes_to_strip.is_empty()
&& message.mentioned_users.contains(&params.bot_user_id)
{
let mut message_text = message.message_text.clone();
if let MatrixMessageContent::Text(message_text) = &message.content {
let mut message_text = message_text.clone();
for prefix in &params.bot_user_prefixes_to_strip {
if let Some(message_text_stripped) = message_text.strip_prefix(prefix) {
message_text = message_text_stripped.to_owned();
for prefix in &params.bot_user_prefixes_to_strip {
if let Some(message_text_stripped) = message_text.strip_prefix(prefix) {
message_text = message_text_stripped.to_owned();
}
}
}
message.message_text = message_text.trim().to_owned();
message.content = MatrixMessageContent::Text(message_text.trim().to_owned());
}
}
messages_filtered.push(message);
@@ -207,23 +212,25 @@ fn is_message_from_allowed_sender(
false
}
pub fn convert_matrix_native_event_to_matrix_message(
pub async fn convert_matrix_native_event_to_matrix_message(
matrix_link: &MatrixLink,
matrix_native_event: &AnySyncMessageLikeEvent,
) -> Option<MatrixMessage> {
) -> Result<Option<MatrixMessage>, mxlink::matrix_sdk::Error> {
let Some(content) = matrix_native_event.original_content() else {
// Redacted message
return None;
return Ok(None);
};
let AnyMessageLikeEventContent::RoomMessage(room_message) = content else {
// Some state event, etc.
return None;
return Ok(None);
};
let (text, is_notice) = match &room_message.msgtype {
MessageType::Text(text_content) => (text_content.body.clone(), false),
MessageType::Notice(notice_content) => (notice_content.body.clone(), true),
_ => return None,
MessageType::Image(image_content) => (image_content.body.clone(), false),
_ => return Ok(None),
};
let is_reply = matches!(room_message.relates_to, Some(Relation::Reply { .. }));
@@ -248,17 +255,45 @@ pub fn convert_matrix_native_event_to_matrix_message(
.map(|m| m.user_ids.iter().map(|u| u.to_owned()).collect())
.unwrap_or(vec![]);
Some(MatrixMessage {
if let MessageType::Image(image_content) = &room_message.msgtype {
let media_request = mxlink::matrix_sdk::media::MediaRequestParameters {
source: image_content.source.to_owned(),
format: mxlink::matrix_sdk::media::MediaFormat::File,
};
let file_name = image_content.filename.clone().unwrap_or(image_content.body.clone());
let mime_type = get_mime_type_from_file_name(&file_name);
tracing::debug!("Determined mime type {} for file {}", mime_type, file_name);
let span = tracing::debug_span!("get_media_content", file_name = %file_name, mime_type = %mime_type);
let media_bytes = matrix_link
.client()
.media()
.get_media_content(&media_request, true)
.instrument(span)
.await?;
return Ok(Some(MatrixMessage {
sender_id: matrix_native_event.sender().to_owned(),
content: MatrixMessageContent::Image(image_content.clone(), mime_type, media_bytes),
mentioned_users,
timestamp,
}));
}
Ok(Some(MatrixMessage {
sender_id: matrix_native_event.sender().to_owned(),
message_type: if is_notice {
MatrixMessageType::Notice
content: if is_notice {
MatrixMessageContent::Notice(text)
} else {
MatrixMessageType::Text
MatrixMessageContent::Text(text)
},
message_text: text,
mentioned_users,
timestamp,
})
}))
}
/// Determines the interaction context for an incoming (new) room event.

View File

@@ -3,7 +3,7 @@ use chrono::{TimeZone, Utc};
use mxlink::matrix_sdk::ruma::OwnedUserId;
use crate::conversation::matrix::{
MatrixMessage, MatrixMessageProcessingParams, MatrixMessageType,
MatrixMessage, MatrixMessageContent, MatrixMessageProcessingParams,
};
#[test]
@@ -17,24 +17,21 @@ fn is_message_from_allowed_sender() {
let bot_message = MatrixMessage {
sender_id: bot_user_id.to_owned(),
message_type: MatrixMessageType::Text,
message_text: "Hello!".to_owned(),
content: MatrixMessageContent::Text("Hello!".to_owned()),
mentioned_users: vec![],
timestamp,
};
let allowed_user_message = MatrixMessage {
sender_id: allowed_user_id.to_owned(),
message_type: MatrixMessageType::Text,
message_text: "Hello!".to_owned(),
content: MatrixMessageContent::Text("Hello!".to_owned()),
mentioned_users: vec![],
timestamp,
};
let unallowed_user_message = MatrixMessage {
sender_id: unallowed_user_id.to_owned(),
message_type: MatrixMessageType::Text,
message_text: "Hello!".to_owned(),
content: MatrixMessageContent::Text("Hello!".to_owned()),
mentioned_users: vec![],
timestamp,
};
@@ -88,48 +85,42 @@ async fn process_matrix_messages() {
let allowed_user_message = MatrixMessage {
sender_id: allowed_user_id.to_owned(),
message_type: MatrixMessageType::Text,
message_text: "Hello from the user!".to_owned(),
content: MatrixMessageContent::Text("Hello from the user!".to_owned()),
mentioned_users: vec![],
timestamp,
};
let allowed_user_message_with_prefix = MatrixMessage {
sender_id: allowed_user_id.to_owned(),
message_type: MatrixMessageType::Text,
message_text: "!bai Hello from the user!".to_owned(),
content: MatrixMessageContent::Text("!bai Hello from the user!".to_owned()),
mentioned_users: vec![],
timestamp,
};
let allowed_user_message_with_prefix_no_space = MatrixMessage {
sender_id: allowed_user_id.to_owned(),
message_type: MatrixMessageType::Text,
message_text: "!baiHello from the user!".to_owned(),
content: MatrixMessageContent::Text("!baiHello from the user!".to_owned()),
mentioned_users: vec![],
timestamp,
};
let allowed_user_message_with_prefix_full_width_space = MatrixMessage {
sender_id: allowed_user_id.to_owned(),
message_type: MatrixMessageType::Text,
message_text: "!bai Hello from the user!".to_owned(),
content: MatrixMessageContent::Text("!bai Hello from the user!".to_owned()),
mentioned_users: vec![],
timestamp,
};
let bot_message = MatrixMessage {
sender_id: bot_user_id.to_owned(),
message_type: MatrixMessageType::Text,
message_text: "Hello from the bot!".to_owned(),
content: MatrixMessageContent::Text("Hello from the bot!".to_owned()),
mentioned_users: vec![],
timestamp,
};
let allowed_user_message_with_bot_mention = MatrixMessage {
sender_id: allowed_user_id.to_owned(),
message_type: MatrixMessageType::Text,
message_text: "@baibot: Hello from the user!".to_owned(),
content: MatrixMessageContent::Text("@baibot: Hello from the user!".to_owned()),
mentioned_users: vec![bot_user_id.to_owned()],
timestamp,
};
@@ -137,16 +128,14 @@ async fn process_matrix_messages() {
// The message text is the same as above - it mentions the bot, but the actually-mentioned user is another user.
let allowed_user_message_with_another_user_mention = MatrixMessage {
sender_id: allowed_user_id.to_owned(),
message_type: MatrixMessageType::Text,
message_text: allowed_user_message_with_bot_mention.message_text.clone(),
content: allowed_user_message_with_bot_mention.content.clone(),
mentioned_users: vec![allowed_user_id.to_owned()],
timestamp,
};
let unallowed_user_message = MatrixMessage {
sender_id: unallowed_user_id.to_owned(),
message_type: MatrixMessageType::Text,
message_text: "Hello from an unallowed user!".to_owned(),
content: MatrixMessageContent::Text("Hello from an unallowed user!".to_owned()),
mentioned_users: vec![],
timestamp,
};
@@ -285,7 +274,10 @@ async fn process_matrix_messages() {
let processed_message_texts = processed_messages
.iter()
.map(|message| message.message_text.clone())
.map(|message| match &message.content {
MatrixMessageContent::Text(text) => text.clone(),
_ => "".to_owned(),
})
.collect::<Vec<String>>();
assert_eq!(

View File

@@ -12,7 +12,7 @@ use super::matrix::{
};
pub async fn create_llm_conversation_for_matrix_thread(
matrix_link: MatrixLink,
matrix_link: &MatrixLink,
room: &mxlink::matrix_sdk::Room,
thread_id: OwnedEventId,
params: &MatrixMessageProcessingParams,
@@ -27,12 +27,13 @@ pub async fn create_llm_conversation_for_matrix_thread(
}
pub async fn create_llm_conversation_for_matrix_reply_chain(
matrix_link: &MatrixLink,
event_fetcher: &Arc<RoomEventFetcher>,
room: &mxlink::matrix_sdk::Room,
event_id: OwnedEventId,
params: &MatrixMessageProcessingParams,
) -> Result<Conversation, mxlink::matrix_sdk::Error> {
let messages = get_matrix_messages_in_reply_chain(event_fetcher, room, event_id).await?;
let messages = get_matrix_messages_in_reply_chain(matrix_link, event_fetcher, room, event_id).await?;
let llm_messages = filter_messages_and_convert_to_llm_messages(messages, params).await;

View File

@@ -1,5 +1,5 @@
use mxlink::matrix_sdk::ruma::events::room::message::{
AudioMessageEventContent, MessageType, TextMessageEventContent,
AudioMessageEventContent, ImageMessageEventContent, MessageType, TextMessageEventContent,
};
use mxlink::matrix_sdk::ruma::{OwnedEventId, OwnedUserId};
@@ -29,6 +29,7 @@ pub enum MessagePayload {
Text(TextMessageEventContent),
Audio(AudioMessageEventContent),
Image(ImageMessageEventContent),
Reaction {
key: String,
@@ -55,6 +56,7 @@ impl TryInto<MessagePayload> for MessageType {
// For this reason, we handle all audio.
MessagePayload::Audio(audio_content)
}
MessageType::Image(image_content) => MessagePayload::Image(image_content),
other => {
return Err(format!("Unsupported message type: {:?}", other));
}

View File

@@ -3,5 +3,5 @@ pub fn heading() -> &'static str {
}
pub fn intro() -> &'static str {
"The bot can perform various tasks, such as 💬 Text Generation, 🗣️ Text-to-Speech, 🦻 Speech-to-Text, 🖌️ Image Generation, and more."
"The bot can perform various tasks, such as 💬 Text Generation, 🗣️ Text-to-Speech, 🦻 Speech-to-Text, 🖌️ Image Creation, and more."
}

15
src/strings/image_edit.rs Normal file
View File

@@ -0,0 +1,15 @@
pub fn guide_how_to_proceed() -> String {
let mut message = String::new();
message.push_str("💡 Respond in this thread (in any order) with:\n");
message.push_str("- one or more images: to use the given images for creating an edit\n");
message.push_str("- more messages: to expand on your original prompt\n");
message.push_str(
"- a message saying `go`: to generate an edit with the current prompt\n",
);
message.push_str(
"- a message saying `again`: to generate one more image edit with the current prompt\n",
);
message
}

View File

@@ -8,7 +8,7 @@ pub fn guide_how_to_proceed() -> String {
message.push_str("💡 Respond in this thread with:\n");
message.push_str("- more messages: to expand on your original prompt\n");
message.push_str(
"- a message saying `again`: to generate one more image with the current prompt.\n",
"- a message saying `again`: to generate one more image with the current prompt\n",
);
message

View File

@@ -5,6 +5,7 @@ pub mod error;
pub mod global_config;
pub mod help;
pub mod image_generation;
pub mod image_edit;
pub mod introduction;
pub mod provider;
pub mod room_config;

View File

@@ -34,26 +34,26 @@ By default, the bot will also perform 💬 Text Generation on the text. This is
If all your messages are in the same language, you can improve accuracy & latency by configuring the language via the **🦻 Speech-to-Text / 🔤 Language** setting.
### 🖌️ Image Generation
### 🖌️ Image Creation
#### Generating images
#### Creating images
Simply send a command like `%command_prefix% image A beautiful sunset over the ocean` and the bot will start a threaded conversation and post an image based on your prompt.
Simply send a command like `%command_prefix% image create A beautiful sunset over the ocean` and the bot will start a threaded conversation and post an image based on your prompt.
You can then, respond in the same message thread with:
- more messages, to add more criteria to your prompt.
- a message saying `again`, to generate one more image with the current prompt.
#### Generating stickers
#### Creating stickers
A variation of **generating images** is to generate "sticker images".
A variation of **creating images** is to create "sticker images".
To generate a sticker, send a command like `%command_prefix% sticker A huge bowl of steaming ramen with a mountain of beansprouts on top`.
To create a sticker, send a command like `%command_prefix% sticker A huge bowl of steaming ramen with a mountain of beansprouts on top`.
The difference from **generating images** is that the bot will:
The difference from **creating images** is that the bot will:
- generate a smaller-resolution image (`256x256`) - smaller/quicker, but still good enough for a sticker
- create a smaller-resolution image (`256x256`) - smaller/quicker, but still good enough for a sticker
- potentially switch to a different (cheaper or otherwise more suitable) model, if available
- post the image directly to the room (as a reply to your message), without starting a threaded conversation
"#;

9
src/utils/base64.rs Normal file
View File

@@ -0,0 +1,9 @@
use base64::{Engine as _, engine::general_purpose::STANDARD};
pub(crate) fn base64_decode(base64_string: &str) -> Result<Vec<u8>, base64::DecodeError> {
STANDARD.decode(base64_string)
}
pub(crate) fn base64_encode(data: &[u8]) -> String {
STANDARD.encode(data)
}

34
src/utils/mime.rs Normal file
View File

@@ -0,0 +1,34 @@
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()
}
pub fn get_mime_type_from_file_name(file_name: &str) -> mime::Mime {
let extension = file_name.rsplit('.').next().unwrap_or("");
match extension.to_lowercase().as_str() {
"jpg" | "jpeg" => mime::IMAGE_JPEG,
"png" => mime::IMAGE_PNG,
"gif" => mime::IMAGE_GIF,
"webp" => "image/webp".parse().unwrap(),
"svg" => mime::IMAGE_SVG,
"tiff" | "tif" => "image/tiff".parse().unwrap(),
"bmp" => "image/bmp".parse().unwrap(),
"heic" | "heif" => "image/heic".parse().unwrap(),
"avif" => "image/avif".parse().unwrap(),
_ => mime::APPLICATION_OCTET_STREAM,
}
}

View File

@@ -1,3 +1,5 @@
pub mod status;
pub mod text;
pub mod text_to_speech;
pub(crate) mod mime;
pub(crate) mod base64;