Initial commit
This commit is contained in:
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))
|
||||
}
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user