Initial commit

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

View File

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

View File

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

View File

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

View File

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