Initial commit
This commit is contained in:
55
src/agent/definition.rs
Normal file
55
src/agent/definition.rs
Normal 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
101
src/agent/identifier.rs
Normal 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
154
src/agent/instantiation.rs
Normal 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
68
src/agent/manager.rs
Normal 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
21
src/agent/mod.rs
Normal 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;
|
||||
71
src/agent/provider/anthropic/config.rs
Normal file
71
src/agent/provider/anthropic/config.rs
Normal 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()
|
||||
}
|
||||
251
src/agent/provider/anthropic/controller.rs
Normal file
251
src/agent/provider/anthropic/controller.rs
Normal 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
|
||||
}
|
||||
}
|
||||
43
src/agent/provider/anthropic/mod.rs
Normal file
43
src/agent/provider/anthropic/mod.rs
Normal 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()
|
||||
}
|
||||
32
src/agent/provider/anthropic/utils.rs
Normal file
32
src/agent/provider/anthropic/utils.rs
Normal 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()
|
||||
}
|
||||
}
|
||||
3
src/agent/provider/config.rs
Normal file
3
src/agent/provider/config.rs
Normal file
@@ -0,0 +1,3 @@
|
||||
pub trait ConfigTrait {
|
||||
fn validate(&self) -> Result<(), String>;
|
||||
}
|
||||
172
src/agent/provider/controller.rs
Normal file
172
src/agent/provider/controller.rs
Normal 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,
|
||||
}
|
||||
}
|
||||
}
|
||||
198
src/agent/provider/entity/agent_provider.rs
Normal file
198
src/agent/provider/entity/agent_provider.rs
Normal 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>,
|
||||
}
|
||||
31
src/agent/provider/entity/image_generation.rs
Normal file
31
src/agent/provider/entity/image_generation.rs
Normal 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>,
|
||||
}
|
||||
13
src/agent/provider/entity/mod.rs
Normal file
13
src/agent/provider/entity/mod.rs
Normal 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};
|
||||
4
src/agent/provider/entity/ping.rs
Normal file
4
src/agent/provider/entity/ping.rs
Normal file
@@ -0,0 +1,4 @@
|
||||
pub enum PingResult {
|
||||
Inconclusive,
|
||||
Successful,
|
||||
}
|
||||
8
src/agent/provider/entity/speech_to_text.rs
Normal file
8
src/agent/provider/entity/speech_to_text.rs
Normal file
@@ -0,0 +1,8 @@
|
||||
#[derive(Default)]
|
||||
pub struct SpeechToTextParams {
|
||||
pub language_override: Option<String>,
|
||||
}
|
||||
|
||||
pub struct SpeechToTextResult {
|
||||
pub text: String,
|
||||
}
|
||||
10
src/agent/provider/entity/text_generation.rs
Normal file
10
src/agent/provider/entity/text_generation.rs
Normal 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,
|
||||
}
|
||||
10
src/agent/provider/entity/text_to_speech.rs
Normal file
10
src/agent/provider/entity/text_to_speech.rs
Normal 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,
|
||||
}
|
||||
26
src/agent/provider/groq/mod.rs
Normal file
26
src/agent/provider/groq/mod.rs
Normal 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
|
||||
}
|
||||
33
src/agent/provider/localai/mod.rs
Normal file
33
src/agent/provider/localai/mod.rs
Normal 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
|
||||
}
|
||||
22
src/agent/provider/mistral/mod.rs
Normal file
22
src/agent/provider/mistral/mod.rs
Normal 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
25
src/agent/provider/mod.rs
Normal 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,
|
||||
};
|
||||
24
src/agent/provider/ollama/mod.rs
Normal file
24
src/agent/provider/ollama/mod.rs
Normal 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
|
||||
}
|
||||
190
src/agent/provider/openai/config.rs
Normal file
190
src/agent/provider/openai/config.rs
Normal 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
|
||||
}
|
||||
465
src/agent/provider/openai/controller.rs
Normal file
465
src/agent/provider/openai/controller.rs
Normal 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))
|
||||
}
|
||||
46
src/agent/provider/openai/mod.rs
Normal file
46
src/agent/provider/openai/mod.rs
Normal 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()
|
||||
}
|
||||
55
src/agent/provider/openai/utils.rs
Normal file
55
src/agent/provider/openai/utils.rs
Normal 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))
|
||||
}
|
||||
}
|
||||
}
|
||||
261
src/agent/provider/openai_compat/config.rs
Normal file
261
src/agent/provider/openai_compat/config.rs
Normal 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())
|
||||
}
|
||||
445
src/agent/provider/openai_compat/controller.rs
Normal file
445
src/agent/provider/openai_compat/controller.rs
Normal 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)
|
||||
}
|
||||
}
|
||||
70
src/agent/provider/openai_compat/mod.rs
Normal file
70
src/agent/provider/openai_compat/mod.rs
Normal 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
|
||||
}
|
||||
78
src/agent/provider/openai_compat/utils.rs
Normal file
78
src/agent/provider/openai_compat/utils.rs
Normal 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))
|
||||
}
|
||||
}
|
||||
}
|
||||
21
src/agent/provider/openrouter/mod.rs
Normal file
21
src/agent/provider/openrouter/mod.rs
Normal 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
|
||||
}
|
||||
21
src/agent/provider/togetherai/mod.rs
Normal file
21
src/agent/provider/togetherai/mod.rs
Normal 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
67
src/agent/purpose.rs
Normal 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
146
src/agent/utils.rs
Normal 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
419
src/bot/implementation.rs
Normal 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
110
src/bot/load_config.rs
Normal 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
334
src/bot/messaging.rs
Normal 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
8
src/bot/mod.rs
Normal 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
262
src/bot/reacting.rs
Normal 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
145
src/bot/rooms.rs
Normal 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())
|
||||
}
|
||||
}
|
||||
63
src/controller/access/determination/mod.rs
Normal file
63
src/controller/access/determination/mod.rs
Normal 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)
|
||||
}
|
||||
58
src/controller/access/determination/tests.rs
Normal file
58
src/controller/access/determination/tests.rs
Normal 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);
|
||||
}
|
||||
}
|
||||
45
src/controller/access/dispatching.rs
Normal file
45
src/controller/access/dispatching.rs
Normal 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
|
||||
}
|
||||
}
|
||||
}
|
||||
183
src/controller/access/help.rs
Normal file
183
src/controller/access/help.rs
Normal 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
|
||||
}
|
||||
8
src/controller/access/mod.rs
Normal file
8
src/controller/access/mod.rs
Normal 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;
|
||||
67
src/controller/access/room_local_agent_managers.rs
Normal file
67
src/controller/access/room_local_agent_managers.rs
Normal 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(())
|
||||
}
|
||||
63
src/controller/access/users.rs
Normal file
63
src/controller/access/users.rs
Normal 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(())
|
||||
}
|
||||
418
src/controller/agent/create/mod.rs
Normal file
418
src/controller/agent/create/mod.rs
Normal 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;
|
||||
}
|
||||
61
src/controller/agent/create/tests.rs
Normal file
61
src/controller/agent/create/tests.rs
Normal 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);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
195
src/controller/agent/delete/mod.rs
Normal file
195
src/controller/agent/delete/mod.rs
Normal 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(())
|
||||
}
|
||||
84
src/controller/agent/details/mod.rs
Normal file
84
src/controller/agent/details/mod.rs
Normal 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(())
|
||||
}
|
||||
99
src/controller/agent/determination/mod.rs
Normal file
99
src/controller/agent/determination/mod.rs
Normal 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)
|
||||
}
|
||||
131
src/controller/agent/determination/tests.rs
Normal file
131
src/controller/agent/determination/tests.rs
Normal 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);
|
||||
}
|
||||
}
|
||||
79
src/controller/agent/help/mod.rs
Normal file
79
src/controller/agent/help/mod.rs
Normal 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(())
|
||||
}
|
||||
39
src/controller/agent/list/mod.rs
Normal file
39
src/controller/agent/list/mod.rs
Normal 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(())
|
||||
}
|
||||
54
src/controller/agent/mod.rs
Normal file
54
src/controller/agent/mod.rs
Normal 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,
|
||||
}
|
||||
}
|
||||
35
src/controller/cfg/common/generic_setting.rs
Normal file
35
src/controller/cfg/common/generic_setting.rs
Normal 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(())
|
||||
}
|
||||
1
src/controller/cfg/common/mod.rs
Normal file
1
src/controller/cfg/common/mod.rs
Normal file
@@ -0,0 +1 @@
|
||||
pub(super) mod generic_setting;
|
||||
74
src/controller/cfg/controller_type.rs
Normal file
74
src/controller/cfg/controller_type.rs
Normal 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>),
|
||||
}
|
||||
141
src/controller/cfg/determination/mod.rs
Normal file
141
src/controller/cfg/determination/mod.rs
Normal 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)
|
||||
}
|
||||
87
src/controller/cfg/determination/speech_to_text/mod.rs
Normal file
87
src/controller/cfg/determination/speech_to_text/mod.rs
Normal 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)
|
||||
}
|
||||
136
src/controller/cfg/determination/speech_to_text/tests.rs
Normal file
136
src/controller/cfg/determination/speech_to_text/tests.rs
Normal 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);
|
||||
}
|
||||
}
|
||||
212
src/controller/cfg/determination/tests.rs
Normal file
212
src/controller/cfg/determination/tests.rs
Normal 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);
|
||||
}
|
||||
}
|
||||
201
src/controller/cfg/determination/text_generation/mod.rs
Normal file
201
src/controller/cfg/determination/text_generation/mod.rs
Normal 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)
|
||||
}
|
||||
332
src/controller/cfg/determination/text_generation/tests.rs
Normal file
332
src/controller/cfg/determination/text_generation/tests.rs
Normal 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);
|
||||
}
|
||||
}
|
||||
158
src/controller/cfg/determination/text_to_speech/mod.rs
Normal file
158
src/controller/cfg/determination/text_to_speech/mod.rs
Normal 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)
|
||||
}
|
||||
252
src/controller/cfg/determination/text_to_speech/tests.rs
Normal file
252
src/controller/cfg/determination/text_to_speech/tests.rs
Normal 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);
|
||||
}
|
||||
}
|
||||
124
src/controller/cfg/dispatching/mod.rs
Normal file
124
src/controller/cfg/dispatching/mod.rs
Normal 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
|
||||
}
|
||||
}
|
||||
}
|
||||
78
src/controller/cfg/dispatching/speech_to_text.rs
Normal file
78
src/controller/cfg/dispatching/speech_to_text.rs
Normal 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
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
155
src/controller/cfg/dispatching/text_generation.rs
Normal file
155
src/controller/cfg/dispatching/text_generation.rs
Normal 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
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
134
src/controller/cfg/dispatching/text_to_speech.rs
Normal file
134
src/controller/cfg/dispatching/text_to_speech.rs
Normal 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
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
38
src/controller/cfg/global_config/generic_setting.rs
Normal file
38
src/controller/cfg/global_config/generic_setting.rs
Normal 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(())
|
||||
}
|
||||
162
src/controller/cfg/global_config/handler.rs
Normal file
162
src/controller/cfg/global_config/handler.rs
Normal 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(())
|
||||
}
|
||||
2
src/controller/cfg/global_config/mod.rs
Normal file
2
src/controller/cfg/global_config/mod.rs
Normal file
@@ -0,0 +1,2 @@
|
||||
pub(super) mod generic_setting;
|
||||
pub(super) mod handler;
|
||||
547
src/controller/cfg/help.rs
Normal file
547
src/controller/cfg/help.rs
Normal 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
12
src/controller/cfg/mod.rs
Normal 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;
|
||||
38
src/controller/cfg/room_config/generic_setting.rs
Normal file
38
src/controller/cfg/room_config/generic_setting.rs
Normal 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(())
|
||||
}
|
||||
141
src/controller/cfg/room_config/handler.rs
Normal file
141
src/controller/cfg/room_config/handler.rs
Normal 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(())
|
||||
}
|
||||
2
src/controller/cfg/room_config/mod.rs
Normal file
2
src/controller/cfg/room_config/mod.rs
Normal file
@@ -0,0 +1,2 @@
|
||||
pub(super) mod generic_setting;
|
||||
pub(super) mod handler;
|
||||
759
src/controller/cfg/status.rs
Normal file
759
src/controller/cfg/status.rs
Normal 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
|
||||
}
|
||||
599
src/controller/chat_completion/mod.rs
Normal file
599
src/controller/chat_completion/mod.rs
Normal 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(),
|
||||
¶ms,
|
||||
)
|
||||
.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
|
||||
}
|
||||
27
src/controller/controller_type.rs
Normal file
27
src/controller/controller_type.rs
Normal 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),
|
||||
}
|
||||
154
src/controller/determination/mod.rs
Normal file
154
src/controller/determination/mod.rs
Normal 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![],
|
||||
})
|
||||
}
|
||||
}
|
||||
}
|
||||
196
src/controller/determination/tests.rs
Normal file
196
src/controller/determination/tests.rs
Normal 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);
|
||||
}
|
||||
}
|
||||
108
src/controller/dispatching.rs
Normal file
108
src/controller/dispatching.rs
Normal 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;
|
||||
}
|
||||
}
|
||||
90
src/controller/help/mod.rs
Normal file
90
src/controller/help/mod.rs
Normal 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(())
|
||||
}
|
||||
179
src/controller/image/generation.rs
Normal file
179
src/controller/image/generation.rs
Normal file
@@ -0,0 +1,179 @@
|
||||
use mxlink::{MatrixLink, MessageResponseType};
|
||||
|
||||
use tracing::Instrument;
|
||||
|
||||
use crate::agent::provider::ImageGenerationParams;
|
||||
use crate::agent::AgentPurpose;
|
||||
use crate::agent::ControllerTrait;
|
||||
use crate::controller::utils::agent::get_effective_agent_for_purpose_or_complain;
|
||||
use crate::conversation::create_llm_conversation_for_matrix_thread;
|
||||
use crate::conversation::matrix::MatrixMessageProcessingParams;
|
||||
use crate::strings;
|
||||
use crate::{entity::MessageContext, Bot};
|
||||
|
||||
// We may make this configurable (per room, etc.) in the future, but for now it's hardcoded.
|
||||
const STICKER_SIZE: &str = "256x256";
|
||||
|
||||
pub async fn handle_image(
|
||||
bot: &Bot,
|
||||
matrix_link: MatrixLink,
|
||||
message_context: &MessageContext,
|
||||
original_prompt: &str,
|
||||
) -> anyhow::Result<()> {
|
||||
let response_type = MessageResponseType::InThread(message_context.thread_info().clone());
|
||||
|
||||
let Some(agent) = get_effective_agent_for_purpose_or_complain(
|
||||
bot,
|
||||
message_context,
|
||||
AgentPurpose::ImageGeneration,
|
||||
response_type.clone(),
|
||||
true,
|
||||
)
|
||||
.await
|
||||
else {
|
||||
return Ok(());
|
||||
};
|
||||
|
||||
let params = MatrixMessageProcessingParams::new(
|
||||
bot.user_id().as_str().to_owned(),
|
||||
message_context.combined_admin_and_user_regexes(),
|
||||
);
|
||||
|
||||
let conversation = create_llm_conversation_for_matrix_thread(
|
||||
matrix_link.clone(),
|
||||
message_context.room(),
|
||||
message_context.thread_info().root_event_id.clone(),
|
||||
¶ms,
|
||||
)
|
||||
.await?;
|
||||
|
||||
let prompt = if conversation.messages.len() >= 2 {
|
||||
// Skip the first message, which contains the original prompt (which we already have)
|
||||
let other_messages = conversation.messages.iter().skip(1).cloned().collect();
|
||||
|
||||
super::prompt::build(original_prompt, other_messages)
|
||||
} else {
|
||||
original_prompt.to_owned()
|
||||
};
|
||||
|
||||
message_context.room().typing_notice(true).await?;
|
||||
|
||||
let span = tracing::debug_span!(
|
||||
"image_generation",
|
||||
agent_id = agent.identifier().as_string()
|
||||
);
|
||||
|
||||
let response = agent
|
||||
.controller()
|
||||
.generate_image(&prompt, ImageGenerationParams::default())
|
||||
.instrument(span)
|
||||
.await?;
|
||||
|
||||
let actual_prompt = response.revised_prompt.as_deref().unwrap_or(&prompt);
|
||||
|
||||
if *actual_prompt.trim() != *prompt.trim() {
|
||||
bot.messaging()
|
||||
.send_notice_markdown_no_fail(
|
||||
message_context.room(),
|
||||
strings::image_generation::revised_prompt(actual_prompt),
|
||||
response_type.clone(),
|
||||
)
|
||||
.await;
|
||||
}
|
||||
|
||||
let attachment_body_text = format!("Generated image based on: {}", actual_prompt);
|
||||
|
||||
let mut event_content = matrix_link
|
||||
.media()
|
||||
.upload_and_prepare_event_content(
|
||||
message_context.room(),
|
||||
&response.mime_type,
|
||||
response.bytes,
|
||||
&attachment_body_text,
|
||||
)
|
||||
.await
|
||||
.map_err(|e| anyhow::anyhow!("Failed to upload and prepare event: {}", e))?;
|
||||
|
||||
matrix_link
|
||||
.messaging()
|
||||
.send_event(
|
||||
message_context.room(),
|
||||
&mut event_content,
|
||||
response_type.clone(),
|
||||
)
|
||||
.await?;
|
||||
|
||||
if conversation.messages.len() == 1 {
|
||||
// If this is the beginning of the thread, send helpful instructions
|
||||
bot.messaging()
|
||||
.send_notice_markdown_no_fail(
|
||||
message_context.room(),
|
||||
strings::image_generation::guide_how_to_proceed(),
|
||||
response_type.clone(),
|
||||
)
|
||||
.await;
|
||||
}
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub async fn handle_sticker(
|
||||
bot: &Bot,
|
||||
matrix_link: MatrixLink,
|
||||
message_context: &MessageContext,
|
||||
original_prompt: &str,
|
||||
) -> anyhow::Result<()> {
|
||||
// Stickers are always sent directly to the room - no threading.
|
||||
let response_type =
|
||||
MessageResponseType::Reply(message_context.thread_info().root_event_id.clone());
|
||||
|
||||
let Some(agent) = get_effective_agent_for_purpose_or_complain(
|
||||
bot,
|
||||
message_context,
|
||||
AgentPurpose::ImageGeneration,
|
||||
response_type.clone(),
|
||||
true,
|
||||
)
|
||||
.await
|
||||
else {
|
||||
return Ok(());
|
||||
};
|
||||
|
||||
message_context.room().typing_notice(true).await?;
|
||||
|
||||
let span = tracing::debug_span!(
|
||||
"sticker_generation",
|
||||
agent_id = agent.identifier().as_string()
|
||||
);
|
||||
|
||||
let params = ImageGenerationParams::default()
|
||||
.with_size_override(Some(STICKER_SIZE.to_owned()))
|
||||
.with_cheaper_model_switching_allowed(true)
|
||||
.with_cheaper_quality_switching_allowed(true);
|
||||
|
||||
let response = agent
|
||||
.controller()
|
||||
.generate_image(original_prompt, params)
|
||||
.instrument(span)
|
||||
.await?;
|
||||
|
||||
let attachment_body_text = format!("Generated sticker image based on: {}", original_prompt);
|
||||
|
||||
let mut event_content = matrix_link
|
||||
.media()
|
||||
.upload_and_prepare_event_content(
|
||||
message_context.room(),
|
||||
&response.mime_type,
|
||||
response.bytes,
|
||||
&attachment_body_text,
|
||||
)
|
||||
.await
|
||||
.map_err(|e| anyhow::anyhow!("Failed to upload and prepare event: {}", e))?;
|
||||
|
||||
matrix_link
|
||||
.messaging()
|
||||
.send_event(message_context.room(), &mut event_content, response_type)
|
||||
.await?;
|
||||
|
||||
Ok(())
|
||||
}
|
||||
2
src/controller/image/mod.rs
Normal file
2
src/controller/image/mod.rs
Normal file
@@ -0,0 +1,2 @@
|
||||
pub mod generation;
|
||||
mod prompt;
|
||||
111
src/controller/image/prompt.rs
Normal file
111
src/controller/image/prompt.rs
Normal file
@@ -0,0 +1,111 @@
|
||||
use crate::conversation::llm::{Author, Message};
|
||||
|
||||
/// Builds a prompt from the original prompt and other messages in the conversation.
|
||||
///
|
||||
/// Only messages authored by the user are considered.
|
||||
///
|
||||
/// Messages that say "Again" (regardless of casing) are ignored. They are considered special messages
|
||||
/// which trigger re-generation, but do not need to be included in the prompt criteria.
|
||||
pub fn build(original_prompt: &str, other_messages: Vec<Message>) -> String {
|
||||
let mut prompt = original_prompt.to_owned();
|
||||
|
||||
// Make a new messages vector that only contains messages we care about
|
||||
let other_messages: Vec<Message> = other_messages
|
||||
.into_iter()
|
||||
.filter(|message| {
|
||||
if let Author::User = message.author {
|
||||
message.message_text.to_lowercase() != "again"
|
||||
} else {
|
||||
false
|
||||
}
|
||||
})
|
||||
.collect();
|
||||
|
||||
if !other_messages.is_empty() {
|
||||
prompt.push_str("\nOther criteria:");
|
||||
for message in other_messages {
|
||||
prompt.push_str(
|
||||
format!("\n- {}", message.message_text.replace("\n", ". ").as_str()).as_str(),
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
prompt
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::build;
|
||||
use super::{Author, Message};
|
||||
|
||||
struct TestCase {
|
||||
original_prompt: &'static str,
|
||||
messages: Vec<Message>,
|
||||
expected_prompt: &'static str,
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_build_prompt() {
|
||||
let test_cases = vec![
|
||||
// Simple case
|
||||
TestCase {
|
||||
original_prompt: "Generate a picture of a cat",
|
||||
messages: vec![],
|
||||
expected_prompt: "Generate a picture of a cat",
|
||||
},
|
||||
// Only a single user message
|
||||
TestCase {
|
||||
original_prompt: "Generate a picture of a dog",
|
||||
messages: vec![Message {
|
||||
author: Author::User,
|
||||
message_text: "Must be blue".to_owned(),
|
||||
}],
|
||||
expected_prompt: "Generate a picture of a dog\nOther criteria:\n- Must be blue",
|
||||
},
|
||||
// Multiple complex user messages dispersed with assistant messages
|
||||
TestCase {
|
||||
original_prompt: "Generate a picture of an elephant",
|
||||
messages: vec![Message {
|
||||
author: Author::User,
|
||||
message_text: "Must be blue".to_owned(),
|
||||
},
|
||||
Message {
|
||||
author: Author::Assistant,
|
||||
message_text: "Whatever".to_owned(),
|
||||
},
|
||||
Message {
|
||||
author: Author::User,
|
||||
message_text: "Must be 3-legged.\nMust be flying.".to_owned(),
|
||||
}],
|
||||
expected_prompt: "Generate a picture of an elephant\nOther criteria:\n- Must be blue\n- Must be 3-legged.. Must be flying.",
|
||||
},
|
||||
// "Again" is ignored.
|
||||
TestCase {
|
||||
original_prompt: "Generate a picture of a grizzly bear",
|
||||
messages: vec![Message {
|
||||
author: Author::User,
|
||||
message_text: "Must be blue".to_owned(),
|
||||
},
|
||||
Message {
|
||||
author: Author::Assistant,
|
||||
message_text: "Whatever".to_owned(),
|
||||
},
|
||||
Message {
|
||||
author: Author::User,
|
||||
message_text: "Again".to_owned(),
|
||||
},
|
||||
Message {
|
||||
author: Author::User,
|
||||
message_text: "again".to_owned(),
|
||||
}],
|
||||
expected_prompt: "Generate a picture of a grizzly bear\nOther criteria:\n- Must be blue",
|
||||
},
|
||||
];
|
||||
|
||||
for test_case in test_cases {
|
||||
let actual_prompt = build(test_case.original_prompt, test_case.messages);
|
||||
|
||||
assert_eq!(actual_prompt, test_case.expected_prompt);
|
||||
}
|
||||
}
|
||||
}
|
||||
28
src/controller/join/mod.rs
Normal file
28
src/controller/join/mod.rs
Normal 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
19
src/controller/mod.rs
Normal 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;
|
||||
97
src/controller/provider/mod.rs
Normal file
97
src/controller/provider/mod.rs
Normal 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(())
|
||||
}
|
||||
50
src/controller/reaction/mod.rs
Normal file
50
src/controller/reaction/mod.rs
Normal 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"
|
||||
)),
|
||||
}
|
||||
}
|
||||
100
src/controller/reaction/text_to_speech.rs
Normal file
100
src/controller/reaction/text_to_speech.rs
Normal 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,
|
||||
}
|
||||
}
|
||||
}
|
||||
21
src/controller/usage/mod.rs
Normal file
21
src/controller/usage/mod.rs
Normal 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(())
|
||||
}
|
||||
63
src/controller/utils/agent.rs
Normal file
63
src/controller/utils/agent.rs
Normal 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
|
||||
}
|
||||
}
|
||||
}
|
||||
29
src/controller/utils/mod.rs
Normal file
29
src/controller/utils/mod.rs
Normal 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
|
||||
}
|
||||
}
|
||||
}
|
||||
168
src/controller/utils/text_to_speech.rs
Normal file
168
src/controller/utils/text_to_speech.rs
Normal 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
|
||||
}
|
||||
17
src/conversation/llm/entity.rs
Normal file
17
src/conversation/llm/entity.rs
Normal 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
Reference in New Issue
Block a user