2024-09-12 13:44:06 +03:00
|
|
|
use serde::{Deserialize, Serialize};
|
|
|
|
|
|
2024-09-21 14:05:19 +00:00
|
|
|
use crate::agent::default_prompt;
|
2024-09-12 13:44:06 +03:00
|
|
|
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)]
|
2024-10-03 10:36:01 +03:00
|
|
|
pub max_response_tokens: Option<u32>,
|
2024-09-12 13:44:06 +03:00
|
|
|
|
|
|
|
|
#[serde(default)]
|
|
|
|
|
pub max_context_tokens: u32,
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
impl Default for TextGenerationConfig {
|
|
|
|
|
fn default() -> Self {
|
|
|
|
|
Self {
|
|
|
|
|
model_id: default_text_model_id(),
|
2024-09-21 14:05:19 +00:00
|
|
|
prompt: Some(default_prompt().to_owned()),
|
2024-09-12 13:44:06 +03:00
|
|
|
temperature: super::super::default_temperature(),
|
2024-10-03 10:36:01 +03:00
|
|
|
max_response_tokens: Some(4096),
|
2024-09-12 13:44:06 +03:00
|
|
|
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,
|
2025-02-27 09:52:24 +02:00
|
|
|
max_completion_tokens: None,
|
2024-09-12 13:44:06 +03:00
|
|
|
max_context_tokens: self.max_context_tokens,
|
2026-01-23 22:05:53 +01:00
|
|
|
tools: Default::default(),
|
2024-09-12 13:44:06 +03:00
|
|
|
})
|
|
|
|
|
}
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
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> {
|
2026-02-10 14:33:19 +02:00
|
|
|
let model_id =
|
|
|
|
|
convert_string_to_enum::<async_openai::types::audio::SpeechModel>(&self.model_id)?;
|
2024-09-12 13:44:06 +03:00
|
|
|
|
2025-11-30 10:31:06 +02:00
|
|
|
let voice = convert_string_to_enum::<async_openai::types::audio::Voice>(&self.voice)?;
|
2024-09-12 13:44:06 +03:00
|
|
|
|
2026-02-10 14:33:19 +02:00
|
|
|
let response_format = convert_string_to_enum::<
|
|
|
|
|
async_openai::types::audio::SpeechResponseFormat,
|
|
|
|
|
>(&self.response_format)?;
|
2024-09-12 13:44:06 +03:00
|
|
|
|
|
|
|
|
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 {
|
2026-02-10 14:33:19 +02:00
|
|
|
Some(convert_string_to_enum::<
|
|
|
|
|
async_openai::types::images::ImageSize,
|
|
|
|
|
>(size)?)
|
2024-09-12 13:44:06 +03:00
|
|
|
} else {
|
2025-05-11 23:18:19 +03:00
|
|
|
None
|
2024-09-12 13:44:06 +03:00
|
|
|
};
|
|
|
|
|
|
|
|
|
|
let style = if let Some(style) = &self.style {
|
2026-02-10 14:33:19 +02:00
|
|
|
Some(convert_string_to_enum::<
|
|
|
|
|
async_openai::types::images::ImageStyle,
|
|
|
|
|
>(style)?)
|
2024-09-12 13:44:06 +03:00
|
|
|
} else {
|
2025-05-03 09:37:10 +03:00
|
|
|
None
|
2024-09-12 13:44:06 +03:00
|
|
|
};
|
|
|
|
|
|
|
|
|
|
let quality = if let Some(quality) = &self.quality {
|
2026-02-10 14:33:19 +02:00
|
|
|
Some(convert_string_to_enum::<
|
|
|
|
|
async_openai::types::images::ImageQuality,
|
|
|
|
|
>(quality)?)
|
2024-09-12 13:44:06 +03:00
|
|
|
} else {
|
2025-05-03 09:37:10 +03:00
|
|
|
None
|
2024-09-12 13:44:06 +03:00
|
|
|
};
|
|
|
|
|
|
|
|
|
|
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())
|
|
|
|
|
}
|