Initial commit

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

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

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

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

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

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

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

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

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

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

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

@@ -0,0 +1,13 @@
mod agent_provider;
mod image_generation;
mod ping;
mod speech_to_text;
mod text_generation;
mod text_to_speech;
pub use agent_provider::{AgentProvider, AgentProviderInfo};
pub use image_generation::{ImageGenerationParams, ImageGenerationResult};
pub use ping::PingResult;
pub use speech_to_text::{SpeechToTextParams, SpeechToTextResult};
pub use text_generation::{TextGenerationParams, TextGenerationResult};
pub use text_to_speech::{TextToSpeechParams, TextToSpeechResult};

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

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

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

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

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

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

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