diff --git a/Cargo.lock b/Cargo.lock index 71428e6..6a3767f 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -227,9 +227,9 @@ dependencies = [ [[package]] name = "async-openai" -version = "0.28.0" +version = "0.28.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "36c566b15aa847e60a9e6c9b9b4b9d4be94bbf776804624279afa69559fea7e1" +checksum = "14d76e2f5af19477d6254415acc95ba97c6cc6f3b1e3cb4676b7f0fab8194298" dependencies = [ "async-openai-macros", "backoff", @@ -353,7 +353,7 @@ version = "1.6.0" dependencies = [ "anthropic", "anyhow", - "async-openai 0.28.0", + "async-openai 0.28.1", "base64 0.22.1", "chrono", "etke_openai_api_rust", diff --git a/docs/agents.md b/docs/agents.md index 6ae1c77..fb6e218 100644 --- a/docs/agents.md +++ b/docs/agents.md @@ -35,7 +35,7 @@ Depending on where the agent is defined (within a room, globally, or [statically When creating an agent, you will be given some sample [YAML](https://en.wikipedia.org/wiki/YAML) configuration which you can use to customize the agent's behavior. -This configuration varies depending on the [โ˜๏ธ provider](./providers.md) used and the capabilities of the agent. Based on the configuration keys you pass, certain features will be enabled or disabled. For example, if you skip the `image_generation` key for an [OpenAI](./providers.md#openai) agent, it won't be able to generate images (see [๐Ÿ–Œ๏ธ Image Generation](./features.md#-image-generation)). +This configuration varies depending on the [โ˜๏ธ provider](./providers.md) used and the capabilities of the agent. Based on the configuration keys you pass, certain features will be enabled or disabled. For example, if you skip the `image_generation` key for an [OpenAI](./providers.md#openai) agent, it won't be able to generate images (see [๐Ÿ–Œ๏ธ Image Creation](./features.md#-image-creation), [๐ŸŽจ Image Editing](./features.md#-image-editing), [๐Ÿซต Sticker Creation](./features.md#-sticker-creation)). After making your modifications to the sample YAML, you submit it back to the bot and the new agent will be created. diff --git a/docs/configuration/README.md b/docs/configuration/README.md index 0ca2555..e941b1f 100644 --- a/docs/configuration/README.md +++ b/docs/configuration/README.md @@ -40,7 +40,7 @@ You can adjust the following settings per room and/or globally: - [๐Ÿ’ฌ Text Generation](text-generation.md) - [๐Ÿฆป Speech-to-Text](speech-to-text.md) - [๐Ÿ—ฃ๏ธ Text-to-Speech](text-to-speech.md) -- [๐Ÿ–Œ๏ธ Image Generation](image-generation.md) +- [๐Ÿ–Œ๏ธ Image Creation](image-generation.md) - [๐Ÿค Handlers](handlers.md) Refer to the bot's help messages (as a response to a `!bai config` help command) for the most up-to-date information on what Room Settings can be configured. diff --git a/docs/configuration/handlers.md b/docs/configuration/handlers.md index cddc6fa..e12ee16 100644 --- a/docs/configuration/handlers.md +++ b/docs/configuration/handlers.md @@ -11,7 +11,7 @@ The bot supports the following use-purposes: - [๐Ÿ’ฌ text-generation](../features.md#-text-generation): communicating with you via text - [๐Ÿฆป speech-to-text](../features.md#-speech-to-text): turning your voice messages into text - [๐Ÿ—ฃ๏ธ text-to-speech](../features.md#๏ธ-text-to-speech): turning bot or users text messages into voice messages -- [๐Ÿ–Œ๏ธ image-generation](../features.md#-image-generation): generating images based on instructions +- [๐Ÿ–Œ๏ธ image-generation](../features.md#image-generation): generating images based on instructions In a given room, each different purpose can be served by a different [provider](../providers.md) and model. This combination of provider and model configuration is called an [๐Ÿค– agent](../agents.md). Each purpose can be served by a different **handler** agent. diff --git a/docs/configuration/image-generation.md b/docs/configuration/image-generation.md index a3237ab..2d019f4 100644 --- a/docs/configuration/image-generation.md +++ b/docs/configuration/image-generation.md @@ -1,9 +1,11 @@ -## ๐Ÿ–Œ๏ธ Image Generation +## Image Generation -The Image Generation feature is not configurable at this moment. +The Image Creation and Image Editing features are not configurable at this moment. You may also wish to see: -- [๐ŸŒŸ Features / ๐Ÿ–Œ๏ธ Image Generation](../features.md#-image-generation) for a higher-level introduction to the Image Generation features -- [๐Ÿ“– Usage / ๐Ÿ–Œ๏ธ Image Generation](../usage.md#-image-generation) section for more details on how to use the bot for Image Generation in a room +- [๐ŸŒŸ Features / Image Generation / ๐Ÿ–Œ๏ธ Image Creation](../features.md#-image-creation) for a higher-level introduction to the Image Creation features +- [๐ŸŒŸ Features / Image Generation / ๐ŸŽจ Image Editing](../features.md#-image-editing) for a higher-level introduction to the Image Editing features +- [๐Ÿ“– Usage / Image Generation / ๐Ÿ–Œ๏ธ Creating Images](../usage.md#-creating-images) section for more details on how to use the bot for Image Creation in a room +- [๐Ÿ“– Usage / Image Generation / ๐ŸŽจ Editing images](../usage.md#-editing-images) section for more details on how to use the bot for Image Editing in a room diff --git a/docs/development.md b/docs/development.md index 4a9fe50..6386954 100644 --- a/docs/development.md +++ b/docs/development.md @@ -93,7 +93,7 @@ For getting started most quickly (and locally), we recommend using [LocalAI](#lo **Ollama is most lightweight** (~2GB for the container image + ~1.6GB for the model), but supports only [๐Ÿ’ฌ text-generation](./features.md#-text-generation). -**LocalAI requires 4x more disk space** (~6GB for the container image + ~12GB for the models), but supports [๐Ÿ’ฌ text-generation](./features.md#-text-generation), [๐Ÿ—ฃ๏ธ text-to-speech](./features.md#๏ธ-text-to-speech), [๐Ÿฆป speech-to-text](./features.md#-speech-to-text) and [๐Ÿ–ผ๏ธ image-generation](./features.md#๏ธ-image-generation). +**LocalAI requires 4x more disk space** (~6GB for the container image + ~12GB for the models), but supports [๐Ÿ’ฌ text-generation](./features.md#-text-generation), [๐Ÿ—ฃ๏ธ text-to-speech](./features.md#๏ธ-text-to-speech), [๐Ÿฆป speech-to-text](./features.md#-speech-to-text) and [๐Ÿ–ผ๏ธ image-generation](./features.md#๏ธ-image-creation). **OpenAI supports all of these capabilities** as well and does not require powerful hardware or lots of disk space. However, it requires signup and an API key. diff --git a/docs/features.md b/docs/features.md index 00d260d..73011f5 100644 --- a/docs/features.md +++ b/docs/features.md @@ -139,26 +139,42 @@ To operate in this mode, you can: - optionally adjust [๐Ÿฆป Speech-to-Text / ๐Ÿช„ Message Type for non-threaded only-transcribed messages](./configuration/speech-to-text.md#-message-type-for-non-threaded-only-transcribed-messages), if you'd like to bot to send messages of type `notice` (for better compatibility with other bots in the room) instead of sending regular `text` messages (default) -### ๐Ÿ–Œ๏ธ Image Generation +### Image Generation -Image generation is the bot's ability to **generate images** based on text prompts. +#### ๐Ÿ–Œ๏ธ Image Creation -See a [๐Ÿ–ผ๏ธ Screenshot of the Image Generation feature](./screenshots/image-generation.webp). +Image creation is the bot's ability to **create images** based on text prompts. + +See a [๐Ÿ–ผ๏ธ Screenshot of the Image Creation feature](./screenshots/image-creation.webp). You may also wish to see: - [๐Ÿ› ๏ธ Configuration / ๐Ÿ–Œ๏ธ Image Generation](./configuration/image-generation.md) for configuration options related to Image Generation -- [๐Ÿ“– Usage / ๐Ÿ–Œ๏ธ Image Generation](./usage.md#-image-generation) section for more details on how to use the bot for Image Generation in a room -- [๐Ÿซต Sticker Generation](#-sticker-generation) - a special case of Image Generation +- [๐Ÿ“– Usage / Image Generation / ๐Ÿ–Œ๏ธ Creating Images](./usage.md#-creating-images) section for more details on how to use the bot for Image Creation in a room +- [๐Ÿ–Œ๏ธ Image Editing](#๏ธ-image-editing) - another image generation feature +- [๐Ÿซต Sticker Creation](#-sticker-creation) - a special case of Image Creation -### ๐Ÿซต Sticker Generation +#### ๐ŸŽจ Image Editing -Sticker generation is the bot's ability to **generate sticker** images based on text prompts. It's a special case of [๐Ÿ–Œ๏ธ Image Generation](#๏ธ-image-generation). +Image editing is the bot's ability to **edit images** based on a prompt and one or more existing images. -See a [๐Ÿ–ผ๏ธ Screenshot of the Sticker Generation feature](./screenshots/sticker-generation.webp). +See a [๐Ÿ–ผ๏ธ Screenshot of the Image Editing feature](./screenshots/image-editing.webp). -See [๐Ÿ“– Usage / ๐Ÿ–Œ๏ธ Image Generation / Generating Stickers](./usage.md#generating-stickers) for details. +You may also wish to see: + +- [๐Ÿ› ๏ธ Configuration / ๐Ÿ–Œ๏ธ Image Generation](./configuration/image-generation.md) for configuration options related to Image Generation +- [๐Ÿ“– Usage / Image Generation / ๐ŸŽจ Editing images](./usage.md#-editing-images) section for more details on how to use the bot for Image Editing in a room +- [๐Ÿ–Œ๏ธ Image Creation](#๏ธ-image-creation) - another image generation feature + + +#### ๐Ÿซต Sticker Creation + +Sticker generation is the bot's ability to **generate sticker** images based on text prompts. It's a special case of [๐Ÿ–Œ๏ธ Image Creation](#๏ธ-image-creation). + +See a [๐Ÿ–ผ๏ธ Screenshot of the Sticker Creation feature](./screenshots/sticker-generation.webp). + +See [๐Ÿ“– Usage / Image Generation / ๐Ÿซต Creating Stickers](./usage.md#-creating-stickers) for details. ### ๐Ÿ”’ Encryption diff --git a/docs/providers.md b/docs/providers.md index 6aa9274..39a08d5 100644 --- a/docs/providers.md +++ b/docs/providers.md @@ -23,7 +23,7 @@ The list of supported providers is below. ### How to choose a provider -If you're not sure which provider to start with, **we recommend [OpenAI](#openai)** as it's the most popular and has the **widest range of capabilities**: [๐Ÿ’ฌ text-generation](./features.md#-text-generation), [๐Ÿ–Œ๏ธ image-generation](./features.md#๏ธ-image-generation), [๐Ÿฆป speech-to-text](./features.md#-speech-to-text), [๐Ÿ—ฃ๏ธ text-to-speech](./features.md#๏ธ-text-to-speech). +If you're not sure which provider to start with, **we recommend [OpenAI](#openai)** as it's the most popular and has the **widest range of capabilities**: [๐Ÿ’ฌ text-generation](./features.md#-text-generation), [๐Ÿ–Œ๏ธ image-generation](./features.md#๏ธimage-generation), [๐Ÿฆป speech-to-text](./features.md#-speech-to-text), [๐Ÿ—ฃ๏ธ text-to-speech](./features.md#๏ธ-text-to-speech). You don't need to choose just one though. The bot supports [mixing & matching models](./features.md#-mixing--matching-models), so you can use multiple providers at the same time. @@ -120,7 +120,7 @@ For services which are not fully compatible with the OpenAI API, consider using - ๐Ÿ†” Identifier: `openai` - ๐Ÿ”— Links: [๐Ÿ  Home page](https://openai.com/), [๐ŸŒ Wiki](https://en.wikipedia.org/wiki/OpenAI), [๐Ÿ‘ค Sign up](https://platform.openai.com/signup), [๐Ÿ“‹ Models list](https://platform.openai.com/docs/models) -- ๐ŸŒŸ Capabilities: [๐Ÿ–Œ๏ธ image-generation](./features.md#๏ธ-image-generation), [๐Ÿ’ฌ text-generation](./features.md#-text-generation), [๐Ÿ—ฃ๏ธ text-to-speech](./features.md#๏ธ-text-to-speech), [๐Ÿฆป speech-to-text](./features.md#-speech-to-text) +- ๐ŸŒŸ Capabilities: [๐Ÿ–Œ๏ธ image-generation](./features.md#๏ธ-image-creation), [๐Ÿ’ฌ text-generation](./features.md#-text-generation), [๐Ÿ—ฃ๏ธ text-to-speech](./features.md#๏ธ-text-to-speech), [๐Ÿฆป speech-to-text](./features.md#-speech-to-text) - ๐Ÿ—ฒ Quick start: - create a room-local agent: `!bai agent create-room-local openai my-openai-agent` - create a global agent: `!bai agent create-global openai my-openai-agent` @@ -140,7 +140,7 @@ Some of these popular services already have **shortcut** providers (leading to t This provider is just as featureful as the [OpenAI](#openai) provider, but is more compatible with services which do not fully adhere to the [OpenAI API spec](https://github.com/openai/openai-openapi/). - ๐Ÿ†” Identifier: `openai-compatible` -- ๐ŸŒŸ Capabilities: [๐Ÿ–Œ๏ธ image-generation](./features.md#๏ธ-image-generation), [๐Ÿ’ฌ text-generation](./features.md#-text-generation), [๐Ÿ—ฃ๏ธ text-to-speech](./features.md#๏ธ-text-to-speech), [๐Ÿฆป speech-to-text](./features.md#-speech-to-text) +- ๐ŸŒŸ Capabilities: [๐Ÿ–Œ๏ธ image-generation](./features.md#๏ธ-image-creation), [๐Ÿ’ฌ text-generation](./features.md#-text-generation), [๐Ÿ—ฃ๏ธ text-to-speech](./features.md#๏ธ-text-to-speech), [๐Ÿฆป speech-to-text](./features.md#-speech-to-text) - ๐Ÿ—ฒ Quick start: - create a room-local agent: `!bai agent create-room-local openai-compatible my-openai-compatible-agent` - create a global agent: `!bai agent create-global openai-compatible my-openai-compatible-agent` diff --git a/docs/screenshots/image-creation.webp b/docs/screenshots/image-creation.webp new file mode 100644 index 0000000..b8355db Binary files /dev/null and b/docs/screenshots/image-creation.webp differ diff --git a/docs/screenshots/image-editing.webp b/docs/screenshots/image-editing.webp new file mode 100644 index 0000000..f5d944c Binary files /dev/null and b/docs/screenshots/image-editing.webp differ diff --git a/docs/screenshots/image-generation.webp b/docs/screenshots/image-generation.webp deleted file mode 100644 index eceb0ea..0000000 Binary files a/docs/screenshots/image-generation.webp and /dev/null differ diff --git a/docs/usage.md b/docs/usage.md index d49c80d..4abe87a 100644 --- a/docs/usage.md +++ b/docs/usage.md @@ -64,18 +64,18 @@ The speech-to-text feature triggers automatically by default, but can be adjuste If all your messages are in the same language, you can improve accuracy & latency by configuring the language (see [๐Ÿฆป Speech-to-Text / ๐Ÿ”ค Language](./configuration/speech-to-text.md#-language)). -### ๐Ÿ–Œ๏ธ Image Generation - -This is related to the [๐Ÿ–Œ๏ธ Image Generation](./features.md#๏ธ-image-generation) feature. +### Image Generation This feature is not configurable at the moment. The configuration (size, quality, style) specified at the [๐Ÿค– agent](./agents.md) level will be used. +Capabilities depend on the [โ˜๏ธ provider](./providers.md) and model used. -#### Generating images -Simply send a command like `!bai image A beautiful sunset over the ocean` and the bot will start a threaded conversation and post an image based on your prompt. +#### ๐Ÿ–Œ๏ธ Creating images -See a [๐Ÿ–ผ๏ธ Screenshot of the Image Generation feature](./screenshots/image-generation.webp). +Simply send a command like `!bai image create A beautiful sunset over the ocean` and the bot will start a threaded conversation and post an image based on your prompt. + +See a [๐Ÿ–ผ๏ธ Screenshot of the Image Creation feature](./screenshots/image-creation.webp). You can then, respond in the same message thread with: @@ -83,15 +83,29 @@ You can then, respond in the same message thread with: - a message saying `again`, to generate one more image with the current prompt. -#### Generating stickers +#### ๐ŸŽจ Editing images -A variation of [generating images](#generating-images) is to generate "sticker images". +Simply send a command like `!bai image edit Turn the following image into an anime-style drawing` and the bot will start a threaded conversation asking for more details. -See a [๐Ÿ–ผ๏ธ Screenshot of the Sticker Generation feature](./screenshots/sticker-generation.webp). +See a [๐Ÿ–ผ๏ธ Screenshot of the Image Editing feature](./screenshots/image-editing.webp). -To generate a sticker, send a command like `!bai sticker A huge ramen bowl with lots of chashu and a mountain of beansprouts on top`. +You can then, respond in the same message thread with: -The difference from [generating images](#generating-images) is that the bot will: +- more messages, to add more criteria to your prompt. +- one or more images, to provide the images that the bot will operate on. +- a message saying `go`, to start the image generation process. +- a message saying `again`, to prompt the bot to generate one more image edit with the current prompt. + + +#### ๐Ÿซต Creating stickers + +A variation of [creating images](#creating-images) is to create "sticker images". + +See a [๐Ÿ–ผ๏ธ Screenshot of the Sticker Creation feature](./screenshots/sticker-generation.webp). + +To create a sticker, send a command like `!bai sticker A huge ramen bowl with lots of chashu and a mountain of beansprouts on top`. + +The difference from [creating images](#creating-images) is that the bot will: - generate a smaller-resolution image (currently hardcoded to `256x256`) - smaller/quicker, but still good enough for a sticker - potentially switch to a different (cheaper or otherwise more suitable) model, if available diff --git a/src/agent/provider/anthropic/controller.rs b/src/agent/provider/anthropic/controller.rs index 2966211..98f3d47 100644 --- a/src/agent/provider/anthropic/controller.rs +++ b/src/agent/provider/anthropic/controller.rs @@ -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, + _params: ImageEditParams, + ) -> anyhow::Result { + Err(anyhow::anyhow!("Image editing is not supported")) + } + async fn text_to_speech( &self, _input: &str, diff --git a/src/agent/provider/anthropic/utils.rs b/src/agent/provider/anthropic/utils.rs index 25a0745..c246d95 100644 --- a/src/agent/provider/anthropic/utils.rs +++ b/src/agent/provider/anthropic/utils.rs @@ -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) -> MessagesRequest { let mut messages = vec![]; @@ -14,13 +14,28 @@ pub(super) fn create_anthropic_message_request(llm_messages: Vec) -> } }; - 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() diff --git a/src/agent/provider/controller.rs b/src/agent/provider/controller.rs index 18b2772..3f8e50c 100644 --- a/src/agent/provider/controller.rs +++ b/src/agent/provider/controller.rs @@ -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> + Send; + fn create_image_edit( + &self, + prompt: &str, + images: Vec, + params: ImageEditParams, + ) -> impl std::future::Future> + 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, + params: ImageEditParams, + ) -> anyhow::Result { + 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, diff --git a/src/agent/provider/entity/image.rs b/src/agent/provider/entity/image.rs new file mode 100644 index 0000000..53975dd --- /dev/null +++ b/src/agent/provider/entity/image.rs @@ -0,0 +1,65 @@ +use mxlink::mime; + +#[derive(Default)] +pub struct ImageGenerationParams { + pub size_override: Option, + + pub cheaper_model_switching_allowed: bool, + + pub cheaper_quality_switching_allowed: bool, +} + +impl ImageGenerationParams { + pub fn with_size_override(mut self, value: Option) -> 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, + pub mime_type: mime::Mime, + pub revised_prompt: Option, +} + +#[derive(Default)] +pub struct ImageEditParams { +} + +pub struct ImageEditResult { + pub bytes: Vec, + pub mime_type: mime::Mime, +} + +pub struct ImageSource { + pub filename: String, + pub bytes: Vec, + pub mime_type: mime::Mime, +} + +impl ImageSource { + pub fn new(filename: String, bytes: Vec, mime_type: mime::Mime) -> Self { + Self { filename, bytes, mime_type } + } +} + +impl Into 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, + }, + } + } +} diff --git a/src/agent/provider/entity/image_generation.rs b/src/agent/provider/entity/image_generation.rs deleted file mode 100644 index f34096a..0000000 --- a/src/agent/provider/entity/image_generation.rs +++ /dev/null @@ -1,31 +0,0 @@ -#[derive(Default)] -pub struct ImageGenerationParams { - pub size_override: Option, - - pub cheaper_model_switching_allowed: bool, - - pub cheaper_quality_switching_allowed: bool, -} - -impl ImageGenerationParams { - pub fn with_size_override(mut self, value: Option) -> 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, - pub mime_type: mxlink::mime::Mime, - pub revised_prompt: Option, -} diff --git a/src/agent/provider/entity/mod.rs b/src/agent/provider/entity/mod.rs index 4cbeffd..89331dc 100644 --- a/src/agent/provider/entity/mod.rs +++ b/src/agent/provider/entity/mod.rs @@ -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::{ diff --git a/src/agent/provider/mod.rs b/src/agent/provider/mod.rs index 495cff7..297814c 100644 --- a/src/agent/provider/mod.rs +++ b/src/agent/provider/mod.rs @@ -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, }; diff --git a/src/agent/provider/openai/controller.rs b/src/agent/provider/openai/controller.rs index d700bb5..809995b 100644 --- a/src/agent/provider/openai/controller.rs +++ b/src/agent/provider/openai/controller.rs @@ -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, + _params: ImageEditParams, + ) -> anyhow::Result { + 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, diff --git a/src/agent/provider/openai/utils.rs b/src/agent/provider/openai/utils.rs index c392e7a..b19bea0 100644 --- a/src/agent/provider/openai/utils.rs +++ b/src/agent/provider/openai/utils.rs @@ -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, @@ -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 { + 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 + } + } + } } } diff --git a/src/agent/provider/openai_compat/controller.rs b/src/agent/provider/openai_compat/controller.rs index 760ee68..d517da5 100644 --- a/src/agent/provider/openai_compat/controller.rs +++ b/src/agent/provider/openai_compat/controller.rs @@ -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, + _params: ImageEditParams, + ) -> anyhow::Result { + Err(anyhow::anyhow!( + "The OpenAI image edit API is not supported by the OpenAI-compat provider" + )) + } + async fn text_to_speech( &self, input: &str, diff --git a/src/agent/provider/openai_compat/utils.rs b/src/agent/provider/openai_compat/utils.rs index b25b072..2efc8f9 100644 --- a/src/agent/provider/openai_compat/utils.rs +++ b/src/agent/provider/openai_compat/utils.rs @@ -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, @@ -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 { 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 + }, } } diff --git a/src/agent/utils.rs b/src/agent/utils.rs index 980eb48..e954da9 100644 --- a/src/agent/utils.rs +++ b/src/agent/utils.rs @@ -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, base64::DecodeError> { - STANDARD.decode(base64_string) -} diff --git a/src/controller/cfg/status.rs b/src/controller/cfg/status.rs index 1a49fe1..c41a3c3 100644 --- a/src/controller/cfg/status.rs +++ b/src/controller/cfg/status.rs @@ -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, diff --git a/src/controller/chat_completion/mod.rs b/src/controller/chat_completion/mod.rs index 1844d6d..3787b39 100644 --- a/src/controller/chat_completion/mod.rs +++ b/src/controller/chat_completion/mod.rs @@ -39,6 +39,8 @@ pub enum ChatCompletionControllerType { Audio, + Image, + ThreadMention, ReplyMention, } @@ -416,7 +418,8 @@ async fn handle_stage_text_generation( ChatCompletionControllerType::TextCommand | ChatCompletionControllerType::TextMention | ChatCompletionControllerType::TextDirect - | ChatCompletionControllerType::Audio => { + | ChatCompletionControllerType::Audio + | ChatCompletionControllerType::Image => { Some(message_context.combined_admin_and_user_regexes()) } @@ -438,6 +441,7 @@ async fn handle_stage_text_generation( // When we're triggered via a reply mention, the context is the whole reply chain upward of the message that triggered us. ChatCompletionControllerType::ReplyMention => { create_llm_conversation_for_matrix_reply_chain( + &matrix_link, &bot.room_event_fetcher().clone(), message_context.room(), message_context.thread_info().last_event_id.clone(), @@ -449,7 +453,7 @@ async fn handle_stage_text_generation( // Everything else is happening in a thread, so the context is the whole thread. _ => { create_llm_conversation_for_matrix_thread( - matrix_link.clone(), + &matrix_link, message_context.room(), message_context.thread_info().root_event_id.clone(), ¶ms, diff --git a/src/controller/controller_type.rs b/src/controller/controller_type.rs index 5d30703..f64861a 100644 --- a/src/controller/controller_type.rs +++ b/src/controller/controller_type.rs @@ -23,5 +23,6 @@ pub enum ControllerType { ChatCompletion(super::chat_completion::ChatCompletionControllerType), ImageGeneration(String), + ImageEdit(String), StickerGeneration(String), } diff --git a/src/controller/determination/mod.rs b/src/controller/determination/mod.rs index 9c1d6de..6a04674 100644 --- a/src/controller/determination/mod.rs +++ b/src/controller/determination/mod.rs @@ -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")) { diff --git a/src/controller/determination/tests.rs b/src/controller/determination/tests.rs index ccd5ed7..d98a008 100644 --- a/src/controller/determination/tests.rs +++ b/src/controller/determination/tests.rs @@ -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()), diff --git a/src/controller/dispatching.rs b/src/controller/dispatching.rs index ed60ad1..088fa37 100644 --- a/src/controller/dispatching.rs +++ b/src/controller/dispatching.rs @@ -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, diff --git a/src/controller/image/determination/mod.rs b/src/controller/image/determination/mod.rs new file mode 100644 index 0000000..3192486 --- /dev/null +++ b/src/controller/image/determination/mod.rs @@ -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 +} diff --git a/src/controller/image/determination/tests.rs b/src/controller/image/determination/tests.rs new file mode 100644 index 0000000..ac2192b --- /dev/null +++ b/src/controller/image/determination/tests.rs @@ -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); + } +} diff --git a/src/controller/image/edit.rs b/src/controller/image/edit.rs new file mode 100644 index 0000000..884db17 --- /dev/null +++ b/src/controller/image/edit.rs @@ -0,0 +1,163 @@ +use mxlink::{MatrixLink, MessageResponseType}; + +use tracing::Instrument; + +use crate::agent::provider::ImageSource; +use crate::agent::AgentPurpose; +use crate::agent::ControllerTrait; +use crate::agent::provider::ImageEditParams; +use crate::controller::utils::agent::get_effective_agent_for_purpose_or_complain; +use crate::utils::mime::get_file_extension; +use crate::conversation::create_llm_conversation_for_matrix_thread; +use crate::conversation::matrix::MatrixMessageProcessingParams; +use crate::strings; +use crate::{Bot, entity::MessageContext}; + +pub async fn handle( + bot: &Bot, + matrix_link: MatrixLink, + message_context: &MessageContext, + original_prompt: &str, +) -> anyhow::Result<()> { + let response_type = MessageResponseType::InThread(message_context.thread_info().clone()); + + let Some(agent) = get_effective_agent_for_purpose_or_complain( + bot, + message_context, + AgentPurpose::ImageGeneration, + response_type.clone(), + true, + ) + .await + else { + return Ok(()); + }; + + if message_context.thread_info().is_thread_root_only() { + return send_guide(bot, message_context).await; + } + + let _typing_notice_guard = bot.start_typing_notice(message_context.room()).await; + + let params = MatrixMessageProcessingParams::new( + bot.user_id().to_owned(), + Some(message_context.combined_admin_and_user_regexes()), + ); + + let conversation = create_llm_conversation_for_matrix_thread( + &matrix_link, + message_context.room(), + message_context.thread_info().root_event_id.clone(), + ¶ms, + ) + .await?; + + let prompt = if conversation.messages.len() >= 2 { + // Skip the first message, which contains the original prompt (which we already have) + let other_messages = conversation.messages.iter().skip(1).cloned().collect(); + + super::prompt::build(original_prompt, other_messages) + } else { + original_prompt.to_owned() + }; + + let got_go_signal = conversation.messages.iter().any(|message| { + if let crate::conversation::llm::MessageContent::Text(text) = &message.content { + text.to_lowercase() == "go" + } else { + false + } + }); + + let image_sources: Vec = 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(()) +} diff --git a/src/controller/image/generation.rs b/src/controller/image/generation.rs index b3c58a7..eb527c7 100644 --- a/src/controller/image/generation.rs +++ b/src/controller/image/generation.rs @@ -6,7 +6,7 @@ use crate::agent::AgentPurpose; use crate::agent::ControllerTrait; use crate::agent::provider::ImageGenerationParams; use crate::controller::utils::agent::get_effective_agent_for_purpose_or_complain; -use crate::controller::utils::mime::get_file_extension; +use crate::utils::mime::get_file_extension; use crate::conversation::create_llm_conversation_for_matrix_thread; use crate::conversation::matrix::MatrixMessageProcessingParams; use crate::strings; @@ -43,7 +43,7 @@ pub async fn handle_image( ); let conversation = create_llm_conversation_for_matrix_thread( - matrix_link.clone(), + &matrix_link, message_context.room(), message_context.thread_info().root_event_id.clone(), ¶ms, diff --git a/src/controller/image/mod.rs b/src/controller/image/mod.rs index fcfd033..cadd5b4 100644 --- a/src/controller/image/mod.rs +++ b/src/controller/image/mod.rs @@ -1,2 +1,6 @@ pub mod generation; +pub mod edit; mod prompt; +mod determination; + +pub use determination::determine_controller; diff --git a/src/controller/image/prompt.rs b/src/controller/image/prompt.rs index 4e5bfc1..b4ef691 100644 --- a/src/controller/image/prompt.rs +++ b/src/controller/image/prompt.rs @@ -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) -> String { let mut prompt = original_prompt.to_owned(); @@ -14,7 +14,11 @@ pub fn build(original_prompt: &str, other_messages: Vec) -> 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) -> 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) -> 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, }, ], diff --git a/src/controller/utils/mime.rs b/src/controller/utils/mime.rs deleted file mode 100644 index 9efe4c4..0000000 --- a/src/controller/utils/mime.rs +++ /dev/null @@ -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() -} diff --git a/src/controller/utils/mod.rs b/src/controller/utils/mod.rs index 6b35a9d..795db6e 100644 --- a/src/controller/utils/mod.rs +++ b/src/controller/utils/mod.rs @@ -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>( diff --git a/src/controller/utils/text_to_speech.rs b/src/controller/utils/text_to_speech.rs index 9b6091e..73e3146 100644 --- a/src/controller/utils/text_to_speech.rs +++ b/src/controller/utils/text_to_speech.rs @@ -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}, diff --git a/src/conversation/llm/entity.rs b/src/conversation/llm/entity.rs index 76569b7..cd3c410 100644 --- a/src/conversation/llm/entity.rs +++ b/src/conversation/llm/entity.rs @@ -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, + pub content: MessageContent, } +#[derive(Debug, Clone)] +pub struct ImageDetails { + pub event_content: ImageMessageEventContent, + pub mime: Mime, + pub data: Vec, +} + +impl ImageDetails { + pub fn new(event_content: ImageMessageEventContent, mime: Mime, data: Vec) -> 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 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, @@ -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 = None; + let mut last_seen_text_from_author: Option = 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); } } diff --git a/src/conversation/llm/tests.rs b/src/conversation/llm/tests.rs index 7428557..5ec360f 100644 --- a/src/conversation/llm/tests.rs +++ b/src/conversation/llm/tests.rs @@ -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(), }; diff --git a/src/conversation/llm/tokenization.rs b/src/conversation/llm/tokenization.rs index 0d70806..4d9b66a 100644 --- a/src/conversation/llm/tokenization.rs +++ b/src/conversation/llm/tokenization.rs @@ -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 ); } } diff --git a/src/conversation/llm/utils.rs b/src/conversation/llm/utils.rs index 3c071ad..eae0e12 100644 --- a/src/conversation/llm/utils.rs +++ b/src/conversation/llm/utils.rs @@ -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 { - 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 { 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 { - 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(), + }) + } + } } diff --git a/src/conversation/matrix/entity.rs b/src/conversation/matrix/entity.rs index be4b3f5..28c3f14 100644 --- a/src/conversation/matrix/entity.rs +++ b/src/conversation/matrix/entity.rs @@ -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, pub timestamp: DateTime, } #[derive(Clone)] -pub enum MatrixMessageType { - Text, - Notice, +pub enum MatrixMessageContent { + Text(String), + Notice(String), + Image(ImageMessageEventContent, Mime, Vec), } #[derive(Clone)] diff --git a/src/conversation/matrix/mod.rs b/src/conversation/matrix/mod.rs index 10a930f..34f651f 100644 --- a/src/conversation/matrix/mod.rs +++ b/src/conversation/matrix/mod.rs @@ -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::*; diff --git a/src/conversation/matrix/utils/mod.rs b/src/conversation/matrix/utils/mod.rs index 576a04b..4bd5048 100644 --- a/src/conversation/matrix/utils/mod.rs +++ b/src/conversation/matrix/utils/mod.rs @@ -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, mxlink::matrix_sdk::Error> { @@ -40,18 +42,18 @@ pub async fn get_matrix_messages_in_thread( let mut messages: Vec = 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, room: &Room, event_id: OwnedEventId, @@ -62,12 +64,11 @@ pub async fn get_matrix_messages_in_reply_chain( let mut messages: Vec = 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 ¶ms.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 ¶ms.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(¶ms.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 ¶ms.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 ¶ms.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 { +) -> Result, 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. diff --git a/src/conversation/matrix/utils/tests.rs b/src/conversation/matrix/utils/tests.rs index f479cf2..95f095a 100644 --- a/src/conversation/matrix/utils/tests.rs +++ b/src/conversation/matrix/utils/tests.rs @@ -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::>(); assert_eq!( diff --git a/src/conversation/matrix_llm_bridge.rs b/src/conversation/matrix_llm_bridge.rs index e042043..fbf48ae 100644 --- a/src/conversation/matrix_llm_bridge.rs +++ b/src/conversation/matrix_llm_bridge.rs @@ -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, room: &mxlink::matrix_sdk::Room, event_id: OwnedEventId, params: &MatrixMessageProcessingParams, ) -> Result { - 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; diff --git a/src/entity/message_payload.rs b/src/entity/message_payload.rs index 027f69f..12a4463 100644 --- a/src/entity/message_payload.rs +++ b/src/entity/message_payload.rs @@ -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 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)); } diff --git a/src/strings/help/usage.rs b/src/strings/help/usage.rs index 2408795..455f9c4 100644 --- a/src/strings/help/usage.rs +++ b/src/strings/help/usage.rs @@ -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." } diff --git a/src/strings/image_edit.rs b/src/strings/image_edit.rs new file mode 100644 index 0000000..8386329 --- /dev/null +++ b/src/strings/image_edit.rs @@ -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 +} diff --git a/src/strings/image_generation.rs b/src/strings/image_generation.rs index 4c13734..f67701e 100644 --- a/src/strings/image_generation.rs +++ b/src/strings/image_generation.rs @@ -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 diff --git a/src/strings/mod.rs b/src/strings/mod.rs index c2aa17a..6f06808 100644 --- a/src/strings/mod.rs +++ b/src/strings/mod.rs @@ -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; diff --git a/src/strings/usage.rs b/src/strings/usage.rs index 0115bb4..ad0c5fb 100644 --- a/src/strings/usage.rs +++ b/src/strings/usage.rs @@ -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 "#; diff --git a/src/utils/base64.rs b/src/utils/base64.rs new file mode 100644 index 0000000..1ead2d7 --- /dev/null +++ b/src/utils/base64.rs @@ -0,0 +1,9 @@ +use base64::{Engine as _, engine::general_purpose::STANDARD}; + +pub(crate) fn base64_decode(base64_string: &str) -> Result, base64::DecodeError> { + STANDARD.decode(base64_string) +} + +pub(crate) fn base64_encode(data: &[u8]) -> String { + STANDARD.encode(data) +} diff --git a/src/utils/mime.rs b/src/utils/mime.rs new file mode 100644 index 0000000..96b2998 --- /dev/null +++ b/src/utils/mime.rs @@ -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, + } +} diff --git a/src/utils/mod.rs b/src/utils/mod.rs index b362efb..5ab8e94 100644 --- a/src/utils/mod.rs +++ b/src/utils/mod.rs @@ -1,3 +1,5 @@ pub mod status; pub mod text; pub mod text_to_speech; +pub(crate) mod mime; +pub(crate) mod base64; \ No newline at end of file