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)
|
||||
}
|
||||
Reference in New Issue
Block a user