fmt
This commit is contained in:
@@ -7,15 +7,15 @@ use anthropic::types::ContentBlock;
|
||||
use super::super::ControllerTrait;
|
||||
use crate::agent::AgentPurpose;
|
||||
use crate::agent::provider::entity::{
|
||||
ImageGenerationResult, ImageEditResult, ImageSource, PingResult, TextGenerationParams,
|
||||
ImageEditResult, ImageGenerationResult, ImageSource, PingResult, TextGenerationParams,
|
||||
TextGenerationResult, TextToSpeechParams, TextToSpeechResult,
|
||||
};
|
||||
use crate::agent::provider::{
|
||||
ImageGenerationParams, ImageEditParams, SpeechToTextParams, SpeechToTextResult,
|
||||
ImageEditParams, ImageGenerationParams, SpeechToTextParams, SpeechToTextResult,
|
||||
};
|
||||
use crate::conversation::llm::{
|
||||
Author as LLMAuthor, Conversation as LLMConversation, Message as LLMMessage, MessageContent as LLMMessageContent,
|
||||
shorten_messages_list_to_context_size,
|
||||
Author as LLMAuthor, Conversation as LLMConversation, Message as LLMMessage,
|
||||
MessageContent as LLMMessageContent, shorten_messages_list_to_context_size,
|
||||
};
|
||||
use crate::strings;
|
||||
|
||||
|
||||
@@ -1,6 +1,10 @@
|
||||
use anthropic::types::{ContentBlock, ImageSource, Message, MessagesRequest, MessagesRequestBuilder, Role};
|
||||
use anthropic::types::{
|
||||
ContentBlock, ImageSource, Message, MessagesRequest, MessagesRequestBuilder, Role,
|
||||
};
|
||||
|
||||
use crate::conversation::llm::{Author as LLMAuthor, Message as LLMMessage, MessageContent as LLMMessageContent};
|
||||
use crate::conversation::llm::{
|
||||
Author as LLMAuthor, Message as LLMMessage, MessageContent as LLMMessageContent,
|
||||
};
|
||||
|
||||
pub(super) fn create_anthropic_message_request(llm_messages: Vec<LLMMessage>) -> MessagesRequest {
|
||||
let mut messages = vec![];
|
||||
|
||||
@@ -1,9 +1,9 @@
|
||||
use crate::{agent::AgentPurpose, conversation::llm::Conversation};
|
||||
|
||||
use super::{
|
||||
ImageGenerationParams, ImageEditParams, SpeechToTextParams, SpeechToTextResult,
|
||||
ImageEditParams, ImageGenerationParams, SpeechToTextParams, SpeechToTextResult,
|
||||
entity::{
|
||||
ImageGenerationResult, ImageEditResult, ImageSource, PingResult, TextGenerationParams,
|
||||
ImageEditResult, ImageGenerationResult, ImageSource, PingResult, TextGenerationParams,
|
||||
TextGenerationResult, TextToSpeechParams, TextToSpeechResult,
|
||||
},
|
||||
};
|
||||
@@ -180,7 +180,9 @@ impl ControllerTrait for ControllerType {
|
||||
params: ImageEditParams,
|
||||
) -> anyhow::Result<ImageEditResult> {
|
||||
match &self {
|
||||
ControllerType::OpenAI(controller) => controller.create_image_edit(prompt, images, params).await,
|
||||
ControllerType::OpenAI(controller) => {
|
||||
controller.create_image_edit(prompt, images, params).await
|
||||
}
|
||||
ControllerType::OpenAICompat(controller) => {
|
||||
controller.create_image_edit(prompt, images, params).await
|
||||
}
|
||||
|
||||
@@ -33,8 +33,7 @@ pub struct ImageGenerationResult {
|
||||
}
|
||||
|
||||
#[derive(Default)]
|
||||
pub struct ImageEditParams {
|
||||
}
|
||||
pub struct ImageEditParams {}
|
||||
|
||||
pub struct ImageEditResult {
|
||||
pub bytes: Vec<u8>,
|
||||
@@ -49,7 +48,11 @@ pub struct ImageSource {
|
||||
|
||||
impl ImageSource {
|
||||
pub fn new(filename: String, bytes: Vec<u8>, mime_type: mime::Mime) -> Self {
|
||||
Self { filename, bytes, mime_type }
|
||||
Self {
|
||||
filename,
|
||||
bytes,
|
||||
mime_type,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -6,7 +6,9 @@ mod text_generation;
|
||||
mod text_to_speech;
|
||||
|
||||
pub use agent_provider::{AgentProvider, AgentProviderInfo};
|
||||
pub use image::{ImageGenerationParams, ImageGenerationResult, ImageEditParams, ImageEditResult, ImageSource};
|
||||
pub use image::{
|
||||
ImageEditParams, ImageEditResult, ImageGenerationParams, ImageGenerationResult, ImageSource,
|
||||
};
|
||||
pub use ping::PingResult;
|
||||
pub use speech_to_text::{SpeechToTextParams, SpeechToTextResult};
|
||||
pub use text_generation::{
|
||||
|
||||
@@ -20,6 +20,7 @@ pub use controller::{ControllerTrait, ControllerType};
|
||||
pub use config::ConfigTrait;
|
||||
|
||||
pub use entity::{
|
||||
AgentProvider, AgentProviderInfo, ImageGenerationParams, ImageEditParams, ImageSource, PingResult, SpeechToTextParams,
|
||||
SpeechToTextResult, TextGenerationParams, TextGenerationPromptVariables, TextToSpeechParams,
|
||||
AgentProvider, AgentProviderInfo, ImageEditParams, ImageGenerationParams, ImageSource,
|
||||
PingResult, SpeechToTextParams, SpeechToTextResult, TextGenerationParams,
|
||||
TextGenerationPromptVariables, TextToSpeechParams,
|
||||
};
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
use serde::{Deserialize, Serialize};
|
||||
|
||||
use crate::agent::{default_prompt, provider::ConfigTrait};
|
||||
use super::OPENAI_IMAGE_MODEL_GPT_IMAGE_1;
|
||||
use crate::agent::{default_prompt, provider::ConfigTrait};
|
||||
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct Config {
|
||||
|
||||
@@ -4,26 +4,16 @@ use async_openai::{
|
||||
Client as OpenAIClient,
|
||||
config::OpenAIConfig,
|
||||
types::{
|
||||
ChatCompletionRequestMessage, CreateChatCompletionRequestArgs, CreateImageRequestArgs,
|
||||
CreateSpeechRequestArgs, CreateTranscriptionRequestArgs, CreateImageEditRequestArgs,
|
||||
ImageModel, DallE2ImageSize, ImageResponseFormat, Image,
|
||||
ChatCompletionRequestMessage, CreateChatCompletionRequestArgs, CreateImageEditRequestArgs,
|
||||
CreateImageRequestArgs, CreateSpeechRequestArgs, CreateTranscriptionRequestArgs,
|
||||
DallE2ImageSize, Image, ImageModel, ImageResponseFormat,
|
||||
},
|
||||
};
|
||||
|
||||
use super::super::ControllerTrait;
|
||||
use crate::{
|
||||
agent::{
|
||||
AgentPurpose,
|
||||
provider::{
|
||||
entity::{ImageGenerationResult, ImageEditResult, ImageSource, PingResult, TextToSpeechParams, TextToSpeechResult},
|
||||
openai::utils::convert_string_to_enum,
|
||||
},
|
||||
},
|
||||
strings,
|
||||
};
|
||||
use crate::{
|
||||
agent::provider::{
|
||||
ImageGenerationParams, ImageEditParams, SpeechToTextParams, SpeechToTextResult,
|
||||
ImageEditParams, ImageGenerationParams, SpeechToTextParams, SpeechToTextResult,
|
||||
entity::{TextGenerationParams, TextGenerationResult},
|
||||
},
|
||||
conversation::llm::{
|
||||
@@ -32,6 +22,19 @@ use crate::{
|
||||
},
|
||||
utils::base64::base64_decode,
|
||||
};
|
||||
use crate::{
|
||||
agent::{
|
||||
AgentPurpose,
|
||||
provider::{
|
||||
entity::{
|
||||
ImageEditResult, ImageGenerationResult, ImageSource, PingResult,
|
||||
TextToSpeechParams, TextToSpeechResult,
|
||||
},
|
||||
openai::utils::convert_string_to_enum,
|
||||
},
|
||||
},
|
||||
strings,
|
||||
};
|
||||
|
||||
use super::config::Config;
|
||||
|
||||
@@ -272,8 +275,8 @@ impl ControllerTrait for Controller {
|
||||
async_openai::types::ImageQuality::HD => {
|
||||
Some(async_openai::types::ImageQuality::Standard)
|
||||
}
|
||||
}
|
||||
None => None
|
||||
},
|
||||
None => None,
|
||||
}
|
||||
} else {
|
||||
image_generation_config.quality.clone()
|
||||
@@ -390,8 +393,12 @@ impl ControllerTrait for Controller {
|
||||
.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::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),
|
||||
@@ -413,7 +420,8 @@ impl ControllerTrait for Controller {
|
||||
request_builder.response_format(response_format);
|
||||
}
|
||||
|
||||
let request = request_builder.build()
|
||||
let request = request_builder
|
||||
.build()
|
||||
.map_err(|e| anyhow::anyhow!("Failed to build CreateImageEditRequest: {}", e))?;
|
||||
|
||||
tracing::trace!(
|
||||
|
||||
@@ -1,8 +1,13 @@
|
||||
use async_openai::types::{
|
||||
ChatCompletionRequestAssistantMessageArgs, ChatCompletionRequestMessage, ChatCompletionRequestMessageContentPartImage, ChatCompletionRequestSystemMessageArgs, ChatCompletionRequestUserMessageArgs, ChatCompletionRequestUserMessageContent, ChatCompletionRequestUserMessageContentPart, ImageUrlArgs
|
||||
ChatCompletionRequestAssistantMessageArgs, ChatCompletionRequestMessage,
|
||||
ChatCompletionRequestMessageContentPartImage, ChatCompletionRequestSystemMessageArgs,
|
||||
ChatCompletionRequestUserMessageArgs, ChatCompletionRequestUserMessageContent,
|
||||
ChatCompletionRequestUserMessageContentPart, ImageUrlArgs,
|
||||
};
|
||||
|
||||
use crate::conversation::llm::{Author as LLMAuthor, Message as LLMMessage, MessageContent as LLMMessageContent};
|
||||
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(
|
||||
@@ -21,47 +26,53 @@ pub fn convert_llm_messages_to_openai_messages(
|
||||
openai_conversation_messages
|
||||
}
|
||||
|
||||
fn convert_llm_message_to_openai_message(llm_message: LLMMessage) -> Option<ChatCompletionRequestMessage> {
|
||||
fn convert_llm_message_to_openai_message(
|
||||
llm_message: LLMMessage,
|
||||
) -> Option<ChatCompletionRequestMessage> {
|
||||
match &llm_message.content {
|
||||
LLMMessageContent::Text(text) => {
|
||||
Some(match llm_message.author {
|
||||
LLMAuthor::Prompt => ChatCompletionRequestSystemMessageArgs::default()
|
||||
.content(text.clone())
|
||||
.build()
|
||||
.expect("Failed building OpenAI system message")
|
||||
.into(),
|
||||
LLMAuthor::Assistant => ChatCompletionRequestAssistantMessageArgs::default()
|
||||
.content(text.clone())
|
||||
.build()
|
||||
.expect("Failed building OpenAI assistant message")
|
||||
.into(),
|
||||
LLMAuthor::User => ChatCompletionRequestUserMessageArgs::default()
|
||||
.content(text.clone())
|
||||
.build()
|
||||
.expect("Failed building OpenAI user message")
|
||||
.into(),
|
||||
})
|
||||
}
|
||||
LLMMessageContent::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 image_url = format!(
|
||||
"data:{};base64,{}",
|
||||
image_details.mime,
|
||||
base64_encode(&image_details.data)
|
||||
);
|
||||
|
||||
let part = ChatCompletionRequestUserMessageContentPart::ImageUrl(
|
||||
ChatCompletionRequestMessageContentPartImage{
|
||||
ChatCompletionRequestMessageContentPartImage {
|
||||
image_url: ImageUrlArgs::default()
|
||||
.url(image_url)
|
||||
.build()
|
||||
.expect("Failed building OpenAI image url")
|
||||
}
|
||||
.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()),
|
||||
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.",
|
||||
|
||||
@@ -230,13 +230,17 @@ impl TryInto<OpenAIImageGenerationConfig> for ImageGenerationConfig {
|
||||
};
|
||||
|
||||
let style = if let Some(style) = &self.style {
|
||||
Some(convert_string_to_enum::<async_openai::types::ImageStyle>(style)?)
|
||||
Some(convert_string_to_enum::<async_openai::types::ImageStyle>(
|
||||
style,
|
||||
)?)
|
||||
} else {
|
||||
None
|
||||
};
|
||||
|
||||
let quality = if let Some(quality) = &self.quality {
|
||||
Some(convert_string_to_enum::<async_openai::types::ImageQuality>(quality)?)
|
||||
Some(convert_string_to_enum::<async_openai::types::ImageQuality>(
|
||||
quality,
|
||||
)?)
|
||||
} else {
|
||||
None
|
||||
};
|
||||
|
||||
@@ -7,7 +7,8 @@ use super::super::ControllerTrait;
|
||||
use crate::utils::base64::base64_decode;
|
||||
use crate::{
|
||||
agent::provider::{
|
||||
ImageGenerationParams, ImageEditParams, ImageSource, SpeechToTextParams, SpeechToTextResult,
|
||||
ImageEditParams, ImageGenerationParams, ImageSource, SpeechToTextParams,
|
||||
SpeechToTextResult,
|
||||
entity::{TextGenerationParams, TextGenerationResult},
|
||||
},
|
||||
conversation::llm::{
|
||||
@@ -19,7 +20,7 @@ use crate::{
|
||||
agent::{
|
||||
AgentPurpose,
|
||||
provider::entity::{
|
||||
ImageGenerationResult, ImageEditResult, PingResult, TextToSpeechParams,
|
||||
ImageEditResult, ImageGenerationResult, PingResult, TextToSpeechParams,
|
||||
TextToSpeechResult,
|
||||
},
|
||||
},
|
||||
|
||||
@@ -2,7 +2,9 @@ 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, MessageContent as LLMMessageContent};
|
||||
use crate::conversation::llm::{
|
||||
Author as LLMAuthor, Message as LLMMessage, MessageContent as LLMMessageContent,
|
||||
};
|
||||
|
||||
pub fn convert_llm_messages_to_openai_messages(
|
||||
conversation_messages: Vec<LLMMessage>,
|
||||
@@ -28,16 +30,16 @@ fn convert_llm_message_to_openai_message(llm_message: LLMMessage) -> Option<Mess
|
||||
};
|
||||
|
||||
match &llm_message.content {
|
||||
LLMMessageContent::Text(text) => {
|
||||
Some(Message {
|
||||
role,
|
||||
content: text.clone(),
|
||||
})
|
||||
},
|
||||
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.");
|
||||
tracing::warn!(
|
||||
"The OpenAI-compat provider's library does not support image content. This image message will be skipped."
|
||||
);
|
||||
None
|
||||
},
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -42,9 +42,7 @@ pub fn determine_controller(
|
||||
.text_generation_prefix_requirement_type();
|
||||
|
||||
match prefix_requirement_type {
|
||||
TextGenerationPrefixRequirementType::CommandPrefix => {
|
||||
ControllerType::Ignore
|
||||
}
|
||||
TextGenerationPrefixRequirementType::CommandPrefix => ControllerType::Ignore,
|
||||
TextGenerationPrefixRequirementType::No => {
|
||||
ControllerType::ChatCompletion(ChatCompletionControllerType::Image)
|
||||
}
|
||||
|
||||
@@ -78,13 +78,8 @@ pub async fn dispatch_controller(
|
||||
.await
|
||||
}
|
||||
ControllerType::ImageEdit(prompt) => {
|
||||
super::image::edit::handle(
|
||||
bot,
|
||||
bot.matrix_link().clone(),
|
||||
message_context,
|
||||
prompt,
|
||||
)
|
||||
.await
|
||||
super::image::edit::handle(bot, bot.matrix_link().clone(), message_context, prompt)
|
||||
.await
|
||||
}
|
||||
ControllerType::StickerGeneration(prompt) => {
|
||||
super::image::generation::handle_sticker(
|
||||
|
||||
@@ -1,9 +1,7 @@
|
||||
use crate::controller::ControllerType;
|
||||
mod tests;
|
||||
|
||||
pub fn determine_controller(
|
||||
text: &str,
|
||||
) -> ControllerType {
|
||||
pub fn determine_controller(text: &str) -> ControllerType {
|
||||
let text = text.trim();
|
||||
|
||||
if let Some(prompt) = text.strip_prefix("create") {
|
||||
|
||||
@@ -22,11 +22,12 @@ fn determine_controller() {
|
||||
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()),
|
||||
expected: super::ControllerType::ImageEdit(
|
||||
"Turn this into an anime-style image".to_owned(),
|
||||
),
|
||||
},
|
||||
];
|
||||
|
||||
|
||||
@@ -2,15 +2,15 @@ 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::agent::provider::ImageSource;
|
||||
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::utils::mime::get_file_extension;
|
||||
use crate::{Bot, entity::MessageContext};
|
||||
|
||||
pub async fn handle(
|
||||
@@ -69,23 +69,25 @@ pub async fn handle(
|
||||
}
|
||||
});
|
||||
|
||||
let image_sources: Vec<ImageSource> = conversation.messages.iter().filter_map(|message| {
|
||||
if let crate::conversation::llm::MessageContent::Image(image_content) = &message.content {
|
||||
Some(image_content.clone().into())
|
||||
} else {
|
||||
None
|
||||
}
|
||||
}).collect();
|
||||
let image_sources: Vec<ImageSource> = conversation
|
||||
.messages
|
||||
.iter()
|
||||
.filter_map(|message| {
|
||||
if let crate::conversation::llm::MessageContent::Image(image_content) = &message.content
|
||||
{
|
||||
Some(image_content.clone().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 span = tracing::debug_span!("image_edit", agent_id = agent.identifier().as_string());
|
||||
|
||||
let result = agent
|
||||
.controller()
|
||||
@@ -147,10 +149,7 @@ pub async fn handle(
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn send_guide(
|
||||
bot: &Bot,
|
||||
message_context: &MessageContext,
|
||||
) -> anyhow::Result<()> {
|
||||
async fn send_guide(bot: &Bot, message_context: &MessageContext) -> anyhow::Result<()> {
|
||||
bot.messaging()
|
||||
.send_text_markdown_no_fail(
|
||||
message_context.room(),
|
||||
|
||||
@@ -6,10 +6,10 @@ 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::utils::mime::get_file_extension;
|
||||
use crate::conversation::create_llm_conversation_for_matrix_thread;
|
||||
use crate::conversation::matrix::MatrixMessageProcessingParams;
|
||||
use crate::strings;
|
||||
use crate::utils::mime::get_file_extension;
|
||||
use crate::{Bot, entity::MessageContext};
|
||||
|
||||
// We may make this configurable (per room, etc.) in the future, but for now it's hardcoded.
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
pub mod generation;
|
||||
pub mod edit;
|
||||
mod prompt;
|
||||
mod determination;
|
||||
pub mod edit;
|
||||
pub mod generation;
|
||||
mod prompt;
|
||||
|
||||
pub use determination::determine_controller;
|
||||
|
||||
@@ -29,9 +29,7 @@ pub fn build(original_prompt: &str, other_messages: Vec<Message>) -> String {
|
||||
prompt.push_str("\nOther criteria:");
|
||||
for message in other_messages {
|
||||
if let MessageContent::Text(text) = &message.content {
|
||||
prompt.push_str(
|
||||
format!("\n- {}", text.replace("\n", ". ").as_str()).as_str(),
|
||||
);
|
||||
prompt.push_str(format!("\n- {}", text.replace("\n", ". ").as_str()).as_str());
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -87,7 +85,9 @@ mod tests {
|
||||
},
|
||||
Message {
|
||||
author: Author::User,
|
||||
content: MessageContent::Text("Must be 3-legged.\nMust be flying.".to_owned()),
|
||||
content: MessageContent::Text(
|
||||
"Must be 3-legged.\nMust be flying.".to_owned(),
|
||||
),
|
||||
timestamp,
|
||||
},
|
||||
],
|
||||
|
||||
@@ -27,11 +27,18 @@ pub struct ImageDetails {
|
||||
|
||||
impl ImageDetails {
|
||||
pub fn new(event_content: ImageMessageEventContent, mime: Mime, data: Vec<u8>) -> Self {
|
||||
Self { event_content, mime, data }
|
||||
Self {
|
||||
event_content,
|
||||
mime,
|
||||
data,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn filename(&self) -> String {
|
||||
self.event_content.filename.clone().unwrap_or(self.event_content.body.clone())
|
||||
self.event_content
|
||||
.filename
|
||||
.clone()
|
||||
.unwrap_or(self.event_content.body.clone())
|
||||
}
|
||||
}
|
||||
|
||||
@@ -54,7 +61,7 @@ impl PartialEq for MessageContent {
|
||||
(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,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -20,11 +20,15 @@ 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.content, MessageContent::Text("Hello!".to_string()));
|
||||
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");
|
||||
|
||||
@@ -43,7 +47,10 @@ 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.content, MessageContent::Text(source_message_text.to_string()));
|
||||
assert_eq!(
|
||||
llm_message.content,
|
||||
MessageContent::Text(source_message_text.to_string())
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
|
||||
@@ -150,8 +150,7 @@ pub mod test {
|
||||
let third = super::Message {
|
||||
author: super::Author::User,
|
||||
content: super::MessageContent::Text(
|
||||
"This is the 3rd message in this conversation. It shall be preserved."
|
||||
.to_owned(),
|
||||
"This is the 3rd message in this conversation. It shall be preserved.".to_owned(),
|
||||
),
|
||||
timestamp: chrono::Utc::now(),
|
||||
};
|
||||
|
||||
@@ -23,17 +23,15 @@ fn convert_bot_message(matrix_message: &MatrixMessage) -> Option<Message> {
|
||||
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(),
|
||||
})
|
||||
}
|
||||
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(),
|
||||
}),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -73,30 +71,24 @@ fn convert_bot_notice_message(
|
||||
|
||||
fn convert_user_message(matrix_message: &MatrixMessage) -> Option<Message> {
|
||||
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(),
|
||||
})
|
||||
}
|
||||
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(),
|
||||
}),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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, MatrixMessageContent};
|
||||
pub(crate) use entity::{MatrixMessage, MatrixMessageContent, MatrixMessageProcessingParams};
|
||||
|
||||
pub(crate) use utils::*;
|
||||
|
||||
@@ -42,7 +42,9 @@ pub async fn get_matrix_messages_in_thread(
|
||||
let mut messages: Vec<MatrixMessage> = Vec::new();
|
||||
|
||||
for matrix_native_message in messages_native {
|
||||
let message_result = convert_matrix_native_event_to_matrix_message(matrix_link, &matrix_native_message).await?;
|
||||
let message_result =
|
||||
convert_matrix_native_event_to_matrix_message(matrix_link, &matrix_native_message)
|
||||
.await?;
|
||||
|
||||
if let Some(message) = message_result {
|
||||
messages.push(message);
|
||||
@@ -64,7 +66,9 @@ pub async fn get_matrix_messages_in_reply_chain(
|
||||
let mut messages: Vec<MatrixMessage> = Vec::new();
|
||||
|
||||
for matrix_native_message in messages_native {
|
||||
let message_result = convert_matrix_native_event_to_matrix_message(matrix_link, &matrix_native_message).await?;
|
||||
let message_result =
|
||||
convert_matrix_native_event_to_matrix_message(matrix_link, &matrix_native_message)
|
||||
.await?;
|
||||
|
||||
if let Some(message) = message_result {
|
||||
messages.push(message);
|
||||
@@ -261,7 +265,10 @@ pub async fn convert_matrix_native_event_to_matrix_message(
|
||||
format: mxlink::matrix_sdk::media::MediaFormat::File,
|
||||
};
|
||||
|
||||
let file_name = image_content.filename.clone().unwrap_or(image_content.body.clone());
|
||||
let file_name = image_content
|
||||
.filename
|
||||
.clone()
|
||||
.unwrap_or(image_content.body.clone());
|
||||
|
||||
let mime_type = get_mime_type_from_file_name(&file_name);
|
||||
|
||||
|
||||
@@ -33,7 +33,8 @@ pub async fn create_llm_conversation_for_matrix_reply_chain(
|
||||
event_id: OwnedEventId,
|
||||
params: &MatrixMessageProcessingParams,
|
||||
) -> Result<Conversation, mxlink::matrix_sdk::Error> {
|
||||
let messages = get_matrix_messages_in_reply_chain(matrix_link, 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;
|
||||
|
||||
|
||||
@@ -4,9 +4,7 @@ pub fn guide_how_to_proceed() -> String {
|
||||
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 `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",
|
||||
);
|
||||
|
||||
@@ -4,8 +4,8 @@ pub mod cfg;
|
||||
pub mod error;
|
||||
pub mod global_config;
|
||||
pub mod help;
|
||||
pub mod image_generation;
|
||||
pub mod image_edit;
|
||||
pub mod image_generation;
|
||||
pub mod introduction;
|
||||
pub mod provider;
|
||||
pub mod room_config;
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
pub(crate) mod base64;
|
||||
pub(crate) mod mime;
|
||||
pub mod status;
|
||||
pub mod text;
|
||||
pub mod text_to_speech;
|
||||
pub(crate) mod mime;
|
||||
pub(crate) mod base64;
|
||||
Reference in New Issue
Block a user