Initial commit

This commit is contained in:
Slavi Pantaleev
2024-09-12 13:44:06 +03:00
commit 946aa9d9e9
220 changed files with 26033 additions and 0 deletions

55
src/agent/definition.rs Normal file
View File

@@ -0,0 +1,55 @@
use serde::de::Error as DeError;
use serde::{Deserialize, Deserializer, Serialize, Serializer};
use super::provider::AgentProvider;
// Custom serialization for AgentProvider
pub fn serialize_provider_to_string<S>(
value: &AgentProvider,
serializer: S,
) -> Result<S::Ok, S::Error>
where
S: Serializer,
{
serializer.serialize_str(value.to_static_str())
}
// Custom deserialization for AgentProvider
pub fn deserialize_provider_from_string<'de, D>(deserializer: D) -> Result<AgentProvider, D::Error>
where
D: Deserializer<'de>,
{
let s = String::deserialize(deserializer)?;
AgentProvider::from_string(&s).map_err(DeError::custom)
}
#[derive(Debug, Clone, Deserialize, Serialize)]
pub struct AgentDefinition {
pub id: String,
#[serde(
serialize_with = "serialize_provider_to_string",
deserialize_with = "deserialize_provider_from_string"
)]
pub provider: AgentProvider,
pub config: serde_yaml::Value,
}
impl AgentDefinition {
pub fn new(id: String, provider: AgentProvider, config: serde_yaml::Value) -> Self {
Self {
id,
provider,
config,
}
}
}
impl PartialEq for AgentDefinition {
fn eq(&self, other: &Self) -> bool {
self.id == other.id
}
}
impl Eq for AgentDefinition {}

101
src/agent/identifier.rs Normal file
View File

@@ -0,0 +1,101 @@
use std::fmt;
#[derive(Debug, PartialEq, Eq, Clone)]
pub enum PublicIdentifier {
Static(String),
DynamicGlobal(String),
DynamicRoomLocal(String),
}
impl PublicIdentifier {
pub fn from_str(s: &str) -> Option<Self> {
if let Some(rest) = s.strip_prefix("static/") {
return Some(PublicIdentifier::Static(rest.to_string()));
} else if let Some(rest) = s.strip_prefix("global/") {
return Some(PublicIdentifier::DynamicGlobal(rest.to_string()));
} else if let Some(rest) = s.strip_prefix("room-local/") {
return Some(PublicIdentifier::DynamicRoomLocal(rest.to_string()));
}
None
}
pub fn as_string(&self) -> String {
match self {
PublicIdentifier::Static(s) => format!("static/{}", s),
PublicIdentifier::DynamicGlobal(s) => format!("global/{}", s),
PublicIdentifier::DynamicRoomLocal(s) => format!("room-local/{}", s),
}
}
pub fn prefixless(&self) -> String {
match self {
PublicIdentifier::Static(s) => s.to_owned(),
PublicIdentifier::DynamicGlobal(s) => s.to_owned(),
PublicIdentifier::DynamicRoomLocal(s) => s.to_owned(),
}
}
pub fn validate(&self) -> Result<(), String> {
let prefixless = self.prefixless();
if prefixless.is_empty() {
return Err("The agent ID must not be empty.".to_owned());
}
// We use a slash to separate the agent type from the agent ID.
if prefixless.contains("/") {
return Err("The agent ID must not contain the `/` character.".to_owned());
}
// Spaces are used for separating command arguments, etc.
if prefixless.contains(" ") {
return Err("The agent ID must not contain spaces.".to_owned());
}
Ok(())
}
}
impl fmt::Display for PublicIdentifier {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(f, "{}", self.as_string())
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_public_identifier_from_str() {
assert_eq!(
PublicIdentifier::from_str("static/abc"),
Some(PublicIdentifier::Static("abc".to_string()))
);
assert_eq!(
PublicIdentifier::from_str("global/abc"),
Some(PublicIdentifier::DynamicGlobal("abc".to_string()))
);
assert_eq!(
PublicIdentifier::from_str("room-local/abc"),
Some(PublicIdentifier::DynamicRoomLocal("abc".to_string()))
);
assert_eq!(PublicIdentifier::from_str("abc"), None);
}
#[test]
fn test_public_identifier_as_string() {
assert_eq!(
PublicIdentifier::Static("abc".to_string()).as_string(),
"static/abc"
);
assert_eq!(
PublicIdentifier::DynamicGlobal("abc".to_string()).as_string(),
"global/abc"
);
assert_eq!(
PublicIdentifier::DynamicRoomLocal("abc".to_string()).as_string(),
"room-local/abc"
);
}
}

154
src/agent/instantiation.rs Normal file
View File

@@ -0,0 +1,154 @@
use super::{
provider::{self, ControllerType},
AgentDefinition, AgentProvider, PublicIdentifier,
};
// Dead-code is allowed. We do not use these enum struct payloads directly,
// but these errors are being print-formatted (`{:?}`) in error messages, so we wish to keep them.
#[derive(Debug)]
#[allow(dead_code)]
pub enum Error {
// Contains the error message from the validation function
ConfigFailsValidation(String),
// Contains the agent ID
ConfigForAgentIsNotAMapping(String),
// Contains the error from the constructor function
ConstructionFailed(anyhow::Error),
// Contains the error from the YAML deserialization function
Yaml(serde_yaml::Error),
}
pub type Result<T> = std::result::Result<T, Error>;
#[derive(Debug, Clone)]
pub struct AgentInstance {
identifier: PublicIdentifier,
definition: AgentDefinition,
controller: ControllerType,
}
impl AgentInstance {
pub fn new(
identifier: PublicIdentifier,
definition: AgentDefinition,
controller: ControllerType,
) -> Self {
Self {
identifier,
definition,
controller,
}
}
pub fn identifier(&self) -> &PublicIdentifier {
&self.identifier
}
pub fn definition(&self) -> &AgentDefinition {
&self.definition
}
pub fn controller(&self) -> &ControllerType {
&self.controller
}
}
pub(super) fn create(
identifier: PublicIdentifier,
definition: AgentDefinition,
) -> Result<AgentInstance> {
let controller = create_controller_from_provider_and_json_value_config(
&definition.id,
&definition.provider,
definition.config.clone(),
)?;
Ok(AgentInstance::new(identifier, definition, controller))
}
pub fn create_from_provider_and_yaml_value_config(
provider: &AgentProvider,
identifier: &PublicIdentifier,
config: serde_yaml::Value,
) -> Result<AgentInstance> {
let definition = AgentDefinition::new(identifier.prefixless(), provider.to_owned(), config);
create(identifier.to_owned(), definition)
}
fn create_controller_from_provider_and_json_value_config(
agent_id: &str,
provider: &AgentProvider,
config: serde_yaml::Value,
) -> Result<ControllerType> {
match provider {
AgentProvider::Anthropic => {
provider::anthropic::create_controller_from_yaml_value_config(agent_id, config)
}
AgentProvider::Groq => {
provider::openai_compat::create_controller_from_yaml_value_config(agent_id, config)
}
AgentProvider::Mistral => {
provider::openai_compat::create_controller_from_yaml_value_config(agent_id, config)
}
AgentProvider::LocalAI => {
provider::openai_compat::create_controller_from_yaml_value_config(agent_id, config)
}
AgentProvider::Ollama => {
provider::openai_compat::create_controller_from_yaml_value_config(agent_id, config)
}
AgentProvider::OpenAI => {
provider::openai::create_controller_from_yaml_value_config(agent_id, config)
}
AgentProvider::OpenAICompat => {
provider::openai_compat::create_controller_from_yaml_value_config(agent_id, config)
}
AgentProvider::OpenRouter => {
provider::openai_compat::create_controller_from_yaml_value_config(agent_id, config)
}
AgentProvider::TogetherAI => {
provider::openai_compat::create_controller_from_yaml_value_config(agent_id, config)
}
}
}
pub fn default_config_for_provider(provider: &AgentProvider) -> serde_yaml::Value {
match provider {
AgentProvider::Anthropic => {
let config = super::provider::anthropic::default_config();
serde_yaml::to_value(config).expect("Failed to serialize config")
}
AgentProvider::Groq => {
let config = super::provider::groq::default_config();
serde_yaml::to_value(config).expect("Failed to serialize config")
}
AgentProvider::LocalAI => {
let config = super::provider::localai::default_config();
serde_yaml::to_value(config).expect("Failed to serialize config")
}
AgentProvider::Mistral => {
let config = super::provider::mistral::default_config();
serde_yaml::to_value(config).expect("Failed to serialize config")
}
AgentProvider::Ollama => {
let config = super::provider::ollama::default_config();
serde_yaml::to_value(config).expect("Failed to serialize config")
}
AgentProvider::OpenAI => {
let config = super::provider::openai::default_config();
serde_yaml::to_value(config).expect("Failed to serialize config")
}
AgentProvider::OpenAICompat => {
let config = super::provider::openai_compat::default_config();
serde_yaml::to_value(config).expect("Failed to serialize config")
}
AgentProvider::OpenRouter => {
let config = super::provider::openrouter::default_config();
serde_yaml::to_value(config).expect("Failed to serialize config")
}
AgentProvider::TogetherAI => {
let config = super::provider::togetherai::default_config();
serde_yaml::to_value(config).expect("Failed to serialize config")
}
}
}

68
src/agent/manager.rs Normal file
View File

@@ -0,0 +1,68 @@
use super::instantiation;
use super::instantiation::AgentInstance;
use super::AgentDefinition;
use super::PublicIdentifier;
use crate::entity::RoomConfigContext;
#[derive(Debug)]
pub struct Manager {
static_agents: Vec<AgentInstance>,
}
impl Manager {
pub fn new(static_agent_definitions: Vec<AgentDefinition>) -> anyhow::Result<Self> {
let mut static_agents = Vec::with_capacity(static_agent_definitions.len());
for definition in static_agent_definitions {
let identifier = PublicIdentifier::Static(definition.id.clone());
match instantiation::create(identifier.clone(), definition.to_owned()) {
Ok(instance) => static_agents.push(instance),
Err(e) => {
return Err(anyhow::anyhow!(
"Failed to create static agent {}: {:?}",
identifier,
e
));
}
}
}
Ok(Self { static_agents })
}
pub fn available_room_agents_by_room_config_context(
&self,
room_config_context: &RoomConfigContext,
) -> Vec<AgentInstance> {
let mut agents: Vec<AgentInstance> = vec![];
for agent in &self.static_agents {
agents.push(agent.clone());
}
for definition in &room_config_context.global_config.agents {
let identifier = PublicIdentifier::DynamicGlobal(definition.id.clone());
match instantiation::create(identifier.clone(), definition.to_owned()) {
Ok(instance) => agents.push(instance),
Err(e) => {
tracing::warn!("Failed to create {} agent: {:?}. Skipping.", identifier, e);
}
}
}
for definition in &room_config_context.room_config.agents {
let identifier = PublicIdentifier::DynamicRoomLocal(definition.id.clone());
match instantiation::create(identifier.clone(), definition.to_owned()) {
Ok(instance) => agents.push(instance),
Err(e) => {
tracing::warn!("Failed to create {} agent: {:?}. Skipping.", identifier, e);
}
}
}
agents
}
}

21
src/agent/mod.rs Normal file
View File

@@ -0,0 +1,21 @@
mod definition;
mod identifier;
mod instantiation;
mod manager;
pub mod provider;
mod purpose;
pub mod utils;
pub use identifier::PublicIdentifier;
pub use manager::Manager;
pub use definition::AgentDefinition;
pub use instantiation::create_from_provider_and_yaml_value_config;
pub use instantiation::default_config_for_provider;
pub use instantiation::AgentInstance;
pub use instantiation::Error as AgentInstantiationError;
pub use instantiation::Result as AgentInstantiationResult;
pub use provider::{AgentProvider, AgentProviderInfo, ControllerTrait};
pub use purpose::AgentPurpose;

View File

@@ -0,0 +1,71 @@
use serde::{Deserialize, Serialize};
use anthropic_rs::models::claude::ClaudeModel;
use crate::agent::provider::ConfigTrait;
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct Config {
pub base_url: String,
pub api_key: String,
pub text_generation: Option<TextGenerationConfig>,
}
impl Default for Config {
fn default() -> Self {
Self {
base_url: "https://api.anthropic.com/v1".to_owned(),
api_key: "YOUR_API_KEY_HERE".to_owned(),
text_generation: Some(TextGenerationConfig::default()),
}
}
}
impl ConfigTrait for Config {
fn validate(&self) -> Result<(), String> {
if self.base_url.is_empty() {
return Err("The base URL must not be empty.".to_owned());
}
if self.api_key.is_empty() {
return Err("The API key must not be empty.".to_owned());
}
Ok(())
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct TextGenerationConfig {
#[serde(default = "default_text_model_id")]
pub model_id: String,
#[serde(default)]
pub prompt: Option<String>,
#[serde(default = "super::super::default_temperature")]
pub temperature: f32,
#[serde(default)]
pub max_response_tokens: u32,
#[serde(default)]
pub max_context_tokens: u32,
}
impl Default for TextGenerationConfig {
fn default() -> Self {
Self {
model_id: default_text_model_id(),
prompt: Some("You are a brief, but helpful bot.".to_owned()),
temperature: super::super::default_temperature(),
max_response_tokens: 8192,
max_context_tokens: 204_800,
}
}
}
fn default_text_model_id() -> String {
ClaudeModel::Claude35Sonnet.as_str().to_owned()
}

View File

@@ -0,0 +1,251 @@
use std::fmt::Debug;
use std::str::FromStr;
use std::sync::Arc;
use anthropic_rs::completion::message::ContentType;
use anthropic_rs::{
client::Client as AnthropicClient, config::Config as AnthropicConfig,
models::claude::ClaudeModel,
};
use super::super::ControllerTrait;
use crate::agent::provider::entity::{
ImageGenerationResult, PingResult, TextGenerationParams, TextGenerationResult,
TextToSpeechParams, TextToSpeechResult,
};
use crate::agent::provider::{ImageGenerationParams, SpeechToTextParams, SpeechToTextResult};
use crate::agent::AgentPurpose;
use crate::conversation::llm::{
shorten_messages_list_to_context_size, Author as LLMAuthor, Conversation as LLMConversation,
Message as LLMMessage,
};
use crate::strings;
use super::config::Config;
struct ControllerInner {
client: AnthropicClient,
}
#[derive(Clone)]
pub struct Controller {
config: Config,
inner: Arc<ControllerInner>,
}
impl Debug for Controller {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("Controller")
.field("config", &self.config)
.finish()
}
}
impl Controller {
pub fn new(config: Config) -> anyhow::Result<Self> {
let anthropic_config =
AnthropicConfig::new(config.api_key.clone()).with_base_url(config.base_url.clone());
let client = match AnthropicClient::new(anthropic_config) {
Ok(client) => client,
Err(err) => {
return Err(anyhow::anyhow!(
"Failed to create Anthropic client: {}",
err.to_string()
));
}
};
Ok(Self {
config,
inner: Arc::new(ControllerInner { client }),
})
}
}
impl ControllerTrait for Controller {
async fn ping(&self) -> anyhow::Result<PingResult> {
if !self.supports_purpose(AgentPurpose::TextGeneration) {
return Ok(PingResult::Inconclusive);
}
let messages = vec![LLMMessage {
author: LLMAuthor::User,
message_text: "Hello!".to_string(),
}];
let conversation = LLMConversation { messages };
self.generate_text(conversation, TextGenerationParams::default())
.await?;
Ok(PingResult::Successful)
}
async fn generate_text(
&self,
conversation: LLMConversation,
params: TextGenerationParams,
) -> anyhow::Result<TextGenerationResult> {
let Some(text_generation_config) = &self.config.text_generation else {
return Err(anyhow::anyhow!(
strings::agent::no_configuration_for_purpose_so_cannot_be_used(
&AgentPurpose::TextGeneration
),
));
};
let prompt_text = params
.prompt_override
.unwrap_or(self.text_generation_prompt().unwrap_or("".to_owned()))
.trim()
.to_owned();
let prompt_message = if prompt_text.is_empty() {
None
} else {
Some(LLMMessage {
author: LLMAuthor::Prompt,
message_text: prompt_text,
})
};
let mut conversation_messages = conversation.messages;
if params.context_management_enabled {
tracing::trace!("Shortening messages list to context size");
conversation_messages = shorten_messages_list_to_context_size(
&text_generation_config.model_id,
&prompt_message,
conversation_messages,
text_generation_config.max_response_tokens,
text_generation_config.max_context_tokens,
);
tracing::trace!("Finished shortening messages list to context size");
};
let messages_count = conversation_messages.len();
let mut request = super::utils::create_anthropic_message_request(conversation_messages);
let model = match ClaudeModel::from_str(&text_generation_config.model_id) {
Ok(model) => model,
Err(err) => {
tracing::debug!(?err, "Failed to parse model ID");
return Err(anyhow::anyhow!(
"Failed to parse model ID: {}",
&text_generation_config.model_id
));
}
};
let temperature = params
.temperature_override
.unwrap_or(text_generation_config.temperature);
if let Some(prompt_message) = prompt_message {
request.system = Some(prompt_message.message_text);
}
request.model = model;
request.temperature = Some(temperature);
request.max_tokens = text_generation_config.max_response_tokens;
if let Ok(request_as_json) = serde_json::to_string(&request) {
tracing::trace!(
model = format!("{:?}", request.model),
?messages_count,
request = request_as_json,
"Sending Anthropic create message API request"
);
}
let response = self.inner.client.create_message(request).await?;
tracing::trace!(?response, "Got response from Anthropic create message API");
// response.content usually contains a single element, but we support handling multiple to account for all possibilities
let mut text_parts = vec![];
for content in response.content {
let content_type = content.content_type;
match content_type {
ContentType::Text => {
text_parts.push(content.text);
} // There are no other content types to handle yet, but there may be in the future
}
}
if text_parts.is_empty() {
return Err(anyhow::anyhow!(
"No text content in response from the Anthropic create message API"
));
}
Ok(TextGenerationResult {
text: text_parts.join("\n\n"),
})
}
async fn speech_to_text(
&self,
_mime_type: &mxlink::mime::Mime,
_media: Vec<u8>,
_params: SpeechToTextParams,
) -> anyhow::Result<SpeechToTextResult> {
Err(anyhow::anyhow!("Speech-to-Text not supported"))
}
async fn generate_image(
&self,
_prompt: &str,
_params: ImageGenerationParams,
) -> anyhow::Result<ImageGenerationResult> {
Err(anyhow::anyhow!("Image generation not supported"))
}
async fn text_to_speech(
&self,
_input: &str,
_params: TextToSpeechParams,
) -> anyhow::Result<TextToSpeechResult> {
Err(anyhow::anyhow!("Speech generation not supported"))
}
fn supports_purpose(&self, purpose: AgentPurpose) -> bool {
match purpose {
AgentPurpose::TextGeneration => self.config.text_generation.is_some(),
AgentPurpose::SpeechToText => false,
AgentPurpose::TextToSpeech => false,
AgentPurpose::ImageGeneration => false,
AgentPurpose::CatchAll => true,
}
}
fn text_generation_prompt(&self) -> Option<String> {
let Some(text_generation_config) = &self.config.text_generation else {
return None;
};
text_generation_config.prompt.clone()
}
fn text_generation_temperature(&self) -> Option<f32> {
let Some(text_generation_config) = &self.config.text_generation else {
return None;
};
Some(text_generation_config.temperature)
}
fn text_to_speech_voice(&self) -> Option<String> {
None
}
fn text_to_speech_speed(&self) -> Option<f32> {
None
}
}

View File

@@ -0,0 +1,43 @@
mod config;
mod controller;
mod utils;
pub use config::Config;
pub use controller::Controller;
use super::super::AgentInstantiationError;
use super::super::AgentInstantiationResult;
use super::controller::ControllerType;
use super::ConfigTrait;
pub fn create_controller_from_yaml_value_config(
agent_id: &str,
config: serde_yaml::Value,
) -> AgentInstantiationResult<ControllerType> {
let config = match &config {
serde_yaml::Value::Mapping(_) => {
let config: Config =
serde_yaml::from_value(config).map_err(AgentInstantiationError::Yaml)?;
config
.validate()
.map_err(AgentInstantiationError::ConfigFailsValidation)?;
config
}
_ => {
return Err(AgentInstantiationError::ConfigForAgentIsNotAMapping(
agent_id.to_owned(),
));
}
};
let controller =
Controller::new(config).map_err(AgentInstantiationError::ConstructionFailed)?;
Ok(ControllerType::Anthropic(Box::new(controller)))
}
pub fn default_config() -> Config {
Config::default()
}

View File

@@ -0,0 +1,32 @@
use anthropic_rs::completion::message::{Content, ContentType, Message, MessageRequest, Role};
use crate::conversation::llm::{Author as LLMAuthor, Message as LLMMessage};
pub(super) fn create_anthropic_message_request(llm_messages: Vec<LLMMessage>) -> MessageRequest {
let mut messages = vec![];
for message in llm_messages {
let role = match message.author {
LLMAuthor::User => Role::User,
LLMAuthor::Assistant => Role::Assistant,
LLMAuthor::Prompt => {
continue;
}
};
let content = vec![Content {
content_type: ContentType::Text,
text: message.message_text,
}];
let message = Message { role, content };
messages.push(message);
}
MessageRequest {
stream: false,
messages,
..Default::default()
}
}

View File

@@ -0,0 +1,3 @@
pub trait ConfigTrait {
fn validate(&self) -> Result<(), String>;
}

View File

@@ -0,0 +1,172 @@
use crate::{agent::AgentPurpose, conversation::llm::Conversation};
use super::{
entity::{
ImageGenerationResult, PingResult, TextGenerationParams, TextGenerationResult,
TextToSpeechParams, TextToSpeechResult,
},
ImageGenerationParams, SpeechToTextParams, SpeechToTextResult,
};
pub trait ControllerTrait {
fn supports_purpose(&self, purpose: AgentPurpose) -> bool;
fn ping(&self) -> impl std::future::Future<Output = anyhow::Result<PingResult>> + Send;
fn text_generation_prompt(&self) -> Option<String>;
fn text_generation_temperature(&self) -> Option<f32>;
fn text_to_speech_voice(&self) -> Option<String>;
fn text_to_speech_speed(&self) -> Option<f32>;
fn generate_text(
&self,
conversation: Conversation,
params: TextGenerationParams,
) -> impl std::future::Future<Output = anyhow::Result<TextGenerationResult>> + Send;
fn speech_to_text(
&self,
mime_type: &mxlink::mime::Mime,
media: Vec<u8>,
params: SpeechToTextParams,
) -> impl std::future::Future<Output = anyhow::Result<SpeechToTextResult>> + Send;
fn generate_image(
&self,
prompt: &str,
params: ImageGenerationParams,
) -> impl std::future::Future<Output = anyhow::Result<ImageGenerationResult>> + Send;
fn text_to_speech(
&self,
text: &str,
params: TextToSpeechParams,
) -> impl std::future::Future<Output = anyhow::Result<TextToSpeechResult>> + Send;
}
#[derive(Debug, Clone)]
pub enum ControllerType {
OpenAI(Box<super::openai::Controller>),
OpenAICompat(Box<super::openai_compat::Controller>),
Anthropic(Box<super::anthropic::Controller>),
}
impl ControllerTrait for ControllerType {
fn supports_purpose(&self, purpose: AgentPurpose) -> bool {
match &self {
ControllerType::OpenAI(controller) => controller.supports_purpose(purpose),
ControllerType::OpenAICompat(controller) => controller.supports_purpose(purpose),
ControllerType::Anthropic(controller) => controller.supports_purpose(purpose),
}
}
fn text_generation_prompt(&self) -> Option<String> {
match &self {
ControllerType::OpenAI(controller) => controller.text_generation_prompt(),
ControllerType::OpenAICompat(controller) => controller.text_generation_prompt(),
ControllerType::Anthropic(controller) => controller.text_generation_prompt(),
}
}
fn text_to_speech_voice(&self) -> Option<String> {
match &self {
ControllerType::OpenAI(controller) => controller.text_to_speech_voice(),
ControllerType::OpenAICompat(controller) => controller.text_to_speech_voice(),
ControllerType::Anthropic(controller) => controller.text_to_speech_voice(),
}
}
fn text_to_speech_speed(&self) -> Option<f32> {
match &self {
ControllerType::OpenAI(controller) => controller.text_to_speech_speed(),
ControllerType::OpenAICompat(controller) => controller.text_to_speech_speed(),
ControllerType::Anthropic(controller) => controller.text_to_speech_speed(),
}
}
fn text_generation_temperature(&self) -> Option<f32> {
match &self {
ControllerType::OpenAI(controller) => controller.text_generation_temperature(),
ControllerType::OpenAICompat(controller) => controller.text_generation_temperature(),
ControllerType::Anthropic(controller) => controller.text_generation_temperature(),
}
}
async fn ping(&self) -> anyhow::Result<PingResult> {
match &self {
ControllerType::OpenAI(controller) => controller.ping().await,
ControllerType::OpenAICompat(controller) => controller.ping().await,
ControllerType::Anthropic(controller) => controller.ping().await,
}
}
async fn generate_text(
&self,
conversation: Conversation,
params: TextGenerationParams,
) -> anyhow::Result<TextGenerationResult> {
match &self {
ControllerType::OpenAI(controller) => {
controller.generate_text(conversation, params).await
}
ControllerType::OpenAICompat(controller) => {
controller.generate_text(conversation, params).await
}
ControllerType::Anthropic(controller) => {
controller.generate_text(conversation, params).await
}
}
}
async fn speech_to_text(
&self,
mime_type: &mxlink::mime::Mime,
media: Vec<u8>,
params: SpeechToTextParams,
) -> anyhow::Result<SpeechToTextResult> {
match &self {
ControllerType::OpenAI(controller) => {
controller.speech_to_text(mime_type, media, params).await
}
ControllerType::OpenAICompat(controller) => {
controller.speech_to_text(mime_type, media, params).await
}
ControllerType::Anthropic(controller) => {
controller.speech_to_text(mime_type, media, params).await
}
}
}
async fn generate_image(
&self,
prompt: &str,
params: ImageGenerationParams,
) -> anyhow::Result<ImageGenerationResult> {
match &self {
ControllerType::OpenAI(controller) => controller.generate_image(prompt, params).await,
ControllerType::OpenAICompat(controller) => {
controller.generate_image(prompt, params).await
}
ControllerType::Anthropic(controller) => {
controller.generate_image(prompt, params).await
}
}
}
async fn text_to_speech(
&self,
text: &str,
params: TextToSpeechParams,
) -> anyhow::Result<TextToSpeechResult> {
match &self {
ControllerType::OpenAI(controller) => controller.text_to_speech(text, params).await,
ControllerType::OpenAICompat(controller) => {
controller.text_to_speech(text, params).await
}
ControllerType::Anthropic(controller) => controller.text_to_speech(text, params).await,
}
}
}

View File

@@ -0,0 +1,198 @@
use crate::agent::AgentPurpose;
#[derive(Debug, Clone)]
pub enum AgentProvider {
Anthropic,
Groq,
LocalAI,
Mistral,
Ollama,
OpenAI,
OpenAICompat,
OpenRouter,
TogetherAI,
}
impl AgentProvider {
pub fn choices() -> Vec<&'static Self> {
vec![
&Self::Anthropic,
&Self::Groq,
&Self::LocalAI,
&Self::Mistral,
&Self::Ollama,
&Self::OpenAI,
&Self::OpenAICompat,
&Self::OpenRouter,
&Self::TogetherAI,
]
}
pub fn to_static_str(&self) -> &'static str {
match &self {
Self::Anthropic => "anthropic",
Self::Groq => "groq",
Self::LocalAI => "localai",
Self::Mistral => "mistral",
Self::Ollama => "ollama",
Self::OpenAI => "openai",
Self::OpenAICompat => "openai-compatible",
Self::OpenRouter => "openrouter",
Self::TogetherAI => "together-ai",
}
}
pub fn from_string(s: &str) -> Result<Self, &'static str> {
match s {
"anthropic" => Ok(Self::Anthropic),
"groq" => Ok(Self::Groq),
"localai" => Ok(Self::LocalAI),
"mistral" => Ok(Self::Mistral),
"ollama" => Ok(Self::Ollama),
"openai" => Ok(Self::OpenAI),
"openai-compatible" => Ok(Self::OpenAICompat),
"openrouter" => Ok(Self::OpenRouter),
"together-ai" => Ok(Self::TogetherAI),
_ => Err("Unexpected string value"),
}
}
pub fn info(&self) -> AgentProviderInfo {
match &self {
Self::Anthropic => AgentProviderInfo {
id: Self::Anthropic.to_static_str(),
name: "Anthropic",
description: "Anthropic is an American AI company founded by former OpenAI engineers and providing powerful language models.",
homepage_url: Some("https://www.anthropic.com/"),
wiki_url: Some("https://en.wikipedia.org/wiki/Anthropic"),
sign_up_url: Some("https://console.anthropic.com/"),
models_list_url: Some("https://docs.anthropic.com/en/docs/about-claude/models"),
supported_purposes: vec![
AgentPurpose::TextGeneration,
],
},
Self::Groq => AgentProviderInfo {
id: Self::Groq.to_static_str(),
name: "Groq",
description: "Groq is an American company developing optimized Language Processing Units (LPU) and offering cloud service which runs various models (built by others) with very high performance.",
homepage_url: Some("https://groq.com/"),
wiki_url: Some("https://en.wikipedia.org/wiki/Groq"),
sign_up_url: Some("https://console.groq.com/login"),
models_list_url: Some("https://console.groq.com/docs/models"),
supported_purposes: vec![
AgentPurpose::TextGeneration,
AgentPurpose::SpeechToText,
],
},
Self::LocalAI => AgentProviderInfo {
id: Self::LocalAI.to_static_str(),
name: "LocalAI",
description: "LocalAI is the free, Open Source OpenAI alternative. LocalAI act as a drop-in replacement REST API that’s compatible with OpenAI API specifications for local inferencing. It allows you to run LLMs, generate images, audio (and not only) locally or on-prem with consumer grade hardware, supporting multiple model families and architectures.",
homepage_url: Some("https://localai.io/"),
wiki_url: None,
sign_up_url: None,
models_list_url: Some("https://localai.io/gallery.html"),
supported_purposes: vec![
AgentPurpose::TextGeneration,
AgentPurpose::TextToSpeech,
AgentPurpose::SpeechToText,
],
},
Self::Mistral => AgentProviderInfo {
id: Self::Mistral.to_static_str(),
name: "Mistral",
description: "Mistral AI is a research lab based in Europe (France) which produces their own language models.",
homepage_url: Some("https://mistral.ai/"),
wiki_url: Some("https://en.wikipedia.org/wiki/Mistral_AI"),
sign_up_url: Some("https://auth.mistral.ai/ui/registration"),
models_list_url: Some("https://docs.mistral.ai/getting-started/models/"),
supported_purposes: vec![
AgentPurpose::TextGeneration,
],
},
Self::Ollama => AgentProviderInfo {
id: Self::Ollama.to_static_str(),
name: "Ollama",
description: "Ollama lets you run various models in a [self-hosted](https://github.com/ollama/ollama?tab=readme-ov-file#ollama) way. This is more advanced and requires powerful hardware for running some of the better models, but ensures your data stays with you.",
homepage_url: Some("https://ollama.com/"),
wiki_url: None,
sign_up_url: None,
models_list_url: Some("https://ollama.com/library"),
supported_purposes: vec![
AgentPurpose::TextGeneration,
],
},
Self::OpenAI => AgentProviderInfo {
id: Self::OpenAI.to_static_str(),
name: "OpenAI",
description: "OpenAI is an American AI company providing powerful language models.\n\nUse this provider either with the OpenAI API or with other OpenAI-compatible API services which **fully** adhere to the [OpenAI API spec](https://github.com/openai/openai-openapi/).\nFor services which are not fully compatible with the OpenAI API, consider using the **OpenAI Compatible** provider.",
homepage_url: Some("https://openai.com/"),
wiki_url: Some("https://en.wikipedia.org/wiki/OpenAI"),
sign_up_url: Some("https://platform.openai.com/signup"),
models_list_url: Some("https://platform.openai.com/docs/models"),
supported_purposes: vec![
AgentPurpose::ImageGeneration,
AgentPurpose::TextGeneration,
AgentPurpose::TextToSpeech,
AgentPurpose::SpeechToText,
],
},
Self::OpenAICompat => AgentProviderInfo {
id: Self::OpenAICompat.to_static_str(),
name: "OpenAI Compatible",
description: "This provider allows you to use OpenAI-compatible API services like [OpenRouter](https://openrouter.ai/), [Together AI](https://www.together.ai/), etc.\n\nSome of these popular services already have **shortcut** providers (leading to this one behind the scenes) - this make it easier to get started.\n\nThis provider just as featureful as the **OpenAI** provider, but is more compatible with services which do not fully adhere to the [OpenAI API spec](https://github.com/openai/openai-openapi/).",
homepage_url: None,
wiki_url: None,
sign_up_url: None,
models_list_url: None,
supported_purposes: vec![
AgentPurpose::ImageGeneration,
AgentPurpose::TextGeneration,
AgentPurpose::TextToSpeech,
AgentPurpose::SpeechToText,
],
},
Self::OpenRouter => AgentProviderInfo {
id: Self::OpenRouter.to_static_str(),
name: "OpenRouter",
description: "OpenRouter is a unified interface for LLMs. The platform scouts for the lowest prices and best latencies/throughputs across dozens of providers, and lets you choose how to [prioritize](https://openrouter.ai/docs/provider-routing) them.",
homepage_url: Some("https://openrouter.ai/"),
wiki_url: None,
sign_up_url: Some("https://openrouter.ai/"),
models_list_url: Some("https://openrouter.ai/models"),
supported_purposes: vec![
AgentPurpose::TextGeneration,
],
},
Self::TogetherAI => AgentProviderInfo {
id: Self::TogetherAI.to_static_str(),
name: "Together AI",
description: "Together AI makes it easy to run or [fine-tune](https://docs.together.ai/docs/fine-tuning-overview) leading open source models with only a few lines of code.",
homepage_url: Some("https://www.together.ai/"),
wiki_url: None,
sign_up_url: Some("https://api.together.ai/signup"),
models_list_url: Some("https://api.together.xyz/models"),
supported_purposes: vec![
AgentPurpose::TextGeneration,
],
},
}
}
}
impl std::fmt::Display for AgentProvider {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(f, "{}", self.to_static_str())
}
}
pub struct AgentProviderInfo {
pub id: &'static str,
pub name: &'static str,
pub description: &'static str,
pub homepage_url: Option<&'static str>,
pub wiki_url: Option<&'static str>,
pub sign_up_url: Option<&'static str>,
pub models_list_url: Option<&'static str>,
pub supported_purposes: Vec<AgentPurpose>,
}

View File

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

View File

@@ -0,0 +1,13 @@
mod agent_provider;
mod image_generation;
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 ping::PingResult;
pub use speech_to_text::{SpeechToTextParams, SpeechToTextResult};
pub use text_generation::{TextGenerationParams, TextGenerationResult};
pub use text_to_speech::{TextToSpeechParams, TextToSpeechResult};

View File

@@ -0,0 +1,4 @@
pub enum PingResult {
Inconclusive,
Successful,
}

View File

@@ -0,0 +1,8 @@
#[derive(Default)]
pub struct SpeechToTextParams {
pub language_override: Option<String>,
}
pub struct SpeechToTextResult {
pub text: String,
}

View File

@@ -0,0 +1,10 @@
#[derive(Default)]
pub struct TextGenerationParams {
pub context_management_enabled: bool,
pub prompt_override: Option<String>,
pub temperature_override: Option<f32>,
}
pub struct TextGenerationResult {
pub text: String,
}

View File

@@ -0,0 +1,10 @@
#[derive(Default)]
pub struct TextToSpeechParams {
pub speed_override: Option<f32>,
pub voice_override: Option<String>,
}
pub struct TextToSpeechResult {
pub bytes: Vec<u8>,
pub mime_type: mxlink::mime::Mime,
}

View File

@@ -0,0 +1,26 @@
// Groq is based on openai_compat, because it's not fully compatible with async-openai.
use super::openai_compat::Config;
pub fn default_config() -> Config {
let mut config = Config {
base_url: "https://api.groq.com/openai/v1".to_owned(),
text_to_speech: None,
image_generation: None,
..Default::default()
};
if let Some(ref mut config) = config.text_generation.as_mut() {
config.model_id = "llama3-70b-8192".to_owned();
config.max_context_tokens = 131_072;
config.max_response_tokens = 4096;
}
if let Some(ref mut config) = config.speech_to_text.as_mut() {
config.model_id = "whisper-large-v3".to_owned();
}
config
}

View File

@@ -0,0 +1,33 @@
// LocalAI is based on OpenAI (async-openai), because it seems to be fully compatible.
// Moreover, openai_api_rust does not support speech-to-text, so if we wish to use this feature
// we need to stick to async-openai.
use super::openai_compat::Config;
pub fn default_config() -> Config {
let mut config = Config {
base_url: "http://my-localai-self-hosted-service:8080/v1".to_owned(),
..Default::default()
};
if let Some(ref mut config) = config.text_generation.as_mut() {
config.model_id = "gpt-4".to_owned();
config.max_context_tokens = 128_000;
config.max_response_tokens = 4096;
}
if let Some(ref mut config) = config.text_to_speech.as_mut() {
config.model_id = "tts-1".to_owned();
}
if let Some(ref mut config) = config.speech_to_text.as_mut() {
config.model_id = "whisper-1".to_owned();
}
if let Some(ref mut config) = config.image_generation.as_mut() {
config.model_id = "stablediffusion".to_owned();
}
config
}

View File

@@ -0,0 +1,22 @@
// Mistral is based on openai_compat, because it's not fully compatible with async-openai.
use super::openai_compat::Config;
pub fn default_config() -> Config {
let mut config = Config {
base_url: "https://api.mistral.ai/v1".to_owned(),
speech_to_text: None,
text_to_speech: None,
image_generation: None,
..Default::default()
};
if let Some(ref mut config) = config.text_generation.as_mut() {
config.model_id = "mistral-large-latest".to_owned();
config.max_context_tokens = 128_000;
}
config
}

25
src/agent/provider/mod.rs Normal file
View File

@@ -0,0 +1,25 @@
pub mod anthropic;
mod config;
mod controller;
mod entity;
pub(super) mod groq;
pub mod localai;
pub(super) mod mistral;
pub mod ollama;
pub mod openai;
pub mod openai_compat;
pub(super) mod openrouter;
pub(super) mod togetherai;
fn default_temperature() -> f32 {
1.0
}
pub use controller::{ControllerTrait, ControllerType};
pub use config::ConfigTrait;
pub use entity::{
AgentProvider, AgentProviderInfo, ImageGenerationParams, PingResult, SpeechToTextParams,
SpeechToTextResult, TextGenerationParams, TextToSpeechParams,
};

View File

@@ -0,0 +1,24 @@
// At the time of testing, Ollama can be powered by `openai`, but we use `openai_compat` for better reliability
// in the event of future updates to `async-openai`.
use super::openai_compat::Config;
pub fn default_config() -> Config {
let mut config = Config {
base_url: "http://my-ollama-self-hosted-service:11434/v1".to_owned(),
text_to_speech: None,
image_generation: None,
speech_to_text: None,
..Default::default()
};
if let Some(ref mut config) = config.text_generation.as_mut() {
config.model_id = "gemma2:2b".to_owned();
config.max_context_tokens = 128_000;
config.max_response_tokens = 4096;
}
config
}

View File

@@ -0,0 +1,190 @@
use serde::{Deserialize, Serialize};
use crate::agent::provider::ConfigTrait;
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct Config {
pub base_url: String,
pub api_key: String,
#[serde(skip_serializing_if = "Option::is_none")]
pub text_generation: Option<TextGenerationConfig>,
#[serde(skip_serializing_if = "Option::is_none")]
pub speech_to_text: Option<SpeechToTextConfig>,
#[serde(skip_serializing_if = "Option::is_none")]
pub text_to_speech: Option<TextToSpeechConfig>,
#[serde(skip_serializing_if = "Option::is_none")]
pub image_generation: Option<ImageGenerationConfig>,
}
impl Default for Config {
fn default() -> Self {
Self {
base_url: "https://api.openai.com/v1".to_owned(),
api_key: "YOUR_API_KEY_HERE".to_owned(),
text_generation: Some(TextGenerationConfig::default()),
speech_to_text: Some(SpeechToTextConfig::default()),
text_to_speech: Some(TextToSpeechConfig::default()),
image_generation: Some(ImageGenerationConfig::default()),
}
}
}
impl ConfigTrait for Config {
fn validate(&self) -> Result<(), String> {
if self.base_url.is_empty() {
return Err("The base URL must not be empty.".to_owned());
}
Ok(())
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct TextGenerationConfig {
#[serde(default = "default_text_model_id")]
pub model_id: String,
#[serde(default)]
pub prompt: Option<String>,
#[serde(default = "super::super::default_temperature")]
pub temperature: f32,
#[serde(default)]
pub max_response_tokens: u32,
#[serde(default)]
pub max_context_tokens: u32,
}
impl Default for TextGenerationConfig {
fn default() -> Self {
Self {
model_id: default_text_model_id(),
prompt: Some("You are a brief, but helpful bot.".to_owned()),
temperature: super::super::default_temperature(),
max_response_tokens: 16_384,
max_context_tokens: 128_000,
}
}
}
fn default_text_model_id() -> String {
"gpt-4o-2024-08-06".to_owned()
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct SpeechToTextConfig {
#[serde(default = "default_speech_to_text_model_id")]
pub model_id: String,
}
impl Default for SpeechToTextConfig {
fn default() -> Self {
Self {
model_id: default_speech_to_text_model_id(),
}
}
}
fn default_speech_to_text_model_id() -> String {
"whisper-1".to_owned()
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct TextToSpeechConfig {
#[serde(default = "default_text_to_speech_model_id")]
pub model_id: async_openai::types::SpeechModel,
#[serde(default = "default_text_to_speech_voice")]
pub voice: async_openai::types::Voice,
#[serde(default = "default_text_to_speech_speed")]
pub speed: f32,
#[serde(default = "default_text_to_speech_response_format")]
pub response_format: async_openai::types::SpeechResponseFormat,
}
impl Default for TextToSpeechConfig {
fn default() -> Self {
Self {
model_id: default_text_to_speech_model_id(),
voice: default_text_to_speech_voice(),
speed: default_text_to_speech_speed(),
response_format: default_text_to_speech_response_format(),
}
}
}
fn default_text_to_speech_model_id() -> async_openai::types::SpeechModel {
async_openai::types::SpeechModel::Tts1Hd
}
fn default_text_to_speech_voice() -> async_openai::types::Voice {
async_openai::types::Voice::Onyx
}
fn default_text_to_speech_speed() -> f32 {
1.0
}
fn default_text_to_speech_response_format() -> async_openai::types::SpeechResponseFormat {
// The API defaults to mp3, but we prefer Opus because it's smaller.
// Our clients should all have support for it.
async_openai::types::SpeechResponseFormat::Opus
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ImageGenerationConfig {
pub model_id: String,
#[serde(default = "default_image_style")]
pub style: async_openai::types::ImageStyle,
#[serde(default = "default_image_size")]
pub size: async_openai::types::ImageSize,
#[serde(default = "default_image_quality")]
pub quality: async_openai::types::ImageQuality,
}
impl Default for ImageGenerationConfig {
fn default() -> Self {
Self {
model_id: "dall-e-3".to_owned(),
style: default_image_style(),
size: default_image_size(),
quality: default_image_quality(),
}
}
}
impl ImageGenerationConfig {
pub fn model_id_as_openai_image_model(
&self,
) -> Result<async_openai::types::ImageModel, String> {
match self.model_id.as_str() {
"dall-e-2" => Ok(async_openai::types::ImageModel::DallE2),
"dall-e-3" => Ok(async_openai::types::ImageModel::DallE3),
other => Ok(async_openai::types::ImageModel::Other(other.to_owned())),
}
}
}
fn default_image_style() -> async_openai::types::ImageStyle {
async_openai::types::ImageStyle::Vivid
}
fn default_image_size() -> async_openai::types::ImageSize {
async_openai::types::ImageSize::S1024x1024
}
fn default_image_quality() -> async_openai::types::ImageQuality {
async_openai::types::ImageQuality::Standard
}

View File

@@ -0,0 +1,465 @@
use std::ops::Deref;
use async_openai::{
config::OpenAIConfig,
types::{
ChatCompletionRequestMessage, CreateChatCompletionRequestArgs, CreateImageRequestArgs,
CreateSpeechRequestArgs, CreateTranscriptionRequestArgs,
},
Client as OpenAIClient,
};
use super::super::ControllerTrait;
use crate::{
agent::{
provider::{
entity::{ImageGenerationResult, PingResult, TextToSpeechParams, TextToSpeechResult},
openai::utils::convert_string_to_enum,
},
AgentPurpose,
},
strings,
};
use crate::{
agent::{
provider::{
entity::{TextGenerationParams, TextGenerationResult},
ImageGenerationParams, SpeechToTextParams, SpeechToTextResult,
},
utils::base64_decode,
},
conversation::llm::{
shorten_messages_list_to_context_size, Author as LLMAuthor,
Conversation as LLMConversation, Message as LLMMessage,
},
};
use super::config::Config;
#[derive(Debug, Clone)]
pub struct Controller {
config: Config,
client: OpenAIClient<OpenAIConfig>,
}
impl Controller {
pub fn new(config: Config) -> Self {
let openai_config = OpenAIConfig::new()
.with_api_base(config.base_url.clone())
.with_api_key(config.api_key.clone());
let client = OpenAIClient::with_config(openai_config);
Self { config, client }
}
}
impl ControllerTrait for Controller {
async fn ping(&self) -> anyhow::Result<PingResult> {
if !self.supports_purpose(AgentPurpose::TextGeneration) {
return Ok(PingResult::Inconclusive);
}
let messages = vec![LLMMessage {
author: LLMAuthor::User,
message_text: "Hello!".to_string(),
}];
let conversation = LLMConversation { messages };
self.generate_text(conversation, TextGenerationParams::default())
.await?;
Ok(PingResult::Successful)
}
async fn generate_text(
&self,
conversation: LLMConversation,
params: TextGenerationParams,
) -> anyhow::Result<TextGenerationResult> {
let Some(text_generation_config) = &self.config.text_generation else {
return Err(anyhow::anyhow!(
strings::agent::no_configuration_for_purpose_so_cannot_be_used(
&AgentPurpose::TextGeneration
),
));
};
let prompt_text = params
.prompt_override
.unwrap_or(self.text_generation_prompt().unwrap_or("".to_owned()))
.trim()
.to_owned();
let prompt_message = if prompt_text.is_empty() {
None
} else {
Some(LLMMessage {
author: LLMAuthor::Prompt,
message_text: prompt_text,
})
};
let mut conversation_messages = conversation.messages;
if params.context_management_enabled {
tracing::trace!("Shortening messages list to context size");
conversation_messages = shorten_messages_list_to_context_size(
&text_generation_config.model_id,
&prompt_message,
conversation_messages,
text_generation_config.max_response_tokens,
text_generation_config.max_context_tokens,
);
tracing::trace!("Finished shortening messages list to context size");
};
if let Some(prompt_message) = prompt_message {
conversation_messages.insert(0, prompt_message);
}
let openai_conversation_messages: Vec<ChatCompletionRequestMessage> =
super::utils::convert_llm_messages_to_openai_messages(conversation_messages);
let messages_count = openai_conversation_messages.len();
let temperature = params
.temperature_override
.unwrap_or(text_generation_config.temperature);
let request = CreateChatCompletionRequestArgs::default()
.max_tokens(text_generation_config.max_response_tokens)
.model(&text_generation_config.model_id)
.temperature(temperature)
.messages(openai_conversation_messages)
.build()?;
if let Ok(request_as_json) = serde_json::to_string(&request) {
tracing::trace!(
model = format!("{:?}", request.model),
?messages_count,
request = request_as_json,
"Sending OpenAI chat completion API request"
);
}
let response = self.client.chat().create(request).await?;
tracing::trace!(
?response,
"Got response from the OpenAI chat completion API"
);
// We only request 1 result, so there should only be 1 choice.
if let Some(choice) = response.choices.into_iter().next() {
match choice.message.content {
Some(text) => {
return Ok(TextGenerationResult { text });
}
None => {
return Err(anyhow::anyhow!(
"No content was found in the response choice from the OpenAI chat completion API"
));
}
}
}
Err(anyhow::anyhow!(
"No response messages choices were returned from the OpenAI chat completion API"
))
}
async fn speech_to_text(
&self,
mime_type: &mxlink::mime::Mime,
media: Vec<u8>,
params: SpeechToTextParams,
) -> anyhow::Result<SpeechToTextResult> {
let Some(speech_to_text_config) = &self.config.speech_to_text else {
return Err(anyhow::anyhow!(
strings::agent::no_configuration_for_purpose_so_cannot_be_used(
&AgentPurpose::SpeechToText
),
));
};
let filename = audio_mime_type_to_file_name(mime_type).unwrap_or("audio.ogg".to_string());
let language = params.language_override.unwrap_or("".to_string());
let request = CreateTranscriptionRequestArgs::default()
.model(&speech_to_text_config.model_id)
.file(async_openai::types::AudioInput {
source: async_openai::types::InputSource::VecU8 {
filename,
vec: media,
},
})
.language(language.clone())
.build()?;
tracing::trace!(
model_id = speech_to_text_config.model_id,
?language,
"Sending OpenAI speech-to-text API request"
);
let response = self.client.audio().transcribe(request).await?;
tracing::trace!(
?response,
"Got response from the OpenAI audio transcription API"
);
Ok(SpeechToTextResult {
text: response.text,
})
}
async fn generate_image(
&self,
prompt: &str,
params: ImageGenerationParams,
) -> anyhow::Result<ImageGenerationResult> {
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 original_model = image_generation_config
.model_id_as_openai_image_model()
.map_err(|err| anyhow::anyhow!(err))?;
let model = if params.cheaper_model_switching_allowed {
// Switch to a cheaper model
match original_model {
async_openai::types::ImageModel::DallE2 => async_openai::types::ImageModel::DallE2,
async_openai::types::ImageModel::DallE3 => async_openai::types::ImageModel::DallE2,
async_openai::types::ImageModel::Other(_) => {
async_openai::types::ImageModel::DallE2
}
}
} else {
original_model
};
let quality = if params.cheaper_quality_switching_allowed {
// Switch to a cheaper quality
match &image_generation_config.quality {
async_openai::types::ImageQuality::Standard => {
async_openai::types::ImageQuality::Standard
}
async_openai::types::ImageQuality::HD => {
async_openai::types::ImageQuality::Standard
}
}
} else {
image_generation_config.quality.clone()
};
let size = params
.size_override
.map(|s| {
convert_string_to_enum::<async_openai::types::ImageSize>(&s)
.unwrap_or(image_generation_config.size)
})
.unwrap_or(image_generation_config.size);
let request = CreateImageRequestArgs::default()
.model(model)
.prompt(prompt.to_owned())
.response_format(async_openai::types::ImageResponseFormat::B64Json)
.size(size)
.style(image_generation_config.style.clone())
.quality(quality)
.build()?;
tracing::trace!(
?prompt,
model = format!("{:?}", request.model),
size = format!("{:?}", request.size),
style = format!("{:?}", request.style),
quality = format!("{:?}", request.quality),
"Sending OpenAI image generation API request"
);
let response = self.client.images().create(request).await?;
if let Some(image) = response.data.into_iter().next() {
match image.deref() {
async_openai::types::Image::B64Json {
b64_json,
revised_prompt,
} => {
let bytes = base64_decode(b64_json)?;
return Ok(ImageGenerationResult {
bytes,
mime_type: mxlink::mime::IMAGE_PNG,
revised_prompt: revised_prompt.clone(),
});
}
_ => {
return Err(anyhow::anyhow!("Unexpected image type"));
}
}
}
Err(anyhow::anyhow!(
"The OpenAI image generation API returned no images"
))
}
async fn text_to_speech(
&self,
input: &str,
params: TextToSpeechParams,
) -> anyhow::Result<TextToSpeechResult> {
let Some(text_to_speech_config) = &self.config.text_to_speech else {
return Err(anyhow::anyhow!(
strings::agent::no_configuration_for_purpose_so_cannot_be_used(
&AgentPurpose::TextToSpeech
),
));
};
let speed = params.speed_override.unwrap_or(text_to_speech_config.speed);
let voice = if let Some(voice_string) = params.voice_override {
// This is a hacky way to construct a Voice enum from the string we have.
let voice: serde_json::Result<async_openai::types::Voice> =
serde_json::from_str(&format!("\"{}\"", voice_string));
match voice {
Ok(voice) => voice,
Err(err) => {
tracing::debug!(?voice_string, ?err, "Failed to parse voice");
return Err(anyhow::anyhow!(
"The configured voice ({}) is not supported.",
voice_string
));
}
}
} else {
text_to_speech_config.voice.clone()
};
let response_format = text_to_speech_config.response_format;
let mime_type = response_format_to_mime_type(&response_format).unwrap_or(
"audio/mp3"
.parse()
.expect("Failed parsing default mime type"),
);
let request = CreateSpeechRequestArgs::default()
.model(text_to_speech_config.model_id.clone())
.voice(voice)
.speed(speed)
.response_format(response_format)
.input(input)
.build()?;
tracing::trace!(
model = format!("{:?}", request.model),
voice = format!("{:?}", request.voice),
speed = format!("{:?}", request.speed),
"Sending OpenAI text-to-speech API request"
);
let result = self.client.audio().speech(request).await?;
Ok(TextToSpeechResult {
bytes: result.bytes.into(),
mime_type,
})
}
fn supports_purpose(&self, purpose: AgentPurpose) -> bool {
match purpose {
AgentPurpose::TextGeneration => self.config.text_generation.is_some(),
AgentPurpose::SpeechToText => self.config.speech_to_text.is_some(),
AgentPurpose::TextToSpeech => self.config.text_to_speech.is_some(),
AgentPurpose::ImageGeneration => self.config.image_generation.is_some(),
AgentPurpose::CatchAll => true,
}
}
fn text_generation_prompt(&self) -> Option<String> {
let Some(text_generation_config) = &self.config.text_generation else {
return None;
};
text_generation_config.prompt.clone()
}
fn text_generation_temperature(&self) -> Option<f32> {
let Some(text_generation_config) = &self.config.text_generation else {
return None;
};
Some(text_generation_config.temperature)
}
fn text_to_speech_voice(&self) -> Option<String> {
let Some(text_to_speech_config) = &self.config.text_to_speech else {
return None;
};
// A hacky way to turn this enum to a string
let voice_as_string = serde_json::to_string(&text_to_speech_config.voice).ok()?;
Some(voice_as_string.replace("\"", ""))
}
fn text_to_speech_speed(&self) -> Option<f32> {
let Some(text_to_speech_config) = &self.config.text_to_speech else {
return None;
};
Some(text_to_speech_config.speed)
}
}
fn response_format_to_mime_type(
response_format: &async_openai::types::SpeechResponseFormat,
) -> Option<mxlink::mime::Mime> {
let content_type = match response_format {
async_openai::types::SpeechResponseFormat::Mp3 => "audio/mp3".to_owned(),
async_openai::types::SpeechResponseFormat::Wav => "audio/wav".to_owned(),
async_openai::types::SpeechResponseFormat::Opus => "audio/ogg".to_owned(),
async_openai::types::SpeechResponseFormat::Aac => "audio/aac".to_owned(),
async_openai::types::SpeechResponseFormat::Flac => "audio/flac".to_owned(),
async_openai::types::SpeechResponseFormat::Pcm => "audio/L8".to_owned(),
};
match content_type.parse() {
Ok(content_type) => Some(content_type),
Err(err) => {
tracing::error!(?err, "Failed to parse content type");
None
}
}
}
fn audio_mime_type_to_file_name(mime_type: &mxlink::mime::Mime) -> Option<String> {
let mime_type_string = mime_type.to_string();
let file_extension = match mime_type_string.as_str() {
"audio/flac" => "flac",
"audio/x-m4a" | "audio/m4a" => "m4a",
"audio/mp3" | "audio/mpeg" => "mp3",
"audio/mp4" => "mp4",
"application/ogg" | "audio/ogg" => "ogg",
"audio/wav" | "audio/x-wav" => "wav",
"audio/webm" => "webm",
_ => return None,
};
Some(format!("audio.{}", file_extension))
}

View File

@@ -0,0 +1,46 @@
mod config;
mod controller;
mod utils;
pub use config::Config;
pub use controller::Controller;
// openai_compat needs these, so it can convert from its own config types to these
pub(super) use config::ImageGenerationConfig;
pub(super) use config::SpeechToTextConfig;
pub(super) use config::TextGenerationConfig;
pub(super) use config::TextToSpeechConfig;
use super::super::AgentInstantiationError;
use super::super::AgentInstantiationResult;
use super::controller::ControllerType;
use super::ConfigTrait;
pub fn create_controller_from_yaml_value_config(
agent_id: &str,
config: serde_yaml::Value,
) -> AgentInstantiationResult<ControllerType> {
let config = match &config {
serde_yaml::Value::Mapping(_) => {
let config: Config =
serde_yaml::from_value(config).map_err(AgentInstantiationError::Yaml)?;
config
.validate()
.map_err(AgentInstantiationError::ConfigFailsValidation)?;
config
}
_ => {
return Err(AgentInstantiationError::ConfigForAgentIsNotAMapping(
agent_id.to_owned(),
));
}
};
Ok(ControllerType::OpenAI(Box::new(Controller::new(config))))
}
pub fn default_config() -> Config {
Config::default()
}

View File

@@ -0,0 +1,55 @@
use async_openai::types::{
ChatCompletionRequestAssistantMessageArgs, ChatCompletionRequestMessage,
ChatCompletionRequestSystemMessageArgs, ChatCompletionRequestUserMessageArgs,
};
use crate::conversation::llm::{Author as LLMAuthor, Message as LLMMessage};
pub fn convert_llm_messages_to_openai_messages(
conversation_messages: Vec<LLMMessage>,
) -> Vec<ChatCompletionRequestMessage> {
let mut openai_conversation_messages: Vec<ChatCompletionRequestMessage> =
Vec::with_capacity(conversation_messages.len());
for message in conversation_messages {
openai_conversation_messages.push(convert_llm_message_to_openai_message(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(),
}
}
pub(super) fn convert_string_to_enum<T>(value: &str) -> Result<T, String>
where
T: serde::de::DeserializeOwned,
{
// This is a hacky way to construct an enum from the string we have.
let enum_result: serde_json::Result<T> = serde_json::from_str(&format!("\"{}\"", value));
match enum_result {
Ok(enum_result) => Ok(enum_result),
Err(err) => {
tracing::debug!(?err, "Failed to parse into enum");
Err(format!("The value ({}) is not supported.", value))
}
}
}

View File

@@ -0,0 +1,261 @@
use serde::{Deserialize, Serialize};
use crate::agent::provider::openai::{
ImageGenerationConfig as OpenAIImageGenerationConfig,
SpeechToTextConfig as OpenAISpeechToTextConfig,
TextGenerationConfig as OpenAITextGenerationConfig,
TextToSpeechConfig as OpenAITextToSpeechConfig,
};
use crate::agent::provider::ConfigTrait;
use super::utils::convert_string_to_enum;
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct Config {
pub base_url: String,
pub api_key: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub text_generation: Option<TextGenerationConfig>,
#[serde(skip_serializing_if = "Option::is_none")]
pub speech_to_text: Option<SpeechToTextConfig>,
#[serde(skip_serializing_if = "Option::is_none")]
pub text_to_speech: Option<TextToSpeechConfig>,
#[serde(skip_serializing_if = "Option::is_none")]
pub image_generation: Option<ImageGenerationConfig>,
}
impl Default for Config {
fn default() -> Self {
Self {
base_url: "".to_owned(),
api_key: Some("YOUR_API_KEY_HERE".to_owned()),
text_generation: Some(TextGenerationConfig::default()),
speech_to_text: Some(SpeechToTextConfig::default()),
text_to_speech: Some(TextToSpeechConfig::default()),
image_generation: Some(ImageGenerationConfig::default()),
}
}
}
impl ConfigTrait for Config {
fn validate(&self) -> Result<(), String> {
if self.base_url.is_empty() {
return Err("The base URL must not be empty.".to_owned());
}
Ok(())
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct TextGenerationConfig {
#[serde(default = "default_text_model_id")]
pub model_id: String,
#[serde(default)]
pub prompt: Option<String>,
#[serde(default = "super::super::default_temperature")]
pub temperature: f32,
#[serde(default)]
pub max_response_tokens: u32,
#[serde(default)]
pub max_context_tokens: u32,
}
impl Default for TextGenerationConfig {
fn default() -> Self {
Self {
model_id: default_text_model_id(),
prompt: Some("You are a brief, but helpful bot.".to_owned()),
temperature: super::super::default_temperature(),
max_response_tokens: 4096,
max_context_tokens: 128_000,
}
}
}
impl TryInto<OpenAITextGenerationConfig> for TextGenerationConfig {
type Error = anyhow::Error;
fn try_into(self) -> Result<OpenAITextGenerationConfig, Self::Error> {
Ok(OpenAITextGenerationConfig {
model_id: self.model_id,
prompt: self.prompt,
temperature: self.temperature,
max_response_tokens: self.max_response_tokens,
max_context_tokens: self.max_context_tokens,
})
}
}
fn default_text_model_id() -> String {
"some-model".to_owned()
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct SpeechToTextConfig {
#[serde(default = "default_speech_to_text_model_id")]
pub model_id: String,
}
impl Default for SpeechToTextConfig {
fn default() -> Self {
Self {
model_id: default_speech_to_text_model_id(),
}
}
}
impl TryInto<OpenAISpeechToTextConfig> for SpeechToTextConfig {
type Error = anyhow::Error;
fn try_into(self) -> Result<OpenAISpeechToTextConfig, Self::Error> {
Ok(OpenAISpeechToTextConfig {
model_id: self.model_id,
})
}
}
fn default_speech_to_text_model_id() -> String {
"whisper-1".to_owned()
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct TextToSpeechConfig {
#[serde(default = "default_text_to_speech_model_id")]
pub model_id: String,
#[serde(default = "default_text_to_speech_voice")]
pub voice: String,
#[serde(default = "default_text_to_speech_speed")]
pub speed: f32,
#[serde(default = "default_text_to_speech_response_format")]
pub response_format: String,
}
impl Default for TextToSpeechConfig {
fn default() -> Self {
Self {
model_id: default_text_to_speech_model_id(),
voice: default_text_to_speech_voice(),
speed: default_text_to_speech_speed(),
response_format: default_text_to_speech_response_format(),
}
}
}
impl TryInto<OpenAITextToSpeechConfig> for TextToSpeechConfig {
type Error = String;
fn try_into(self) -> Result<OpenAITextToSpeechConfig, Self::Error> {
let model_id = convert_string_to_enum::<async_openai::types::SpeechModel>(&self.model_id)?;
let voice = convert_string_to_enum::<async_openai::types::Voice>(&self.voice)?;
let response_format = convert_string_to_enum::<async_openai::types::SpeechResponseFormat>(
&self.response_format,
)?;
Ok(OpenAITextToSpeechConfig {
model_id,
voice,
speed: self.speed,
response_format,
})
}
}
fn default_text_to_speech_model_id() -> String {
"tts-1".to_owned()
}
fn default_text_to_speech_voice() -> String {
"onyx".to_owned()
}
fn default_text_to_speech_speed() -> f32 {
1.0
}
fn default_text_to_speech_response_format() -> String {
"opus".to_owned()
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ImageGenerationConfig {
pub model_id: String,
#[serde(default = "default_image_style")]
pub style: Option<String>,
#[serde(default = "default_image_size")]
pub size: Option<String>,
#[serde(default = "default_image_quality")]
pub quality: Option<String>,
}
impl Default for ImageGenerationConfig {
fn default() -> Self {
Self {
model_id: "stablediffusion".to_owned(),
style: default_image_style(),
size: default_image_size(),
quality: default_image_quality(),
}
}
}
impl TryInto<OpenAIImageGenerationConfig> for ImageGenerationConfig {
type Error = String;
fn try_into(self) -> Result<OpenAIImageGenerationConfig, Self::Error> {
let size = if let Some(size) = &self.size {
convert_string_to_enum::<async_openai::types::ImageSize>(size)?
} else {
async_openai::types::ImageSize::S1024x1024
};
let style = if let Some(style) = &self.style {
convert_string_to_enum::<async_openai::types::ImageStyle>(style)?
} else {
async_openai::types::ImageStyle::Vivid
};
let quality = if let Some(quality) = &self.quality {
convert_string_to_enum::<async_openai::types::ImageQuality>(quality)?
} else {
async_openai::types::ImageQuality::Standard
};
Ok(OpenAIImageGenerationConfig {
model_id: self.model_id,
style,
size,
quality,
})
}
}
fn default_image_style() -> Option<String> {
Some("vivid".to_owned())
}
fn default_image_size() -> Option<String> {
Some("1024x1024".to_owned())
}
fn default_image_quality() -> Option<String> {
Some("standard".to_owned())
}

View File

@@ -0,0 +1,445 @@
use openai_api_rust::audio::{AudioApi, AudioBody};
use openai_api_rust::chat::{ChatApi, ChatBody};
use openai_api_rust::images::{ImagesApi, ImagesBody};
use openai_api_rust::{Auth, Message, OpenAI};
use super::super::ControllerTrait;
use crate::agent::utils::base64_decode;
use crate::{
agent::provider::{
entity::{TextGenerationParams, TextGenerationResult},
ImageGenerationParams, SpeechToTextParams, SpeechToTextResult,
},
conversation::llm::{
shorten_messages_list_to_context_size, Author as LLMAuthor,
Conversation as LLMConversation, Message as LLMMessage,
},
};
use crate::{
agent::{
provider::entity::{
ImageGenerationResult, PingResult, TextToSpeechParams, TextToSpeechResult,
},
AgentPurpose,
},
strings,
};
use super::Config;
#[derive(Debug, Clone)]
pub struct Controller {
config: Config,
client: OpenAI,
}
impl Controller {
pub fn new(config: Config) -> Self {
let api_key = config.api_key.clone().unwrap_or("".to_owned());
let auth = Auth::new(&api_key);
// The library we use chokes if there's no trailing slash
let base_url = if config.base_url.ends_with("/") {
config.base_url.clone()
} else {
format!("{}/", config.base_url)
};
let client = OpenAI::new(auth, &base_url);
Self { config, client }
}
}
impl ControllerTrait for Controller {
async fn ping(&self) -> anyhow::Result<PingResult> {
if !self.supports_purpose(AgentPurpose::TextGeneration) {
return Ok(PingResult::Inconclusive);
}
let messages = vec![LLMMessage {
author: LLMAuthor::User,
message_text: "Hello!".to_string(),
}];
let conversation = LLMConversation { messages };
self.generate_text(conversation, TextGenerationParams::default())
.await?;
Ok(PingResult::Successful)
}
async fn generate_text(
&self,
conversation: LLMConversation,
params: TextGenerationParams,
) -> anyhow::Result<TextGenerationResult> {
let Some(text_generation_config) = &self.config.text_generation else {
return Err(anyhow::anyhow!(
strings::agent::no_configuration_for_purpose_so_cannot_be_used(
&AgentPurpose::TextGeneration
),
));
};
let prompt_text = params
.prompt_override
.unwrap_or(self.text_generation_prompt().unwrap_or("".to_owned()))
.trim()
.to_owned();
let prompt_message = if prompt_text.is_empty() {
None
} else {
Some(LLMMessage {
author: LLMAuthor::Prompt,
message_text: prompt_text,
})
};
let mut conversation_messages = conversation.messages;
if params.context_management_enabled {
tracing::trace!("Shortening messages list to context size");
conversation_messages = shorten_messages_list_to_context_size(
&text_generation_config.model_id,
&prompt_message,
conversation_messages,
text_generation_config.max_response_tokens,
text_generation_config.max_context_tokens,
);
tracing::trace!("Finished shortening messages list to context size");
};
if let Some(prompt_message) = prompt_message {
conversation_messages.insert(0, prompt_message);
}
let openai_conversation_messages: Vec<Message> =
super::utils::convert_llm_messages_to_openai_messages(conversation_messages);
let messages_count = openai_conversation_messages.len();
let temperature = params
.temperature_override
.unwrap_or(text_generation_config.temperature);
let max_tokens = text_generation_config
.max_response_tokens
.try_into()
.expect("Failed converting max_response_tokens from u32 to i32");
let request = ChatBody {
model: text_generation_config.model_id.clone(),
max_tokens: Some(max_tokens),
temperature: Some(temperature),
top_p: None,
n: Some(1),
stream: Some(false),
stop: None,
presence_penalty: None,
frequency_penalty: None,
logit_bias: None,
user: None,
messages: openai_conversation_messages,
};
if let Ok(request_as_json) = serde_json::to_string(&request) {
tracing::trace!(
model = format!("{:?}", request.model),
?messages_count,
request = request_as_json,
"Sending OpenAI-compat chat completion API request"
);
}
// This library is not async-aware, so we need to use `spawn_blocking` to run the request on a separate thread.
let client = self.client.clone();
let response =
tokio::task::spawn_blocking(move || client.chat_completion_create(&request)).await?;
let response = match response {
Ok(response) => response,
Err(err) => {
return Err(anyhow::anyhow!(
"Failed to get response from the OpenAI-compat chat completion API: {:?}",
err
));
}
};
tracing::trace!(
?response,
"Got response from the OpenAI-compat chat completion API"
);
// We only request 1 result, so there should only be 1 choice.
if let Some(choice) = response.choices.into_iter().next() {
let Some(message) = choice.message else {
return Err(anyhow::anyhow!(
"No response message in choice was returned from the OpenAI-compat chat completion API"
));
};
return Ok(TextGenerationResult {
text: message.content,
});
}
Err(anyhow::anyhow!(
"No response messages choices were returned from the OpenAI-compat chat completion API"
))
}
async fn speech_to_text(
&self,
_mime_type: &mxlink::mime::Mime,
media: Vec<u8>,
params: SpeechToTextParams,
) -> anyhow::Result<SpeechToTextResult> {
let Some(speech_to_text_config) = &self.config.speech_to_text else {
return Err(anyhow::anyhow!(
strings::agent::no_configuration_for_purpose_so_cannot_be_used(
&AgentPurpose::SpeechToText
),
));
};
// This library does not support passing the audio data as a byte slice, so we need to write it to a temporary file :/
//
// This temporary file will get auto-deleted when the variable goes out of scope.
let temp_file = tokio::task::spawn_blocking(move || {
let mut temp_file = match tempfile::NamedTempFile::new() {
Ok(file) => file,
Err(e) => return Err(e),
};
match std::io::Write::write_all(&mut temp_file, &media) {
Ok(_) => (),
Err(e) => return Err(e),
}
Ok(temp_file)
})
.await??;
let file_path = temp_file
.path()
.to_str()
.ok_or_else(|| anyhow::anyhow!("Failed to get temporary file path"))?;
let language = params.language_override.clone();
let request = AudioBody {
file: std::fs::File::open(file_path)?,
model: speech_to_text_config.model_id.to_owned(),
prompt: None,
response_format: None,
temperature: None,
language: language.clone(),
};
tracing::trace!(
model_id = speech_to_text_config.model_id,
?language,
"Sending OpenAI-compat speech-to-text API request"
);
// This library is not async-aware, so we need to use `spawn_blocking` to run the request on a separate thread.
let client = self.client.clone();
let response =
tokio::task::spawn_blocking(move || client.audio_transcription_create(request)).await?;
let response = match response {
Ok(response) => response,
Err(err) => {
return Err(anyhow::anyhow!(
"Failed to get response from the OpenAI-compat audio transcription API: {:?}",
err
));
}
};
tracing::trace!(
?response,
"Got response from the OpenAI-compat audio transcription API"
);
let Some(text) = response.text else {
return Err(anyhow::anyhow!(
"No response text was returned from the OpenAI-compat audio transcription API"
));
};
Ok(SpeechToTextResult { text })
}
async fn generate_image(
&self,
prompt: &str,
params: ImageGenerationParams,
) -> anyhow::Result<ImageGenerationResult> {
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
),
));
};
// It seems like some OpenAI-compatible providers (e.g. LocalAI with StableDiffusion) skip some requirements
// when they span multiple lines.
let prompt = prompt.replace("\n", " ");
let size: Option<String> = params
.size_override
.or_else(|| image_generation_config.size.clone());
let request = ImagesBody {
model: Some(image_generation_config.model_id.to_owned()),
prompt: prompt.to_owned(),
n: Some(1),
quality: image_generation_config.quality.clone(),
size,
style: image_generation_config.style.clone(),
response_format: Some("b64_json".to_string()),
user: None,
};
tracing::trace!(
?prompt,
model = format!("{:?}", request.model),
size = format!("{:?}", request.size),
style = format!("{:?}", request.style),
quality = format!("{:?}", request.quality),
"Sending OpenAI-compat image generation API request"
);
// This library is not async-aware, so we need to use `spawn_blocking` to run the request on a separate thread.
let client = self.client.clone();
let response = tokio::task::spawn_blocking(move || client.image_create(&request)).await?;
let response = match response {
Ok(response) => response,
Err(err) => {
return Err(anyhow::anyhow!(
"Failed to get response from the OpenAI-compat image creation API: {:?}",
err
));
}
};
let Some(data) = response.data else {
return Err(anyhow::anyhow!(
"The OpenAI-compat image generationAPI returned no image data"
));
};
if let Some(image) = data.into_iter().next() {
let Some(b64_json) = &image.b64_json else {
return Err(anyhow::anyhow!(
"The OpenAI-compat image generation API returned no b64_json image data"
));
};
let bytes = base64_decode(b64_json)?;
return Ok(ImageGenerationResult {
bytes,
mime_type: mxlink::mime::IMAGE_PNG,
revised_prompt: image.revised_prompt,
});
}
Err(anyhow::anyhow!(
"The OpenAI image generation API returned no images"
))
}
async fn text_to_speech(
&self,
input: &str,
params: TextToSpeechParams,
) -> anyhow::Result<TextToSpeechResult> {
// openai_api_rust does not support text-to-speech, so our only bet is to do it via async-openai and hope it works.
// At the time of testing (2024-09-09), providers like LocalAI can be used for text-to-speech via async-openai.
//
// So.. below we try to convert our Config struct to the Config struct from the openai module
// and invoke the openai controller.
// Quick check to make sure doing work below is worth it
let Some(_text_to_speech_config) = &self.config.text_to_speech else {
return Err(anyhow::anyhow!(
strings::agent::no_configuration_for_purpose_so_cannot_be_used(
&AgentPurpose::TextToSpeech
),
));
};
tracing::debug!("Converting OpenAI-compact config to OpenAI config..");
let openai_config = super::utils::convert_config_to_openai_config_lossy(&self.config);
let Some(_text_to_speech_config) = &openai_config.text_to_speech else {
return Err(anyhow::anyhow!(
strings::agent::no_configuration_for_purpose_after_conversion_so_cannot_be_used(
&AgentPurpose::TextToSpeech
),
));
};
let openai_controller = super::super::openai::Controller::new(openai_config);
tracing::error!("Invoking text-to-speech via the OpenAI controller..");
openai_controller.text_to_speech(input, params).await
}
fn supports_purpose(&self, purpose: AgentPurpose) -> bool {
match purpose {
AgentPurpose::ImageGeneration => self.config.image_generation.is_some(),
AgentPurpose::TextGeneration => self.config.text_generation.is_some(),
AgentPurpose::SpeechToText => self.config.speech_to_text.is_some(),
AgentPurpose::TextToSpeech => self.config.text_to_speech.is_some(),
AgentPurpose::CatchAll => true,
}
}
fn text_generation_prompt(&self) -> Option<String> {
let Some(text_generation_config) = &self.config.text_generation else {
return None;
};
text_generation_config.prompt.clone()
}
fn text_generation_temperature(&self) -> Option<f32> {
let Some(text_generation_config) = &self.config.text_generation else {
return None;
};
Some(text_generation_config.temperature)
}
fn text_to_speech_voice(&self) -> Option<String> {
let Some(text_to_speech_config) = &self.config.text_to_speech else {
return None;
};
// A hacky way to turn this enum to a string
let voice_as_string = serde_json::to_string(&text_to_speech_config.voice).ok()?;
Some(voice_as_string.replace("\"", ""))
}
fn text_to_speech_speed(&self) -> Option<f32> {
let Some(text_to_speech_config) = &self.config.text_to_speech else {
return None;
};
Some(text_to_speech_config.speed)
}
}

View File

@@ -0,0 +1,70 @@
// The openai_compat provider aims to support a wider ranger of OpenAI-compatible providers.
//
// The `openai` provider is based on `async-openai`, which only aims to support the OpenAI API spec. See:
// - https://github.com/64bit/async-openai/issues/266
// - https://github.com/64bit/async-openai/blob/05d5a1b4fa6476829dd1a34447b80279cf89d4f8/async-openai/README.md#contributing
//
// This module uses its own configuration, which avoids using strict types tied to OpenAI,
// and thus allows for more flexibility.
//
// Communication with the OpenAI-compatible API is handled by the `openai_api_rust` crate.
// Since this crate is not async-aware, we need to use tokio's `spawn_blocking` to invoke it.
//
// Certain features (e.g. text-to-speech) are not supported by `openai_api_rust` yet, so we may try to delegate them to the `openai` provider.
mod config;
mod controller;
mod utils;
pub use config::Config;
pub use controller::Controller;
use super::super::AgentInstantiationError;
use super::super::AgentInstantiationResult;
use super::controller::ControllerType;
use super::ConfigTrait;
pub fn create_controller_from_yaml_value_config(
agent_id: &str,
config: serde_yaml::Value,
) -> AgentInstantiationResult<ControllerType> {
let config = match &config {
serde_yaml::Value::Mapping(_) => {
let config: Config =
serde_yaml::from_value(config).map_err(AgentInstantiationError::Yaml)?;
config
.validate()
.map_err(AgentInstantiationError::ConfigFailsValidation)?;
config
}
_ => {
return Err(AgentInstantiationError::ConfigForAgentIsNotAMapping(
agent_id.to_owned(),
));
}
};
Ok(ControllerType::OpenAICompat(Box::new(Controller::new(
config,
))))
}
pub fn default_config() -> Config {
let mut config = Config::default();
if let Some(text_generation) = &mut config.text_generation {
text_generation.model_id = "some-model".to_string();
text_generation.max_response_tokens = 4096;
text_generation.max_context_tokens = 128_000;
}
// We don't support these, so let's remove them from the configuration.
config.text_to_speech = None;
config.image_generation = None;
config.base_url = "".to_owned();
config
}

View File

@@ -0,0 +1,78 @@
use openai_api_rust::{Message, Role};
use crate::agent::provider::openai::Config as OpenAIConfig;
use crate::conversation::llm::{Author as LLMAuthor, Message as LLMMessage};
pub fn convert_llm_messages_to_openai_messages(
conversation_messages: Vec<LLMMessage>,
) -> Vec<Message> {
let mut openai_conversation_messages: Vec<Message> =
Vec::with_capacity(conversation_messages.len());
for message in conversation_messages {
openai_conversation_messages.push(convert_llm_message_to_openai_message(message));
}
openai_conversation_messages
}
fn convert_llm_message_to_openai_message(llm_message: LLMMessage) -> Message {
let role = match llm_message.author {
LLMAuthor::Prompt => Role::System,
LLMAuthor::Assistant => Role::Assistant,
LLMAuthor::User => Role::User,
};
Message {
role,
content: llm_message.message_text,
}
}
pub(super) fn convert_config_to_openai_config_lossy(config: &super::Config) -> OpenAIConfig {
let text_generation = config
.text_generation
.as_ref()
.and_then(|tg| tg.clone().try_into().ok());
let speech_to_text = config
.speech_to_text
.as_ref()
.and_then(|stt| stt.clone().try_into().ok());
let text_to_speech = config
.text_to_speech
.as_ref()
.and_then(|tts| tts.clone().try_into().ok());
let image_generation = config
.image_generation
.as_ref()
.and_then(|ig| ig.clone().try_into().ok());
OpenAIConfig {
api_key: config.api_key.clone().unwrap_or("".to_string()),
text_generation,
speech_to_text,
text_to_speech,
image_generation,
base_url: config.base_url.clone(),
}
}
pub(super) fn convert_string_to_enum<T>(value: &str) -> Result<T, String>
where
T: serde::de::DeserializeOwned,
{
// This is a hacky way to construct an enum from the string we have.
let enum_result: serde_json::Result<T> = serde_json::from_str(&format!("\"{}\"", value));
match enum_result {
Ok(enum_result) => Ok(enum_result),
Err(err) => {
tracing::debug!(?err, "Failed to parse into enum");
Err(format!("The value ({}) is not supported.", value))
}
}
}

View File

@@ -0,0 +1,21 @@
use super::openai_compat::Config;
pub fn default_config() -> Config {
let mut config = Config {
base_url: "https://openrouter.ai/api/v1".to_owned(),
text_to_speech: None,
image_generation: None,
speech_to_text: None,
..Default::default()
};
if let Some(ref mut config) = config.text_generation.as_mut() {
config.model_id = "mattshumer/reflection-70b:free".to_owned();
config.max_context_tokens = 8192;
config.max_response_tokens = 2048;
}
config
}

View File

@@ -0,0 +1,21 @@
use super::openai_compat::Config;
pub fn default_config() -> Config {
let mut config = Config {
base_url: "https://api.together.xyz/v1".to_owned(),
text_to_speech: None,
image_generation: None,
speech_to_text: None,
..Default::default()
};
if let Some(ref mut config) = config.text_generation.as_mut() {
config.model_id = "meta-llama/Meta-Llama-3.1-405B-Instruct-Turbo".to_owned();
config.max_context_tokens = 8192;
config.max_response_tokens = 2048;
}
config
}

67
src/agent/purpose.rs Normal file
View File

@@ -0,0 +1,67 @@
#[derive(Debug, Clone, Copy, PartialEq)]
pub enum AgentPurpose {
CatchAll,
ImageGeneration,
TextGeneration,
TextToSpeech,
SpeechToText,
}
impl AgentPurpose {
pub fn from_str(s: &str) -> Option<Self> {
match s {
"catch-all" => Some(Self::CatchAll),
"image-generation" => Some(Self::ImageGeneration),
"text-generation" => Some(Self::TextGeneration),
"text-to-speech" => Some(Self::TextToSpeech),
"speech-to-text" => Some(Self::SpeechToText),
_ => None,
}
}
pub fn as_str(&self) -> &'static str {
match self {
Self::CatchAll => "catch-all",
Self::ImageGeneration => "image-generation",
Self::TextGeneration => "text-generation",
Self::TextToSpeech => "text-to-speech",
Self::SpeechToText => "speech-to-text",
}
}
pub fn choices() -> Vec<&'static Self> {
vec![
&Self::TextGeneration,
&Self::SpeechToText,
&Self::TextToSpeech,
&Self::ImageGeneration,
&Self::CatchAll,
]
}
pub fn emoji(&self) -> &'static str {
match self {
Self::CatchAll => "❓",
Self::TextGeneration => "💬",
Self::SpeechToText => "🦻",
Self::TextToSpeech => "🗣️",
Self::ImageGeneration => "🖌️",
}
}
pub fn heading(&self) -> &'static str {
match self {
Self::CatchAll => "Catch-All",
Self::TextGeneration => "Text Generation",
Self::SpeechToText => "Speech-to-Text",
Self::TextToSpeech => "Text-to-Speech",
Self::ImageGeneration => "Image Generation",
}
}
}
impl std::fmt::Display for AgentPurpose {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str(self.as_str())
}
}

146
src/agent/utils.rs Normal file
View File

@@ -0,0 +1,146 @@
use base64::{engine::general_purpose::STANDARD, Engine as _};
use crate::{
agent::{
AgentInstance, AgentPurpose, ControllerTrait, Manager as AgentManager, PublicIdentifier,
},
entity::RoomConfigContext,
strings,
};
#[derive(Debug)]
pub struct AgentForPurposeDeterminationInfo {
pub instance: AgentInstance,
pub configuration_source: AgentForPurposeDeterminationInfoConfigurationSource,
}
#[derive(Debug)]
pub enum AgentForPurposeDeterminationInfoConfigurationSource {
Room,
Global,
}
#[derive(Debug)]
pub enum AgentForPurposeDeterminationError {
Unknown(String),
NoneConfigured,
ConfiguredButMissing(PublicIdentifier),
ConfiguredButLacksSupport(PublicIdentifier),
}
pub async fn get_effective_agent_for_purpose(
agent_manager: &AgentManager,
room_config_context: &RoomConfigContext,
agent_purpose: AgentPurpose,
) -> Result<AgentForPurposeDeterminationInfo, AgentForPurposeDeterminationError> {
let (agent_identifier, configuration_source) =
match get_effective_room_agent_identifier_for_purpose(room_config_context, agent_purpose)
.await
{
Ok((agent_identifier, configuration_source)) => {
(agent_identifier, configuration_source)
}
Err(err) => {
return Err(AgentForPurposeDeterminationError::Unknown(err));
}
};
let Some(agent_identifier) = agent_identifier else {
return Err(AgentForPurposeDeterminationError::NoneConfigured);
};
let agents = agent_manager.available_room_agents_by_room_config_context(room_config_context);
let Some(agent_instance) = agents.iter().find(|a| *a.identifier() == agent_identifier) else {
return Err(AgentForPurposeDeterminationError::ConfiguredButMissing(
agent_identifier,
));
};
let agent_instance = agent_instance.clone();
let supports_purpose = agent_instance.controller().supports_purpose(agent_purpose);
if !supports_purpose {
return Err(AgentForPurposeDeterminationError::ConfiguredButLacksSupport(agent_identifier));
}
Ok(AgentForPurposeDeterminationInfo {
instance: agent_instance,
configuration_source,
})
}
async fn get_effective_room_agent_identifier_for_purpose(
room_config_context: &RoomConfigContext,
purpose: AgentPurpose,
) -> Result<
(
Option<PublicIdentifier>,
AgentForPurposeDeterminationInfoConfigurationSource,
),
String,
> {
let (agent_id, configuration_source) =
get_effective_room_agent_raw_id_for_purpose(room_config_context, purpose).await;
let Some(agent_id) = agent_id else {
return Ok((None, configuration_source));
};
let agent_identifier = match PublicIdentifier::from_str(agent_id.as_str()) {
Some(agent_identifier) => agent_identifier,
None => return Err(strings::agent::invalid_id_generic()),
};
Ok((Some(agent_identifier), configuration_source))
}
async fn get_effective_room_agent_raw_id_for_purpose(
room_config_context: &RoomConfigContext,
purpose: AgentPurpose,
) -> (
Option<String>,
AgentForPurposeDeterminationInfoConfigurationSource,
) {
let agent_id = room_config_context
.room_config
.settings
.handler
.get_by_purpose_with_catch_all_fallback(purpose);
if let Some(agent_id) = agent_id {
return (
Some(agent_id),
AgentForPurposeDeterminationInfoConfigurationSource::Room,
);
}
tracing::trace!(
?purpose,
"No specific agent found for purpose in room, falling back to global.",
);
(
get_global_agent_id_for_purpose(room_config_context, purpose).await,
AgentForPurposeDeterminationInfoConfigurationSource::Global,
)
}
async fn get_global_agent_id_for_purpose(
room_config_context: &RoomConfigContext,
purpose: AgentPurpose,
) -> Option<String> {
room_config_context
.global_config
.fallback_room_settings
.handler
.get_by_purpose_with_catch_all_fallback(purpose)
}
pub(crate) fn base64_decode(base64_string: &str) -> Result<Vec<u8>, base64::DecodeError> {
STANDARD.decode(base64_string)
}

419
src/bot/implementation.rs Normal file
View File

@@ -0,0 +1,419 @@
use std::sync::Arc;
use std::{future::Future, pin::Pin};
use mxlink::matrix_sdk::media::{MediaFormat, MediaRequest};
use mxlink::matrix_sdk::ruma::{
events::room::MediaSource, MilliSecondsSinceUnixEpoch, OwnedUserId,
};
use mxlink::matrix_sdk::Room;
use mxlink::{
InitConfig, LoginConfig, LoginCredentials, LoginEncryption, MatrixLink, PersistenceConfig,
};
use mxlink::helpers::account_data_config::{
ConfigError, GlobalConfigManager as AccountDataGlobalConfigManager,
RoomConfigManager as AccountDataRoomConfigManager,
};
use mxlink::helpers::encryption::Manager as EncryptionManager;
use crate::agent::Manager as AgentManager;
use crate::entity::catch_up_marker::{
CatchUpMarker, CatchUpMarkerManager, DelayedCatchUpMarkerManager,
};
use crate::entity::cfg::Config;
use crate::entity::globalconfig::{GlobalConfig, GlobalConfigurationManager};
use crate::entity::roomconfig::{RoomConfig, RoomConfigurationManager};
use crate::agent::Manager;
use crate::conversation::matrix::{RoomDisplayNameFetcher, RoomEventFetcher};
const ROOM_EVENT_FETCHER_LRU_CACHE_SIZE: usize = 1000;
const ROOM_DISPLAY_NAME_FETCHER_LRU_CACHE_SIZE: usize = 1000;
const ROOM_CONFIG_MANAGER_LRU_CACHE_SIZE: usize = 1000;
const LOGO_BYTES: &[u8] = include_bytes!("../../etc/assets/baibot-torso-768.png");
const LOGO_MIME_TYPE: &str = "image/png";
/// Controls how often we persist the catch-up marker to Account Data.
/// Consult the `DelayedCatchUpMarkerManager` documentation for more information.
const DELAYED_CATCH_UP_MARKER_MANAGER_PERSIST_INTERVAL_DURATION: std::time::Duration =
std::time::Duration::from_secs(10);
/// Controls what federation delay we will tolerate. The timestamp that gets persisted
/// will be based on the last seen event's `origin_server_ts` minus this duration.
/// Consult the `DelayedCatchUpMarkerManager` documentation for more information.
const DELAYED_CATCH_UP_MARKER_MANAGER_FEDERATION_DELAY_TOLERANCE_DURATION: std::time::Duration =
std::time::Duration::from_secs(90);
struct BotInner {
config: Config,
matrix_link: MatrixLink,
delayed_catch_up_marker_manager: DelayedCatchUpMarkerManager,
global_config_manager: tokio::sync::Mutex<GlobalConfigurationManager>,
room_config_manager: tokio::sync::Mutex<RoomConfigurationManager>,
room_event_fetcher: Arc<RoomEventFetcher>,
room_display_name_fetcher: Arc<RoomDisplayNameFetcher>,
agent_manager: Manager,
admin_pattern_regexes: Vec<regex::Regex>,
}
/// Bot represents a bot instance.
///
/// All of the state is held in an `Arc` so the `Bot` can be cloned freely.
#[derive(Clone)]
pub struct Bot {
inner: Arc<BotInner>,
}
impl Bot {
pub async fn new(config: Config) -> anyhow::Result<Self> {
// Take some potentially problematic configuration values out of the config early on.
// If we'd be failing, we'd like it to happen early, before we log in, etc.
let initial_global_config: GlobalConfig =
config.initial_global_config.clone().try_into()?;
let admin_pattern_regexes = config.access.admin_pattern_regexes()?;
let persistence_config_encryption_key = config.persistence.config_encryption_key()?;
let agent_manager = AgentManager::new(config.agents.static_definitions.clone())?;
let encryption_manager = EncryptionManager::new(persistence_config_encryption_key);
let matrix_link = create_matrix_link(&config).await?;
let catch_up_marker_manager = create_catch_up_marker_manager(matrix_link.clone());
let delayed_catch_up_marker_manager = DelayedCatchUpMarkerManager::new(
catch_up_marker_manager,
DELAYED_CATCH_UP_MARKER_MANAGER_PERSIST_INTERVAL_DURATION,
DELAYED_CATCH_UP_MARKER_MANAGER_FEDERATION_DELAY_TOLERANCE_DURATION,
);
let global_config_manager = tokio::sync::Mutex::new(create_global_configuration_manager(
matrix_link.clone(),
encryption_manager.clone(),
initial_global_config,
));
let room_config_manager = tokio::sync::Mutex::new(create_room_configuration_manager(
matrix_link.clone(),
encryption_manager.clone(),
));
let room_event_fetcher = RoomEventFetcher::new(Some(ROOM_EVENT_FETCHER_LRU_CACHE_SIZE));
let room_display_name_fetcher = RoomDisplayNameFetcher::new(
matrix_link.clone(),
Some(ROOM_DISPLAY_NAME_FETCHER_LRU_CACHE_SIZE),
);
Ok(Self {
inner: Arc::new(BotInner {
config,
matrix_link,
delayed_catch_up_marker_manager,
global_config_manager,
room_config_manager,
room_event_fetcher: Arc::new(room_event_fetcher),
room_display_name_fetcher: Arc::new(room_display_name_fetcher),
agent_manager,
admin_pattern_regexes,
}),
})
}
pub(crate) fn admin_patterns(&self) -> &Vec<String> {
&self.inner.config.access.admin_patterns
}
pub(crate) fn name(&self) -> &str {
&self.inner.config.user.name
}
pub(crate) fn command_prefix(&self) -> &str {
&self.inner.config.command_prefix
}
pub(crate) fn homeserver_name(&self) -> &str {
&self.inner.config.homeserver.server_name
}
pub(crate) fn global_config_manager(&self) -> &tokio::sync::Mutex<GlobalConfigurationManager> {
&self.inner.global_config_manager
}
pub(crate) fn room_config_manager(&self) -> &tokio::sync::Mutex<RoomConfigurationManager> {
&self.inner.room_config_manager
}
pub(crate) fn room_event_fetcher(&self) -> Arc<RoomEventFetcher> {
self.inner.room_event_fetcher.clone()
}
pub(crate) fn room_display_name_fetcher(&self) -> Arc<RoomDisplayNameFetcher> {
self.inner.room_display_name_fetcher.clone()
}
pub(crate) fn agent_manager(&self) -> &Manager {
&self.inner.agent_manager
}
pub(crate) fn matrix_link(&self) -> &MatrixLink {
&self.inner.matrix_link
}
pub(crate) fn user_id(&self) -> &OwnedUserId {
self.matrix_link().user_id()
}
pub(crate) fn reacting(&self) -> super::reacting::Reacting {
super::reacting::Reacting::new(self.clone())
}
pub(crate) fn rooms(&self) -> super::rooms::Rooms {
super::rooms::Rooms::new(self.clone())
}
pub(crate) fn messaging(&self) -> super::messaging::Messaging {
super::messaging::Messaging::new(self.clone())
}
pub(crate) fn admin_pattern_regexes(&self) -> &Vec<regex::Regex> {
&self.inner.admin_pattern_regexes
}
pub(crate) async fn global_config(&self) -> Result<GlobalConfig, ConfigError> {
let mut global_config_manager_guard = self.inner.global_config_manager.lock().await;
global_config_manager_guard.get_or_create().await
}
pub(crate) async fn is_caught_up(
&self,
event_origin_server_ts: MilliSecondsSinceUnixEpoch,
) -> Result<bool, ConfigError> {
self.inner
.delayed_catch_up_marker_manager
.is_caught_up(event_origin_server_ts.0.into())
.await
}
pub(crate) async fn catch_up(&self, event_origin_server_ts: MilliSecondsSinceUnixEpoch) {
self.inner
.delayed_catch_up_marker_manager
.catch_up(event_origin_server_ts.0.into())
.await
}
pub async fn start(&self) -> anyhow::Result<()> {
self.rooms().attach_event_handlers().await;
self.messaging().attach_event_handlers().await;
self.reacting().attach_event_handlers().await;
self.inner.delayed_catch_up_marker_manager.start().await;
self.prepare_profile().await?;
self.inner
.matrix_link
.start()
.await
.map_err(|e| anyhow::anyhow!("Failed to sync: {:?}", e))
}
async fn prepare_profile(&self) -> anyhow::Result<()> {
use std::time::Duration;
use tokio::time::sleep;
let mut delay = Duration::from_secs(3);
let max_delay = Duration::from_secs(30);
loop {
match self.do_prepare_profile().await {
Ok(_) => return Ok(()),
Err(err) => {
tracing::warn!(
?err,
?delay,
"Failed to prepare profile.. Will retry after delay..."
);
sleep(delay).await;
delay = std::cmp::min(delay * 2, max_delay);
}
}
}
}
async fn do_prepare_profile(&self) -> anyhow::Result<()> {
tracing::debug!("Preparing profile..");
let account = self.inner.matrix_link.client().account();
let media = self.inner.matrix_link.client().media();
let desired_display_name = self.inner.config.user.name.clone();
let profile = account
.get_profile()
.await
.map_err(|e| anyhow::anyhow!("Failed fetching profile: {:?}", e))?;
let should_update_display_name = match &profile.displayname {
Some(displayname) => displayname != &desired_display_name,
None => true,
};
if should_update_display_name {
tracing::info!(
?profile.displayname,
?desired_display_name,
"Updating display name.."
);
if let Err(err) = account.set_display_name(Some(&desired_display_name)).await {
return Err(anyhow::anyhow!("Failed setting display name: {:?}", err));
}
}
let should_update_avatar = match &profile.avatar_url {
Some(avatar_url) => {
let request = MediaRequest {
source: MediaSource::Plain(avatar_url.to_owned()),
format: MediaFormat::File,
};
let content = media
.get_media_content(&request, true)
.await
.map_err(|e| anyhow::anyhow!("Failed fetching existing avatar: {:?}", e))?;
content.as_slice() != LOGO_BYTES
}
None => true,
};
if should_update_avatar {
tracing::info!("Updating avatar..");
let mime_type = LOGO_MIME_TYPE
.parse()
.expect("Failed parsing mime type for logo");
account
.upload_avatar(&mime_type, LOGO_BYTES.to_vec())
.await
.map_err(|e| anyhow::anyhow!("Failed uploading avatar: {:?}", e))?;
}
Ok(())
}
}
async fn create_matrix_link(config: &Config) -> anyhow::Result<MatrixLink> {
let session_file_path = config.persistence.session_file_path()?;
let session_encryption_key = config.persistence.session_encryption_key()?;
let db_dir_path: std::path::PathBuf = config.persistence.db_dir_path()?;
let login_creds = LoginCredentials::UserPassword(
config.user.mxid_localpart.to_owned(),
config.user.password.to_owned(),
);
let login_encryption = LoginEncryption::new(
config.user.encryption.recovery_passphrase.clone(),
config.user.encryption.recovery_reset_allowed,
);
let login_config = LoginConfig::new(
config.homeserver.url.to_owned(),
login_creds,
Some(login_encryption),
config.user.name.to_owned(),
);
let persistence_config =
PersistenceConfig::new(session_file_path, session_encryption_key, db_dir_path);
let init_config = InitConfig::new(login_config, persistence_config);
mxlink::init(&init_config).await.map_err(|e| e.into())
}
pub fn create_global_configuration_manager(
matrix_link: MatrixLink,
encryption_manager: EncryptionManager,
initial_global_config: GlobalConfig,
) -> GlobalConfigurationManager {
let initial_global_config_callback = move || {
let initial_global_config = initial_global_config.clone();
let future = create_initial_global_config(initial_global_config);
// Explicitly box the future to match the expected type
Box::pin(future) as Pin<Box<dyn Future<Output = GlobalConfig> + Send>>
};
AccountDataGlobalConfigManager::new(
matrix_link,
encryption_manager,
initial_global_config_callback,
)
}
async fn create_initial_global_config(initial_global_config: GlobalConfig) -> GlobalConfig {
initial_global_config
}
pub fn create_room_configuration_manager(
matrix_link: MatrixLink,
encryption_manager: EncryptionManager,
) -> RoomConfigurationManager {
let initial_room_config_callback = |room: Room| {
let future = create_initial_room_config(room);
// Explicitly box the future to match the expected type
Box::pin(future) as Pin<Box<dyn Future<Output = RoomConfig> + Send>>
};
AccountDataRoomConfigManager::new(
matrix_link.user_id().clone(),
encryption_manager,
initial_room_config_callback,
Some(ROOM_CONFIG_MANAGER_LRU_CACHE_SIZE),
)
}
async fn create_initial_room_config(room: Room) -> RoomConfig {
RoomConfig::default().with_room(room).await
}
pub fn create_catch_up_marker_manager(matrix_link: MatrixLink) -> CatchUpMarkerManager {
let initial_global_config_callback = || {
let future = create_initial_catch_up_marker();
// Explicitly box the future to match the expected type
Box::pin(future) as Pin<Box<dyn Future<Output = CatchUpMarker> + Send>>
};
// Intentionally not using encryption, to make this resilient even if we lose our encryption key.
// We're not worried about the catch-up marker being read or tampered with, as it's not sensitive data.
let encryption_manager = EncryptionManager::new(None);
let catch_up_marker_manager: CatchUpMarkerManager = AccountDataGlobalConfigManager::new(
matrix_link.clone(),
encryption_manager,
initial_global_config_callback,
);
catch_up_marker_manager
}
async fn create_initial_catch_up_marker() -> CatchUpMarker {
CatchUpMarker::new(0)
}

110
src/bot/load_config.rs Normal file
View File

@@ -0,0 +1,110 @@
use std::env;
use std::path::PathBuf;
use anyhow::anyhow;
use crate::agent::AgentPurpose;
pub use crate::entity::cfg::{defaults as cfg_defaults, env as cfg_env, Config};
pub fn load() -> anyhow::Result<Config> {
let config_file_path = env::var(cfg_env::BAIBOT_CONFIG_FILE_PATH)
.unwrap_or_else(|_| cfg_defaults::config_file_path().to_owned());
let config_file_path = PathBuf::from(config_file_path);
if !config_file_path.exists() {
return Err(anyhow!(
"Config file ({}) not found. Adjust the {} environment variable to use another config file.",
config_file_path.display(),
cfg_env::BAIBOT_CONFIG_FILE_PATH,
));
}
let config_str = std::fs::read_to_string(config_file_path)?;
let mut config: Config = serde_yaml::from_str(&config_str)?;
// Allow environment variables to override some configuration keys
for (key, value) in env::vars() {
match key.as_str() {
cfg_env::BAIBOT_HOMESERVER_SERVER_NAME => config.homeserver.server_name = value,
cfg_env::BAIBOT_HOMESERVER_URL => config.homeserver.url = value,
cfg_env::BAIBOT_USER_MXID_LOCALPART => config.user.mxid_localpart = value,
cfg_env::BAIBOT_USER_PASSWORD => config.user.password = value,
cfg_env::BAIBOT_USER_ENCRYPTION_RECOVERY_PASSPHRASE => {
config.user.encryption.recovery_passphrase = Some(value);
}
cfg_env::BAIBOT_USER_NAME => config.user.name = value,
cfg_env::BAIBOT_COMMAND_PREFIX => config.command_prefix = value,
cfg_env::BAIBOT_LOGGING => {
config.logging = value;
}
cfg_env::BAIBOT_ACCESS_ADMIN_PATTERNS => {
config.access.admin_patterns = value
.split(' ')
.map(|s| s.trim().to_string())
.filter(|s| !s.is_empty())
.collect();
}
cfg_env::BAIBOT_PERSISTENCE_DATA_DIR_PATH => {
config.persistence.data_dir_path = Some(value);
}
cfg_env::BAIBOT_PERSISTENCE_CONFIG_ENCRYPTION_KEY => {
config.persistence.config_encryption_key = Some(value);
}
cfg_env::BAIBOT_INITIAL_GLOBAL_CONFIG_HANDLER_CATCH_ALL => {
let value = if value.is_empty() { None } else { Some(value) };
config
.initial_global_config
.handler
.set_by_purpose(AgentPurpose::CatchAll, value);
}
cfg_env::BAIBOT_INITIAL_GLOBAL_CONFIG_HANDLER_TEXT_GENERATION => {
let value = if value.is_empty() { None } else { Some(value) };
config
.initial_global_config
.handler
.set_by_purpose(AgentPurpose::TextGeneration, value);
}
cfg_env::BAIBOT_INITIAL_GLOBAL_CONFIG_HANDLER_TEXT_TO_SPEECH => {
let value = if value.is_empty() { None } else { Some(value) };
config
.initial_global_config
.handler
.set_by_purpose(AgentPurpose::TextToSpeech, value);
}
cfg_env::BAIBOT_INITIAL_GLOBAL_CONFIG_HANDLER_SPEECH_TO_TEXT => {
let value = if value.is_empty() { None } else { Some(value) };
config
.initial_global_config
.handler
.set_by_purpose(AgentPurpose::SpeechToText, value);
}
cfg_env::BAIBOT_INITIAL_GLOBAL_CONFIG_HANDLER_IMAGE_GENERATION => {
let value = if value.is_empty() { None } else { Some(value) };
config
.initial_global_config
.handler
.set_by_purpose(AgentPurpose::ImageGeneration, value);
}
cfg_env::BAIBOT_INITIAL_GLOBAL_CONFIG_USER_PATTERNS => {
config.initial_global_config.user_patterns = Some(
value
.split(' ')
.map(|s| s.trim().to_string())
.filter(|s| !s.is_empty())
.collect(),
);
}
_ => {}
}
}
config.validate().map_err(|s| anyhow!(s))?;
Ok(config)
}

334
src/bot/messaging.rs Normal file
View File

@@ -0,0 +1,334 @@
use mxlink::matrix_sdk::{
ruma::{
api::client::receipt::create_receipt::v3::ReceiptType,
events::room::message::OriginalSyncRoomMessageEvent, OwnedEventId,
},
Room,
};
use mxlink::{CallbackError, MessageResponseType};
use tracing::Instrument;
use crate::{
conversation::matrix::determine_thread_context_for_room_event,
entity::{MessageContext, MessagePayload, RoomConfigContext, TriggerEventInfo},
};
#[derive(Clone)]
pub struct Messaging {
bot: super::Bot,
}
impl Messaging {
pub fn new(bot: super::Bot) -> Self {
Self { bot }
}
pub async fn send_text_markdown_no_fail(
&self,
room: &Room,
message: String,
response_type: MessageResponseType,
) -> Option<mxlink::matrix_sdk::ruma::api::client::message::send_message_event::v3::Response>
{
let result = self
.bot
.matrix_link()
.messaging()
.send_text_markdown(room, message, response_type)
.await;
match result {
Ok(result) => Some(result),
Err(err) => {
tracing::error!(
room_id = format!("{:?}", room.room_id()),
?err,
"Failed to send text message to room",
);
None
}
}
}
pub async fn send_notice_markdown_no_fail(
&self,
room: &Room,
message: String,
response_type: MessageResponseType,
) -> Option<mxlink::matrix_sdk::ruma::api::client::message::send_message_event::v3::Response>
{
let result = self
.bot
.matrix_link()
.messaging()
.send_notice_markdown(room, message, response_type)
.await;
match result {
Ok(result) => Some(result),
Err(err) => {
tracing::error!(
room_id = format!("{:?}", room.room_id()),
?err,
"Failed to send notice message to room",
);
None
}
}
}
pub async fn send_tooltip_markdown_no_fail(
&self,
room: &Room,
message: &str,
response_type: MessageResponseType,
) -> Option<mxlink::matrix_sdk::ruma::api::client::message::send_message_event::v3::Response>
{
self.send_notice_markdown_no_fail(
room,
crate::utils::status::create_tooltip_message_text(message),
response_type,
)
.await
}
pub async fn send_success_markdown_no_fail(
&self,
room: &Room,
message: &str,
response_type: MessageResponseType,
) -> Option<mxlink::matrix_sdk::ruma::api::client::message::send_message_event::v3::Response>
{
self.send_notice_markdown_no_fail(
room,
crate::utils::status::create_success_message_text(message),
response_type,
)
.await
}
pub async fn send_error_markdown_no_fail(
&self,
room: &Room,
err: &str,
response_type: MessageResponseType,
) -> Option<mxlink::matrix_sdk::ruma::api::client::message::send_message_event::v3::Response>
{
self.send_notice_markdown_no_fail(
room,
crate::utils::status::create_error_message_text(err),
response_type,
)
.await
}
pub async fn redact_event_no_fail(
&self,
room: &Room,
target_event_id: OwnedEventId,
reason: Option<String>,
) -> Option<mxlink::matrix_sdk::ruma::api::client::redact::redact_event::v3::Response> {
let result = self
.bot
.matrix_link()
.messaging()
.redact_event(room, target_event_id.clone(), reason)
.await;
match result {
Ok(result) => Some(result),
Err(err) => {
tracing::error!(
room_id = format!("{:?}", room.room_id()),
?target_event_id,
?err,
"Failed to send redaction to room",
);
None
}
}
}
pub(super) async fn attach_event_handlers(&self) {
let matrix_link_messaging = self.bot.matrix_link().messaging();
let this = self.clone();
matrix_link_messaging.on_actionable_room_message(|event, room| async move {
this.on_actionable_message(event, room).await
});
}
#[tracing::instrument(name = "bot_on_actionable_message", skip_all, fields(room_id = room.room_id().as_str(), event_id = event.event_id.as_str()))]
async fn on_actionable_message(
&self,
event: OriginalSyncRoomMessageEvent,
room: Room,
) -> Result<(), CallbackError> {
if self
.bot
.is_caught_up(event.origin_server_ts)
.await
.map_err(|e| {
CallbackError::Unknown(
format!("Failed to determine catch-up state: {:?}", e).into(),
)
})?
{
tracing::debug!(
event_origin_server_ts = format!("{:?}", event.origin_server_ts),
"Ignoring old message event",
);
return Ok(());
}
tracing::info!("Processing message");
let global_config = self
.bot
.global_config()
.await
.map_err(|err| CallbackError::Unknown(err.into()))?;
tracing::trace!(?global_config, "Global config");
let room_config = self
.bot
.room_config_manager()
.lock()
.await
.get_or_create_for_room(&room)
.await
.map_err(|err| CallbackError::Unknown(err.into()))?;
tracing::trace!(?room_config, "Room config");
let trigger_event_sender_is_admin = mxidwc::match_user_id(
event.sender.clone().as_str(),
self.bot.admin_pattern_regexes(),
);
let trigger_event_sender_is_allowed_user = match &global_config.access.user_patterns {
Some(user_patterns) => {
let allowed_user_regexes = mxidwc::parse_patterns_vector(user_patterns)
.map_err(|err| CallbackError::Unknown(err.into()))?;
mxidwc::match_user_id(event.sender.clone().as_str(), &allowed_user_regexes)
}
None => false,
};
if !trigger_event_sender_is_admin && !trigger_event_sender_is_allowed_user {
tracing::debug!("Ignoring message from non-admin/non-allowed user");
return Ok(());
}
let payload: Result<MessagePayload, String> = event.content.msgtype.clone().try_into();
let payload = match payload {
Ok(payload) => payload,
Err(err) => {
tracing::debug!(
msg_type = event.content.msgtype(),
?err,
"Ignoring message not supported by us",
);
return Ok(());
}
};
let thread_context = determine_thread_context_for_room_event(
self.bot.user_id(),
&room,
&event,
&payload,
&self.bot.room_event_fetcher(),
)
.await;
let thread_context = match thread_context {
Ok(value) => value,
Err(err) => {
tracing::error!(?err, "Failed to determine thread context for event");
return Ok(());
}
};
let Some(thread_context) = thread_context else {
tracing::debug!("Ignoring message with unknown thread context (likely not a threaded message or a top-level message)");
return Ok(());
};
let room_config_context =
RoomConfigContext::new(global_config.clone(), room_config.clone());
let trigger_event_info = TriggerEventInfo::new(
event.event_id.clone(),
event.sender.clone(),
payload,
trigger_event_sender_is_admin,
);
let message_context = MessageContext::new(
room.clone(),
room_config_context,
self.bot.admin_pattern_regexes().clone(),
trigger_event_info,
thread_context.info.clone(),
);
let bot_display_name = self
.bot
.room_display_name_fetcher()
.own_display_name_in_room(message_context.room())
.await;
let bot_display_name = match bot_display_name {
Ok(value) => value,
Err(err) => {
tracing::warn!(
?err,
"Failed to fetch bot display name. Proceeding without it"
);
None
}
};
// The first event in the thread determines which handler processes the current event.
let controller_type = crate::controller::determine_controller(
self.bot.command_prefix(),
&thread_context.first_message,
&message_context,
self.bot.user_id(),
&bot_display_name,
);
tracing::info!(?controller_type, "Determined controller");
let _ = room
.send_single_receipt(
ReceiptType::Read,
thread_context.info.clone().into(),
event.event_id.clone(),
)
.await;
let start_time = std::time::Instant::now();
let event_span = tracing::error_span!("message_controller", ?controller_type);
crate::controller::dispatch_controller(&controller_type, &message_context, &self.bot)
.instrument(event_span)
.await;
let duration = std::time::Instant::now().duration_since(start_time);
tracing::debug!(?duration, "Controller finished");
self.bot.catch_up(event.origin_server_ts).await;
return Ok(());
}
}

8
src/bot/mod.rs Normal file
View File

@@ -0,0 +1,8 @@
mod implementation;
mod load_config;
mod messaging;
mod reacting;
mod rooms;
pub use implementation::Bot;
pub use load_config::load as load_config;

262
src/bot/reacting.rs Normal file
View File

@@ -0,0 +1,262 @@
use mxlink::matrix_sdk::{
ruma::{
events::{
room::message::Relation, AnyMessageLikeEvent, AnySyncTimelineEvent, AnyTimelineEvent,
MessageLikeEvent,
},
OwnedEventId, OwnedUserId,
},
Room,
};
use mxlink::CallbackError;
use mxlink::ThreadInfo;
use tracing::Instrument;
use crate::entity::{MessageContext, MessagePayload, RoomConfigContext, TriggerEventInfo};
#[derive(Clone)]
pub struct Reacting {
bot: super::Bot,
}
impl Reacting {
pub fn new(bot: super::Bot) -> Self {
Self { bot }
}
pub async fn react_no_fail(
&self,
room: &Room,
target_event_id: OwnedEventId,
reaction_key: String,
) -> Option<mxlink::matrix_sdk::ruma::api::client::message::send_message_event::v3::Response>
{
let result = self
.bot
.matrix_link()
.reacting()
.react(room, target_event_id.clone(), reaction_key)
.await;
match result {
Ok(result) => Some(result),
Err(err) => {
tracing::error!(
"Failed to send reaction to {} in room {:?}: {:?}",
target_event_id,
room.room_id(),
err
);
None
}
}
}
pub(super) async fn attach_event_handlers(&self) {
let matrix_link_reacting = self.bot.matrix_link().reacting();
let this = self.clone();
matrix_link_reacting.on_actionable_reaction(
|event, room, reaction_event_content| async move {
this.on_actionable_reaction(event, room, reaction_event_content)
.await
},
);
}
#[tracing::instrument(name = "bot_on_actionable_reaction", skip_all, fields(room_id = room.room_id().as_str(), event_id = event.event_id().as_str()))]
async fn on_actionable_reaction(
&self,
event: AnySyncTimelineEvent,
room: Room,
reaction_event_content: mxlink::matrix_sdk::ruma::events::reaction::ReactionEventContent,
) -> Result<(), CallbackError> {
if self
.bot
.is_caught_up(event.origin_server_ts())
.await
.map_err(|e| {
CallbackError::Unknown(
format!("Failed to determine catch-up state: {:?}", e).into(),
)
})?
{
tracing::debug!(
event_origin_server_ts = format!("{:?}", event.origin_server_ts()),
"Ignoring old reaction event",
);
return Ok(());
}
tracing::info!("Handling reaction");
let global_config = self
.bot
.global_config()
.await
.map_err(|err| CallbackError::Unknown(err.into()))?;
tracing::trace!(?global_config, "Global config");
let trigger_event_sender_is_admin =
mxidwc::match_user_id(event.sender().as_str(), self.bot.admin_pattern_regexes());
let trigger_event_sender_is_allowed_user = match &global_config.access.user_patterns {
Some(user_patterns) => {
let allowed_user_regexes = mxidwc::parse_patterns_vector(user_patterns)
.map_err(|err| CallbackError::Unknown(err.into()))?;
mxidwc::match_user_id(event.sender().as_str(), &allowed_user_regexes)
}
None => false,
};
if !trigger_event_sender_is_admin && !trigger_event_sender_is_allowed_user {
tracing::debug!("Ignoring reaction from non-admin/non-allowed user");
return Ok(());
}
let reacted_to_event_id = &reaction_event_content.relates_to.event_id;
let reacted_to_event = self
.bot
.room_event_fetcher()
.fetch_event_in_room(reacted_to_event_id, &room)
.await;
let reacted_to_event = match reacted_to_event {
Ok(value) => value,
Err(err) => {
tracing::error!(
?reacted_to_event_id,
?err,
"Failed to fetch reacted-to event",
);
return Ok(());
}
};
let reacted_to_event_any_timeline_event = match reacted_to_event.event.deserialize() {
Ok(value) => value,
Err(err) => {
tracing::error!(
?reacted_to_event_id,
?err,
"Failed to deserialize reacted-to event event",
);
return Ok(());
}
};
let reacted_to_event_sender_id: OwnedUserId =
reacted_to_event_any_timeline_event.sender().to_owned();
let AnyTimelineEvent::MessageLike(reacted_to_event_message_like) =
reacted_to_event_any_timeline_event
else {
tracing::debug!(
?reacted_to_event_id,
"Ignoring non-MessageLike reacted-to event",
);
return Ok(());
};
let AnyMessageLikeEvent::RoomMessage(reacted_to_event_room_message) =
reacted_to_event_message_like
else {
tracing::debug!(
?reacted_to_event_id,
"Ignoring non-RoomMessage reacted-to event",
);
return Ok(());
};
let MessageLikeEvent::Original(reacted_to_event_room_message_original) =
reacted_to_event_room_message
else {
tracing::debug!(?reacted_to_event_id, "Ignoring redacted reacted-to event",);
return Ok(());
};
let reacted_to_event_payload: Result<MessagePayload, String> =
reacted_to_event_room_message_original
.content
.msgtype
.clone()
.try_into();
let Ok(reacted_to_event_payload) = reacted_to_event_payload else {
tracing::debug!(
msg_type = reacted_to_event_room_message_original.content.msgtype(),
"Ignoring reaction to message of unknown type",
);
return Ok(());
};
let thread_root_event_id = match reacted_to_event_room_message_original.content.relates_to {
Some(relation) => {
if let Relation::Thread(thread_id) = relation {
thread_id.event_id.clone()
} else {
reacted_to_event_id.clone()
}
}
None => reacted_to_event_id.clone(),
};
let thread_info = ThreadInfo::new(thread_root_event_id, reacted_to_event_id.clone());
let room_config = self
.bot
.room_config_manager()
.lock()
.await
.get_or_create_for_room(&room)
.await
.map_err(|err| CallbackError::Unknown(err.into()))?;
tracing::trace!(?room_config, "Room config");
let room_config_context =
RoomConfigContext::new(global_config.clone(), room_config.clone());
let trigger_event_info = TriggerEventInfo::new(
event.event_id().to_owned(),
event.sender().to_owned(),
MessagePayload::Reaction {
key: reaction_event_content.relates_to.key,
reacted_to_event_payload: Box::new(reacted_to_event_payload),
reacted_to_event_id: reaction_event_content.relates_to.event_id.clone(),
reacted_to_event_sender_id,
},
trigger_event_sender_is_admin,
);
let message_context = MessageContext::new(
room,
room_config_context,
self.bot.admin_pattern_regexes().clone(),
trigger_event_info,
thread_info,
);
tracing::info!("Handling reaction via reaction controller");
let event_span = tracing::error_span!("reaction_controller");
crate::controller::reaction::handle(
&self.bot,
self.bot.matrix_link().clone(),
&message_context,
)
.instrument(event_span)
.await
.map_err(|err| CallbackError::Unknown(err.into()))?;
self.bot.catch_up(event.origin_server_ts()).await;
Ok(())
}
}

145
src/bot/rooms.rs Normal file
View File

@@ -0,0 +1,145 @@
use mxlink::{
matrix_sdk::{
ruma::events::{room::member::StrippedRoomMemberEvent, AnySyncTimelineEvent},
Room,
},
InvitationDecision,
};
use mxlink::CallbackError;
use tracing::Instrument;
use crate::entity::RoomConfigContext;
#[derive(Clone)]
pub struct Rooms {
bot: super::Bot,
}
impl Rooms {
pub fn new(bot: super::Bot) -> Self {
Self { bot }
}
pub(super) async fn attach_event_handlers(&self) {
let matrix_link_rooms = self.bot.matrix_link().rooms();
let this = self.clone();
matrix_link_rooms.on_being_last_member(|event, room| async move {
this.on_being_last_member(event, room).await
});
let this = self.clone();
matrix_link_rooms
.on_invitation(|event, room| async move { this.on_invitation(event, room).await });
let this = self.clone();
matrix_link_rooms.on_joined(|event, room| async move { this.on_joined(event, room).await });
}
async fn on_invitation(
&self,
room_member: StrippedRoomMemberEvent,
_room: Room,
) -> Result<InvitationDecision, CallbackError> {
tracing::debug!("Deciding on room invitation");
let global_config = self
.bot
.global_config()
.await
.map_err(|e| CallbackError::Unknown(e.into()))?;
let sender_is_admin = mxidwc::match_user_id(
room_member.sender.clone().as_str(),
self.bot.admin_pattern_regexes(),
);
let sender_is_allowed_user = match &global_config.access.user_patterns {
Some(user_patterns) => {
let allowed_user_regexes = mxidwc::parse_patterns_vector(user_patterns)
.map_err(|e| CallbackError::Unknown(e.into()))?;
mxidwc::match_user_id(room_member.sender.clone().as_str(), &allowed_user_regexes)
}
None => false,
};
if !(sender_is_admin || sender_is_allowed_user) {
return Ok(InvitationDecision::Reject);
}
Ok(InvitationDecision::Join)
}
#[tracing::instrument(name = "bot_on_joined", skip_all, fields(room_id = room.room_id().as_str(), event_id = event.event_id().as_str()))]
async fn on_joined(
&self,
event: AnySyncTimelineEvent,
room: Room,
) -> Result<(), CallbackError> {
if self
.bot
.is_caught_up(event.origin_server_ts())
.await
.map_err(|e| {
CallbackError::Unknown(
format!("Failed to determine catch-up state: {:?}", e).into(),
)
})?
{
tracing::debug!(
event_origin_server_ts = format!("{:?}", event.origin_server_ts()),
"Ignoring old room join event",
);
return Ok(());
}
tracing::info!("Handling room join");
let global_config = self
.bot
.global_config()
.await
.map_err(|e| CallbackError::Unknown(e.into()))?;
let room_config_manager = self.bot.room_config_manager().lock().await;
// We force-create a new config when we join anew to ensure we:
// - always start from a known clean state
// - record the last join timestamp, so we can accurately service the room (ignoring past messages, etc.)
let room_config = room_config_manager
.create_new_for_room(&room)
.await
.map_err(|e| CallbackError::Unknown(e.into()))?;
let room_config_context = RoomConfigContext::new(global_config, room_config);
let event_span = tracing::error_span!("join_controller");
let result = crate::controller::join::handle(&self.bot, &room, &room_config_context)
.instrument(event_span)
.await
.map_err(|e| CallbackError::Unknown(e.into()));
self.bot.catch_up(event.origin_server_ts()).await;
result
}
async fn on_being_last_member(
&self,
_event: AnySyncTimelineEvent,
room: mxlink::matrix_sdk::Room,
) -> Result<(), CallbackError> {
tracing::info!(
"Leaving room {} because we are the last member",
room.room_id()
);
// We are last in this room. Let's just leave
room.leave().await.map_err(|e| e.into())
}
}

View File

@@ -0,0 +1,63 @@
#[cfg(test)]
mod tests;
use super::super::ControllerType;
#[derive(Debug, PartialEq)]
pub enum AccessControllerType {
Help,
GetUsers,
SetUsers(Option<Vec<String>>),
GetRoomLocalAgentManagers,
SetRoomLocalAgentManagers(Option<Vec<String>>),
}
pub fn determine_controller(text: &str) -> ControllerType {
if text.starts_with("users") {
return ControllerType::Access(AccessControllerType::GetUsers);
}
if let Some(patterns_string) = text.strip_prefix("set-users") {
let patterns_string = patterns_string.trim().to_owned();
let patterns_option = if patterns_string.is_empty() {
None
} else {
let patterns_vector = patterns_string
.split(" ")
.map(|s| s.to_string())
.collect::<Vec<String>>();
Some(patterns_vector)
};
return ControllerType::Access(AccessControllerType::SetUsers(patterns_option));
}
if text.starts_with("room-local-agent-managers") {
return ControllerType::Access(AccessControllerType::GetRoomLocalAgentManagers);
}
if let Some(patterns_string) = text.strip_prefix("set-room-local-agent-managers") {
let patterns_string = patterns_string.trim().to_owned();
let patterns_option = if patterns_string.is_empty() {
None
} else {
let patterns_vector = patterns_string
.split(" ")
.map(|s| s.to_string())
.collect::<Vec<String>>();
Some(patterns_vector)
};
return ControllerType::Access(AccessControllerType::SetRoomLocalAgentManagers(
patterns_option,
));
}
ControllerType::Access(AccessControllerType::Help)
}

View File

@@ -0,0 +1,58 @@
#[test]
fn determine_controller() {
struct TestCase {
name: &'static str,
input: &'static str,
expected: super::ControllerType,
}
let test_cases = vec![
TestCase {
name: "Top-level is help",
input: "",
expected: super::ControllerType::Access(super::AccessControllerType::Help),
},
TestCase {
name: "Anything else goes to top-level",
input: "whatever",
expected: super::ControllerType::Access(super::AccessControllerType::Help),
},
TestCase {
name: "Users",
input: "users",
expected: super::ControllerType::Access(super::AccessControllerType::GetUsers),
},
TestCase {
name: "Set-users",
input: "set-users @user:example.com @bot.*:example.org",
expected: super::ControllerType::Access(super::AccessControllerType::SetUsers(Some(
vec![
"@user:example.com".to_owned(),
"@bot.*:example.org".to_owned(),
],
))),
},
TestCase {
name: "Room-local-agent-managers",
input: "room-local-agent-managers",
expected: super::ControllerType::Access(
super::AccessControllerType::GetRoomLocalAgentManagers,
),
},
TestCase {
name: "Set-room-local-agent-managers",
input: "set-room-local-agent-managers @user:example.com @bot.*:example.org",
expected: super::ControllerType::Access(
super::AccessControllerType::SetRoomLocalAgentManagers(Some(vec![
"@user:example.com".to_owned(),
"@bot.*:example.org".to_owned(),
])),
),
},
];
for test_case in test_cases {
let result = super::determine_controller(test_case.input);
assert_eq!(result, test_case.expected, "Test case: {}", test_case.name);
}
}

View File

@@ -0,0 +1,45 @@
use mxlink::MessageResponseType;
use crate::{entity::MessageContext, strings, Bot};
use super::AccessControllerType;
pub async fn dispatch_controller(
handler: &AccessControllerType,
message_context: &MessageContext,
bot: &Bot,
) -> anyhow::Result<()> {
// Only the help command is available without access control, so that all users can get familiar with how the bot's access system works.
match handler {
AccessControllerType::Help => {}
_ => {
if !message_context.sender_can_manage_global_config()? {
bot.messaging()
.send_error_markdown_no_fail(
message_context.room(),
strings::global_config::no_permissions_to_administrate(),
MessageResponseType::Reply(
message_context.thread_info().root_event_id.clone(),
),
)
.await;
return Ok(());
}
}
};
match handler {
AccessControllerType::Help => super::help::handle(bot, message_context).await,
AccessControllerType::GetUsers => super::users::handle_get(bot, message_context).await,
AccessControllerType::SetUsers(patterns) => {
super::users::handle_set(bot, message_context, patterns).await
}
AccessControllerType::GetRoomLocalAgentManagers => {
super::room_local_agent_managers::handle_get(bot, message_context).await
}
AccessControllerType::SetRoomLocalAgentManagers(patterns) => {
super::room_local_agent_managers::handle_set(bot, message_context, patterns).await
}
}
}

View File

@@ -0,0 +1,183 @@
use mxlink::MessageResponseType;
use crate::{entity::MessageContext, strings, Bot};
pub async fn handle(bot: &Bot, message_context: &MessageContext) -> anyhow::Result<()> {
let mut message = String::new();
message.push_str(&build_section_intro());
message.push_str("\n\n");
message.push_str(&build_section_joining_rooms());
message.push_str("\n\n");
message.push_str(&build_section_users(
bot.command_prefix(),
bot.homeserver_name(),
message_context,
));
message.push_str("\n\n");
message.push_str(&build_section_administrators(bot.admin_patterns()));
message.push_str("\n\n");
message.push_str(&build_section_room_local_agent_managers(
bot.command_prefix(),
bot.homeserver_name(),
message_context,
));
bot.messaging()
.send_text_markdown_no_fail(
message_context.room(),
message,
MessageResponseType::Reply(message_context.thread_info().root_event_id.clone()),
)
.await;
Ok(())
}
fn build_section_intro() -> String {
let mut message = String::new();
message.push_str(&format!("## {}", strings::help::access::heading()));
message.push_str("\n\n");
message.push_str(&strings::help::access::intro());
message
}
fn build_section_joining_rooms() -> String {
let mut message = String::new();
message.push_str(&format!(
"### {}",
strings::help::access::room_auto_join_heading()
));
message.push_str("\n\n");
message.push_str(&strings::help::access::room_auto_join_intro());
message.push_str("\n\n");
message
}
fn build_section_users(
command_prefix: &str,
homeserver_name: &str,
message_context: &MessageContext,
) -> String {
let mut message = String::new();
message.push_str(&format!("### {}", strings::help::access::users_heading()));
message.push_str("\n\n");
message.push_str(&strings::help::access::users_intro());
message.push('\n');
message.push_str(&strings::help::access::users_access());
message.push_str("\n\n");
if let Some(user_patterns) = &message_context.global_config().access.user_patterns {
if user_patterns.is_empty() {
message.push_str(&strings::access::users_no_patterns());
} else {
message.push_str(&strings::access::users_now_match_patterns(user_patterns));
}
} else {
message.push_str(&strings::access::users_no_patterns());
}
let can_manage_global_config = message_context.sender_can_manage_global_config();
if let Ok(can_manage_global_config) = can_manage_global_config {
if can_manage_global_config {
message.push_str("\n\n");
message.push_str(strings::the_following_commands_are_available());
message.push('\n');
message.push_str(&strings::help::access::users_command_get(command_prefix));
message.push('\n');
message.push_str(&strings::help::access::users_command_set(command_prefix));
message.push_str("\n\n");
message.push_str(&strings::help::access::example_user_patterns(
homeserver_name,
));
}
}
message
}
fn build_section_administrators(admin_patterns: &[String]) -> String {
let mut message = String::new();
message.push_str(&format!(
"### {}",
strings::help::access::administrators_heading()
));
message.push_str("\n\n");
message.push_str(&strings::help::access::administrators_intro());
message.push_str("\n\n");
message.push_str(&strings::help::access::administrators_now_match_patterns(
admin_patterns,
));
message.push_str("\n\n");
message.push_str(&strings::help::access::administrators_outro());
message
}
fn build_section_room_local_agent_managers(
command_prefix: &str,
homeserver_name: &str,
message_context: &MessageContext,
) -> String {
let mut message = String::new();
message.push_str(&format!(
"### {}",
strings::help::access::room_local_agent_managers_heading()
));
message.push_str("\n\n");
message.push_str(&strings::help::access::room_local_agent_managers_intro(
command_prefix,
));
message.push('\n');
message.push_str(&strings::help::access::room_local_agent_managers_security_warning());
message.push_str("\n\n");
if let Some(user_patterns) = &message_context
.global_config()
.access
.room_local_agent_manager_patterns
{
if user_patterns.is_empty() {
message.push_str(&strings::access::room_local_agent_managers_no_patterns());
} else {
message.push_str(
&strings::access::room_local_agent_managers_now_match_patterns(user_patterns),
);
}
} else {
message.push_str(&strings::access::room_local_agent_managers_no_patterns());
}
let can_manage_global_config = message_context.sender_can_manage_global_config();
if let Ok(can_manage_global_config) = can_manage_global_config {
if can_manage_global_config {
message.push_str("\n\n");
message.push_str(strings::the_following_commands_are_available());
message.push('\n');
message.push_str(
&strings::help::access::room_local_agent_managers_command_get(command_prefix),
);
message.push('\n');
message.push_str(
&strings::help::access::room_local_agent_managers_command_set(command_prefix),
);
message.push_str("\n\n");
message.push_str(&strings::help::access::example_user_patterns(
homeserver_name,
));
}
}
message
}

View File

@@ -0,0 +1,8 @@
mod determination;
mod dispatching;
pub mod help;
mod room_local_agent_managers;
mod users;
pub use determination::{determine_controller, AccessControllerType};
pub use dispatching::dispatch_controller;

View File

@@ -0,0 +1,67 @@
use mxlink::MessageResponseType;
use crate::{entity::MessageContext, strings, Bot};
pub async fn handle_get(bot: &Bot, message_context: &MessageContext) -> anyhow::Result<()> {
let message = match &message_context
.global_config()
.access
.room_local_agent_manager_patterns
{
Some(patterns) => strings::access::room_local_agent_managers_now_match_patterns(patterns),
None => strings::access::room_local_agent_managers_no_patterns(),
};
bot.messaging()
.send_text_markdown_no_fail(
message_context.room(),
message,
MessageResponseType::Reply(message_context.thread_info().root_event_id.clone()),
)
.await;
Ok(())
}
pub async fn handle_set(
bot: &Bot,
message_context: &MessageContext,
patterns: &Option<Vec<String>>,
) -> anyhow::Result<()> {
if let Some(patterns) = patterns {
if let Err(err) = mxidwc::parse_patterns_vector(patterns) {
bot.messaging()
.send_error_markdown_no_fail(
message_context.room(),
&strings::access::failed_to_parse_patterns(&err.to_string()),
MessageResponseType::Reply(message_context.thread_info().root_event_id.clone()),
)
.await;
return Ok(());
}
}
let mut global_config_manager_guard = bot.global_config_manager().lock().await;
let mut global_config = global_config_manager_guard.get_or_create().await?;
global_config.access.room_local_agent_manager_patterns = patterns.clone();
global_config_manager_guard.persist(&global_config).await?;
let message = match patterns {
Some(patterns) => strings::access::room_local_agent_managers_now_match_patterns(patterns),
None => strings::access::room_local_agent_managers_no_patterns(),
};
bot.messaging()
.send_success_markdown_no_fail(
message_context.room(),
&message,
MessageResponseType::Reply(message_context.thread_info().root_event_id.clone()),
)
.await;
Ok(())
}

View File

@@ -0,0 +1,63 @@
use mxlink::MessageResponseType;
use crate::{entity::MessageContext, strings, Bot};
pub async fn handle_get(bot: &Bot, message_context: &MessageContext) -> anyhow::Result<()> {
let message = match &message_context.global_config().access.user_patterns {
Some(patterns) => strings::access::users_now_match_patterns(patterns),
None => strings::access::users_no_patterns(),
};
bot.messaging()
.send_text_markdown_no_fail(
message_context.room(),
message,
MessageResponseType::Reply(message_context.thread_info().root_event_id.clone()),
)
.await;
Ok(())
}
pub async fn handle_set(
bot: &Bot,
message_context: &MessageContext,
patterns: &Option<Vec<String>>,
) -> anyhow::Result<()> {
if let Some(patterns) = patterns {
if let Err(err) = mxidwc::parse_patterns_vector(patterns) {
bot.messaging()
.send_error_markdown_no_fail(
message_context.room(),
&strings::access::failed_to_parse_patterns(&err.to_string()),
MessageResponseType::Reply(message_context.thread_info().root_event_id.clone()),
)
.await;
return Ok(());
}
}
let mut global_config_manager_guard = bot.global_config_manager().lock().await;
let mut global_config = global_config_manager_guard.get_or_create().await?;
global_config.access.user_patterns = patterns.clone();
global_config_manager_guard.persist(&global_config).await?;
let message = match patterns {
Some(patterns) => strings::access::users_now_match_patterns(patterns),
None => strings::access::users_no_patterns(),
};
bot.messaging()
.send_success_markdown_no_fail(
message_context.room(),
&message,
MessageResponseType::Reply(message_context.thread_info().root_event_id.clone()),
)
.await;
Ok(())
}

View File

@@ -0,0 +1,418 @@
#[cfg(test)]
mod tests;
use mxlink::MessageResponseType;
use crate::agent::provider::{ControllerTrait, PingResult};
use crate::agent::PublicIdentifier;
use crate::agent::{create_from_provider_and_yaml_value_config, AgentDefinition};
use crate::agent::{AgentInstance, AgentProvider};
use crate::controller::utils::get_text_body_or_complain;
use crate::entity::globalconfig::GlobalConfigurationManager;
use crate::entity::roomconfig::RoomConfigurationManager;
use crate::strings;
use crate::{entity::MessageContext, Bot};
struct ParsedAgentConfig {
agent: AgentInstance,
config: serde_yaml::Value,
}
pub async fn handle_room_local(
bot: &Bot,
room_config_manager: &tokio::sync::Mutex<RoomConfigurationManager>,
message_context: &MessageContext,
provider: &str,
agent_id_prefixless: &str,
) -> anyhow::Result<()> {
if !message_context.sender_can_manage_room_local_agents()? {
bot.messaging()
.send_error_markdown_no_fail(
message_context.room(),
&strings::agent::not_allowed_to_manage_room_local_agents_in_room(),
MessageResponseType::Reply(message_context.thread_info().root_event_id.clone()),
)
.await;
return Ok(());
}
let Ok(provider) = AgentProvider::from_string(provider) else {
bot.messaging()
.send_error_markdown_no_fail(
message_context.room(),
&strings::provider::invalid(provider),
MessageResponseType::Reply(message_context.thread_info().root_event_id.clone()),
)
.await;
return Ok(());
};
let agent_identifier = PublicIdentifier::DynamicRoomLocal(agent_id_prefixless.to_owned());
if let Err(err) = agent_identifier.validate() {
bot.messaging()
.send_error_markdown_no_fail(
message_context.room(),
&strings::agent::invalid_id_validation_error(err),
MessageResponseType::Reply(message_context.thread_info().root_event_id.clone()),
)
.await;
return Ok(());
}
let agent_exists = bot
.agent_manager()
.available_room_agents_by_room_config_context(message_context.room_config_context())
.iter()
.any(|agent| *agent.identifier() == agent_identifier);
if agent_exists {
bot.messaging()
.send_error_markdown_no_fail(
message_context.room(),
&strings::agent::already_exists_see_help(agent_id_prefixless, bot.command_prefix()),
MessageResponseType::InThread(message_context.thread_info().clone()),
)
.await;
return Ok(());
}
if message_context.thread_info().is_thread_root_only() {
return send_guide(bot, message_context, &agent_identifier, &provider).await;
}
let Some(text_message_content) = get_text_body_or_complain(bot, message_context).await else {
return Ok(());
};
let parsed_config = parse_agent_config_from_message_or_complain(
bot,
message_context,
&provider,
&agent_identifier,
text_message_content,
)
.await;
let Some(parsed_config) = parsed_config else {
return Ok(());
};
message_context.room().typing_notice(true).await?;
if !try_to_ping_agent_or_complain(bot, message_context, &parsed_config.agent).await {
return Ok(());
}
let agent_definition = AgentDefinition::new(
agent_identifier.prefixless(),
provider,
parsed_config.config.clone(),
);
let mut room_config = message_context.room_config().clone();
room_config.agents.push(agent_definition.clone());
room_config_manager
.lock()
.await
.persist(message_context.room(), &room_config)
.await?;
send_completion_wrap_up(
bot,
message_context,
&agent_identifier,
&parsed_config.agent,
)
.await;
Ok(())
}
pub async fn handle_global(
bot: &Bot,
global_config_manager: &tokio::sync::Mutex<GlobalConfigurationManager>,
message_context: &MessageContext,
provider: &str,
agent_id_prefixless: &str,
) -> anyhow::Result<()> {
if !message_context.sender_can_manage_global_config()? {
bot.messaging()
.send_error_markdown_no_fail(
message_context.room(),
strings::global_config::no_permissions_to_administrate(),
MessageResponseType::Reply(message_context.thread_info().root_event_id.clone()),
)
.await;
return Ok(());
}
let Ok(provider) = AgentProvider::from_string(provider) else {
bot.messaging()
.send_error_markdown_no_fail(
message_context.room(),
&strings::provider::invalid(provider),
MessageResponseType::Reply(message_context.thread_info().root_event_id.clone()),
)
.await;
return Ok(());
};
let agent_identifier = PublicIdentifier::DynamicGlobal(agent_id_prefixless.to_owned());
if let Err(err) = agent_identifier.validate() {
bot.messaging()
.send_error_markdown_no_fail(
message_context.room(),
&strings::agent::invalid_id_validation_error(err),
MessageResponseType::Reply(message_context.thread_info().root_event_id.clone()),
)
.await;
return Ok(());
}
let agent_exists = bot
.agent_manager()
.available_room_agents_by_room_config_context(message_context.room_config_context())
.iter()
.any(|agent| *agent.identifier() == agent_identifier);
if agent_exists {
bot.messaging()
.send_error_markdown_no_fail(
message_context.room(),
&strings::agent::already_exists_see_help(agent_id_prefixless, bot.command_prefix()),
MessageResponseType::InThread(message_context.thread_info().clone()),
)
.await;
return Ok(());
}
if message_context.thread_info().is_thread_root_only() {
return send_guide(bot, message_context, &agent_identifier, &provider).await;
}
let Some(text_message_content) = get_text_body_or_complain(bot, message_context).await else {
return Ok(());
};
let parsed_config = parse_agent_config_from_message_or_complain(
bot,
message_context,
&provider,
&agent_identifier,
text_message_content,
)
.await;
let Some(parsed_config) = parsed_config else {
return Ok(());
};
message_context.room().typing_notice(true).await?;
if !try_to_ping_agent_or_complain(bot, message_context, &parsed_config.agent).await {
return Ok(());
}
let agent_definition = AgentDefinition::new(
agent_identifier.prefixless(),
provider,
parsed_config.config.clone(),
);
let mut global_config = message_context.global_config().clone();
global_config.agents.push(agent_definition.clone());
global_config_manager
.lock()
.await
.persist(&global_config)
.await?;
send_completion_wrap_up(
bot,
message_context,
&agent_identifier,
&parsed_config.agent,
)
.await;
Ok(())
}
async fn send_guide(
bot: &Bot,
message_context: &MessageContext,
agent_identifier: &PublicIdentifier,
provider: &AgentProvider,
) -> anyhow::Result<()> {
let sample_config = crate::agent::default_config_for_provider(provider);
let sample_config_pretty_yaml = serde_yaml::to_string(&sample_config)?;
bot.messaging()
.send_text_markdown_no_fail(
message_context.room(),
strings::agent::creation_guide(agent_identifier, provider, &sample_config_pretty_yaml),
MessageResponseType::InThread(message_context.thread_info().clone()),
)
.await;
Ok(())
}
fn parse_from_message_to_yaml_value(text: &str) -> Result<serde_yaml::Value, String> {
let mut text = text.trim();
if text.starts_with("```") {
// Try to strip ```yml and ```yaml first and fall back to the generic ``` later.
text = text.trim_start_matches("```yml");
text = text.trim_start_matches("```yaml");
text = text.trim_start_matches("```");
text = text.trim_end_matches("```");
}
let config: serde_yaml::Value = serde_yaml::from_str(text).map_err(|e| e.to_string())?;
match config {
serde_yaml::Value::Mapping(_) => {}
_ => {
return Err("Not a valid YAML hashmap".to_owned());
}
};
Ok(config)
}
async fn parse_agent_config_from_message_or_complain(
bot: &Bot,
message_context: &MessageContext,
provider: &AgentProvider,
agent_identifier: &PublicIdentifier,
text: &str,
) -> Option<ParsedAgentConfig> {
let config_yaml_value = parse_from_message_to_yaml_value(text);
let config_yaml_value = match config_yaml_value {
Ok(config_yaml_value) => config_yaml_value,
Err(err) => {
bot.messaging()
.send_error_markdown_no_fail(
message_context.room(),
&strings::agent::configuration_not_a_valid_yaml_hashmap(err),
MessageResponseType::InThread(message_context.thread_info().clone()),
)
.await;
return None;
}
};
let agent = create_from_provider_and_yaml_value_config(
provider,
agent_identifier,
config_yaml_value.clone(),
);
let agent = match agent {
Ok(agent) => ParsedAgentConfig {
agent,
config: config_yaml_value,
},
Err(err) => {
bot.messaging()
.send_error_markdown_no_fail(
message_context.room(),
&strings::provider::invalid_configuration_for_provider(provider, err),
MessageResponseType::InThread(message_context.thread_info().clone()),
)
.await;
return None;
}
};
Some(agent)
}
async fn try_to_ping_agent_or_complain(
bot: &Bot,
message_context: &MessageContext,
agent_instance: &AgentInstance,
) -> bool {
bot.messaging()
.send_notice_markdown_no_fail(
message_context.room(),
format!("⏳ {}", strings::agent::configuration_agent_will_ping()),
MessageResponseType::InThread(message_context.thread_info().clone()),
)
.await;
match agent_instance.controller().ping().await {
Ok(ping_result) => {
let message = match ping_result {
PingResult::Inconclusive => format!(
"❓ {}",
strings::agent::configuration_agent_ping_inconclusive()
),
PingResult::Successful => {
format!("✅ {}", strings::agent::configuration_agent_ping_ok())
}
};
bot.messaging()
.send_notice_markdown_no_fail(
message_context.room(),
message,
MessageResponseType::InThread(message_context.thread_info().clone()),
)
.await;
true
}
Err(err) => {
bot.messaging()
.send_error_markdown_no_fail(
message_context.room(),
&strings::agent::configuration_does_not_result_in_a_working_agent(err),
MessageResponseType::InThread(message_context.thread_info().clone()),
)
.await;
false
}
}
}
async fn send_completion_wrap_up(
bot: &Bot,
message_context: &MessageContext,
agent_identifier: &PublicIdentifier,
agent_instance: &AgentInstance,
) {
bot.messaging()
.send_success_markdown_no_fail(
message_context.room(),
&strings::agent::created(agent_identifier),
MessageResponseType::InThread(message_context.thread_info().clone()),
)
.await;
bot.messaging()
.send_tooltip_markdown_no_fail(
message_context.room(),
&strings::agent::post_creation_helpful_commands(
agent_identifier,
agent_instance,
bot.command_prefix(),
),
MessageResponseType::InThread(message_context.thread_info().clone()),
)
.await;
}

View File

@@ -0,0 +1,61 @@
#[test]
fn agent_config_parsing_works() {
struct TestCase {
input: String,
expected: Option<serde_yaml::Value>,
}
let provider = crate::agent::AgentProvider::OpenAI;
let sample_config = crate::agent::default_config_for_provider(&provider);
let sample_config_pretty_yaml = serde_yaml::to_string(&sample_config).unwrap();
let test_cases = vec![
// Invalid input
TestCase {
input: r#"Hello"#.to_owned(),
expected: None,
},
// Plain text
TestCase {
input: sample_config_pretty_yaml.clone(),
expected: Some(sample_config.clone()),
},
// Generic code block
TestCase {
input: format!("```\n{}```", sample_config_pretty_yaml),
expected: Some(sample_config.clone()),
},
// YAML code block (yaml)
TestCase {
input: format!("```yaml\n{}```", sample_config_pretty_yaml),
expected: Some(sample_config.clone()),
},
// YAML code block (yml)
TestCase {
input: format!("```yml\n{}```", sample_config_pretty_yaml),
expected: Some(sample_config.clone()),
},
// JSON code block
TestCase {
input: format!("```json\n{}```", sample_config_pretty_yaml),
expected: None,
},
];
for (i, test_case) in test_cases.iter().enumerate() {
let result = super::parse_from_message_to_yaml_value(&test_case.input);
match result {
Ok(config) => {
assert_eq!(
config,
test_case.expected.clone().unwrap(),
"Test case {} failed",
i
);
}
Err(_) => {
assert_eq!(test_case.expected, None, "Test case {} failed", i);
}
}
}
}

View File

@@ -0,0 +1,195 @@
use mxlink::MessageResponseType;
use crate::entity::{
globalconfig::GlobalConfigurationManager, roomconfig::RoomConfigurationManager, MessageContext,
};
use crate::{agent::PublicIdentifier, strings, Bot};
pub async fn handle(
bot: &Bot,
room_config_manager: &tokio::sync::Mutex<RoomConfigurationManager>,
global_config_manager: &tokio::sync::Mutex<GlobalConfigurationManager>,
message_context: &MessageContext,
agent_identifier: &PublicIdentifier,
) -> anyhow::Result<()> {
let agents = bot
.agent_manager()
.available_room_agents_by_room_config_context(message_context.room_config_context());
let agent = agents.iter().find(|a| a.identifier() == agent_identifier);
let Some(_) = agent else {
bot.messaging()
.send_error_markdown_no_fail(
message_context.room(),
&strings::agent::agent_with_given_identifier_missing(agent_identifier),
MessageResponseType::Reply(message_context.thread_info().root_event_id.clone()),
)
.await;
return Ok(());
};
match &agent_identifier {
PublicIdentifier::DynamicRoomLocal(_) => {
if !message_context.sender_can_manage_room_local_agents()? {
bot.messaging()
.send_error_markdown_no_fail(
message_context.room(),
&strings::agent::not_allowed_to_manage_room_local_agents_in_room(),
MessageResponseType::Reply(
message_context.thread_info().root_event_id.clone(),
),
)
.await;
return Ok(());
}
delete_room_local_agent(bot, room_config_manager, message_context, agent_identifier)
.await
}
PublicIdentifier::DynamicGlobal(_) => {
if !message_context.sender_can_manage_global_config()? {
bot.messaging()
.send_error_markdown_no_fail(
message_context.room(),
strings::global_config::no_permissions_to_administrate(),
MessageResponseType::Reply(
message_context.thread_info().root_event_id.clone(),
),
)
.await;
return Ok(());
}
delete_global_agent(
bot,
global_config_manager,
message_context,
agent_identifier,
)
.await
}
PublicIdentifier::Static(_) => {
bot.messaging()
.send_error_markdown_no_fail(
message_context.room(),
&strings::agent::not_allowed_to_manage_static_agents(),
MessageResponseType::Reply(message_context.thread_info().root_event_id.clone()),
)
.await;
Ok(())
}
}
}
async fn delete_room_local_agent(
bot: &Bot,
room_config_manager: &tokio::sync::Mutex<RoomConfigurationManager>,
message_context: &MessageContext,
agent_id: &PublicIdentifier,
) -> anyhow::Result<()> {
let mut room_config = message_context.room_config().clone();
let mut was_deleted = false;
let agent_id_prefixless = agent_id.prefixless();
let mut agents = Vec::new();
for agent_config in room_config.agents {
if agent_config.id == agent_id_prefixless {
was_deleted = true;
} else {
agents.push(agent_config.clone());
}
}
if !was_deleted {
bot.messaging()
.send_error_markdown_no_fail(
message_context.room(),
&strings::agent::agent_with_given_identifier_missing(agent_id),
MessageResponseType::Reply(message_context.thread_info().root_event_id.clone()),
)
.await;
return Ok(());
}
room_config.agents = agents;
let room_config_manager = room_config_manager.lock().await;
// We may unset all handlers in the room config which refer to this agent.
// We intentionally do not do this, because we do not support "agent edit" yet and ask people to do "agent delete" and "agent create" instead.
// We'd rather not magically reconfigure the room on agent deletion and obstruct this use case.
room_config_manager
.persist(message_context.room(), &room_config)
.await?;
bot.messaging()
.send_success_markdown_no_fail(
message_context.room(),
&strings::agent::removed_room_local(agent_id, bot.command_prefix()),
MessageResponseType::Reply(message_context.thread_info().root_event_id.clone()),
)
.await;
Ok(())
}
async fn delete_global_agent(
bot: &Bot,
global_config_manager: &tokio::sync::Mutex<GlobalConfigurationManager>,
message_context: &MessageContext,
agent_id: &PublicIdentifier,
) -> anyhow::Result<()> {
let mut global_config = message_context.global_config().clone();
let mut was_deleted = false;
let agent_id_prefixless = agent_id.prefixless();
let mut agents = Vec::new();
for agent_config in global_config.agents {
if agent_config.id == agent_id_prefixless {
was_deleted = true;
} else {
agents.push(agent_config.clone());
}
}
if !was_deleted {
bot.messaging()
.send_error_markdown_no_fail(
message_context.room(),
&strings::agent::agent_with_given_identifier_missing(agent_id),
MessageResponseType::Reply(message_context.thread_info().root_event_id.clone()),
)
.await;
return Ok(());
}
global_config.agents = agents;
global_config_manager
.lock()
.await
.persist(&global_config)
.await?;
bot.messaging()
.send_success_markdown_no_fail(
message_context.room(),
&strings::agent::removed_global(agent_id, bot.command_prefix()),
MessageResponseType::Reply(message_context.thread_info().root_event_id.clone()),
)
.await;
Ok(())
}

View File

@@ -0,0 +1,84 @@
use mxlink::MessageResponseType;
use crate::{agent::PublicIdentifier, entity::MessageContext, strings, Bot};
pub async fn handle(
bot: &Bot,
message_context: &MessageContext,
agent_identifier: &PublicIdentifier,
) -> anyhow::Result<()> {
let agents = bot
.agent_manager()
.available_room_agents_by_room_config_context(message_context.room_config_context());
let agent = agents.iter().find(|a| a.identifier() == agent_identifier);
let agent = match agent {
Some(agent) => agent,
None => {
bot.messaging()
.send_error_markdown_no_fail(
message_context.room(),
&strings::agent::agent_with_given_identifier_missing(agent_identifier),
MessageResponseType::Reply(message_context.thread_info().root_event_id.clone()),
)
.await;
return Ok(());
}
};
// Access checks
match &agent_identifier {
PublicIdentifier::DynamicRoomLocal(_) => {
if !message_context.sender_can_manage_room_local_agents()? {
bot.messaging()
.send_error_markdown_no_fail(
message_context.room(),
&strings::agent::not_allowed_to_manage_room_local_agents_in_room(),
MessageResponseType::Reply(
message_context.thread_info().root_event_id.clone(),
),
)
.await;
return Ok(());
}
}
PublicIdentifier::DynamicGlobal(_) => {
if !message_context.sender_can_manage_global_config()? {
bot.messaging()
.send_error_markdown_no_fail(
message_context.room(),
strings::global_config::no_permissions_to_administrate(),
MessageResponseType::Reply(
message_context.thread_info().root_event_id.clone(),
),
)
.await;
return Ok(());
}
}
PublicIdentifier::Static(_) => {}
};
let config_yaml_pretty = serde_yaml::to_string(&agent.definition().config)?;
bot.messaging()
.send_text_markdown_no_fail(
message_context.room(),
format!(
"Configuration for agent `{}` (powered by the `{}` provider):\n```yml\n{}\n```",
agent_identifier,
agent.definition().provider.to_static_str(),
config_yaml_pretty.trim(),
)
.to_owned(),
MessageResponseType::Reply(message_context.thread_info().root_event_id.clone()),
)
.await;
Ok(())
}

View File

@@ -0,0 +1,99 @@
#[cfg(test)]
mod tests;
use crate::{agent::PublicIdentifier, controller::ControllerType, strings};
#[derive(Debug, PartialEq)]
pub enum AgentControllerType {
List,
Details(PublicIdentifier),
CreateRoomLocal { provider: String, agent_id: String },
CreateGlobal { provider: String, agent_id: String },
Delete(PublicIdentifier),
Help,
}
pub fn determine_controller(command_prefix: &str, text: &str) -> ControllerType {
if text.starts_with("list") {
return ControllerType::Agent(AgentControllerType::List);
}
if let Some(agent_id_string) = text.strip_prefix("details") {
let agent_id_string = agent_id_string.trim();
if agent_id_string.is_empty() || agent_id_string.contains(" ") {
return ControllerType::Error(
strings::agent::incorrect_invocation_expects_agent_id_arg(command_prefix),
);
}
let Some(agent_identifier) = PublicIdentifier::from_str(agent_id_string) else {
return ControllerType::Error(strings::agent::invalid_id_generic());
};
return ControllerType::Agent(AgentControllerType::Details(agent_identifier));
}
if let Some(remaining_text) = text.strip_prefix("create-room-local") {
// `remaining_text` should be something like: `PROVIDER ID`
let remaining_text = remaining_text.trim();
let parts = remaining_text.split_once(' ');
let Some((provider, agent_id_string)) = parts else {
return ControllerType::Error(strings::agent::incorrect_creation_invocation(
command_prefix,
));
};
if agent_id_string.contains(" ") {
return ControllerType::Error(strings::agent::incorrect_creation_invocation(
command_prefix,
));
}
return ControllerType::Agent(AgentControllerType::CreateRoomLocal {
provider: provider.to_owned(),
agent_id: agent_id_string.trim().to_owned(),
});
}
if let Some(remaining_text) = text.strip_prefix("create-global") {
// `remaining_text` should be something like: `PROVIDER ID`
let remaining_text = remaining_text.trim();
let parts = remaining_text.split_once(' ');
let Some((provider, agent_id_string)) = parts else {
return ControllerType::Error(strings::agent::incorrect_creation_invocation(
command_prefix,
));
};
if agent_id_string.contains(" ") {
return ControllerType::Error(strings::agent::incorrect_creation_invocation(
command_prefix,
));
}
return ControllerType::Agent(AgentControllerType::CreateGlobal {
provider: provider.to_owned(),
agent_id: agent_id_string.trim().to_owned(),
});
}
if let Some(agent_id_string) = text.strip_prefix("delete") {
let agent_id_string = agent_id_string.trim();
if agent_id_string.is_empty() || agent_id_string.contains(" ") {
return ControllerType::Error(
strings::agent::incorrect_invocation_expects_agent_id_arg(command_prefix),
);
}
let Some(agent_identifier) = PublicIdentifier::from_str(agent_id_string) else {
return ControllerType::Error(strings::agent::invalid_id_generic());
};
return ControllerType::Agent(AgentControllerType::Delete(agent_identifier));
}
ControllerType::Agent(AgentControllerType::Help)
}

View File

@@ -0,0 +1,131 @@
#[test]
fn determine_controller() {
use crate::agent::PublicIdentifier;
struct TestCase {
name: &'static str,
input: &'static str,
expected: super::ControllerType,
}
let command_prefix = "!bai";
let test_cases = vec![
TestCase {
name: "Top-level is help",
input: "",
expected: super::ControllerType::Agent(super::AgentControllerType::Help),
},
TestCase {
name: "Anything else goes to top-level",
input: "whatever",
expected: super::ControllerType::Agent(super::AgentControllerType::Help),
},
TestCase {
name: "List",
input: "list",
expected: super::ControllerType::Agent(super::AgentControllerType::List),
},
TestCase {
name: "details",
input: "details static/agent-id",
expected: super::ControllerType::Agent(super::AgentControllerType::Details(
PublicIdentifier::Static("agent-id".to_owned()),
)),
},
TestCase {
name: "details with invalid agent identifier",
input: "details agent-id",
expected: super::ControllerType::Error(crate::strings::agent::invalid_id_generic()),
},
TestCase {
name: "create-room-local no arguments",
input: "create-room-local",
expected: super::ControllerType::Error(
crate::strings::agent::incorrect_creation_invocation(command_prefix),
),
},
TestCase {
name: "create-room-local only with provider",
input: "create-room-local openai",
expected: super::ControllerType::Error(
crate::strings::agent::incorrect_creation_invocation(command_prefix),
),
},
TestCase {
name: "create-room-local correct",
input: "create-room-local openai my-agent-id",
expected: super::ControllerType::Agent(super::AgentControllerType::CreateRoomLocal {
provider: "openai".to_owned(),
agent_id: "my-agent-id".trim().to_owned(),
}),
},
TestCase {
name: "create-global extra arguments",
input: "create-global openai my-agent-id more arguments here",
expected: super::ControllerType::Error(
crate::strings::agent::incorrect_creation_invocation(command_prefix),
),
},
TestCase {
name: "create-global no arguments",
input: "create-global",
expected: super::ControllerType::Error(
crate::strings::agent::incorrect_creation_invocation(command_prefix),
),
},
TestCase {
name: "create-global only with provider",
input: "create-global openai",
expected: super::ControllerType::Error(
crate::strings::agent::incorrect_creation_invocation(command_prefix),
),
},
TestCase {
name: "create-global correct",
input: "create-global openai my-agent-id",
expected: super::ControllerType::Agent(super::AgentControllerType::CreateGlobal {
provider: "openai".to_owned(),
agent_id: "my-agent-id".trim().to_owned(),
}),
},
TestCase {
name: "create-global extra arguments",
input: "create-global openai my-agent-id more arguments here",
expected: super::ControllerType::Error(
crate::strings::agent::incorrect_creation_invocation(command_prefix),
),
},
TestCase {
name: "delete no arguments",
input: "delete",
expected: super::ControllerType::Error(
crate::strings::agent::incorrect_invocation_expects_agent_id_arg(command_prefix),
),
},
TestCase {
name: "delete too many arguments",
input: "delete agent-id extra arguments",
expected: super::ControllerType::Error(
crate::strings::agent::incorrect_invocation_expects_agent_id_arg(command_prefix),
),
},
TestCase {
name: "delete",
input: "delete static/agent-id",
expected: super::ControllerType::Agent(super::AgentControllerType::Delete(
PublicIdentifier::Static("agent-id".to_owned()),
)),
},
TestCase {
name: "delete with invalid agent identifier",
input: "delete agent-id",
expected: super::ControllerType::Error(crate::strings::agent::invalid_id_generic()),
},
];
for test_case in test_cases {
let result = super::determine_controller(command_prefix, test_case.input);
assert_eq!(result, test_case.expected, "Test case: {}", test_case.name);
}
}

View File

@@ -0,0 +1,79 @@
use mxlink::MessageResponseType;
use crate::{entity::MessageContext, strings, Bot};
pub async fn handle(bot: &Bot, message_context: &MessageContext) -> anyhow::Result<()> {
// Anyone can access this help command, because certain subcommands ("list")
// are also useful to regular users and it'd be great for them to learn about them.
let mut message = String::new();
let can_manage_agents = message_context.sender_can_manage_room_local_agents()?;
message.push_str(&format!("## {}", strings::help::agent::heading()));
message.push_str("\n\n");
message.push_str(&strings::help::agent::intro(
bot.command_prefix(),
can_manage_agents,
));
message.push('\n');
message.push_str(&strings::help::agent::intro_capabilities());
message.push_str("\n\n");
message.push_str(&strings::help::agent::intro_handler_relation(
bot.command_prefix(),
));
if can_manage_agents {
message.push_str("\n\n");
message.push_str(strings::help::available_commands_intro());
message.push('\n');
message.push_str(&strings::help::agent::list_agents(bot.command_prefix()));
message.push('\n');
message.push_str(strings::help::agent::create_agent_intro());
message.push('\n');
message.push_str(&strings::help::agent::create_agent_room_local(
bot.command_prefix(),
));
message.push('\n');
if message_context.sender_can_manage_global_config()? {
message.push_str(&strings::help::agent::create_agent_global(
bot.command_prefix(),
));
message.push('\n');
}
message.push_str(&strings::help::agent::create_agent_example(
bot.command_prefix(),
));
message.push('\n');
message.push_str(&strings::help::agent::show_agent_details(
bot.command_prefix(),
));
message.push('\n');
message.push_str(&strings::help::agent::delete_agent(bot.command_prefix()));
message.push_str("\n\n");
message.push_str(strings::help::agent::available_commands_outro_update_note());
} else {
message.push_str("\n\n");
message.push_str(strings::help::agent::no_permission_to_create_agents());
}
bot.messaging()
.send_text_markdown_no_fail(
message_context.room(),
message,
MessageResponseType::Reply(message_context.thread_info().root_event_id.clone()),
)
.await;
Ok(())
}

View File

@@ -0,0 +1,39 @@
use mxlink::MessageResponseType;
use crate::agent::AgentPurpose;
use crate::strings;
use crate::{entity::MessageContext, Bot};
pub async fn handle(bot: &Bot, message_context: &MessageContext) -> anyhow::Result<()> {
let agents = bot
.agent_manager()
.available_room_agents_by_room_config_context(message_context.room_config_context());
let mut message = String::new();
if agents.is_empty() {
message.push_str(strings::agent::agent_list_empty().as_str());
} else {
message.push_str(&strings::agent::non_empty_agent_list_block(&agents));
message.push_str("\n\n");
message.push_str(strings::agent::agent_list_legend_intro().as_str());
for purpose in AgentPurpose::choices() {
message.push_str(&format!(
"\n- {} `{}` ({})",
purpose.emoji(),
purpose.as_str(),
strings::agent::purpose_howto(purpose),
));
}
}
bot.messaging()
.send_text_markdown_no_fail(
message_context.room(),
message,
MessageResponseType::Reply(message_context.thread_info().root_event_id.clone()),
)
.await;
Ok(())
}

View File

@@ -0,0 +1,54 @@
use crate::{entity::MessageContext, Bot};
pub mod create;
pub mod delete;
pub mod details;
pub mod determination;
pub mod help;
pub mod list;
pub use determination::{determine_controller, AgentControllerType};
pub async fn dispatch_controller(
handler: &AgentControllerType,
message_context: &MessageContext,
bot: &Bot,
) -> anyhow::Result<()> {
match handler {
AgentControllerType::CreateRoomLocal { provider, agent_id } => {
create::handle_room_local(
bot,
bot.room_config_manager(),
message_context,
provider,
agent_id,
)
.await
}
AgentControllerType::CreateGlobal { provider, agent_id } => {
create::handle_global(
bot,
bot.global_config_manager(),
message_context,
provider,
agent_id,
)
.await
}
AgentControllerType::List => list::handle(bot, message_context).await,
AgentControllerType::Details(agent_identifier) => {
details::handle(bot, message_context, agent_identifier).await
}
AgentControllerType::Delete(agent_identifier) => {
delete::handle(
bot,
bot.room_config_manager(),
bot.global_config_manager(),
message_context,
agent_identifier,
)
.await
}
AgentControllerType::Help => help::handle(bot, message_context).await,
}
}

View File

@@ -0,0 +1,35 @@
use mxlink::MessageResponseType;
use crate::{entity::MessageContext, strings, Bot};
pub async fn handle_get<T>(
bot: &Bot,
message_context: &MessageContext,
value: &Option<T>,
) -> anyhow::Result<()>
where
T: std::fmt::Display,
{
match value {
Some(value) => {
bot.messaging()
.send_text_markdown_no_fail(
message_context.room(),
strings::cfg::value_currently_set_to(value),
MessageResponseType::Reply(message_context.thread_info().root_event_id.clone()),
)
.await;
}
None => {
bot.messaging()
.send_text_markdown_no_fail(
message_context.room(),
strings::cfg::value_currently_unset(),
MessageResponseType::Reply(message_context.thread_info().root_event_id.clone()),
)
.await;
}
}
Ok(())
}

View File

@@ -0,0 +1 @@
pub(super) mod generic_setting;

View File

@@ -0,0 +1,74 @@
use crate::{
agent::{AgentPurpose, PublicIdentifier},
entity::roomconfig::{
SpeechToTextFlowType, TextGenerationAutoUsage, TextGenerationPrefixRequirementType,
TextToSpeechBotMessagesFlowType, TextToSpeechUserMessagesFlowType,
},
};
#[derive(Debug, PartialEq)]
pub enum SettingsStorageSource {
Room,
Global,
}
#[derive(Debug, PartialEq)]
pub enum ConfigControllerType {
Help,
Status,
SettingsRelated(SettingsStorageSource, ConfigSettingRelatedControllerType),
}
#[derive(Debug, PartialEq)]
pub enum ConfigSettingRelatedControllerType {
GetHandler(AgentPurpose),
SetHandler(AgentPurpose, Option<PublicIdentifier>),
TextGeneration(ConfigTextGenerationSettingRelatedControllerType),
SpeechToText(ConfigSpeechToTextSettingRelatedControllerType),
TextToSpeech(ConfigTextToSpeechSettingRelatedControllerType),
}
#[derive(Debug, PartialEq)]
pub enum ConfigTextGenerationSettingRelatedControllerType {
GetContextManagementEnabled,
SetContextManagementEnabled(Option<bool>),
GetPrefixRequirementType,
SetPrefixRequirementType(Option<TextGenerationPrefixRequirementType>),
GetAutoUsage,
SetAutoUsage(Option<TextGenerationAutoUsage>),
GetPromptOverride,
SetPromptOverride(Option<String>),
GetTemperatureOverride,
SetTemperatureOverride(Option<f32>),
}
#[derive(Debug, PartialEq)]
pub enum ConfigSpeechToTextSettingRelatedControllerType {
GetFlowType,
SetFlowType(Option<SpeechToTextFlowType>),
GetLanguage,
SetLanguage(Option<String>),
}
#[derive(Debug, PartialEq)]
pub enum ConfigTextToSpeechSettingRelatedControllerType {
GetBotMessagesFlowType,
SetBotMessagesFlowType(Option<TextToSpeechBotMessagesFlowType>),
GetUserMessagesFlowType,
SetUserMessagesFlowType(Option<TextToSpeechUserMessagesFlowType>),
GetSpeedOverride,
SetSpeedOverride(Option<f32>),
GetVoiceOverride,
SetVoiceOverride(Option<String>),
}

View File

@@ -0,0 +1,141 @@
#[cfg(test)]
mod tests;
mod speech_to_text;
mod text_generation;
mod text_to_speech;
use crate::{
agent::{AgentPurpose, PublicIdentifier},
controller::ControllerType,
strings,
};
use super::controller_type::{
ConfigControllerType, ConfigSettingRelatedControllerType, SettingsStorageSource,
};
pub fn determine_controller(text: &str) -> ControllerType {
if text.starts_with("status") {
return ControllerType::Config(ConfigControllerType::Status);
}
// Someone pasted our instructions verbatim.
if text.strip_prefix("CONFIG_TYPE").is_some() {
return ControllerType::Error(strings::cfg::error_config_type_not_replaced());
}
if let Some(remaining_text) = text.strip_prefix("room ") {
return match do_determine_controller(remaining_text.trim()) {
Ok(handler) => ControllerType::Config(ConfigControllerType::SettingsRelated(
SettingsStorageSource::Room,
handler,
)),
Err(controller_type) => controller_type,
};
}
if let Some(remaining_text) = text.strip_prefix("global ") {
return match do_determine_controller(remaining_text.trim()) {
Ok(handler) => ControllerType::Config(ConfigControllerType::SettingsRelated(
SettingsStorageSource::Global,
handler,
)),
Err(controller_type) => controller_type,
};
}
ControllerType::Config(ConfigControllerType::Help)
}
fn do_determine_controller(
text: &str,
) -> Result<ConfigSettingRelatedControllerType, ControllerType> {
if let Some(purpose_str) = text.strip_prefix("handler") {
let purpose_str = purpose_str.trim();
if purpose_str.is_empty() {
return Err(ControllerType::Error(
strings::cfg::configuration_invocation_incorrect_more_values_expected().to_owned(),
));
}
let Some(purpose) = AgentPurpose::from_str(purpose_str) else {
return Err(ControllerType::Error(
strings::agent::purpose_unrecognized(purpose_str).to_owned(),
));
};
return Ok(ConfigSettingRelatedControllerType::GetHandler(purpose));
}
if let Some(remaining_text) = text.strip_prefix("set-handler") {
// Something like:
// - `PURPOSE ID`
// - `PURPOSE`
let remaining_text = remaining_text.trim();
if remaining_text.is_empty() {
return Err(ControllerType::Error(
strings::cfg::configuration_invocation_incorrect_more_values_expected().to_owned(),
));
}
// This will be None if we're just dealing with `PURPOSE` and lack an `ID`.
// In such cases, the whole thing is the purpose string.
let parts = remaining_text.split_once(' ');
let (purpose_str, agent_id_string_option) = if let Some(parts) = parts {
(parts.0, Some(parts.1.to_owned()))
} else {
(remaining_text, None)
};
let Some(purpose) = AgentPurpose::from_str(purpose_str) else {
return Err(ControllerType::Error(
strings::agent::purpose_unrecognized(purpose_str).to_owned(),
));
};
let agent_identifier = match agent_id_string_option {
Some(agent_id_string) => {
let Some(agent_identifier) = PublicIdentifier::from_str(&agent_id_string) else {
return Err(ControllerType::Error(
strings::agent::invalid_id_generic().to_owned(),
));
};
Some(agent_identifier)
}
None => None,
};
return Ok(ConfigSettingRelatedControllerType::SetHandler(
purpose,
agent_identifier,
));
}
if let Some(remaining_text) = text.strip_prefix("text-generation") {
return match text_generation::determine(remaining_text.trim()) {
Ok(handler) => Ok(ConfigSettingRelatedControllerType::TextGeneration(handler)),
Err(controller_type) => Err(controller_type),
};
}
if let Some(remaining_text) = text.strip_prefix("text-to-speech") {
return match text_to_speech::determine(remaining_text.trim()) {
Ok(handler) => Ok(ConfigSettingRelatedControllerType::TextToSpeech(handler)),
Err(controller_type) => Err(controller_type),
};
}
if let Some(remaining_text) = text.strip_prefix("speech-to-text") {
return match speech_to_text::determine(remaining_text.trim()) {
Ok(handler) => Ok(ConfigSettingRelatedControllerType::SpeechToText(handler)),
Err(controller_type) => Err(controller_type),
};
}
Err(ControllerType::Unknown)
}

View File

@@ -0,0 +1,87 @@
#[cfg(test)]
mod tests;
use crate::{controller::ControllerType, entity::roomconfig::SpeechToTextFlowType, strings};
use super::super::controller_type::ConfigSpeechToTextSettingRelatedControllerType;
pub(super) fn determine(
text: &str,
) -> Result<ConfigSpeechToTextSettingRelatedControllerType, ControllerType> {
// Flow Type
if let Some(remaining_text) = text.strip_prefix("flow-type") {
let remaining_text = remaining_text.trim();
if !remaining_text.is_empty() {
return Err(ControllerType::Error(
strings::cfg::configuration_getter_used_with_extra_text(
"flow-type",
remaining_text,
)
.to_owned(),
));
}
return Ok(ConfigSpeechToTextSettingRelatedControllerType::GetFlowType);
}
if let Some(value_string) = text.strip_prefix("set-flow-type") {
let value_string = value_string.trim().to_owned();
let value_choice = if value_string.is_empty() {
None
} else {
let value_choice = SpeechToTextFlowType::from_str(&value_string.to_lowercase());
if value_choice.is_none() {
return Err(ControllerType::Error(
strings::cfg::configuration_value_unrecognized(&value_string).to_owned(),
));
}
value_choice
};
return Ok(ConfigSpeechToTextSettingRelatedControllerType::SetFlowType(
value_choice,
));
}
// Language
if let Some(remaining_text) = text.strip_prefix("language") {
let remaining_text = remaining_text.trim();
if !remaining_text.is_empty() {
return Err(ControllerType::Error(
strings::cfg::configuration_getter_used_with_extra_text("language", remaining_text)
.to_owned(),
));
}
return Ok(ConfigSpeechToTextSettingRelatedControllerType::GetLanguage);
}
if let Some(value_string) = text.strip_prefix("set-language") {
let value_string = value_string.trim().to_owned();
let value_string = if value_string.is_empty() {
None
} else {
if value_string.len() != 2 {
return Err(ControllerType::Error(
strings::speech_to_text::language_code_invalid(&value_string).to_owned(),
));
}
Some(value_string)
};
return Ok(ConfigSpeechToTextSettingRelatedControllerType::SetLanguage(
value_string,
));
}
Err(ControllerType::Unknown)
}

View File

@@ -0,0 +1,136 @@
#[test]
fn determine_controller_other() {
use super::ConfigSpeechToTextSettingRelatedControllerType;
use super::ControllerType;
struct TestCase {
name: &'static str,
input: &'static str,
expected: Result<ConfigSpeechToTextSettingRelatedControllerType, ControllerType>,
}
let test_cases = vec![TestCase {
name: "Unknown",
input: "whatever",
expected: Err(ControllerType::Unknown),
}];
for test_case in test_cases {
let result = super::determine(test_case.input);
assert_eq!(result, test_case.expected, "Test case: {}", test_case.name);
}
}
#[test]
fn determine_controller_flow_type() {
use super::ConfigSpeechToTextSettingRelatedControllerType;
use super::ControllerType;
use crate::entity::roomconfig::SpeechToTextFlowType;
struct TestCase {
name: &'static str,
input: &'static str,
expected: Result<ConfigSpeechToTextSettingRelatedControllerType, ControllerType>,
}
let test_cases = vec![
TestCase {
name: "flow-type getter ok",
input: "flow-type",
expected: Ok(ConfigSpeechToTextSettingRelatedControllerType::GetFlowType),
},
TestCase {
name: "flow-type getter extra args",
input: "flow-type some values here",
expected: Err(ControllerType::Error(
crate::strings::cfg::configuration_getter_used_with_extra_text(
"flow-type",
"some values here",
),
)),
},
TestCase {
name: "flow-type setter",
input: "set-flow-type transcribe_and_generate_text",
expected: Ok(ConfigSpeechToTextSettingRelatedControllerType::SetFlowType(
Some(SpeechToTextFlowType::TranscribeAndGenerateText),
)),
},
TestCase {
name: "flow-type setter",
input: "set-flow-type unknown-Value",
expected: Err(ControllerType::Error(
crate::strings::cfg::configuration_value_unrecognized("unknown-Value"),
)),
},
TestCase {
name: "flow-type unsetter",
input: "set-flow-type",
expected: Ok(ConfigSpeechToTextSettingRelatedControllerType::SetFlowType(
None,
)),
},
];
for test_case in test_cases {
let result = super::determine(test_case.input);
assert_eq!(result, test_case.expected, "Test case: {}", test_case.name);
}
}
#[test]
fn determine_controller_language() {
use super::ConfigSpeechToTextSettingRelatedControllerType;
use super::ControllerType;
struct TestCase {
name: &'static str,
input: &'static str,
expected: Result<ConfigSpeechToTextSettingRelatedControllerType, ControllerType>,
}
let test_cases = vec![
TestCase {
name: "language getter ok",
input: "language",
expected: Ok(ConfigSpeechToTextSettingRelatedControllerType::GetLanguage),
},
TestCase {
name: "language getter extra args",
input: "language some values here",
expected: Err(ControllerType::Error(
crate::strings::cfg::configuration_getter_used_with_extra_text(
"language",
"some values here",
),
)),
},
TestCase {
name: "language setter 2-letter code (ja)",
input: "set-language ja",
expected: Ok(ConfigSpeechToTextSettingRelatedControllerType::SetLanguage(
Some("ja".to_owned()),
)),
},
// OpenAI does not support 3-letter codes, so we won't be allowing it either
TestCase {
name: "language setter 3-letter code (jpn) fails",
input: "set-language jpn",
expected: Err(ControllerType::Error(
crate::strings::speech_to_text::language_code_invalid("jpn"),
)),
},
TestCase {
name: "language unsetter",
input: "set-language",
expected: Ok(ConfigSpeechToTextSettingRelatedControllerType::SetLanguage(
None,
)),
},
];
for test_case in test_cases {
let result = super::determine(test_case.input);
assert_eq!(result, test_case.expected, "Test case: {}", test_case.name);
}
}

View File

@@ -0,0 +1,212 @@
#[test]
fn determine_controller() {
use super::super::controller_type;
use crate::agent::{AgentPurpose, PublicIdentifier};
struct TestCase {
name: &'static str,
input: &'static str,
expected: super::ControllerType,
}
let test_cases = vec![
TestCase {
name: "Top-level is help",
input: "",
expected: super::ControllerType::Config(controller_type::ConfigControllerType::Help),
},
TestCase {
name: "unknown commands is help",
input: "whatever",
expected: super::ControllerType::Config(controller_type::ConfigControllerType::Help),
},
TestCase {
name: "Status",
input: "status",
expected: super::ControllerType::Config(controller_type::ConfigControllerType::Status),
},
TestCase {
name: "per-room handler getter - catch-all",
input: "room handler catch-all",
expected: super::ControllerType::Config(controller_type::ConfigControllerType::SettingsRelated(
controller_type::SettingsStorageSource::Room,
controller_type::ConfigSettingRelatedControllerType::GetHandler(AgentPurpose::CatchAll),
)),
},
TestCase {
name: "per-room handler getter - text-generation",
input: "room handler text-generation",
expected: super::ControllerType::Config(controller_type::ConfigControllerType::SettingsRelated(
controller_type::SettingsStorageSource::Room,
controller_type::ConfigSettingRelatedControllerType::GetHandler(AgentPurpose::TextGeneration),
)),
},
TestCase {
name: "per-room handler getter - invalid purpose",
input: "room handler invalid-purpose-here",
expected: super::ControllerType::Error(
crate::strings::agent::purpose_unrecognized("invalid-purpose-here").to_owned()
),
},
TestCase {
name: "per-room handler getter - invalid purpose with spaces",
input: "room handler invalid purpose here",
expected: super::ControllerType::Error(
crate::strings::agent::purpose_unrecognized("invalid purpose here").to_owned()
),
},
TestCase {
name: "per-room handler setter - too few values",
input: "room set-handler",
expected: super::ControllerType::Error(
crate::strings::cfg::configuration_invocation_incorrect_more_values_expected().to_owned()
),
},
TestCase {
name: "per-room handler setter - catch-all",
input: "room set-handler catch-all static/agent-id",
expected: super::ControllerType::Config(controller_type::ConfigControllerType::SettingsRelated(
controller_type::SettingsStorageSource::Room,
controller_type::ConfigSettingRelatedControllerType::SetHandler(AgentPurpose::CatchAll, Some(
PublicIdentifier::Static("agent-id".to_owned())
)),
)),
},
TestCase {
name: "per-room handler setter - catch-all with bare agent id",
input: "room set-handler catch-all agent-id",
expected: super::ControllerType::Error(
crate::strings::agent::invalid_id_generic().to_owned()
),
},
TestCase {
name: "per-room handler setter - catch-all unsetter",
input: "room set-handler catch-all",
expected: super::ControllerType::Config(controller_type::ConfigControllerType::SettingsRelated(
controller_type::SettingsStorageSource::Room,
controller_type::ConfigSettingRelatedControllerType::SetHandler(AgentPurpose::CatchAll, None),
)),
},
TestCase {
name: "per-room handler setter - text-generation",
input: "room set-handler text-generation room-local/agent-id",
expected: super::ControllerType::Config(controller_type::ConfigControllerType::SettingsRelated(
controller_type::SettingsStorageSource::Room,
controller_type::ConfigSettingRelatedControllerType::SetHandler(AgentPurpose::TextGeneration, Some(
PublicIdentifier::DynamicRoomLocal("agent-id".to_owned())
)),
)),
},
TestCase {
name: "per-room handler setter - too many values",
input: "room set-handler text-generation agent-id more values here",
expected: super::ControllerType::Error(
crate::strings::agent::invalid_id_generic().to_owned()
),
},
// We have few global handler test cases. We've exercised the per-room handlers enough.
// These share the same code path, so we don't need to test all the permutations again.
TestCase {
name: "global handler getter - catch-all",
input: "global handler catch-all",
expected: super::ControllerType::Config(controller_type::ConfigControllerType::SettingsRelated(
controller_type::SettingsStorageSource::Global,
controller_type::ConfigSettingRelatedControllerType::GetHandler(AgentPurpose::CatchAll),
)),
},
TestCase {
name: "global handler setter - text-generation with global agent",
input: "global set-handler text-generation global/agent-id",
expected: super::ControllerType::Config(controller_type::ConfigControllerType::SettingsRelated(
controller_type::SettingsStorageSource::Global,
controller_type::ConfigSettingRelatedControllerType::SetHandler(AgentPurpose::TextGeneration, Some(
PublicIdentifier::DynamicGlobal("agent-id".to_owned())
)),
)),
},
// This test case passes, even though the handler function will subsequently reject using room-local agents for global handlers.
TestCase {
name: "global handler setter - text-generation with room-local agent",
input: "global set-handler text-generation room-local/agent-id",
expected: super::ControllerType::Config(controller_type::ConfigControllerType::SettingsRelated(
controller_type::SettingsStorageSource::Global,
controller_type::ConfigSettingRelatedControllerType::SetHandler(AgentPurpose::TextGeneration, Some(
PublicIdentifier::DynamicRoomLocal("agent-id".to_owned())
)),
)),
},
// We'll only test one handler per sub-category to ensure proper routing is done here.
// Extensive tests for each sub-category are done in their respective modules.
TestCase {
name: "per-room text-generation/context-management-enabled getter",
input: "room text-generation context-management-enabled",
expected: super::ControllerType::Config(controller_type::ConfigControllerType::SettingsRelated(
controller_type::SettingsStorageSource::Room,
controller_type::ConfigSettingRelatedControllerType::TextGeneration(
controller_type::ConfigTextGenerationSettingRelatedControllerType::GetContextManagementEnabled,
),
)),
},
TestCase {
name: "global text-generation/context-management-enabled getter",
input: "global text-generation context-management-enabled",
expected: super::ControllerType::Config(controller_type::ConfigControllerType::SettingsRelated(
controller_type::SettingsStorageSource::Global,
controller_type::ConfigSettingRelatedControllerType::TextGeneration(
controller_type::ConfigTextGenerationSettingRelatedControllerType::GetContextManagementEnabled,
),
)),
},
TestCase {
name: "per-room text-to-speech/speed-override getter",
input: "room text-to-speech speed-override",
expected: super::ControllerType::Config(controller_type::ConfigControllerType::SettingsRelated(
controller_type::SettingsStorageSource::Room,
controller_type::ConfigSettingRelatedControllerType::TextToSpeech(
controller_type::ConfigTextToSpeechSettingRelatedControllerType::GetSpeedOverride,
),
)),
},
TestCase {
name: "global text-to-speech/speed-override getter",
input: "global text-to-speech speed-override",
expected: super::ControllerType::Config(controller_type::ConfigControllerType::SettingsRelated(
controller_type::SettingsStorageSource::Global,
controller_type::ConfigSettingRelatedControllerType::TextToSpeech(
controller_type::ConfigTextToSpeechSettingRelatedControllerType::GetSpeedOverride,
),
)),
},
TestCase {
name: "room speech-to-text/flow-type getter",
input: "room speech-to-text flow-type",
expected: super::ControllerType::Config(controller_type::ConfigControllerType::SettingsRelated(
controller_type::SettingsStorageSource::Room,
controller_type::ConfigSettingRelatedControllerType::SpeechToText(
controller_type::ConfigSpeechToTextSettingRelatedControllerType::GetFlowType,
),
)),
},
TestCase {
name: "global speech-to-text/flow-type getter",
input: "global speech-to-text flow-type",
expected: super::ControllerType::Config(controller_type::ConfigControllerType::SettingsRelated(
controller_type::SettingsStorageSource::Global,
controller_type::ConfigSettingRelatedControllerType::SpeechToText(
controller_type::ConfigSpeechToTextSettingRelatedControllerType::GetFlowType,
),
)),
},
];
for test_case in test_cases {
let result = super::determine_controller(test_case.input);
assert_eq!(result, test_case.expected, "Test case: {}", test_case.name);
}
}

View File

@@ -0,0 +1,201 @@
#[cfg(test)]
mod tests;
use crate::{
controller::ControllerType,
entity::roomconfig::{TextGenerationAutoUsage, TextGenerationPrefixRequirementType},
strings,
};
use super::super::controller_type::ConfigTextGenerationSettingRelatedControllerType;
pub(super) fn determine(
text: &str,
) -> Result<ConfigTextGenerationSettingRelatedControllerType, ControllerType> {
if let Some(remaining_text) = text.strip_prefix("context-management-enabled") {
let remaining_text = remaining_text.trim();
if !remaining_text.is_empty() {
return Err(ControllerType::Error(
strings::cfg::configuration_getter_used_with_extra_text(
"context-management-enabled",
remaining_text,
)
.to_owned(),
));
}
return Ok(ConfigTextGenerationSettingRelatedControllerType::GetContextManagementEnabled);
}
if let Some(value_string) = text.strip_prefix("set-context-management-enabled") {
let value_string = value_string.trim().to_owned();
let value_opt = if value_string.is_empty() {
None
} else {
let value_string_lowercase = value_string.to_lowercase();
Some(match value_string_lowercase.as_str() {
"true" => true,
"false" => false,
_ => {
return Err(ControllerType::Error(
strings::cfg::configuration_value_unrecognized(&value_string).to_owned(),
));
}
})
};
return Ok(
ConfigTextGenerationSettingRelatedControllerType::SetContextManagementEnabled(
value_opt,
),
);
}
if let Some(remaining_text) = text.strip_prefix("prefix-requirement-type") {
let remaining_text = remaining_text.trim();
if !remaining_text.is_empty() {
return Err(ControllerType::Error(
strings::cfg::configuration_getter_used_with_extra_text(
"prefix-requirement-type",
remaining_text,
)
.to_owned(),
));
}
return Ok(ConfigTextGenerationSettingRelatedControllerType::GetPrefixRequirementType);
}
if let Some(value_string) = text.strip_prefix("set-prefix-requirement-type") {
let value_string = value_string.trim().to_owned();
let value_choice = if value_string.is_empty() {
None
} else {
let value_choice =
TextGenerationPrefixRequirementType::from_str(&value_string.to_lowercase());
if value_choice.is_none() {
return Err(ControllerType::Error(
strings::cfg::configuration_value_unrecognized(&value_string).to_owned(),
));
}
value_choice
};
return Ok(
ConfigTextGenerationSettingRelatedControllerType::SetPrefixRequirementType(
value_choice,
),
);
}
if let Some(remaining_text) = text.strip_prefix("auto-usage") {
let remaining_text = remaining_text.trim();
if !remaining_text.is_empty() {
return Err(ControllerType::Error(
strings::cfg::configuration_getter_used_with_extra_text(
"auto-usage",
remaining_text,
)
.to_owned(),
));
}
return Ok(ConfigTextGenerationSettingRelatedControllerType::GetAutoUsage);
}
if let Some(value_string) = text.strip_prefix("set-auto-usage") {
let value_string = value_string.trim().to_owned();
let value_choice = if value_string.is_empty() {
None
} else {
let value_choice = TextGenerationAutoUsage::from_str(&value_string.to_lowercase());
if value_choice.is_none() {
return Err(ControllerType::Error(
strings::cfg::configuration_value_unrecognized(&value_string).to_owned(),
));
}
value_choice
};
return Ok(ConfigTextGenerationSettingRelatedControllerType::SetAutoUsage(value_choice));
}
if let Some(remaining_text) = text.strip_prefix("prompt-override") {
let remaining_text = remaining_text.trim();
if !remaining_text.is_empty() {
return Err(ControllerType::Error(
strings::cfg::configuration_getter_used_with_extra_text(
"prompt-override",
remaining_text,
)
.to_owned(),
));
}
return Ok(ConfigTextGenerationSettingRelatedControllerType::GetPromptOverride);
}
if let Some(value_string) = text.strip_prefix("set-prompt-override") {
let value_string = value_string.trim().to_owned();
if value_string.is_empty() {
return Ok(ConfigTextGenerationSettingRelatedControllerType::SetPromptOverride(None));
}
return Ok(
ConfigTextGenerationSettingRelatedControllerType::SetPromptOverride(Some(value_string)),
);
}
if let Some(remaining_text) = text.strip_prefix("temperature-override") {
let remaining_text = remaining_text.trim();
if !remaining_text.is_empty() {
return Err(ControllerType::Error(
strings::cfg::configuration_getter_used_with_extra_text(
"temperature-override",
remaining_text,
)
.to_owned(),
));
}
return Ok(ConfigTextGenerationSettingRelatedControllerType::GetTemperatureOverride);
}
if let Some(value_string) = text.strip_prefix("set-temperature-override") {
let value_string = value_string.trim().to_owned();
if value_string.is_empty() {
return Ok(
ConfigTextGenerationSettingRelatedControllerType::SetTemperatureOverride(None),
);
}
let value_f32 = value_string.parse::<f32>();
let Ok(value_f32) = value_f32 else {
return Err(ControllerType::Error(
strings::cfg::configuration_value_not_f32(&value_string).to_owned(),
));
};
return Ok(
ConfigTextGenerationSettingRelatedControllerType::SetTemperatureOverride(Some(
value_f32,
)),
);
}
Err(ControllerType::Unknown)
}

View File

@@ -0,0 +1,332 @@
#[test]
fn determine_controller_other() {
use super::ConfigTextGenerationSettingRelatedControllerType;
use super::ControllerType;
struct TestCase {
name: &'static str,
input: &'static str,
expected: Result<ConfigTextGenerationSettingRelatedControllerType, ControllerType>,
}
let test_cases = vec![TestCase {
name: "Unknown",
input: "whatever",
expected: Err(ControllerType::Unknown),
}];
for test_case in test_cases {
let result = super::determine(test_case.input);
assert_eq!(result, test_case.expected, "Test case: {}", test_case.name);
}
}
#[test]
fn determine_controller_context_management() {
use super::ConfigTextGenerationSettingRelatedControllerType;
use super::ControllerType;
struct TestCase {
name: &'static str,
input: &'static str,
expected: Result<ConfigTextGenerationSettingRelatedControllerType, ControllerType>,
}
let test_cases = vec![
TestCase {
name: "context-management-enabled getter ok",
input: "context-management-enabled",
expected: Ok(
ConfigTextGenerationSettingRelatedControllerType::GetContextManagementEnabled,
),
},
TestCase {
name: "context-management-enabled getter extra args",
input: "context-management-enabled some values here",
expected: Err(ControllerType::Error(
crate::strings::cfg::configuration_getter_used_with_extra_text(
"context-management-enabled",
"some values here",
),
)),
},
TestCase {
name: "context-management-enabled setter",
input: "set-context-management-enabled true",
expected: Ok(
ConfigTextGenerationSettingRelatedControllerType::SetContextManagementEnabled(
Some(true),
),
),
},
TestCase {
name: "context-management-enabled setter uppercase",
input: "set-context-management-enabled TRUE",
expected: Ok(
ConfigTextGenerationSettingRelatedControllerType::SetContextManagementEnabled(
Some(true),
),
),
},
TestCase {
name: "context-management-enabled setter non-bool",
input: "set-context-management-enabled non-Bool-Value",
expected: Err(ControllerType::Error(
crate::strings::cfg::configuration_value_unrecognized("non-Bool-Value"),
)),
},
TestCase {
name: "context-management-enabled unsetter",
input: "set-context-management-enabled",
expected: Ok(
ConfigTextGenerationSettingRelatedControllerType::SetContextManagementEnabled(None),
),
},
];
for test_case in test_cases {
let result = super::determine(test_case.input);
assert_eq!(result, test_case.expected, "Test case: {}", test_case.name);
}
}
#[test]
fn determine_controller_prefix_requirement_type() {
use super::ConfigTextGenerationSettingRelatedControllerType;
use super::ControllerType;
use crate::entity::roomconfig::TextGenerationPrefixRequirementType;
struct TestCase {
name: &'static str,
input: &'static str,
expected: Result<ConfigTextGenerationSettingRelatedControllerType, ControllerType>,
}
let test_cases = vec![
TestCase {
name: "prefix-requirement-type getter ok",
input: "prefix-requirement-type",
expected: Ok(
ConfigTextGenerationSettingRelatedControllerType::GetPrefixRequirementType,
),
},
TestCase {
name: "prefix-requirement-type getter extra args",
input: "prefix-requirement-type some values here",
expected: Err(ControllerType::Error(
crate::strings::cfg::configuration_getter_used_with_extra_text(
"prefix-requirement-type",
"some values here",
),
)),
},
TestCase {
name: "prefix-requirement-type setter (command_prefix)",
input: "set-prefix-requirement-type command_prefix",
expected: Ok(
ConfigTextGenerationSettingRelatedControllerType::SetPrefixRequirementType(Some(
TextGenerationPrefixRequirementType::CommandPrefix,
)),
),
},
TestCase {
name: "prefix-requirement-type setter (no)",
input: "set-prefix-requirement-type no",
expected: Ok(
ConfigTextGenerationSettingRelatedControllerType::SetPrefixRequirementType(Some(
TextGenerationPrefixRequirementType::No,
)),
),
},
TestCase {
name: "prefix-requirement-type setter",
input: "set-prefix-requirement-type unknown-Value",
expected: Err(ControllerType::Error(
crate::strings::cfg::configuration_value_unrecognized("unknown-Value"),
)),
},
TestCase {
name: "prefix-requirement-type unsetter",
input: "set-prefix-requirement-type",
expected: Ok(
ConfigTextGenerationSettingRelatedControllerType::SetPrefixRequirementType(None),
),
},
];
for test_case in test_cases {
let result = super::determine(test_case.input);
assert_eq!(result, test_case.expected, "Test case: {}", test_case.name);
}
}
#[test]
fn determine_controller_auto_usage() {
use super::ConfigTextGenerationSettingRelatedControllerType;
use super::ControllerType;
use crate::entity::roomconfig::TextGenerationAutoUsage;
struct TestCase {
name: &'static str,
input: &'static str,
expected: Result<ConfigTextGenerationSettingRelatedControllerType, ControllerType>,
}
let test_cases = vec![
TestCase {
name: "auto-usage getter ok",
input: "auto-usage",
expected: Ok(ConfigTextGenerationSettingRelatedControllerType::GetAutoUsage),
},
TestCase {
name: "auto-usage getter extra args",
input: "auto-usage some values here",
expected: Err(ControllerType::Error(
crate::strings::cfg::configuration_getter_used_with_extra_text(
"auto-usage",
"some values here",
),
)),
},
TestCase {
name: "auto-usage setter",
input: "set-auto-usage only_for_voice",
expected: Ok(
ConfigTextGenerationSettingRelatedControllerType::SetAutoUsage(Some(
TextGenerationAutoUsage::OnlyForVoice,
)),
),
},
TestCase {
name: "auto-usage setter",
input: "set-auto-usage unknown-Value",
expected: Err(ControllerType::Error(
crate::strings::cfg::configuration_value_unrecognized("unknown-Value"),
)),
},
TestCase {
name: "auto-usage unsetter",
input: "set-auto-usage",
expected: Ok(ConfigTextGenerationSettingRelatedControllerType::SetAutoUsage(None)),
},
];
for test_case in test_cases {
let result = super::determine(test_case.input);
assert_eq!(result, test_case.expected, "Test case: {}", test_case.name);
}
}
#[test]
fn determine_controller_prompt_override() {
use super::ConfigTextGenerationSettingRelatedControllerType;
use super::ControllerType;
struct TestCase {
name: &'static str,
input: &'static str,
expected: Result<ConfigTextGenerationSettingRelatedControllerType, ControllerType>,
}
let test_cases = vec![
TestCase {
name: "prompt-override getter ok",
input: "prompt-override",
expected: Ok(ConfigTextGenerationSettingRelatedControllerType::GetPromptOverride),
},
TestCase {
name: "prompt-override getter extra args",
input: "prompt-override some values here",
expected: Err(ControllerType::Error(
crate::strings::cfg::configuration_getter_used_with_extra_text(
"prompt-override",
"some values here",
),
)),
},
TestCase {
name: "prompt-override setter with multiple words",
input: "set-prompt-override Hello! You are a bot",
expected: Ok(
ConfigTextGenerationSettingRelatedControllerType::SetPromptOverride(Some(
"Hello! You are a bot".to_owned(),
)),
),
},
TestCase {
name: "prompt-override setter with multi-line",
input: "set-prompt-override Hello!\n\nYou are a bot",
expected: Ok(
ConfigTextGenerationSettingRelatedControllerType::SetPromptOverride(Some(
"Hello!\n\nYou are a bot".to_owned(),
)),
),
},
TestCase {
name: "prompt-override unsetter",
input: "set-prompt-override",
expected: Ok(ConfigTextGenerationSettingRelatedControllerType::SetPromptOverride(None)),
},
];
for test_case in test_cases {
let result = super::determine(test_case.input);
assert_eq!(result, test_case.expected, "Test case: {}", test_case.name);
}
}
#[test]
fn determine_controller_temperature_override() {
use super::ConfigTextGenerationSettingRelatedControllerType;
use super::ControllerType;
struct TestCase {
name: &'static str,
input: &'static str,
expected: Result<ConfigTextGenerationSettingRelatedControllerType, ControllerType>,
}
let test_cases = vec![
TestCase {
name: "temperature-override getter ok",
input: "temperature-override",
expected: Ok(ConfigTextGenerationSettingRelatedControllerType::GetTemperatureOverride),
},
TestCase {
name: "temperature-override getter extra args",
input: "temperature-override some values here",
expected: Err(ControllerType::Error(
crate::strings::cfg::configuration_getter_used_with_extra_text(
"temperature-override",
"some values here",
),
)),
},
TestCase {
name: "temperature-override setter",
input: "set-temperature-override 0.5",
expected: Ok(
ConfigTextGenerationSettingRelatedControllerType::SetTemperatureOverride(Some(0.5)),
),
},
TestCase {
name: "temperature-override setter",
input: "set-temperature-override unknown-Value",
expected: Err(ControllerType::Error(
crate::strings::cfg::configuration_value_not_f32("unknown-Value"),
)),
},
TestCase {
name: "temperature-override unsetter",
input: "set-temperature-override",
expected: Ok(
ConfigTextGenerationSettingRelatedControllerType::SetTemperatureOverride(None),
),
},
];
for test_case in test_cases {
let result = super::determine(test_case.input);
assert_eq!(result, test_case.expected, "Test case: {}", test_case.name);
}
}

View File

@@ -0,0 +1,158 @@
#[cfg(test)]
mod tests;
use crate::{
controller::ControllerType,
entity::roomconfig::{TextToSpeechBotMessagesFlowType, TextToSpeechUserMessagesFlowType},
strings,
};
use super::super::controller_type::ConfigTextToSpeechSettingRelatedControllerType;
pub(super) fn determine(
text: &str,
) -> Result<ConfigTextToSpeechSettingRelatedControllerType, ControllerType> {
if let Some(remaining_text) = text.strip_prefix("bot-msgs-flow-type") {
let remaining_text = remaining_text.trim();
if !remaining_text.is_empty() {
return Err(ControllerType::Error(
strings::cfg::configuration_getter_used_with_extra_text(
"bot-msgs-flow-type",
remaining_text,
)
.to_owned(),
));
}
return Ok(ConfigTextToSpeechSettingRelatedControllerType::GetBotMessagesFlowType);
}
if let Some(value_string) = text.strip_prefix("set-bot-msgs-flow-type") {
let value_string = value_string.trim().to_owned();
let value_choice = if value_string.is_empty() {
None
} else {
let value_choice =
TextToSpeechBotMessagesFlowType::from_str(&value_string.to_lowercase());
if value_choice.is_none() {
return Err(ControllerType::Error(
strings::cfg::configuration_value_unrecognized(&value_string).to_owned(),
));
}
value_choice
};
return Ok(
ConfigTextToSpeechSettingRelatedControllerType::SetBotMessagesFlowType(value_choice),
);
}
if let Some(remaining_text) = text.strip_prefix("user-msgs-flow-type") {
let remaining_text = remaining_text.trim();
if !remaining_text.is_empty() {
return Err(ControllerType::Error(
strings::cfg::configuration_getter_used_with_extra_text(
"user-msgs-flow-type",
remaining_text,
)
.to_owned(),
));
}
return Ok(ConfigTextToSpeechSettingRelatedControllerType::GetUserMessagesFlowType);
}
if let Some(value_string) = text.strip_prefix("set-user-msgs-flow-type") {
let value_string = value_string.trim().to_owned();
let value_choice = if value_string.is_empty() {
None
} else {
let value_choice =
TextToSpeechUserMessagesFlowType::from_str(&value_string.to_lowercase());
if value_choice.is_none() {
return Err(ControllerType::Error(
strings::cfg::configuration_value_unrecognized(&value_string).to_owned(),
));
}
value_choice
};
return Ok(
ConfigTextToSpeechSettingRelatedControllerType::SetUserMessagesFlowType(value_choice),
);
}
if let Some(remaining_text) = text.strip_prefix("speed-override") {
let remaining_text = remaining_text.trim();
if !remaining_text.is_empty() {
return Err(ControllerType::Error(
strings::cfg::configuration_getter_used_with_extra_text(
"speed-override",
remaining_text,
)
.to_owned(),
));
}
return Ok(ConfigTextToSpeechSettingRelatedControllerType::GetSpeedOverride);
}
if let Some(value_string) = text.strip_prefix("set-speed-override") {
let value_string = value_string.trim().to_owned();
if value_string.is_empty() {
return Ok(ConfigTextToSpeechSettingRelatedControllerType::SetSpeedOverride(None));
}
let value_f32 = value_string.parse::<f32>();
let Ok(value_f32) = value_f32 else {
return Err(ControllerType::Error(
strings::cfg::configuration_value_not_f32(&value_string).to_owned(),
));
};
return Ok(
ConfigTextToSpeechSettingRelatedControllerType::SetSpeedOverride(Some(value_f32)),
);
}
if let Some(remaining_text) = text.strip_prefix("voice-override") {
let remaining_text = remaining_text.trim();
if !remaining_text.is_empty() {
return Err(ControllerType::Error(
strings::cfg::configuration_getter_used_with_extra_text(
"voice-override",
remaining_text,
)
.to_owned(),
));
}
return Ok(ConfigTextToSpeechSettingRelatedControllerType::GetVoiceOverride);
}
if let Some(value_string) = text.strip_prefix("set-voice-override") {
let value_string = value_string.trim().to_owned();
if value_string.is_empty() {
return Ok(ConfigTextToSpeechSettingRelatedControllerType::SetVoiceOverride(None));
}
return Ok(
ConfigTextToSpeechSettingRelatedControllerType::SetVoiceOverride(Some(value_string)),
);
}
Err(ControllerType::Unknown)
}

View File

@@ -0,0 +1,252 @@
#[test]
fn determine_controller_other() {
use super::ConfigTextToSpeechSettingRelatedControllerType;
use super::ControllerType;
struct TestCase {
name: &'static str,
input: &'static str,
expected: Result<ConfigTextToSpeechSettingRelatedControllerType, ControllerType>,
}
let test_cases = vec![TestCase {
name: "Unknown",
input: "whatever",
expected: Err(ControllerType::Unknown),
}];
for test_case in test_cases {
let result = super::determine(test_case.input);
assert_eq!(result, test_case.expected, "Test case: {}", test_case.name);
}
}
#[test]
fn determine_controller_bot_msgs_flow_type() {
use super::ConfigTextToSpeechSettingRelatedControllerType;
use super::ControllerType;
use crate::entity::roomconfig::TextToSpeechBotMessagesFlowType;
struct TestCase {
name: &'static str,
input: &'static str,
expected: Result<ConfigTextToSpeechSettingRelatedControllerType, ControllerType>,
}
let test_cases = vec![
TestCase {
name: "bot-msgs-flow-type getter ok",
input: "bot-msgs-flow-type",
expected: Ok(ConfigTextToSpeechSettingRelatedControllerType::GetBotMessagesFlowType),
},
TestCase {
name: "bot-msgs-flow-type getter extra args",
input: "bot-msgs-flow-type some values here",
expected: Err(ControllerType::Error(
crate::strings::cfg::configuration_getter_used_with_extra_text(
"bot-msgs-flow-type",
"some values here",
),
)),
},
TestCase {
name: "bot-msgs-flow-type setter",
input: "set-bot-msgs-flow-type only_for_voice",
expected: Ok(
ConfigTextToSpeechSettingRelatedControllerType::SetBotMessagesFlowType(Some(
TextToSpeechBotMessagesFlowType::OnlyForVoice,
)),
),
},
TestCase {
name: "bot-msgs-flow-type setter",
input: "set-bot-msgs-flow-type unknown-Value",
expected: Err(ControllerType::Error(
crate::strings::cfg::configuration_value_unrecognized("unknown-Value"),
)),
},
TestCase {
name: "bot-msgs-flow-type unsetter",
input: "set-bot-msgs-flow-type",
expected: Ok(
ConfigTextToSpeechSettingRelatedControllerType::SetBotMessagesFlowType(None),
),
},
];
for test_case in test_cases {
let result = super::determine(test_case.input);
assert_eq!(result, test_case.expected, "Test case: {}", test_case.name);
}
}
#[test]
fn determine_controller_user_msgs_flow_type() {
use super::ConfigTextToSpeechSettingRelatedControllerType;
use super::ControllerType;
use crate::entity::roomconfig::TextToSpeechUserMessagesFlowType;
struct TestCase {
name: &'static str,
input: &'static str,
expected: Result<ConfigTextToSpeechSettingRelatedControllerType, ControllerType>,
}
let test_cases = vec![
TestCase {
name: "user-msgs-flow-type getter ok",
input: "user-msgs-flow-type",
expected: Ok(ConfigTextToSpeechSettingRelatedControllerType::GetUserMessagesFlowType),
},
TestCase {
name: "user-msgs-flow-type getter extra args",
input: "user-msgs-flow-type some values here",
expected: Err(ControllerType::Error(
crate::strings::cfg::configuration_getter_used_with_extra_text(
"user-msgs-flow-type",
"some values here",
),
)),
},
TestCase {
name: "user-msgs-flow-type setter",
input: "set-user-msgs-flow-type on_demand",
expected: Ok(
ConfigTextToSpeechSettingRelatedControllerType::SetUserMessagesFlowType(Some(
TextToSpeechUserMessagesFlowType::OnDemand,
)),
),
},
TestCase {
name: "user-msgs-flow-type setter",
input: "set-user-msgs-flow-type unknown-Value",
expected: Err(ControllerType::Error(
crate::strings::cfg::configuration_value_unrecognized("unknown-Value"),
)),
},
TestCase {
name: "user-msgs-flow-type unsetter",
input: "set-user-msgs-flow-type",
expected: Ok(
ConfigTextToSpeechSettingRelatedControllerType::SetUserMessagesFlowType(None),
),
},
];
for test_case in test_cases {
let result = super::determine(test_case.input);
assert_eq!(result, test_case.expected, "Test case: {}", test_case.name);
}
}
#[test]
fn determine_controller_speed_override() {
use super::ConfigTextToSpeechSettingRelatedControllerType;
use super::ControllerType;
struct TestCase {
name: &'static str,
input: &'static str,
expected: Result<ConfigTextToSpeechSettingRelatedControllerType, ControllerType>,
}
let test_cases = vec![
TestCase {
name: "speed-override getter ok",
input: "speed-override",
expected: Ok(ConfigTextToSpeechSettingRelatedControllerType::GetSpeedOverride),
},
TestCase {
name: "speed-override getter extra args",
input: "speed-override some values here",
expected: Err(ControllerType::Error(
crate::strings::cfg::configuration_getter_used_with_extra_text(
"speed-override",
"some values here",
),
)),
},
TestCase {
name: "speed-override setter",
input: "set-speed-override 0.5",
expected: Ok(
ConfigTextToSpeechSettingRelatedControllerType::SetSpeedOverride(Some(0.5)),
),
},
TestCase {
name: "speed-override setter",
input: "set-speed-override unknown-Value",
expected: Err(ControllerType::Error(
crate::strings::cfg::configuration_value_not_f32("unknown-Value"),
)),
},
TestCase {
name: "speed-override unsetter",
input: "set-speed-override",
expected: Ok(ConfigTextToSpeechSettingRelatedControllerType::SetSpeedOverride(None)),
},
];
for test_case in test_cases {
let result = super::determine(test_case.input);
assert_eq!(result, test_case.expected, "Test case: {}", test_case.name);
}
}
#[test]
fn determine_controller_voice_override() {
use super::ConfigTextToSpeechSettingRelatedControllerType;
use super::ControllerType;
struct TestCase {
name: &'static str,
input: &'static str,
expected: Result<ConfigTextToSpeechSettingRelatedControllerType, ControllerType>,
}
let test_cases = vec![
TestCase {
name: "voice-override getter ok",
input: "voice-override",
expected: Ok(ConfigTextToSpeechSettingRelatedControllerType::GetVoiceOverride),
},
TestCase {
name: "voice-override getter extra args",
input: "voice-override some values here",
expected: Err(ControllerType::Error(
crate::strings::cfg::configuration_getter_used_with_extra_text(
"voice-override",
"some values here",
),
)),
},
TestCase {
name: "voice-override setter",
input: "set-voice-override alex",
expected: Ok(
ConfigTextToSpeechSettingRelatedControllerType::SetVoiceOverride(Some(
"alex".to_owned(),
)),
),
},
TestCase {
name: "voice-override setter preserves case",
input: "set-voice-override Alex",
expected: Ok(
ConfigTextToSpeechSettingRelatedControllerType::SetVoiceOverride(Some(
"Alex".to_owned(),
)),
),
},
TestCase {
name: "voice-override unsetter",
input: "set-voice-override",
expected: Ok(ConfigTextToSpeechSettingRelatedControllerType::SetVoiceOverride(None)),
},
];
for test_case in test_cases {
let result = super::determine(test_case.input);
assert_eq!(result, test_case.expected, "Test case: {}", test_case.name);
}
}

View File

@@ -0,0 +1,124 @@
use crate::strings;
use crate::{entity::MessageContext, Bot};
use mxlink::MessageResponseType;
use super::controller_type::{
ConfigControllerType, ConfigSettingRelatedControllerType, SettingsStorageSource,
};
mod speech_to_text;
mod text_generation;
mod text_to_speech;
pub async fn dispatch_controller(
handler: &ConfigControllerType,
message_context: &MessageContext,
bot: &Bot,
) -> anyhow::Result<()> {
// Anyone can access Help and Status.
// Settings-related access checks are done in dispatch_config_related_handler().
match handler {
ConfigControllerType::Help => super::help::handle(bot, message_context).await,
ConfigControllerType::Status => super::status::handle(bot, message_context).await,
ConfigControllerType::SettingsRelated(config_type, config_related_handler) => {
dispatch_config_related_handler(
config_type,
config_related_handler,
message_context,
bot,
)
.await
}
}
}
async fn dispatch_config_related_handler(
config_type: &SettingsStorageSource,
handler: &ConfigSettingRelatedControllerType,
message_context: &MessageContext,
bot: &Bot,
) -> anyhow::Result<()> {
if let SettingsStorageSource::Global = config_type {
if !message_context.sender_can_manage_global_config()? {
bot.messaging()
.send_error_markdown_no_fail(
message_context.room(),
strings::global_config::no_permissions_to_administrate(),
MessageResponseType::Reply(message_context.thread_info().root_event_id.clone()),
)
.await;
return Ok(());
}
};
let room_settings = match config_type {
SettingsStorageSource::Room => &message_context.room_config().settings,
SettingsStorageSource::Global => &message_context.global_config().fallback_room_settings,
};
match handler {
ConfigSettingRelatedControllerType::GetHandler(purpose) => match config_type {
SettingsStorageSource::Room => {
super::room_config::handler::handle_get(bot, message_context, *purpose).await
}
SettingsStorageSource::Global => {
super::global_config::handler::handle_get(bot, message_context, *purpose).await
}
},
ConfigSettingRelatedControllerType::SetHandler(purpose, agent_identifier) => {
match config_type {
SettingsStorageSource::Room => {
super::room_config::handler::handle_set(
bot,
bot.room_config_manager(),
message_context,
*purpose,
agent_identifier,
)
.await
}
SettingsStorageSource::Global => {
super::global_config::handler::handle_set(
bot,
bot.global_config_manager(),
message_context,
*purpose,
agent_identifier,
)
.await
}
}
}
ConfigSettingRelatedControllerType::TextGeneration(controller_type) => {
text_generation::dispatch(
controller_type,
message_context,
bot,
room_settings,
config_type,
)
.await
}
ConfigSettingRelatedControllerType::SpeechToText(controller_type) => {
speech_to_text::dispatch(
controller_type,
message_context,
bot,
room_settings,
config_type,
)
.await
}
ConfigSettingRelatedControllerType::TextToSpeech(controller_type) => {
text_to_speech::dispatch(
controller_type,
message_context,
bot,
room_settings,
config_type,
)
.await
}
}
}

View File

@@ -0,0 +1,78 @@
use crate::entity::roomconfig::{RoomSettings, SpeechToTextFlowType};
use crate::{entity::MessageContext, Bot};
use super::super::controller_type::{
ConfigSpeechToTextSettingRelatedControllerType, SettingsStorageSource,
};
use super::super::common::generic_setting::handle_get as setting_get;
use super::super::global_config::generic_setting::handle_set as global_setting_set;
use super::super::room_config::generic_setting::handle_set as room_setting_set;
pub(super) async fn dispatch(
handler: &ConfigSpeechToTextSettingRelatedControllerType,
message_context: &MessageContext,
bot: &Bot,
room_settings: &RoomSettings,
config_type: &SettingsStorageSource,
) -> anyhow::Result<()> {
match handler {
ConfigSpeechToTextSettingRelatedControllerType::GetFlowType => {
let value = &room_settings.speech_to_text.flow_type;
setting_get::<SpeechToTextFlowType>(bot, message_context, value).await
}
ConfigSpeechToTextSettingRelatedControllerType::SetFlowType(value) => {
let value = value.to_owned();
let setter_callback = Box::new(move |room_settings: &mut RoomSettings| {
room_settings.speech_to_text.flow_type = value;
});
match config_type {
SettingsStorageSource::Room => {
room_setting_set::<SpeechToTextFlowType>(
bot,
message_context,
&value,
setter_callback,
)
.await
}
SettingsStorageSource::Global => {
global_setting_set::<SpeechToTextFlowType>(
bot,
message_context,
&value,
setter_callback,
)
.await
}
}
}
ConfigSpeechToTextSettingRelatedControllerType::GetLanguage => {
let value = &room_settings.speech_to_text.language;
setting_get::<String>(bot, message_context, value).await
}
ConfigSpeechToTextSettingRelatedControllerType::SetLanguage(value) => {
let value = value.to_owned();
let value_setter = value.clone();
let setter_callback = Box::new(move |room_settings: &mut RoomSettings| {
room_settings.speech_to_text.language = value_setter;
});
match config_type {
SettingsStorageSource::Room => {
room_setting_set::<String>(bot, message_context, &value, setter_callback).await
}
SettingsStorageSource::Global => {
global_setting_set::<String>(bot, message_context, &value, setter_callback)
.await
}
}
}
}
}

View File

@@ -0,0 +1,155 @@
use crate::entity::roomconfig::{
RoomSettings, TextGenerationAutoUsage, TextGenerationPrefixRequirementType,
};
use crate::{entity::MessageContext, Bot};
use super::super::controller_type::{
ConfigTextGenerationSettingRelatedControllerType, SettingsStorageSource,
};
use super::super::common::generic_setting::handle_get as setting_get;
use super::super::global_config::generic_setting::handle_set as global_setting_set;
use super::super::room_config::generic_setting::handle_set as room_setting_set;
pub(super) async fn dispatch(
handler: &ConfigTextGenerationSettingRelatedControllerType,
message_context: &MessageContext,
bot: &Bot,
room_settings: &RoomSettings,
config_type: &SettingsStorageSource,
) -> anyhow::Result<()> {
match handler {
ConfigTextGenerationSettingRelatedControllerType::GetContextManagementEnabled => {
let value = &room_settings.text_generation.context_management_enabled;
setting_get::<bool>(bot, message_context, value).await
}
ConfigTextGenerationSettingRelatedControllerType::SetContextManagementEnabled(value) => {
let value = value.to_owned();
let setter_callback = Box::new(move |room_settings: &mut RoomSettings| {
room_settings.text_generation.context_management_enabled = value;
});
match config_type {
SettingsStorageSource::Room => {
room_setting_set::<bool>(bot, message_context, &value, setter_callback).await
}
SettingsStorageSource::Global => {
global_setting_set::<bool>(bot, message_context, &value, setter_callback).await
}
}
}
ConfigTextGenerationSettingRelatedControllerType::GetPrefixRequirementType => {
let value = &room_settings.text_generation.prefix_requirement_type;
setting_get::<TextGenerationPrefixRequirementType>(bot, message_context, value).await
}
ConfigTextGenerationSettingRelatedControllerType::SetPrefixRequirementType(value) => {
let value = value.to_owned();
let setter_callback = Box::new(move |room_settings: &mut RoomSettings| {
room_settings.text_generation.prefix_requirement_type = value;
});
match config_type {
SettingsStorageSource::Room => {
room_setting_set::<TextGenerationPrefixRequirementType>(
bot,
message_context,
&value,
setter_callback,
)
.await
}
SettingsStorageSource::Global => {
global_setting_set::<TextGenerationPrefixRequirementType>(
bot,
message_context,
&value,
setter_callback,
)
.await
}
}
}
ConfigTextGenerationSettingRelatedControllerType::GetAutoUsage => {
let value = &room_settings.text_generation.auto_usage;
setting_get::<TextGenerationAutoUsage>(bot, message_context, value).await
}
ConfigTextGenerationSettingRelatedControllerType::SetAutoUsage(value) => {
let value = value.to_owned();
let setter_callback = Box::new(move |room_settings: &mut RoomSettings| {
room_settings.text_generation.auto_usage = value;
});
match config_type {
SettingsStorageSource::Room => {
room_setting_set::<TextGenerationAutoUsage>(
bot,
message_context,
&value,
setter_callback,
)
.await
}
SettingsStorageSource::Global => {
global_setting_set::<TextGenerationAutoUsage>(
bot,
message_context,
&value,
setter_callback,
)
.await
}
}
}
ConfigTextGenerationSettingRelatedControllerType::GetPromptOverride => {
let value = &room_settings.text_generation.prompt_override;
setting_get::<String>(bot, message_context, value).await
}
ConfigTextGenerationSettingRelatedControllerType::SetPromptOverride(value) => {
let value = value.to_owned();
let value_setter = value.clone();
let setter_callback = Box::new(move |room_settings: &mut RoomSettings| {
room_settings.text_generation.prompt_override = value_setter;
});
match config_type {
SettingsStorageSource::Room => {
room_setting_set::<String>(bot, message_context, &value, setter_callback).await
}
SettingsStorageSource::Global => {
global_setting_set::<String>(bot, message_context, &value, setter_callback)
.await
}
}
}
ConfigTextGenerationSettingRelatedControllerType::GetTemperatureOverride => {
let value = &room_settings.text_generation.temperature_override;
setting_get::<f32>(bot, message_context, value).await
}
ConfigTextGenerationSettingRelatedControllerType::SetTemperatureOverride(value) => {
let value = value.to_owned();
let setter_callback = Box::new(move |room_settings: &mut RoomSettings| {
room_settings.text_generation.temperature_override = value;
});
match config_type {
SettingsStorageSource::Room => {
room_setting_set::<f32>(bot, message_context, &value, setter_callback).await
}
SettingsStorageSource::Global => {
global_setting_set::<f32>(bot, message_context, &value, setter_callback).await
}
}
}
}
}

View File

@@ -0,0 +1,134 @@
use crate::entity::roomconfig::{
RoomSettings, TextToSpeechBotMessagesFlowType, TextToSpeechUserMessagesFlowType,
};
use crate::{entity::MessageContext, Bot};
use super::super::controller_type::{
ConfigTextToSpeechSettingRelatedControllerType, SettingsStorageSource,
};
use super::super::common::generic_setting::handle_get as setting_get;
use super::super::global_config::generic_setting::handle_set as global_setting_set;
use super::super::room_config::generic_setting::handle_set as room_setting_set;
pub(super) async fn dispatch(
handler: &ConfigTextToSpeechSettingRelatedControllerType,
message_context: &MessageContext,
bot: &Bot,
room_settings: &RoomSettings,
config_type: &SettingsStorageSource,
) -> anyhow::Result<()> {
match handler {
ConfigTextToSpeechSettingRelatedControllerType::GetBotMessagesFlowType => {
let value = &room_settings.text_to_speech.bot_msgs_flow_type;
setting_get::<TextToSpeechBotMessagesFlowType>(bot, message_context, value).await
}
ConfigTextToSpeechSettingRelatedControllerType::SetBotMessagesFlowType(value) => {
let value = value.to_owned();
let setter_callback = Box::new(move |room_settings: &mut RoomSettings| {
room_settings.text_to_speech.bot_msgs_flow_type = value;
});
match config_type {
SettingsStorageSource::Room => {
room_setting_set::<TextToSpeechBotMessagesFlowType>(
bot,
message_context,
&value,
setter_callback,
)
.await
}
SettingsStorageSource::Global => {
global_setting_set::<TextToSpeechBotMessagesFlowType>(
bot,
message_context,
&value,
setter_callback,
)
.await
}
}
}
ConfigTextToSpeechSettingRelatedControllerType::GetUserMessagesFlowType => {
let value = &room_settings.text_to_speech.user_msgs_flow_type;
setting_get::<TextToSpeechUserMessagesFlowType>(bot, message_context, value).await
}
ConfigTextToSpeechSettingRelatedControllerType::SetUserMessagesFlowType(value) => {
let value = value.to_owned();
let setter_callback = Box::new(move |room_settings: &mut RoomSettings| {
room_settings.text_to_speech.user_msgs_flow_type = value;
});
match config_type {
SettingsStorageSource::Room => {
room_setting_set::<TextToSpeechUserMessagesFlowType>(
bot,
message_context,
&value,
setter_callback,
)
.await
}
SettingsStorageSource::Global => {
global_setting_set::<TextToSpeechUserMessagesFlowType>(
bot,
message_context,
&value,
setter_callback,
)
.await
}
}
}
ConfigTextToSpeechSettingRelatedControllerType::GetSpeedOverride => {
let value = &room_settings.text_to_speech.speed_override;
setting_get::<f32>(bot, message_context, value).await
}
ConfigTextToSpeechSettingRelatedControllerType::SetSpeedOverride(value) => {
let value = value.to_owned();
let setter_callback = Box::new(move |room_settings: &mut RoomSettings| {
room_settings.text_to_speech.speed_override = value;
});
match config_type {
SettingsStorageSource::Room => {
room_setting_set::<f32>(bot, message_context, &value, setter_callback).await
}
SettingsStorageSource::Global => {
global_setting_set::<f32>(bot, message_context, &value, setter_callback).await
}
}
}
ConfigTextToSpeechSettingRelatedControllerType::GetVoiceOverride => {
let value = &room_settings.text_to_speech.voice_override;
setting_get::<String>(bot, message_context, value).await
}
ConfigTextToSpeechSettingRelatedControllerType::SetVoiceOverride(value) => {
let value = value.to_owned();
let value_setter = value.clone();
let setter_callback = Box::new(move |room_settings: &mut RoomSettings| {
room_settings.text_to_speech.voice_override = value_setter;
});
match config_type {
SettingsStorageSource::Room => {
room_setting_set::<String>(bot, message_context, &value, setter_callback).await
}
SettingsStorageSource::Global => {
global_setting_set::<String>(bot, message_context, &value, setter_callback)
.await
}
}
}
}
}

View File

@@ -0,0 +1,38 @@
use mxlink::MessageResponseType;
use crate::entity::{roomconfig::RoomSettings, MessageContext};
use crate::{strings, Bot};
pub async fn handle_set<T>(
bot: &Bot,
message_context: &MessageContext,
value: &Option<T>,
setter_callback: Box<dyn FnOnce(&mut RoomSettings) + Send>,
) -> anyhow::Result<()>
where
T: std::fmt::Display,
{
let mut global_config = message_context.global_config().clone();
setter_callback(&mut global_config.fallback_room_settings);
bot.global_config_manager()
.lock()
.await
.persist(&global_config)
.await?;
let message = match value {
Some(value) => strings::global_config::value_was_set_to(value),
None => strings::global_config::value_was_unset(),
};
bot.messaging()
.send_success_markdown_no_fail(
message_context.room(),
&message,
MessageResponseType::Reply(message_context.thread_info().root_event_id.clone()),
)
.await;
Ok(())
}

View File

@@ -0,0 +1,162 @@
use mxlink::MessageResponseType;
use crate::{
agent::{AgentPurpose, PublicIdentifier},
entity::{globalconfig::GlobalConfigurationManager, MessageContext},
strings, Bot,
};
pub async fn handle_get(
bot: &Bot,
message_context: &MessageContext,
purpose: AgentPurpose,
) -> anyhow::Result<()> {
let agent_id = message_context
.global_config()
.fallback_room_settings
.handler
.get_by_purpose(purpose);
let Some(agent_id) = agent_id else {
bot.messaging()
.send_text_markdown_no_fail(
message_context.room(),
strings::global_config::global_config_lacks_specific_agent_for_purpose(purpose),
MessageResponseType::Reply(message_context.thread_info().root_event_id.clone()),
)
.await;
return Ok(());
};
let agent_identifier = match PublicIdentifier::from_str(agent_id.as_str()) {
Some(agent_identifier) => agent_identifier,
None => {
bot.messaging()
.send_error_markdown_no_fail(
message_context.room(),
&strings::agent::invalid_id_generic(),
MessageResponseType::Reply(message_context.thread_info().root_event_id.clone()),
)
.await;
return Ok(());
}
};
let agent_exists = bot
.agent_manager()
.available_room_agents_by_room_config_context(message_context.room_config_context())
.iter()
.any(|agent| *agent.identifier() == agent_identifier);
if agent_exists {
bot.messaging()
.send_text_markdown_no_fail(
message_context.room(),
strings::global_config::configured_to_use_agent_for_purpose(
&agent_identifier,
purpose,
),
MessageResponseType::Reply(message_context.thread_info().root_event_id.clone()),
)
.await;
} else {
bot.messaging()
.send_text_markdown_no_fail(
message_context.room(),
strings::global_config::configures_agent_for_purpose_but_does_not_exist(
&agent_identifier,
purpose,
),
MessageResponseType::Reply(message_context.thread_info().root_event_id.clone()),
)
.await;
return Ok(());
}
Ok(())
}
pub async fn handle_set(
bot: &Bot,
global_config_manager: &tokio::sync::Mutex<GlobalConfigurationManager>,
message_context: &MessageContext,
purpose: AgentPurpose,
agent_identifier: &Option<PublicIdentifier>,
) -> anyhow::Result<()> {
if let Some(agent_identifier) = agent_identifier {
let is_allowed = match &agent_identifier {
PublicIdentifier::Static(_) => true,
PublicIdentifier::DynamicGlobal(_) => true,
PublicIdentifier::DynamicRoomLocal(_) => false,
};
if !is_allowed {
bot.messaging()
.send_error_markdown_no_fail(
message_context.room(),
&strings::global_config::not_allowed_to_use_agent_in_global_config(
agent_identifier,
),
MessageResponseType::Reply(message_context.thread_info().root_event_id.clone()),
)
.await;
return Ok(());
}
let agent_exists = bot
.agent_manager()
.available_room_agents_by_room_config_context(message_context.room_config_context())
.iter()
.any(|agent| agent.identifier() == agent_identifier);
if !agent_exists {
bot.messaging()
.send_error_markdown_no_fail(
message_context.room(),
&strings::agent::agent_with_given_identifier_missing(agent_identifier),
MessageResponseType::Reply(message_context.thread_info().root_event_id.clone()),
)
.await;
return Ok(());
}
}
let agent_id = agent_identifier
.as_ref()
.map(|agent_identifier| agent_identifier.as_string());
let mut global_config = global_config_manager.lock().await.get_or_create().await?;
global_config
.fallback_room_settings
.handler
.set_by_purpose(purpose, agent_id);
global_config_manager
.lock()
.await
.persist(&global_config)
.await?;
let message = match agent_identifier {
Some(agent_identifier) => {
strings::global_config::reconfigured_to_use_agent_for_purpose(agent_identifier, purpose)
}
None => strings::global_config::reconfigured_to_not_specify_agent_for_purpose(purpose),
};
bot.messaging()
.send_success_markdown_no_fail(
message_context.room(),
&message,
MessageResponseType::Reply(message_context.thread_info().root_event_id.clone()),
)
.await;
Ok(())
}

View File

@@ -0,0 +1,2 @@
pub(super) mod generic_setting;
pub(super) mod handler;

547
src/controller/cfg/help.rs Normal file
View File

@@ -0,0 +1,547 @@
use mxlink::MessageResponseType;
use crate::{
entity::{
roomconfig::{
SpeechToTextFlowType, TextGenerationAutoUsage, TextGenerationPrefixRequirementType,
TextToSpeechBotMessagesFlowType, TextToSpeechUserMessagesFlowType,
},
MessageContext,
},
strings, Bot,
};
pub async fn handle(bot: &Bot, message_context: &MessageContext) -> anyhow::Result<()> {
let mut message = String::new();
message.push_str(&build_section_intro());
message.push_str("\n\n");
message.push_str("\n---\n");
message.push_str(&build_section_status(bot.command_prefix()));
message.push_str("\n\n");
message.push_str("\n---\n");
message.push_str(&build_section_handlers(bot.command_prefix()));
message.push_str("\n\n");
message.push_str("\n---\n");
message.push_str(&build_section_text_generation(
bot.command_prefix(),
bot.user_id().localpart(),
));
message.push_str("\n\n");
message.push_str("\n---\n");
message.push_str(&build_section_text_to_speech(bot.command_prefix()));
message.push_str("\n\n");
message.push_str("\n---\n");
message.push_str(&build_section_speech_to_text(bot.command_prefix()));
message.push_str("\n\n");
message.push_str("\n---\n");
message.push_str(&build_section_image_generation());
bot.messaging()
.send_text_markdown_no_fail(
message_context.room(),
message,
MessageResponseType::Reply(message_context.thread_info().root_event_id.clone()),
)
.await;
Ok(())
}
fn build_section_intro() -> String {
let mut message = String::new();
message.push_str(&format!("## {}", strings::help::cfg::heading()));
message.push_str("\n\n");
message.push_str(&strings::help::cfg::intro_long());
message
}
fn build_section_status(command_prefix: &str) -> String {
let mut message = String::new();
message.push_str(&format!("### {}", strings::help::cfg::status_heading()));
message.push_str("\n\n");
message.push_str(&strings::help::cfg::status_intro(command_prefix));
message
}
fn build_section_handlers(command_prefix: &str) -> String {
let mut message = String::new();
message.push_str(&format!("### {}", strings::help::cfg::handlers_heading()));
message.push_str("\n\n");
message.push_str(&strings::help::cfg::handlers_intro_common());
message.push('\n');
message.push_str(&strings::help::cfg::handlers_intro_purposes());
message.push_str("\n\n");
message.push_str(strings::help::available_commands_intro());
message.push('\n');
message.push_str(&format!(
"- {}",
strings::help::cfg::handlers_show(command_prefix)
));
message.push('\n');
message.push_str(&format!(
"- {}",
strings::help::cfg::handlers_set(command_prefix)
));
message.push('\n');
message.push_str(&format!(
"- {}",
strings::help::cfg::handlers_unset(command_prefix)
));
message
}
fn build_section_text_generation(command_prefix: &str, bot_username: &str) -> String {
let mut message = String::new();
message.push_str(&format!(
"### {}",
strings::help::cfg::text_generation_heading()
));
message.push_str("\n\n");
message.push_str(&strings::help::cfg::text_generation_common());
message.push_str("\n\n");
// Prefix requirement type
message.push_str(&format!(
"#### {}",
strings::help::cfg::text_generation_prefix_requirement_type_heading()
));
message.push_str("\n\n");
message.push_str(&strings::help::cfg::text_generation_prefix_requirement_type_intro());
message.push('\n');
message.push_str(
&strings::help::cfg::the_following_configuration_values_are_recognized(
TextGenerationPrefixRequirementType::choices(),
),
);
message.push_str("\n\n");
message
.push_str(&strings::help::cfg::text_generation_prefix_requirement_type_outro(bot_username));
message.push('\n');
message.push_str(&format!(
"- {}",
&strings::help::cfg::current_setting_show(
command_prefix,
"text-generation prefix-requirement-type"
)
));
message.push('\n');
message.push_str(&format!(
"- {}",
&strings::help::cfg::current_setting_set(
command_prefix,
"text-generation set-prefix-requirement-type VALUE"
)
));
message.push('\n');
message.push_str(&format!(
"- {}",
&strings::help::cfg::current_setting_unset(
command_prefix,
"text-generation set-prefix-requirement-type"
)
));
message.push_str("\n\n");
// Auto Usage
message.push_str(&format!(
"#### {}",
strings::help::cfg::text_generation_auto_usage_heading()
));
message.push_str("\n\n");
message.push_str(&strings::help::cfg::text_generation_auto_usage_intro());
message.push('\n');
message.push_str(
&strings::help::cfg::the_following_configuration_values_are_recognized(
TextGenerationAutoUsage::choices(),
),
);
message.push_str("\n\n");
message.push_str(&format!(
"- {}",
&strings::help::cfg::current_setting_show(command_prefix, "text-generation auto-usage")
));
message.push('\n');
message.push_str(&format!(
"- {}",
&strings::help::cfg::current_setting_set(
command_prefix,
"text-generation set-auto-usage VALUE"
)
));
message.push('\n');
message.push_str(&format!(
"- {}",
&strings::help::cfg::current_setting_unset(
command_prefix,
"text-generation set-auto-usage"
)
));
message.push_str("\n\n");
// Context Management
message.push_str(&format!(
"#### {}",
strings::help::cfg::text_generation_context_management_heading()
));
message.push_str("\n\n");
message.push_str(&strings::help::cfg::text_generation_context_management_intro());
message.push('\n');
message.push_str(
&strings::help::cfg::the_following_configuration_values_are_recognized(vec![true, false]),
);
message.push_str("\n\n");
message.push_str(&format!(
"- {}",
&strings::help::cfg::current_setting_show(
command_prefix,
"text-generation context-management-enabled"
)
));
message.push('\n');
message.push_str(&format!(
"- {}",
&strings::help::cfg::current_setting_set(
command_prefix,
"text-generation set-context-management-enabled VALUE"
)
));
message.push('\n');
message.push_str(&format!(
"- {}",
&strings::help::cfg::current_setting_unset(
command_prefix,
"text-generation set-context-management-enabled"
)
));
message.push_str("\n\n");
// Prompt override
message.push_str(&format!(
"#### {}",
strings::help::cfg::text_generation_prompt_override_heading()
));
message.push_str("\n\n");
message.push_str(&strings::help::cfg::text_generation_prompt_override_intro());
message.push_str("\n\n");
message.push_str(&format!(
"- {}",
&strings::help::cfg::current_setting_show(
command_prefix,
"text-generation prompt-override"
)
));
message.push('\n');
message.push_str(&format!(
"- {}",
&strings::help::cfg::current_setting_set(
command_prefix,
"text-generation set-prompt-override VALUE"
)
));
message.push('\n');
message.push_str(&format!(
"- {}",
&strings::help::cfg::current_setting_unset(
command_prefix,
"text-generation set-prompt-override"
)
));
message.push_str("\n\n");
// Speed override
message.push_str(&format!(
"#### {}",
strings::help::cfg::text_generation_temperature_override_heading()
));
message.push_str("\n\n");
message.push_str(&strings::help::cfg::text_generation_temperature_override_intro());
message.push_str("\n\n");
message.push_str(&format!(
"- {}",
&strings::help::cfg::current_setting_show(
command_prefix,
"text-generation temperature-override"
)
));
message.push('\n');
message.push_str(&format!(
"- {}",
&strings::help::cfg::current_setting_set(
command_prefix,
"text-generation set-temperature-override VALUE"
)
));
message.push('\n');
message.push_str(&format!(
"- {}",
&strings::help::cfg::current_setting_unset(
command_prefix,
"text-generation set-temperature-override"
)
));
message
}
fn build_section_speech_to_text(command_prefix: &str) -> String {
let mut message = String::new();
message.push_str(&format!(
"### {}",
strings::help::cfg::speech_to_text_heading()
));
message.push_str("\n\n");
message.push_str(&strings::help::cfg::speech_to_text_common());
message.push_str("\n\n");
// Flow Type
message.push_str(&format!(
"#### {}",
strings::help::cfg::speech_to_text_flow_type_heading()
));
message.push_str("\n\n");
message.push_str(strings::help::cfg::speech_to_text_flow_type_intro());
message.push('\n');
message.push_str(
&strings::help::cfg::the_following_configuration_values_are_recognized(
SpeechToTextFlowType::choices(),
),
);
message.push_str("\n\n");
message.push_str(&format!(
"- {}",
&strings::help::cfg::current_setting_show(command_prefix, "speech-to-text flow-type")
));
message.push('\n');
message.push_str(&format!(
"- {}",
&strings::help::cfg::current_setting_set(
command_prefix,
"speech-to-text set-flow-type VALUE"
)
));
message.push('\n');
message.push_str(&format!(
"- {}",
&strings::help::cfg::current_setting_unset(command_prefix, "speech-to-text set-flow-type")
));
message.push_str("\n\n");
// Language
message.push_str(&format!(
"#### {}",
strings::help::cfg::speech_to_text_language_heading()
));
message.push_str("\n\n");
message.push_str(strings::help::cfg::speech_to_text_language_intro());
message.push_str("\n\n");
message.push_str(&format!(
"- {}",
&strings::help::cfg::current_setting_show(command_prefix, "speech-to-text language")
));
message.push('\n');
message.push_str(&format!(
"- {}",
&strings::help::cfg::current_setting_set(
command_prefix,
"speech-to-text set-language VALUE"
)
));
message.push('\n');
message.push_str(&format!(
"- {}",
&strings::help::cfg::current_setting_unset(command_prefix, "speech-to-text set-language")
));
message.push_str("\n\n");
message
}
fn build_section_text_to_speech(command_prefix: &str) -> String {
let mut message = String::new();
message.push_str(&format!(
"### {}",
strings::help::cfg::text_to_speech_heading()
));
message.push_str("\n\n");
message.push_str(strings::help::cfg::text_to_speech_common());
message.push_str("\n\n");
// Bot Messages Flow Type
message.push_str(&format!(
"#### {}",
strings::help::cfg::text_to_speech_bot_msgs_flow_type_heading()
));
message.push_str("\n\n");
message.push_str(&strings::help::cfg::text_to_speech_bot_msgs_flow_type_intro());
message.push('\n');
message.push_str(
&strings::help::cfg::the_following_configuration_values_are_recognized(
TextToSpeechBotMessagesFlowType::choices(),
),
);
message.push_str("\n\n");
message.push_str(&format!(
"- {}",
&strings::help::cfg::current_setting_show(
command_prefix,
"text-to-speech bot-msgs-flow-type"
)
));
message.push('\n');
message.push_str(&format!(
"- {}",
&strings::help::cfg::current_setting_set(
command_prefix,
"text-to-speech set-bot-msgs-flow-type VALUE"
)
));
message.push('\n');
message.push_str(&format!(
"- {}",
&strings::help::cfg::current_setting_unset(
command_prefix,
"text-to-speech set-bot-msgs-flow-type"
)
));
message.push_str("\n\n");
// User Messages Flow Type
message.push_str(&format!(
"#### {}",
strings::help::cfg::text_to_speech_user_msgs_flow_type_heading()
));
message.push_str("\n\n");
message.push_str(&strings::help::cfg::text_to_speech_user_msgs_flow_type_intro());
message.push('\n');
message.push_str(
&strings::help::cfg::the_following_configuration_values_are_recognized(
TextToSpeechUserMessagesFlowType::choices(),
),
);
message.push_str("\n\n");
message.push_str(&format!(
"- {}",
&strings::help::cfg::current_setting_show(
command_prefix,
"text-to-speech user-msgs-flow-type"
)
));
message.push('\n');
message.push_str(&format!(
"- {}",
&strings::help::cfg::current_setting_set(
command_prefix,
"text-to-speech set-user-msgs-flow-type VALUE"
)
));
message.push('\n');
message.push_str(&format!(
"- {}",
&strings::help::cfg::current_setting_unset(
command_prefix,
"text-to-speech set-user-msgs-flow-type"
)
));
message.push_str("\n\n");
// Speed override
message.push_str(&format!(
"#### {}",
strings::help::cfg::text_to_speech_speed_override_heading()
));
message.push_str("\n\n");
message.push_str(&strings::help::cfg::text_to_speech_speed_override_intro());
message.push_str("\n\n");
message.push_str(&format!(
"- {}",
&strings::help::cfg::current_setting_show(command_prefix, "text-to-speech speed-override")
));
message.push('\n');
message.push_str(&format!(
"- {}",
&strings::help::cfg::current_setting_set(
command_prefix,
"text-to-speech set-speed-override VALUE"
)
));
message.push('\n');
message.push_str(&format!(
"- {}",
&strings::help::cfg::current_setting_unset(
command_prefix,
"text-to-speech set-speed-override"
)
));
message.push_str("\n\n");
// Voice override
message.push_str(&format!(
"#### {}",
strings::help::cfg::text_to_speech_voice_override_heading()
));
message.push_str("\n\n");
message.push_str(&strings::help::cfg::text_to_speech_voice_override_intro());
message.push_str("\n\n");
message.push_str(&format!(
"- {}",
&strings::help::cfg::current_setting_show(command_prefix, "text-to-speech voice-override")
));
message.push('\n');
message.push_str(&format!(
"- {}",
&strings::help::cfg::current_setting_set(
command_prefix,
"text-to-speech set-voice-override VALUE"
)
));
message.push('\n');
message.push_str(&format!(
"- {}",
&strings::help::cfg::current_setting_unset(
command_prefix,
"text-to-speech set-voice-override"
)
));
message.push_str("\n\n");
message
}
fn build_section_image_generation() -> String {
let mut message = String::new();
message.push_str(&format!(
"### {}",
strings::help::cfg::image_generation_heading()
));
message.push_str("\n\n");
message.push_str(strings::help::cfg::image_generation_common());
message
}

12
src/controller/cfg/mod.rs Normal file
View File

@@ -0,0 +1,12 @@
mod common;
mod controller_type;
mod determination;
mod dispatching;
mod global_config;
mod help;
mod room_config;
mod status;
pub use controller_type::ConfigControllerType;
pub use determination::determine_controller;
pub use dispatching::dispatch_controller;

View File

@@ -0,0 +1,38 @@
use mxlink::MessageResponseType;
use crate::entity::{roomconfig::RoomSettings, MessageContext};
use crate::{strings, Bot};
pub async fn handle_set<T>(
bot: &Bot,
message_context: &MessageContext,
value: &Option<T>,
setter_callback: Box<dyn FnOnce(&mut RoomSettings) + Send>,
) -> anyhow::Result<()>
where
T: std::fmt::Display,
{
let mut room_config = message_context.room_config().clone();
setter_callback(&mut room_config.settings);
bot.room_config_manager()
.lock()
.await
.persist(message_context.room(), &room_config)
.await?;
let message = match value {
Some(value) => strings::room_config::value_was_set_to(value),
None => strings::room_config::value_was_unset(),
};
bot.messaging()
.send_success_markdown_no_fail(
message_context.room(),
&message,
MessageResponseType::Reply(message_context.thread_info().root_event_id.clone()),
)
.await;
Ok(())
}

View File

@@ -0,0 +1,141 @@
use mxlink::MessageResponseType;
use crate::{
agent::{AgentPurpose, PublicIdentifier},
entity::MessageContext,
strings, Bot,
};
use crate::entity::roomconfig::RoomConfigurationManager;
pub async fn handle_get(
bot: &Bot,
message_context: &MessageContext,
purpose: AgentPurpose,
) -> anyhow::Result<()> {
let agent_id = message_context
.room_config()
.settings
.handler
.get_by_purpose(purpose);
let Some(agent_id) = agent_id else {
bot.messaging()
.send_error_markdown_no_fail(
message_context.room(),
&strings::room_config::room_not_configured_with_specific_agent_for_purpose(purpose),
MessageResponseType::Reply(message_context.thread_info().root_event_id.clone()),
)
.await;
return Ok(());
};
let Some(agent_identifier) = PublicIdentifier::from_str(agent_id.as_str()) else {
bot.messaging()
.send_error_markdown_no_fail(
message_context.room(),
&strings::agent::invalid_id_generic(),
MessageResponseType::Reply(message_context.thread_info().root_event_id.clone()),
)
.await;
return Ok(());
};
let agent_exists = bot
.agent_manager()
.available_room_agents_by_room_config_context(message_context.room_config_context())
.iter()
.any(|agent| *agent.identifier() == agent_identifier);
if agent_exists {
bot.messaging()
.send_text_markdown_no_fail(
message_context.room(),
strings::room_config::configured_to_use_agent_for_purpose(
&agent_identifier,
purpose,
),
MessageResponseType::Reply(message_context.thread_info().root_event_id.clone()),
)
.await;
} else {
bot.messaging()
.send_text_markdown_no_fail(
message_context.room(),
strings::room_config::configures_agent_for_purpose_but_does_not_exist(
&agent_identifier,
purpose,
),
MessageResponseType::Reply(message_context.thread_info().root_event_id.clone()),
)
.await;
return Ok(());
}
Ok(())
}
pub async fn handle_set(
bot: &Bot,
room_config_manager: &tokio::sync::Mutex<RoomConfigurationManager>,
message_context: &MessageContext,
purpose: AgentPurpose,
agent_identifier: &Option<PublicIdentifier>,
) -> anyhow::Result<()> {
if let Some(agent_identifier) = agent_identifier {
let agent_exists = bot
.agent_manager()
.available_room_agents_by_room_config_context(message_context.room_config_context())
.iter()
.any(|agent| agent.identifier() == agent_identifier);
if !agent_exists {
bot.messaging()
.send_error_markdown_no_fail(
message_context.room(),
&strings::agent::agent_with_given_identifier_missing(agent_identifier),
MessageResponseType::Reply(message_context.thread_info().root_event_id.clone()),
)
.await;
return Ok(());
}
}
let mut new_room_config = message_context.room_config().clone();
let agent_id = agent_identifier
.as_ref()
.map(|agent_identifier| agent_identifier.as_string());
new_room_config
.settings
.handler
.set_by_purpose(purpose, agent_id);
room_config_manager
.lock()
.await
.persist(message_context.room(), &new_room_config)
.await?;
let message = match agent_identifier {
Some(agent_identifier) => {
strings::room_config::reconfigured_to_use_agent_for_purpose(agent_identifier, purpose)
}
None => strings::room_config::reconfigured_to_not_specify_agent_for_purpose(purpose),
};
bot.messaging()
.send_success_markdown_no_fail(
message_context.room(),
&message,
MessageResponseType::Reply(message_context.thread_info().root_event_id.clone()),
)
.await;
Ok(())
}

View File

@@ -0,0 +1,2 @@
pub(super) mod generic_setting;
pub(super) mod handler;

View File

@@ -0,0 +1,759 @@
use mxlink::MessageResponseType;
use crate::{
agent::{
utils::get_effective_agent_for_purpose, AgentInstance, AgentPurpose, ControllerTrait,
Manager as AgentManager, PublicIdentifier,
},
entity::{
roomconfig::{RoomConfig, RoomSettingsHandler},
MessageContext, RoomConfigContext,
},
strings, Bot,
};
pub async fn handle(bot: &Bot, message_context: &MessageContext) -> anyhow::Result<()> {
let mut message = String::new();
let agent_manager = bot.agent_manager();
let agents = agent_manager
.available_room_agents_by_room_config_context(message_context.room_config_context());
// Room handlers
message.push_str(&generate_room_handlers_section(
&message_context.room_config().settings.handler,
&agents,
bot.command_prefix(),
));
message.push_str("\n\n");
// Global handlers
message.push_str(&generate_global_config_handlers_section(
&message_context
.global_config()
.fallback_room_settings
.handler,
&agents,
bot.command_prefix(),
));
message.push_str("\n\n");
// Agents
message.push_str(&generate_room_agents_section(
message_context.room_config(),
&agents,
bot.command_prefix(),
));
message.push_str("\n\n");
// Text Generation
message.push_str(
&generate_text_generation_section(agent_manager, message_context.room_config_context())
.await,
);
message.push_str("\n\n");
// Text-to-Speech
message.push_str(
&generate_text_to_speech_section(agent_manager, message_context.room_config_context())
.await,
);
message.push_str("\n\n");
// Speech-to-Text
message.push_str(
&generate_speech_to_text_section(agent_manager, message_context.room_config_context())
.await,
);
message.push_str("\n\n");
// Image Generation
message.push_str(
&generate_image_generation_section(agent_manager, message_context.room_config_context())
.await,
);
message.push_str("\n\n");
bot.messaging()
.send_text_markdown_no_fail(
message_context.room(),
message,
MessageResponseType::Reply(message_context.thread_info().root_event_id.clone()),
)
.await;
Ok(())
}
fn generate_room_handlers_section(
handler_config: &RoomSettingsHandler,
agents: &[AgentInstance],
command_prefix: &str,
) -> String {
let mut message = String::new();
message.push_str(
format!(
"## {}\n",
strings::cfg::status_room_config_handlers_heading()
)
.as_str(),
);
message.push_str(strings::cfg::status_room_config_handlers_intro());
message.push_str("\n\n");
for purpose in AgentPurpose::choices() {
message.push_str(&generate_handler_line_for_purpose(
purpose,
handler_config,
agents,
false,
));
message.push('\n');
}
message.push_str("\n\n");
message.push_str(strings::cfg::status_room_config_handlers_outro(command_prefix).as_str());
message
}
fn generate_global_config_handlers_section(
handler_config: &RoomSettingsHandler,
agents: &[AgentInstance],
command_prefix: &str,
) -> String {
let mut message = String::new();
message.push_str(
format!(
"## {}\n",
strings::cfg::status_global_config_handlers_heading()
)
.as_str(),
);
message.push_str(strings::cfg::status_global_config_handlers_intro());
message.push_str("\n\n");
for purpose in AgentPurpose::choices() {
message.push_str(&generate_handler_line_for_purpose(
purpose,
handler_config,
agents,
true,
));
message.push('\n');
}
message.push_str("\n\n");
message.push_str(strings::cfg::status_global_config_handlers_outro(command_prefix).as_str());
message
}
fn generate_handler_line_for_purpose(
purpose: &AgentPurpose,
handler_config: &RoomSettingsHandler,
agents: &[AgentInstance],
is_for_global_config: bool,
) -> String {
let agent_id = handler_config.get_by_purpose(*purpose);
match agent_id {
Some(agent_id) => {
let agent = agents
.iter()
.find(|a| *a.identifier().as_string() == agent_id);
strings::cfg::status_handler_line_agent_found(purpose, &agent_id, agent)
}
None => match purpose {
AgentPurpose::CatchAll => {
if is_for_global_config {
return strings::cfg::status_handler_line_catch_all_agent_not_set_globally();
}
strings::cfg::status_handler_line_catch_all_agent_not_set_in_room_default_to_global(
)
}
_ => {
if is_for_global_config {
return strings::cfg::status_handler_line_non_catch_all_agent_not_set_globally(
purpose,
);
}
strings::cfg::status_handler_line_non_catch_all_agent_not_set_in_room_default_to_global(
purpose,
)
}
},
}
}
fn generate_room_agents_section(
room_config: &RoomConfig,
agents: &Vec<AgentInstance>,
command_prefix: &str,
) -> String {
let mut message = String::new();
message.push_str(format!("## {}\n", strings::cfg::status_room_agents_heading()).as_str());
if room_config.agents.is_empty() {
message.push_str(strings::cfg::status_room_agents_empty());
message.push_str("\n\n");
message.push_str(&strings::help::learn_more_send_a_command(
command_prefix,
"agent",
));
return message;
}
message.push_str(strings::cfg::status_room_agents_intro());
message.push_str("\n\n");
for agent in agents {
let PublicIdentifier::DynamicRoomLocal(_) = agent.identifier() else {
continue;
};
message.push_str(&format!(
"- `{}` ({})\n",
agent.identifier(),
strings::agent::create_support_badges_text(agent.controller()),
));
}
message.push_str("\n\n");
message.push_str(strings::cfg::status_room_agents_outro(command_prefix).as_str());
message
}
async fn generate_text_generation_section(
agent_manager: &AgentManager,
room_config_context: &RoomConfigContext,
) -> String {
let mut message = String::new();
message.push_str(format!("## {}\n", strings::cfg::status_text_generation_heading()).as_str());
let text_generation_agent_info = get_effective_agent_for_purpose(
agent_manager,
room_config_context,
AgentPurpose::TextGeneration,
)
.await;
// Effective agent
let text_generation_agent = match text_generation_agent_info {
Ok(text_generation_agent_info) => {
message.push_str(&strings::cfg::status_entry_effective_agent(
text_generation_agent_info.instance.identifier(),
text_generation_agent_info.configuration_source,
));
Some(text_generation_agent_info.instance)
}
Err(err) => {
tracing::error!(?err, "Failed to determine text-generation agent");
message.push_str(&strings::cfg::status_entry_effective_agent_error());
None
}
};
// Prefix requirement type
let effective_prefix_requirement_type =
room_config_context.text_generation_prefix_requirement_type();
let room_config_prefix_requirement_type = room_config_context
.room_config
.settings
.text_generation
.prefix_requirement_type;
let global_config_prefix_requirement_type = room_config_context
.global_config
.fallback_room_settings
.text_generation
.prefix_requirement_type;
let prefix_requirement_type_set_where = if room_config_prefix_requirement_type.is_some() {
strings::cfg::status_badge_set_in_room_config()
} else if global_config_prefix_requirement_type.is_some() {
strings::cfg::status_badge_set_in_global_config()
} else {
strings::cfg::status_badge_using_hardcoded_default()
};
message.push_str(
&strings::cfg::status_text_generation_entry_prefix_requirement_type(
effective_prefix_requirement_type,
prefix_requirement_type_set_where,
),
);
// Auto usage
let effective_auto_usage = room_config_context.auto_text_generation_usage();
let room_config_auto_usage = room_config_context
.room_config
.settings
.text_generation
.auto_usage;
let global_config_auto_usage = room_config_context
.global_config
.fallback_room_settings
.text_generation
.auto_usage;
let auto_usage_set_where = if room_config_auto_usage.is_some() {
strings::cfg::status_badge_set_in_room_config()
} else if global_config_auto_usage.is_some() {
strings::cfg::status_badge_set_in_global_config()
} else {
strings::cfg::status_badge_using_hardcoded_default()
};
message.push_str(&strings::cfg::status_text_generation_entry_auto_usage(
effective_auto_usage,
auto_usage_set_where,
));
// Context Management
let effective_context_management =
room_config_context.text_generation_context_management_enabled();
let room_config_context_management = room_config_context
.room_config
.settings
.text_generation
.context_management_enabled;
let global_config_context_management = room_config_context
.global_config
.fallback_room_settings
.text_generation
.context_management_enabled;
let context_management_set_where = if room_config_context_management.is_some() {
strings::cfg::status_badge_set_in_room_config()
} else if global_config_context_management.is_some() {
strings::cfg::status_badge_set_in_global_config()
} else {
strings::cfg::status_badge_using_hardcoded_default()
};
message.push_str(
&strings::cfg::status_text_generation_entry_context_management(
effective_context_management,
context_management_set_where,
),
);
// Prompt override
let text_agent_prompt = if let Some(text_generation_agent) = &text_generation_agent {
text_generation_agent.controller().text_generation_prompt()
} else {
None
};
let room_config_prompt_override = room_config_context
.room_config
.settings
.text_generation
.prompt_override
.clone();
let global_config_prompt_override = room_config_context
.global_config
.fallback_room_settings
.text_generation
.prompt_override
.clone();
let (prompt, prompt_set_where) =
if let Some(room_config_prompt_override) = room_config_prompt_override {
(
room_config_prompt_override,
strings::cfg::status_badge_set_in_room_config(),
)
} else if let Some(global_config_prompt_override) = global_config_prompt_override {
(
global_config_prompt_override,
strings::cfg::status_badge_set_in_global_config(),
)
} else {
(
text_agent_prompt.unwrap_or("".to_owned()),
strings::cfg::status_badge_set_in_agent_config(),
)
};
message.push_str(&strings::cfg::status_text_generation_entry_prompt(
&prompt,
prompt_set_where,
));
// Temperature
let text_agent_temperature = if let Some(text_generation_agent) = &text_generation_agent {
text_generation_agent
.controller()
.text_generation_temperature()
} else {
None
};
let room_config_temperature_override = room_config_context
.room_config
.settings
.text_generation
.temperature_override;
let global_config_temperature_override = room_config_context
.global_config
.fallback_room_settings
.text_generation
.temperature_override;
let (effective_temperature, set_where) = if let Some(room_config_temperature_override) =
room_config_temperature_override
{
(
Some(room_config_temperature_override),
strings::cfg::status_badge_set_in_room_config(),
)
} else if let Some(global_config_temperature_override) = global_config_temperature_override {
(
Some(global_config_temperature_override),
strings::cfg::status_badge_set_in_global_config(),
)
} else {
(
text_agent_temperature,
strings::cfg::status_badge_set_in_agent_config(),
)
};
message.push_str(&strings::cfg::status_text_generation_entry_temperature(
effective_temperature,
set_where,
));
message
}
async fn generate_speech_to_text_section(
agent_manager: &AgentManager,
room_config_context: &RoomConfigContext,
) -> String {
let mut message = String::new();
message.push_str(format!("## {}\n", strings::cfg::status_speech_to_text_heading()).as_str());
let speech_to_text_agent_info = get_effective_agent_for_purpose(
agent_manager,
room_config_context,
AgentPurpose::SpeechToText,
)
.await;
// Effective agent
match speech_to_text_agent_info {
Ok(speech_to_text_agent_info) => {
message.push_str(&strings::cfg::status_entry_effective_agent(
speech_to_text_agent_info.instance.identifier(),
speech_to_text_agent_info.configuration_source,
));
}
Err(err) => {
tracing::error!(?err, "Failed to determine speech-to-text agent");
message.push_str(&strings::cfg::status_entry_effective_agent_error());
}
};
// Flow type
let effective_flow_type = room_config_context.speech_to_text_flow_type();
let room_config_flow_type = room_config_context
.room_config
.settings
.speech_to_text
.flow_type;
let global_config_flow_type = room_config_context
.global_config
.fallback_room_settings
.speech_to_text
.flow_type;
let flow_type_set_where = if room_config_flow_type.is_some() {
strings::cfg::status_badge_set_in_room_config()
} else if global_config_flow_type.is_some() {
strings::cfg::status_badge_set_in_global_config()
} else {
strings::cfg::status_badge_using_hardcoded_default()
};
message.push_str(&strings::cfg::status_speech_to_text_entry_flow_type(
effective_flow_type,
flow_type_set_where,
));
// Language
let effective_language = room_config_context.speech_to_text_language();
let room_config_language = &room_config_context
.room_config
.settings
.speech_to_text
.language;
let global_config_language = &room_config_context
.global_config
.fallback_room_settings
.speech_to_text
.language;
let language_set_where = if room_config_language.is_some() {
strings::cfg::status_badge_set_in_room_config()
} else if global_config_language.is_some() {
strings::cfg::status_badge_set_in_global_config()
} else {
strings::cfg::status_badge_using_hardcoded_default()
};
message.push_str(&strings::cfg::status_speech_to_text_entry_language(
effective_language,
language_set_where,
));
message
}
async fn generate_text_to_speech_section(
agent_manager: &AgentManager,
room_config_context: &RoomConfigContext,
) -> String {
let mut message = String::new();
message.push_str(format!("## {}\n", strings::cfg::status_text_to_speech_heading()).as_str());
let text_to_speech_agent_info = get_effective_agent_for_purpose(
agent_manager,
room_config_context,
AgentPurpose::TextToSpeech,
)
.await;
// Effective agent
let text_to_speech_agent = match text_to_speech_agent_info {
Ok(text_to_speech_agent_info) => {
message.push_str(&strings::cfg::status_entry_effective_agent(
text_to_speech_agent_info.instance.identifier(),
text_to_speech_agent_info.configuration_source,
));
Some(text_to_speech_agent_info.instance)
}
Err(err) => {
tracing::error!(?err, "Failed to determine text-to-speech agent");
message.push_str(&strings::cfg::status_entry_effective_agent_error());
None
}
};
// Bot messages flow type
let effective_bot_messages_tts_flow_type =
room_config_context.text_to_speech_bot_messages_flow_type();
let room_config_bot_messages_tts_flow_type = room_config_context
.room_config
.settings
.text_to_speech
.bot_msgs_flow_type;
let global_config_bot_messages_tts_flow_type = room_config_context
.global_config
.fallback_room_settings
.text_to_speech
.bot_msgs_flow_type;
let bot_messages_tts_flow_type_set_where = if room_config_bot_messages_tts_flow_type.is_some() {
strings::cfg::status_badge_set_in_room_config()
} else if global_config_bot_messages_tts_flow_type.is_some() {
strings::cfg::status_badge_set_in_global_config()
} else {
strings::cfg::status_badge_using_hardcoded_default()
};
message.push_str(
&strings::cfg::status_text_to_speech_entry_bot_msgs_flow_type(
effective_bot_messages_tts_flow_type,
bot_messages_tts_flow_type_set_where,
),
);
// User messages flow type
let effective_user_messages_tts_flow_type =
room_config_context.text_to_speech_user_messages_flow_type();
let room_config_user_messages_tts_flow_type = room_config_context
.room_config
.settings
.text_to_speech
.user_msgs_flow_type;
let global_config_user_messages_tts_flow_type = room_config_context
.global_config
.fallback_room_settings
.text_to_speech
.user_msgs_flow_type;
let user_messages_tts_flow_type_set_where = if room_config_user_messages_tts_flow_type.is_some()
{
strings::cfg::status_badge_set_in_room_config()
} else if global_config_user_messages_tts_flow_type.is_some() {
strings::cfg::status_badge_set_in_global_config()
} else {
strings::cfg::status_badge_using_hardcoded_default()
};
message.push_str(
&strings::cfg::status_text_to_speech_entry_user_msgs_flow_type(
effective_user_messages_tts_flow_type,
user_messages_tts_flow_type_set_where,
),
);
// Speed
let agent_speed = if let Some(text_to_speech_agent) = &text_to_speech_agent {
text_to_speech_agent.controller().text_to_speech_speed()
} else {
None
};
let room_config_speed_override = room_config_context
.room_config
.settings
.text_to_speech
.speed_override;
let global_config_speed_override = room_config_context
.global_config
.fallback_room_settings
.text_to_speech
.speed_override;
let (effective_speed, set_where) =
if let Some(room_config_speed_override) = room_config_speed_override {
(
Some(room_config_speed_override),
strings::cfg::status_badge_set_in_room_config(),
)
} else if let Some(global_config_speed_override) = global_config_speed_override {
(
Some(global_config_speed_override),
strings::cfg::status_badge_set_in_global_config(),
)
} else if agent_speed.is_some() {
(
agent_speed,
strings::cfg::status_badge_set_in_agent_config(),
)
} else {
(None, strings::cfg::status_badge_using_hardcoded_default())
};
message.push_str(&strings::cfg::status_text_to_speech_entry_speed(
effective_speed,
set_where,
));
// Voice
let agent_voice = if let Some(text_to_speech_agent) = &text_to_speech_agent {
text_to_speech_agent.controller().text_to_speech_voice()
} else {
None
};
let room_config_voice_override = room_config_context
.room_config
.settings
.text_to_speech
.voice_override
.clone();
let global_config_voice_override = room_config_context
.global_config
.fallback_room_settings
.text_to_speech
.voice_override
.clone();
let (effective_voice, set_where) =
if let Some(room_config_voice_override) = room_config_voice_override {
(
Some(room_config_voice_override),
strings::cfg::status_badge_set_in_room_config(),
)
} else if let Some(global_config_voice_override) = global_config_voice_override {
(
Some(global_config_voice_override),
strings::cfg::status_badge_set_in_global_config(),
)
} else if agent_voice.is_some() {
(
agent_voice,
strings::cfg::status_badge_set_in_agent_config(),
)
} else {
(None, strings::cfg::status_badge_using_hardcoded_default())
};
message.push_str(&strings::cfg::status_text_to_speech_entry_voice(
effective_voice,
set_where,
));
message
}
async fn generate_image_generation_section(
agent_manager: &AgentManager,
room_config_context: &RoomConfigContext,
) -> String {
let mut message = String::new();
message.push_str(format!("## {}\n", strings::cfg::status_image_generation_heading()).as_str());
let image_generation_agent_info = get_effective_agent_for_purpose(
agent_manager,
room_config_context,
AgentPurpose::ImageGeneration,
)
.await;
// Effective agent
let _image_generation_agent = match image_generation_agent_info {
Ok(image_generation_agent_info) => {
message.push_str(&strings::cfg::status_entry_effective_agent(
image_generation_agent_info.instance.identifier(),
image_generation_agent_info.configuration_source,
));
Some(image_generation_agent_info.instance)
}
Err(err) => {
tracing::error!(?err, "Failed to determine image generation agent");
message.push_str(&strings::cfg::status_entry_effective_agent_error());
None
}
};
message
}

View File

@@ -0,0 +1,599 @@
use mxlink::matrix_sdk::ruma::events::room::message::AudioMessageEventContent;
use mxlink::matrix_sdk::ruma::OwnedEventId;
use mxlink::{MatrixLink, MessageResponseType};
use tracing::Instrument;
use crate::agent::provider::{SpeechToTextParams, TextGenerationParams};
use crate::agent::AgentInstance;
use crate::agent::AgentPurpose;
use crate::agent::ControllerTrait;
use crate::controller::utils::agent::get_effective_agent_for_purpose_or_complain;
use crate::conversation::matrix::MatrixMessageProcessingParams;
use crate::entity::roomconfig::{
SpeechToTextFlowType, TextToSpeechBotMessagesFlowType, TextToSpeechUserMessagesFlowType,
};
use crate::entity::MessagePayload;
use crate::strings;
use crate::utils::text_to_speech::create_transcribed_message_text;
use crate::{conversation::create_llm_conversation_for_matrix_thread, entity::MessageContext, Bot};
#[derive(Debug, PartialEq)]
pub enum ChatCompletionControllerType {
ViaText { prefixes_to_strip: Vec<String> },
ViaAudio,
}
struct TextToSpeechEligiblePayload {
text: String,
event_id: OwnedEventId,
}
enum TextToSpeechParams {
Perform(TextToSpeechEligiblePayload, MessageResponseType),
Offer(TextToSpeechEligiblePayload, MessageResponseType),
}
pub async fn handle(
bot: &Bot,
matrix_link: MatrixLink,
message_context: &MessageContext,
controller_type: &ChatCompletionControllerType,
) -> anyhow::Result<()> {
let mut original_message_is_audio = false;
let speech_to_text_flow_type = message_context
.room_config_context()
.speech_to_text_flow_type();
let mut speech_to_text_created_event_id: Option<OwnedEventId> = None;
if let MessagePayload::Audio(audio_content) = &message_context.payload() {
original_message_is_audio = true;
let response_type = match speech_to_text_flow_type {
SpeechToTextFlowType::Ignore => {
tracing::debug!("Intentionally ignoring audio message");
return Ok(());
}
SpeechToTextFlowType::TranscribeAndGenerateText => {
tracing::debug!("Will be trascribing and possibly generating text..");
MessageResponseType::InThread(message_context.thread_info().clone())
}
SpeechToTextFlowType::OnlyTranscribe => {
tracing::debug!("Will only be trascribing audio to text..");
if message_context.thread_info().is_thread_root_only() {
MessageResponseType::Reply(message_context.thread_info().root_event_id.clone())
} else {
MessageResponseType::InThread(message_context.thread_info().clone())
}
}
};
let Some(speech_to_text_created_event_id_result) =
handle_stage_speech_to_text(bot, message_context, audio_content, response_type).await
else {
return Ok(());
};
speech_to_text_created_event_id = Some(speech_to_text_created_event_id_result);
if speech_to_text_flow_type == SpeechToTextFlowType::OnlyTranscribe {
tracing::debug!(
"Intentionally not continuing with text generation after transcription"
);
return Ok(());
}
// We've pushed a transcription to the room.
// Let's proceed below where we potentially handle text-generation.
}
let text_to_speech_stage_params: Option<TextToSpeechParams>;
if message_context
.room_config_context()
.should_auto_text_generate(original_message_is_audio)
{
let speech_to_text_created_event_id_reaction_event_id =
if let Some(speech_to_text_created_event_id) = speech_to_text_created_event_id {
let reaction_event_response = bot
.reacting()
.react_no_fail(
message_context.room(),
speech_to_text_created_event_id.clone(),
strings::PROGRESS_INDICATOR_EMOJI.to_owned(),
)
.await;
reaction_event_response
.map(|reaction_event_response| reaction_event_response.event_id)
} else {
None
};
let response_type = MessageResponseType::InThread(message_context.thread_info().clone());
let text_to_speech_eligible_payload = handle_stage_text_generation(
bot,
matrix_link.clone(),
message_context,
controller_type,
response_type.clone(),
)
.await;
if let Some(speech_to_text_created_event_id_reaction_event_id) =
speech_to_text_created_event_id_reaction_event_id
{
bot.messaging()
.redact_event_no_fail(
message_context.room(),
speech_to_text_created_event_id_reaction_event_id,
Some("Done".to_owned()),
)
.await;
}
// If no text was generated (due to some issue), there's no point in continuing.
let Some(text_to_speech_eligible_payload) = text_to_speech_eligible_payload else {
return Ok(());
};
text_to_speech_stage_params = match message_context
.room_config_context()
.text_to_speech_bot_messages_flow_type()
{
TextToSpeechBotMessagesFlowType::Never => None,
TextToSpeechBotMessagesFlowType::OnDemandAlways => Some(TextToSpeechParams::Offer(
text_to_speech_eligible_payload,
response_type,
)),
TextToSpeechBotMessagesFlowType::OnDemandForVoice => {
if original_message_is_audio {
Some(TextToSpeechParams::Offer(
text_to_speech_eligible_payload,
response_type,
))
} else {
None
}
}
TextToSpeechBotMessagesFlowType::OnlyForVoice => {
if original_message_is_audio {
Some(TextToSpeechParams::Perform(
text_to_speech_eligible_payload,
response_type,
))
} else {
None
}
}
TextToSpeechBotMessagesFlowType::Always => Some(TextToSpeechParams::Perform(
text_to_speech_eligible_payload,
response_type,
)),
};
} else {
tracing::debug!("Not generating text due to auto-usage configuration");
let response_type = MessageResponseType::Reply(message_context.event_id().clone());
// If we got text from the user, perhaps it's eligible for text-to-speech.
let MessagePayload::Text(text_payload) = &message_context.payload() else {
// Audio message, or a notice or something else.
// We don't wish to proceed with potential TTS for non-text messages.
return Ok(());
};
let text_to_speech_eligible_payload = TextToSpeechEligiblePayload {
text: text_payload.body.clone(),
event_id: message_context.event_id().clone(),
};
text_to_speech_stage_params = match message_context
.room_config_context()
.text_to_speech_user_messages_flow_type()
{
TextToSpeechUserMessagesFlowType::Never => None,
TextToSpeechUserMessagesFlowType::OnDemand => Some(TextToSpeechParams::Offer(
text_to_speech_eligible_payload,
response_type,
)),
TextToSpeechUserMessagesFlowType::Always => Some(TextToSpeechParams::Perform(
text_to_speech_eligible_payload,
response_type,
)),
};
}
// We're potentially dealing with some text in text_to_speech_eligible_payload - either coming directly from the user or generated by an agent.
match text_to_speech_stage_params {
Some(TextToSpeechParams::Perform(text_to_speech_eligible_payload, response_type)) => {
let _tts_result = generate_and_send_tts_for_message(
bot,
matrix_link.clone(),
message_context,
response_type,
text_to_speech_eligible_payload.event_id,
&text_to_speech_eligible_payload.text,
)
.await;
}
Some(TextToSpeechParams::Offer(text_to_speech_eligible_payload, response_type)) => {
send_tts_offer_for_message(
bot,
message_context,
response_type,
text_to_speech_eligible_payload.event_id,
)
.await;
}
None => {}
}
Ok(())
}
async fn handle_stage_speech_to_text(
bot: &Bot,
message_context: &MessageContext,
audio_content: &AudioMessageEventContent,
response_type: MessageResponseType,
) -> Option<OwnedEventId> {
let agent = get_effective_agent_for_purpose_or_complain(
bot,
message_context,
AgentPurpose::SpeechToText,
response_type.clone(),
true,
)
.await?;
tracing::debug!(
agent_id = agent.identifier().as_string(),
"Handling speech-to-text",
);
let reaction_event_response = bot
.reacting()
.react_no_fail(
message_context.room(),
message_context.event_id().clone(),
AgentPurpose::SpeechToText.emoji().to_owned(),
)
.await;
let speech_to_text_created_event_id = handle_stage_speech_to_text_actual_transcribing(
bot,
message_context,
&agent,
audio_content,
response_type.clone(),
)
.await;
if let Some(reaction_event_response) = reaction_event_response {
let redaction_reason = if speech_to_text_created_event_id.is_ok() {
strings::speech_to_text::redaction_reason_done()
} else {
strings::speech_to_text::redaction_reason_failed()
};
bot.messaging()
.redact_event_no_fail(
message_context.room(),
reaction_event_response.event_id,
Some(redaction_reason.to_owned()),
)
.await;
}
let speech_to_text_created_event_id = match speech_to_text_created_event_id {
Ok(event_id) => event_id,
Err(err) => {
tracing::warn!(
"Error in room {} while trying to transcribe 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::SpeechToText,
&err,
),
response_type,
)
.await;
return None;
}
};
Some(speech_to_text_created_event_id)
}
async fn handle_stage_text_generation(
bot: &Bot,
matrix_link: MatrixLink,
message_context: &MessageContext,
controller_type: &ChatCompletionControllerType,
response_type: MessageResponseType,
) -> Option<TextToSpeechEligiblePayload> {
let agent = get_effective_agent_for_purpose_or_complain(
bot,
message_context,
AgentPurpose::TextGeneration,
response_type.clone(),
true,
)
.await?;
_ = message_context.room().typing_notice(true).await;
let prefixes_to_strip = match controller_type {
ChatCompletionControllerType::ViaText { prefixes_to_strip } => prefixes_to_strip.clone(),
ChatCompletionControllerType::ViaAudio => vec![],
};
let params = MatrixMessageProcessingParams::new(
bot.user_id().as_str().to_owned(),
message_context.combined_admin_and_user_regexes(),
)
.with_first_message_stripped_prefixes(prefixes_to_strip);
let conversation = create_llm_conversation_for_matrix_thread(
matrix_link.clone(),
message_context.room(),
message_context.thread_info().root_event_id.clone(),
&params,
)
.await;
let conversation = match conversation {
Ok(conversation) => conversation,
Err(err) => {
tracing::warn!(?err, "Error while trying to create conversation");
bot.messaging()
.send_error_markdown_no_fail(
message_context.room(),
&strings::agent::error_while_serving_purpose(
agent.identifier(),
&AgentPurpose::TextGeneration,
&err,
),
response_type,
)
.await;
return None;
}
};
tracing::debug!(
agent_id = agent.identifier().as_string(),
provider = format!("{}", agent.definition().provider.clone()),
"Invoking LLM for text generation with conversation.."
);
let span = tracing::debug_span!(
"text_generation",
agent_id = agent.identifier().as_string(),
provider = format!("{}", agent.definition().provider.clone()),
);
let start_time = std::time::Instant::now();
let params = TextGenerationParams {
context_management_enabled: message_context
.room_config_context()
.text_generation_context_management_enabled(),
prompt_override: message_context
.room_config_context()
.text_generation_prompt_override(),
temperature_override: message_context
.room_config_context()
.text_generation_temperature_override(),
};
let result = agent
.controller()
.generate_text(conversation, params)
.instrument(span)
.await;
let duration = std::time::Instant::now().duration_since(start_time);
tracing::debug!(
agent_id = agent.identifier().as_string(),
provider = format!("{}", agent.definition().provider.clone()),
?duration,
"Done with LLM text generation"
);
let result = match result {
Ok(result) => result,
Err(err) => {
tracing::warn!(
"Error in room {} while trying to generate text 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::TextGeneration,
&err,
),
response_type,
)
.await;
return None;
}
};
let text = result.text.clone().trim().to_owned();
if text.is_empty() {
tracing::warn!(
agent_id = agent.identifier().as_string(),
"Agent returned empty text",
);
bot.messaging()
.send_error_markdown_no_fail(
message_context.room(),
&strings::agent::empty_response_returned(agent.identifier()),
response_type,
)
.await;
return None;
}
let send_message_response = bot
.messaging()
.send_text_markdown_no_fail(message_context.room(), text.clone(), response_type)
.await?;
Some(TextToSpeechEligiblePayload {
text,
event_id: send_message_response.event_id,
})
}
async fn handle_stage_speech_to_text_actual_transcribing(
bot: &Bot,
message_context: &MessageContext,
agent: &AgentInstance,
audio_content: &AudioMessageEventContent,
response_type: MessageResponseType,
) -> anyhow::Result<OwnedEventId> {
let src = &audio_content.source;
let media_request = mxlink::matrix_sdk::media::MediaRequest {
source: src.to_owned(),
format: mxlink::matrix_sdk::media::MediaFormat::File,
};
let media = message_context
.room()
.client()
.media()
.get_media_content(&media_request, true)
.await?;
_ = message_context.room().typing_notice(true).await;
let span = tracing::debug_span!(
"speech_to_text_generation",
agent_id = agent.identifier().as_string()
);
let mime_type = audio_content
.info
.as_ref()
.and_then(|info| info.mimetype.clone())
.unwrap_or_else(|| "audio/ogg".to_string())
.parse::<mxlink::mime::Mime>()
.map_err(|err| anyhow::anyhow!("Invalid MIME type: {}", err))?;
let params = SpeechToTextParams {
language_override: message_context
.room_config_context()
.speech_to_text_language(),
};
let speech_to_text_result = agent
.controller()
.speech_to_text(&mime_type, media, params)
.instrument(span)
.await?;
let transcribed_text = create_transcribed_message_text(&speech_to_text_result.text);
let result = bot
.messaging()
.send_notice_markdown_no_fail(message_context.room(), transcribed_text, response_type)
.await;
result
.map(|result| result.event_id)
.ok_or_else(|| anyhow::anyhow!("Failed to send transcribed text"))
}
async fn send_tts_offer_for_message(
bot: &Bot,
message_context: &MessageContext,
response_type: MessageResponseType,
event_id: OwnedEventId,
) {
// Offers may be enabled, but there's no guarantee that whatever agent is configured can actually do TTS.
// So.. do not complain if there's no agent available. Just silently ignore it.
let speech_agent = get_effective_agent_for_purpose_or_complain(
bot,
message_context,
AgentPurpose::TextToSpeech,
response_type,
false,
)
.await;
if speech_agent.is_some() {
bot.reacting()
.react_no_fail(
message_context.room(),
event_id,
AgentPurpose::TextToSpeech.emoji().to_owned(),
)
.await;
}
}
async fn generate_and_send_tts_for_message(
bot: &Bot,
matrix_link: MatrixLink,
message_context: &MessageContext,
response_type: MessageResponseType,
event_id: OwnedEventId,
text: &str,
) -> bool {
let speech_agent = get_effective_agent_for_purpose_or_complain(
bot,
message_context,
AgentPurpose::TextToSpeech,
response_type.clone(),
true,
)
.await;
let Some(speech_agent) = speech_agent else {
return false;
};
crate::controller::utils::text_to_speech::generate_and_send_tts_for_message(
bot,
matrix_link,
message_context,
response_type,
&speech_agent,
&event_id,
text,
)
.await
}

View File

@@ -0,0 +1,27 @@
#[derive(Debug, PartialEq)]
pub enum ControllerType {
// Denotes that the message is to be ignored.
Ignore,
Help,
UsageHelp,
Unknown,
Error(String),
ErrorInThread(String, mxlink::ThreadInfo),
ProviderHelp,
Access(super::access::AccessControllerType),
Agent(super::agent::AgentControllerType),
Config(super::cfg::ConfigControllerType),
ChatCompletion(super::chat_completion::ChatCompletionControllerType),
ImageGeneration(String),
StickerGeneration(String),
}

View File

@@ -0,0 +1,154 @@
#[cfg(test)]
mod tests;
use mxlink::matrix_sdk::ruma::OwnedUserId;
use super::chat_completion::ChatCompletionControllerType;
use crate::{
entity::{
roomconfig::TextGenerationPrefixRequirementType, MessageContext, MessagePayload,
ThreadContextFirstMessage,
},
strings,
};
use super::ControllerType;
pub fn determine_controller(
command_prefix: &str,
first_thread_message: &ThreadContextFirstMessage,
message_context: &MessageContext,
bot_user_id: &OwnedUserId,
bot_display_name: &Option<String>,
) -> ControllerType {
match &first_thread_message.payload {
MessagePayload::Text(text_message_content) => {
let prefix_requirement_type = message_context
.room_config_context()
.text_generation_prefix_requirement_type();
determine_text_controller(
command_prefix,
&text_message_content.body,
prefix_requirement_type,
first_thread_message.is_mentioning_bot,
bot_user_id,
bot_display_name,
)
}
MessagePayload::Encrypted(thread_info) => {
if thread_info.is_thread_root_only() {
ControllerType::Error(strings::error::message_is_encrypted().to_owned())
} else {
ControllerType::ErrorInThread(
strings::error::first_message_in_thread_is_encrypted().to_owned(),
thread_info.clone(),
)
}
}
MessagePayload::Audio(_) => {
ControllerType::ChatCompletion(ChatCompletionControllerType::ViaAudio)
}
MessagePayload::Reaction { .. } => {
panic!("Handling reaction as first message in thread does not make sense")
}
}
}
fn determine_text_controller(
command_prefix: &str,
text: &str,
room_text_generation_prefix_requirement_type: TextGenerationPrefixRequirementType,
is_mentioning_bot: bool,
bot_user_id: &OwnedUserId,
bot_display_name: &Option<String>,
) -> ControllerType {
let text = text.trim();
if text.starts_with(&format!("{command_prefix} help")) || text == command_prefix {
return ControllerType::Help;
}
if let Some(remaining) = text.strip_prefix(&format!("{command_prefix} access")) {
return super::access::determine_controller(remaining.trim());
}
if let Some(remaining) = text.strip_prefix(&format!("{command_prefix} provider")) {
return super::provider::determine_controller(remaining.trim());
}
if let Some(remaining) = text.strip_prefix(&format!("{command_prefix} agent")) {
return super::agent::determine_controller(command_prefix, remaining.trim());
}
if let Some(remaining) = text.strip_prefix(&format!("{command_prefix} config")) {
return super::cfg::determine_controller(remaining.trim());
}
if let Some(prompt) = text.strip_prefix(&format!("{command_prefix} image")) {
return ControllerType::ImageGeneration(prompt.trim().to_owned());
}
if let Some(prompt) = text.strip_prefix(&format!("{command_prefix} sticker")) {
return ControllerType::StickerGeneration(prompt.trim().to_owned());
}
if let Some(remaining) = text.strip_prefix(&format!("{command_prefix} usage")) {
return super::usage::determine_controller(remaining.trim());
}
// Regular text message that does not match any command.
// If it mentions the bot, it's a chat completion.
// Otherwise, it depends on the prefix requirement for text generation - it may be routed for chat completion or ignored.
if is_mentioning_bot {
// Different clients do mentions differently.
// The body text containing the mention usually contains one of:
// - the full user ID (includes a @ prefix by default)
// - the localpart (with a @ prefix)
// - the localpart (without a @ prefix)
// - the display name (with a @ prefix)
// - the display name (without a @ prefix)
//
// Some add a `: ` suffix after the mention.
//
// There's no guarantee that the mention is at the start even.
// It being there is most common and we try to strip it from there
// as best as we can.
let bot_user_id_localpart = bot_user_id.localpart();
let mut prefixes_to_strip = vec![
bot_user_id.as_str().to_owned(),
format!("@{}", bot_user_id_localpart),
bot_user_id_localpart.to_owned(),
];
if let Some(bot_display_name) = bot_display_name {
prefixes_to_strip.push(format!("@{}", bot_display_name));
prefixes_to_strip.push(bot_display_name.to_owned());
}
prefixes_to_strip.push(":".to_owned());
return ControllerType::ChatCompletion(ChatCompletionControllerType::ViaText {
prefixes_to_strip,
});
}
match room_text_generation_prefix_requirement_type {
TextGenerationPrefixRequirementType::CommandPrefix => {
if text.starts_with(command_prefix) {
ControllerType::ChatCompletion(ChatCompletionControllerType::ViaText {
prefixes_to_strip: vec![command_prefix.to_owned()],
})
} else {
ControllerType::Ignore
}
}
TextGenerationPrefixRequirementType::No => {
ControllerType::ChatCompletion(ChatCompletionControllerType::ViaText {
prefixes_to_strip: vec![],
})
}
}
}

View File

@@ -0,0 +1,196 @@
#[test]
fn determine_text_controller() {
use super::super::chat_completion::ChatCompletionControllerType;
use super::ControllerType;
use crate::controller;
let bot_user_id = mxlink::matrix_sdk::ruma::owned_user_id!("@bot:example.com");
let bot_display_name = "Bot";
let command_prefix = "!bai";
struct TestCase {
name: &'static str,
input: &'static str,
is_mentioning_bot: bool,
expected: ControllerType,
// This value only matters for some of the tests.
// We default to using the No variant for most tests where it's irrelevant.
room_text_generation_prefix_requirement_type: super::TextGenerationPrefixRequirementType,
}
// We only have top-level test cases here.
// Each submodule defines its own test cases.
let test_cases = vec![
TestCase {
name: "Help",
input: "!bai help",
is_mentioning_bot: false,
room_text_generation_prefix_requirement_type:
super::TextGenerationPrefixRequirementType::No,
expected: ControllerType::Help,
},
TestCase {
name: "Prefix only leads to help",
input: "!bai",
is_mentioning_bot: false,
room_text_generation_prefix_requirement_type:
super::TextGenerationPrefixRequirementType::No,
expected: ControllerType::Help,
},
TestCase {
name: "Prefix and unknown command leads to chat completion",
input: "!bai something-else",
is_mentioning_bot: false,
room_text_generation_prefix_requirement_type:
super::TextGenerationPrefixRequirementType::No,
expected: ControllerType::ChatCompletion(ChatCompletionControllerType::ViaText {
prefixes_to_strip: vec![],
}),
},
TestCase {
name: "Access top-level",
input: "!bai access",
is_mentioning_bot: false,
room_text_generation_prefix_requirement_type:
super::TextGenerationPrefixRequirementType::No,
expected: ControllerType::Access(controller::access::AccessControllerType::Help),
},
TestCase {
name: "Provider",
input: "!bai provider",
is_mentioning_bot: false,
room_text_generation_prefix_requirement_type:
super::TextGenerationPrefixRequirementType::No,
expected: ControllerType::ProviderHelp,
},
TestCase {
name: "Usage",
input: "!bai usage",
is_mentioning_bot: false,
room_text_generation_prefix_requirement_type:
super::TextGenerationPrefixRequirementType::No,
expected: ControllerType::UsageHelp,
},
TestCase {
name: "Agent top-level",
input: "!bai agent",
is_mentioning_bot: false,
room_text_generation_prefix_requirement_type:
super::TextGenerationPrefixRequirementType::No,
expected: ControllerType::Agent(controller::agent::AgentControllerType::Help),
},
TestCase {
name: "Config top-level",
input: "!bai config",
is_mentioning_bot: false,
room_text_generation_prefix_requirement_type:
super::TextGenerationPrefixRequirementType::No,
expected: ControllerType::Config(controller::cfg::ConfigControllerType::Help),
},
TestCase {
name: "Image generation",
input: "!bai image Draw a cat!",
is_mentioning_bot: false,
room_text_generation_prefix_requirement_type:
super::TextGenerationPrefixRequirementType::No,
expected: ControllerType::ImageGeneration("Draw a cat!".to_owned()),
},
TestCase {
name: "Sticker generation",
input: "!bai sticker A surprised cat",
is_mentioning_bot: false,
room_text_generation_prefix_requirement_type:
super::TextGenerationPrefixRequirementType::No,
expected: ControllerType::StickerGeneration("A surprised cat".to_owned()),
},
TestCase {
name: "Regular text triggers completion when prefix not required",
input: "Regular text goes here",
is_mentioning_bot: false,
room_text_generation_prefix_requirement_type:
super::TextGenerationPrefixRequirementType::No,
expected: ControllerType::ChatCompletion(ChatCompletionControllerType::ViaText {
prefixes_to_strip: vec![],
}),
},
TestCase {
name: "Regular text is ignored when prefix is required",
input: "Regular text goes here",
is_mentioning_bot: false,
room_text_generation_prefix_requirement_type:
super::TextGenerationPrefixRequirementType::CommandPrefix,
expected: ControllerType::Ignore,
},
TestCase {
name: "Command-prefixed text triggers completion when prefix is required",
input: "!bai Regular text goes here",
is_mentioning_bot: false,
room_text_generation_prefix_requirement_type:
super::TextGenerationPrefixRequirementType::CommandPrefix,
expected: ControllerType::ChatCompletion(ChatCompletionControllerType::ViaText {
prefixes_to_strip: vec!["!bai".to_owned()],
}),
},
TestCase {
name: "Command-prefixed text triggers completion even when prefix is not required",
input: "!bai Regular text goes here",
is_mentioning_bot: false,
room_text_generation_prefix_requirement_type:
super::TextGenerationPrefixRequirementType::No,
expected: ControllerType::ChatCompletion(ChatCompletionControllerType::ViaText {
prefixes_to_strip: vec![],
}),
},
TestCase {
name: "Regular message with bot mention triggers completion stripping bot id and display name (no prefix requirement)",
input: "Regular text goes here",
is_mentioning_bot: true,
room_text_generation_prefix_requirement_type:
super::TextGenerationPrefixRequirementType::No,
expected: ControllerType::ChatCompletion(ChatCompletionControllerType::ViaText {
prefixes_to_strip: vec![
"@bot:example.com".to_owned(),
"@bot".to_owned(),
"bot".to_owned(),
"@Bot".to_owned(),
"Bot".to_owned(),
":".to_owned(),
],
}),
},
// This test case is the same as the one above, just with a different prefix requirement.
// We expect the same result.
TestCase {
name: "Regular message with bot mention triggers completion stripping bot id and display name (command_prefix requirement)",
input: "Regular text goes here",
is_mentioning_bot: true,
room_text_generation_prefix_requirement_type:
super::TextGenerationPrefixRequirementType::CommandPrefix,
expected: ControllerType::ChatCompletion(ChatCompletionControllerType::ViaText {
prefixes_to_strip: vec![
"@bot:example.com".to_owned(),
"@bot".to_owned(),
"bot".to_owned(),
"@Bot".to_owned(),
"Bot".to_owned(),
":".to_owned(),
],
}),
},
];
for test_case in test_cases {
let bot_display_name = Some(bot_display_name.to_owned());
let result = super::determine_text_controller(
command_prefix,
test_case.input,
test_case.room_text_generation_prefix_requirement_type,
test_case.is_mentioning_bot,
&bot_user_id,
&bot_display_name,
);
assert_eq!(result, test_case.expected, "Test case: {}", test_case.name);
}
}

View File

@@ -0,0 +1,108 @@
use mxlink::MessageResponseType;
use crate::{entity::MessageContext, strings, Bot};
use super::ControllerType;
pub async fn dispatch_controller(
controller_type: &ControllerType,
message_context: &MessageContext,
bot: &Bot,
) {
let result = match controller_type {
ControllerType::Access(controller_type) => {
super::access::dispatch_controller(controller_type, message_context, bot).await
}
ControllerType::Agent(controller_type) => {
super::agent::dispatch_controller(controller_type, message_context, bot).await
}
ControllerType::Config(controller_type) => {
super::cfg::dispatch_controller(controller_type, message_context, bot).await
}
ControllerType::Help => super::help::handle(bot, message_context).await,
ControllerType::Unknown => {
bot.messaging()
.send_error_markdown_no_fail(
message_context.room(),
&strings::error::unknown_command_see_help(bot.command_prefix()),
MessageResponseType::Reply(message_context.thread_info().root_event_id.clone()),
)
.await;
Ok(())
}
ControllerType::ProviderHelp => super::provider::handle_help(message_context, bot).await,
ControllerType::UsageHelp => super::usage::handle_help(message_context, bot).await,
ControllerType::ChatCompletion(controller_type) => {
super::chat_completion::handle(
bot,
bot.matrix_link().clone(),
message_context,
controller_type,
)
.await
}
ControllerType::Error(message) => {
bot.messaging()
.send_error_markdown_no_fail(
message_context.room(),
message,
MessageResponseType::Reply(message_context.thread_info().root_event_id.clone()),
)
.await;
Ok(())
}
ControllerType::ErrorInThread(message, thread_info) => {
bot.messaging()
.send_error_markdown_no_fail(
message_context.room(),
message,
MessageResponseType::InThread(thread_info.clone()),
)
.await;
Ok(())
}
ControllerType::Ignore => {
tracing::trace!("Ignoring text message");
Ok(())
}
ControllerType::ImageGeneration(prompt) => {
super::image::generation::handle_image(
bot,
bot.matrix_link().clone(),
message_context,
prompt,
)
.await
}
ControllerType::StickerGeneration(prompt) => {
super::image::generation::handle_sticker(
bot,
bot.matrix_link().clone(),
message_context,
prompt,
)
.await
}
};
if let Err(e) = result {
tracing::error!(
"Error handling message {} from sender {} in room {}: {:?}",
message_context.event_id(),
message_context.sender_id(),
message_context.room_id(),
e,
);
bot.messaging()
.send_error_markdown_no_fail(
message_context.room(),
strings::error::error_while_processing_message(),
MessageResponseType::Reply(message_context.thread_info().root_event_id.clone()),
)
.await;
}
}

View File

@@ -0,0 +1,90 @@
use mxlink::MessageResponseType;
use crate::{entity::MessageContext, strings, Bot};
pub async fn handle(bot: &Bot, message_context: &MessageContext) -> anyhow::Result<()> {
let sender_can_manage_global_config = message_context.sender_can_manage_global_config()?;
let sender_can_manage_room_local_agents =
message_context.sender_can_manage_room_local_agents()?;
let mut message = String::from("");
message.push_str(&format!("## {}\n\n", strings::help::heading_introduction()));
message.push_str(&strings::introduction::create_short_introduction(
bot.name(),
));
message.push_str("\n\n");
// Agents
message.push_str(&format!("## {}", strings::help::agent::heading()));
message.push_str("\n\n");
message.push_str(&strings::help::agent::intro(
bot.command_prefix(),
sender_can_manage_room_local_agents,
));
message.push_str("\n\n");
message.push_str(&strings::help::agent::intro_handler_relation(
bot.command_prefix(),
));
message.push_str("\n\n");
message.push_str(&strings::help::learn_more_send_a_command(
bot.command_prefix(),
"agent",
));
message.push_str("\n\n");
// Providers
if sender_can_manage_room_local_agents || sender_can_manage_global_config {
message.push_str(&format!("## {}", strings::help::provider::heading()));
message.push_str("\n\n");
message.push_str(&strings::help::provider::intro());
message.push_str("\n\n");
message.push_str(&strings::help::learn_more_send_a_command(
bot.command_prefix(),
"provider",
));
message.push_str("\n\n");
}
// Access
message.push_str(&format!("## {}", strings::help::access::heading()));
message.push_str("\n\n");
message.push_str(&strings::help::access::intro());
message.push_str("\n\n");
message.push_str(&strings::help::learn_more_send_a_command(
bot.command_prefix(),
"access",
));
message.push_str("\n\n");
// Configuration
message.push_str(&format!("## {}", strings::help::cfg::heading()));
message.push_str("\n\n");
message.push_str(strings::help::cfg::intro_short());
message.push_str("\n\n");
message.push_str(&strings::help::learn_more_send_a_command(
bot.command_prefix(),
"config",
));
message.push_str("\n\n");
// Usage
message.push_str(&format!("## {}", strings::help::usage::heading()));
message.push_str("\n\n");
message.push_str(strings::help::usage::intro());
message.push_str("\n\n");
message.push_str(&strings::help::learn_more_send_a_command(
bot.command_prefix(),
"usage",
));
bot.messaging()
.send_text_markdown_no_fail(
message_context.room(),
message,
MessageResponseType::Reply(message_context.thread_info().last_event_id.clone()),
)
.await;
Ok(())
}

View File

@@ -0,0 +1,179 @@
use mxlink::{MatrixLink, MessageResponseType};
use tracing::Instrument;
use crate::agent::provider::ImageGenerationParams;
use crate::agent::AgentPurpose;
use crate::agent::ControllerTrait;
use crate::controller::utils::agent::get_effective_agent_for_purpose_or_complain;
use crate::conversation::create_llm_conversation_for_matrix_thread;
use crate::conversation::matrix::MatrixMessageProcessingParams;
use crate::strings;
use crate::{entity::MessageContext, Bot};
// We may make this configurable (per room, etc.) in the future, but for now it's hardcoded.
const STICKER_SIZE: &str = "256x256";
pub async fn handle_image(
bot: &Bot,
matrix_link: MatrixLink,
message_context: &MessageContext,
original_prompt: &str,
) -> anyhow::Result<()> {
let response_type = MessageResponseType::InThread(message_context.thread_info().clone());
let Some(agent) = get_effective_agent_for_purpose_or_complain(
bot,
message_context,
AgentPurpose::ImageGeneration,
response_type.clone(),
true,
)
.await
else {
return Ok(());
};
let params = MatrixMessageProcessingParams::new(
bot.user_id().as_str().to_owned(),
message_context.combined_admin_and_user_regexes(),
);
let conversation = create_llm_conversation_for_matrix_thread(
matrix_link.clone(),
message_context.room(),
message_context.thread_info().root_event_id.clone(),
&params,
)
.await?;
let prompt = if conversation.messages.len() >= 2 {
// Skip the first message, which contains the original prompt (which we already have)
let other_messages = conversation.messages.iter().skip(1).cloned().collect();
super::prompt::build(original_prompt, other_messages)
} else {
original_prompt.to_owned()
};
message_context.room().typing_notice(true).await?;
let span = tracing::debug_span!(
"image_generation",
agent_id = agent.identifier().as_string()
);
let response = agent
.controller()
.generate_image(&prompt, ImageGenerationParams::default())
.instrument(span)
.await?;
let actual_prompt = response.revised_prompt.as_deref().unwrap_or(&prompt);
if *actual_prompt.trim() != *prompt.trim() {
bot.messaging()
.send_notice_markdown_no_fail(
message_context.room(),
strings::image_generation::revised_prompt(actual_prompt),
response_type.clone(),
)
.await;
}
let attachment_body_text = format!("Generated image based on: {}", actual_prompt);
let mut event_content = matrix_link
.media()
.upload_and_prepare_event_content(
message_context.room(),
&response.mime_type,
response.bytes,
&attachment_body_text,
)
.await
.map_err(|e| anyhow::anyhow!("Failed to upload and prepare event: {}", e))?;
matrix_link
.messaging()
.send_event(
message_context.room(),
&mut event_content,
response_type.clone(),
)
.await?;
if conversation.messages.len() == 1 {
// If this is the beginning of the thread, send helpful instructions
bot.messaging()
.send_notice_markdown_no_fail(
message_context.room(),
strings::image_generation::guide_how_to_proceed(),
response_type.clone(),
)
.await;
}
Ok(())
}
pub async fn handle_sticker(
bot: &Bot,
matrix_link: MatrixLink,
message_context: &MessageContext,
original_prompt: &str,
) -> anyhow::Result<()> {
// Stickers are always sent directly to the room - no threading.
let response_type =
MessageResponseType::Reply(message_context.thread_info().root_event_id.clone());
let Some(agent) = get_effective_agent_for_purpose_or_complain(
bot,
message_context,
AgentPurpose::ImageGeneration,
response_type.clone(),
true,
)
.await
else {
return Ok(());
};
message_context.room().typing_notice(true).await?;
let span = tracing::debug_span!(
"sticker_generation",
agent_id = agent.identifier().as_string()
);
let params = ImageGenerationParams::default()
.with_size_override(Some(STICKER_SIZE.to_owned()))
.with_cheaper_model_switching_allowed(true)
.with_cheaper_quality_switching_allowed(true);
let response = agent
.controller()
.generate_image(original_prompt, params)
.instrument(span)
.await?;
let attachment_body_text = format!("Generated sticker image based on: {}", original_prompt);
let mut event_content = matrix_link
.media()
.upload_and_prepare_event_content(
message_context.room(),
&response.mime_type,
response.bytes,
&attachment_body_text,
)
.await
.map_err(|e| anyhow::anyhow!("Failed to upload and prepare event: {}", e))?;
matrix_link
.messaging()
.send_event(message_context.room(), &mut event_content, response_type)
.await?;
Ok(())
}

View File

@@ -0,0 +1,2 @@
pub mod generation;
mod prompt;

View File

@@ -0,0 +1,111 @@
use crate::conversation::llm::{Author, Message};
/// Builds a prompt from the original prompt and other messages in the conversation.
///
/// Only messages authored by the user are considered.
///
/// Messages that say "Again" (regardless of casing) are ignored. They are considered special messages
/// which trigger re-generation, but do not need to be included in the prompt criteria.
pub fn build(original_prompt: &str, other_messages: Vec<Message>) -> String {
let mut prompt = original_prompt.to_owned();
// Make a new messages vector that only contains messages we care about
let other_messages: Vec<Message> = other_messages
.into_iter()
.filter(|message| {
if let Author::User = message.author {
message.message_text.to_lowercase() != "again"
} else {
false
}
})
.collect();
if !other_messages.is_empty() {
prompt.push_str("\nOther criteria:");
for message in other_messages {
prompt.push_str(
format!("\n- {}", message.message_text.replace("\n", ". ").as_str()).as_str(),
);
}
}
prompt
}
#[cfg(test)]
mod tests {
use super::build;
use super::{Author, Message};
struct TestCase {
original_prompt: &'static str,
messages: Vec<Message>,
expected_prompt: &'static str,
}
#[test]
fn test_build_prompt() {
let test_cases = vec![
// Simple case
TestCase {
original_prompt: "Generate a picture of a cat",
messages: vec![],
expected_prompt: "Generate a picture of a cat",
},
// Only a single user message
TestCase {
original_prompt: "Generate a picture of a dog",
messages: vec![Message {
author: Author::User,
message_text: "Must be blue".to_owned(),
}],
expected_prompt: "Generate a picture of a dog\nOther criteria:\n- Must be blue",
},
// Multiple complex user messages dispersed with assistant messages
TestCase {
original_prompt: "Generate a picture of an elephant",
messages: vec![Message {
author: Author::User,
message_text: "Must be blue".to_owned(),
},
Message {
author: Author::Assistant,
message_text: "Whatever".to_owned(),
},
Message {
author: Author::User,
message_text: "Must be 3-legged.\nMust be flying.".to_owned(),
}],
expected_prompt: "Generate a picture of an elephant\nOther criteria:\n- Must be blue\n- Must be 3-legged.. Must be flying.",
},
// "Again" is ignored.
TestCase {
original_prompt: "Generate a picture of a grizzly bear",
messages: vec![Message {
author: Author::User,
message_text: "Must be blue".to_owned(),
},
Message {
author: Author::Assistant,
message_text: "Whatever".to_owned(),
},
Message {
author: Author::User,
message_text: "Again".to_owned(),
},
Message {
author: Author::User,
message_text: "again".to_owned(),
}],
expected_prompt: "Generate a picture of a grizzly bear\nOther criteria:\n- Must be blue",
},
];
for test_case in test_cases {
let actual_prompt = build(test_case.original_prompt, test_case.messages);
assert_eq!(actual_prompt, test_case.expected_prompt);
}
}
}

View File

@@ -0,0 +1,28 @@
use mxlink::MessageResponseType;
use crate::entity::RoomConfigContext;
use crate::{strings, Bot};
pub async fn handle(
bot: &Bot,
room: &mxlink::matrix_sdk::Room,
room_config_context: &RoomConfigContext,
) -> anyhow::Result<()> {
let agent_manager = bot.agent_manager();
bot.messaging()
.send_text_markdown_no_fail(
room,
strings::introduction::create_on_join_introduction(
bot.name(),
bot.command_prefix(),
agent_manager,
room_config_context,
)
.await,
MessageResponseType::InRoom,
)
.await;
Ok(())
}

19
src/controller/mod.rs Normal file
View File

@@ -0,0 +1,19 @@
mod controller_type;
pub mod access;
pub mod agent;
pub mod cfg;
pub mod chat_completion;
mod determination;
mod dispatching;
pub mod help;
pub mod image;
pub mod join;
pub mod provider;
pub mod reaction;
pub mod usage;
mod utils;
pub use controller_type::ControllerType;
pub use determination::determine_controller;
pub use dispatching::dispatch_controller;

View File

@@ -0,0 +1,97 @@
use mxlink::MessageResponseType;
use crate::{agent::AgentProvider, entity::MessageContext, strings, Bot};
use super::ControllerType;
pub fn determine_controller(_text: &str) -> ControllerType {
ControllerType::ProviderHelp
}
pub async fn handle_help(message_context: &MessageContext, bot: &Bot) -> anyhow::Result<()> {
if !message_context.sender_can_manage_room_local_agents()? {
bot.messaging()
.send_error_markdown_no_fail(
message_context.room(),
&strings::provider::not_allowed(),
MessageResponseType::Reply(message_context.thread_info().root_event_id.clone()),
)
.await;
return Ok(());
}
let mut message = String::new();
message.push_str(&format!("## {}", strings::help::provider::heading()));
message.push_str("\n\n");
message.push_str(&strings::help::provider::intro());
message.push_str("\n\n");
message.push_str(&strings::provider::providers_list_intro());
message.push_str("\n\n");
// How to choose
message.push_str(&format!(
"### {}",
strings::provider::help_how_to_choose_heading()
));
message.push_str("\n\n");
message.push_str(&strings::provider::help_how_to_choose_description(
bot.command_prefix(),
));
message.push_str("\n\n");
// How to use
message.push_str(&format!(
"### {}",
strings::provider::help_how_to_use_heading()
));
message.push_str("\n\n");
message.push_str(&strings::provider::help_how_to_use_description(
bot.command_prefix(),
));
message.push_str("\n\n");
for provider in AgentProvider::choices() {
let provider_info = provider.info();
message.push_str(&format!(
"### {}",
strings::provider::help_provider_heading(
provider_info.name,
&provider_info.homepage_url.as_ref().map(|s| s.to_string())
)
));
message.push_str("\n\n");
message.push_str(&strings::provider::help_provider_details(
provider.to_static_str(),
&provider_info,
));
message.push_str("- 🗲 Quick start:\n");
message.push_str(&format!(
"\t- create a room-local agent: `{command_prefix} agent create-room-local {provider_id} my-{provider_id}-agent`",
command_prefix = bot.command_prefix(),
provider_id = provider.to_static_str(),
));
message.push('\n');
message.push_str(&format!(
"\t- create a global agent: `{command_prefix} agent create-global {provider_id} my-{provider_id}-agent`",
command_prefix = bot.command_prefix(),
provider_id = provider.to_static_str(),
));
message.push_str("\n\n");
}
bot.messaging()
.send_text_markdown_no_fail(
message_context.room(),
message,
MessageResponseType::Reply(message_context.thread_info().root_event_id.clone()),
)
.await;
Ok(())
}

View File

@@ -0,0 +1,50 @@
use std::ops::Deref;
use mxlink::MatrixLink;
use crate::{
agent::AgentPurpose,
entity::{MessageContext, MessagePayload},
Bot,
};
mod text_to_speech;
pub async fn handle(
bot: &Bot,
matrix_link: MatrixLink,
message_context: &MessageContext,
) -> anyhow::Result<()> {
match &message_context.payload() {
MessagePayload::Reaction {
key,
reacted_to_event_payload,
reacted_to_event_id,
reacted_to_event_sender_id,
} => {
if key == AgentPurpose::TextToSpeech.emoji() {
if let MessagePayload::Text(text_content) = reacted_to_event_payload.deref() {
return text_to_speech::handle(
bot,
matrix_link,
message_context,
reacted_to_event_id,
reacted_to_event_sender_id,
text_content,
)
.await;
}
tracing::debug!("Ignoring text-to-speech reaction to non-text message");
return Ok(());
}
tracing::debug!("Ignoring unknown reaction");
Ok(())
}
_ => Err(anyhow::anyhow!(
"Reaction controller called with a non-reaction message"
)),
}
}

View File

@@ -0,0 +1,100 @@
use mxlink::{MatrixLink, MessageResponseType};
use mxlink::matrix_sdk::ruma::{
events::room::message::TextMessageEventContent, OwnedEventId, OwnedUserId,
};
use crate::entity::roomconfig::{
TextToSpeechBotMessagesFlowType, TextToSpeechUserMessagesFlowType,
};
use crate::{
agent::AgentPurpose, controller::utils::agent::get_effective_agent_for_purpose_or_complain,
entity::MessageContext, Bot,
};
pub(super) async fn handle(
bot: &Bot,
matrix_link: MatrixLink,
message_context: &MessageContext,
reacted_to_event_id: &OwnedEventId,
reacted_to_event_sender_id: &OwnedUserId,
text_content: &TextMessageEventContent,
) -> anyhow::Result<()> {
// If we're in a thread, we're likely dealing with a bot message, so we should start in the thread.
// Otherwise, we're likely operating in "TTS user messages" mode, so we should reply to the reacted-to message and avoid threads.
let response_type = if message_context.thread_info().is_thread_root_only() {
MessageResponseType::Reply(reacted_to_event_id.clone())
} else {
MessageResponseType::InThread(message_context.thread_info().clone())
};
if !is_allowed_to_tts_for_event(
message_context,
reacted_to_event_sender_id,
matrix_link.user_id(),
) {
tracing::debug!("Ignoring request for on-demand text-to-speech (via reaction) due to room configuration");
return Ok(());
}
let speech_agent = get_effective_agent_for_purpose_or_complain(
bot,
message_context,
AgentPurpose::TextToSpeech,
response_type.clone(),
true,
)
.await;
let Some(speech_agent) = speech_agent else {
// We've already complained about this in get_effective_agent_or_complain
return Ok(());
};
crate::controller::utils::text_to_speech::generate_and_send_tts_for_message(
bot,
matrix_link,
message_context,
response_type,
&speech_agent,
reacted_to_event_id,
&text_content.body,
)
.await;
Ok(())
}
fn is_allowed_to_tts_for_event(
message_context: &MessageContext,
sender_id: &OwnedUserId,
bot_user_id: &OwnedUserId,
) -> bool {
// Whether we're allowed depends on who the original message sender is (the bot or some user).
//
// The user may be an allowed bot user or someone else.
// Regardless, we've been invoked by an allowed user, so if the user wants TTS for a foreign message, we should allow it.
if *sender_id == *bot_user_id {
match message_context
.room_config_context()
.text_to_speech_bot_messages_flow_type()
{
TextToSpeechBotMessagesFlowType::Never => false,
TextToSpeechBotMessagesFlowType::OnDemandAlways => true,
TextToSpeechBotMessagesFlowType::OnDemandForVoice => true,
TextToSpeechBotMessagesFlowType::OnlyForVoice => true,
TextToSpeechBotMessagesFlowType::Always => true,
}
} else {
match message_context
.room_config_context()
.text_to_speech_user_messages_flow_type()
{
TextToSpeechUserMessagesFlowType::Never => false,
TextToSpeechUserMessagesFlowType::OnDemand => true,
TextToSpeechUserMessagesFlowType::Always => true,
}
}
}

View File

@@ -0,0 +1,21 @@
use mxlink::MessageResponseType;
use crate::{entity::MessageContext, strings, Bot};
use super::ControllerType;
pub fn determine_controller(_text: &str) -> ControllerType {
ControllerType::UsageHelp
}
pub async fn handle_help(message_context: &MessageContext, bot: &Bot) -> anyhow::Result<()> {
bot.messaging()
.send_text_markdown_no_fail(
message_context.room(),
strings::usage::intro(bot.command_prefix()),
MessageResponseType::Reply(message_context.thread_info().root_event_id.clone()),
)
.await;
Ok(())
}

View File

@@ -0,0 +1,63 @@
use mxlink::MessageResponseType;
use crate::{
agent::{
utils::{get_effective_agent_for_purpose, AgentForPurposeDeterminationError},
AgentInstance, AgentPurpose,
},
entity::MessageContext,
strings, Bot,
};
pub async fn get_effective_agent_for_purpose_or_complain<'a>(
bot: &'a Bot,
message_context: &MessageContext,
agent_purpose: AgentPurpose,
response_type: MessageResponseType,
complain_when_purpose_unsupported: bool,
) -> Option<AgentInstance> {
let agent_info = get_effective_agent_for_purpose(
bot.agent_manager(),
message_context.room_config_context(),
agent_purpose,
)
.await;
match agent_info {
Ok(agent_info) => Some(agent_info.instance),
Err(err) => {
let error_message = match err {
AgentForPurposeDeterminationError::Unknown(err_string) => Some(err_string),
AgentForPurposeDeterminationError::NoneConfigured => None,
AgentForPurposeDeterminationError::ConfiguredButMissing(agent_identifier) => Some(
strings::room_config::configures_agent_for_purpose_but_does_not_exist(
&agent_identifier,
agent_purpose,
),
),
AgentForPurposeDeterminationError::ConfiguredButLacksSupport(agent_identifier) => {
if complain_when_purpose_unsupported {
Some(strings::room_config::configures_agent_for_purpose_but_agent_does_not_support_it(
&agent_identifier,
agent_purpose,
))
} else {
None
}
}
};
if let Some(error_message) = error_message {
bot.messaging()
.send_error_markdown_no_fail(
message_context.room(),
&error_message,
response_type,
)
.await;
};
None
}
}
}

View File

@@ -0,0 +1,29 @@
use mxlink::MessageResponseType;
use crate::{
entity::{MessageContext, MessagePayload},
Bot,
};
pub mod agent;
pub mod text_to_speech;
pub async fn get_text_body_or_complain<'a>(
bot: &Bot,
message_context: &'a MessageContext,
) -> Option<&'a str> {
match &message_context.payload() {
MessagePayload::Text(text_message_content) => Some(&text_message_content.body),
_ => {
bot.messaging()
.send_text_markdown_no_fail(
message_context.room(),
"This command only works with text messages.".to_owned(),
MessageResponseType::InThread(message_context.thread_info().clone()),
)
.await;
None
}
}
}

View File

@@ -0,0 +1,168 @@
use mxlink::matrix_sdk::ruma::OwnedEventId;
use mxlink::{MatrixLink, MessageResponseType};
use tracing::Instrument;
use crate::{
agent::{provider::TextToSpeechParams, AgentInstance, AgentPurpose, ControllerTrait},
entity::MessageContext,
strings, Bot,
};
pub async fn generate_and_send_tts_for_message(
bot: &Bot,
matrix_link: MatrixLink,
message_context: &MessageContext,
response_type: MessageResponseType,
speech_agent: &AgentInstance,
text_message_event_id: &OwnedEventId,
text_content: &str,
) -> bool {
_ = message_context.room().typing_notice(true).await;
let reaction_event_response = bot
.reacting()
.react_no_fail(
message_context.room(),
text_message_event_id.clone(),
strings::PROGRESS_INDICATOR_EMOJI.to_owned(),
)
.await;
let result = do_generate_and_send_tts_for_message(
bot,
matrix_link,
message_context,
response_type,
speech_agent,
text_content,
)
.await;
if let Some(reaction_event_response) = reaction_event_response {
let redaction_reason = if result {
strings::text_to_speech::redaction_reason_done()
} else {
strings::text_to_speech::redaction_reason_failed()
};
bot.messaging()
.redact_event_no_fail(
message_context.room(),
reaction_event_response.event_id,
Some(redaction_reason.to_owned()),
)
.await;
}
result
}
async fn do_generate_and_send_tts_for_message(
bot: &Bot,
matrix_link: MatrixLink,
message_context: &MessageContext,
response_type: MessageResponseType,
speech_agent: &AgentInstance,
text_content: &str,
) -> bool {
let params = TextToSpeechParams {
speed_override: message_context
.room_config_context()
.text_to_speech_speed_override(),
voice_override: message_context
.room_config_context()
.text_to_speech_voice_override(),
};
let text_content = if let Some(text_content) = text_content.strip_prefix(bot.command_prefix()) {
text_content.trim()
} else {
text_content
};
let span = tracing::debug_span!(
"text_to_speech_generation",
agent_id = speech_agent.identifier().as_string()
);
let text_to_speech_result = speech_agent
.controller()
.text_to_speech(text_content, params)
.instrument(span)
.await;
let text_to_speech_result = match text_to_speech_result {
Ok(text_to_speech_result) => text_to_speech_result,
Err(err) => {
tracing::warn!(
"Error in room {} while trying to generate TTS via agent {}: {:?}",
message_context.room_id(),
speech_agent.identifier(),
err,
);
bot.messaging()
.send_error_markdown_no_fail(
message_context.room(),
&strings::agent::error_while_serving_purpose(
speech_agent.identifier(),
&AgentPurpose::SpeechToText,
&err,
),
response_type,
)
.await;
return false;
}
};
let attachment_body_text = strings::text_to_speech::alternate_body_text();
let event_content = matrix_link
.media()
.upload_and_prepare_event_content(
message_context.room(),
&text_to_speech_result.mime_type,
text_to_speech_result.bytes,
&attachment_body_text,
)
.await;
let mut event_content = match event_content {
Ok(event_content) => event_content,
Err(err) => {
tracing::error!(
?err,
"Error in room {} while trying to upload TTS via agent {}",
message_context.room_id(),
speech_agent.identifier(),
);
return false;
}
};
let result = matrix_link
.messaging()
.send_event(
message_context.room(),
&mut event_content,
response_type.clone(),
)
.await;
let Err(err) = result else {
return true;
};
tracing::error!(
?err,
"Error in room {} while trying to send TTS payload",
message_context.room_id(),
);
false
}

View File

@@ -0,0 +1,17 @@
#[derive(Debug, Clone, PartialEq)]
pub enum Author {
Prompt,
Assistant,
User,
}
#[derive(Debug, Clone)]
pub struct Message {
pub author: Author,
pub message_text: String,
}
#[derive(Debug)]
pub struct Conversation {
pub messages: Vec<Message>,
}

Some files were not shown because too many files have changed in this diff Show More