use serde::{Deserialize, Serialize};
use crate::{MediaSource, ModelCallContext, ModelError, ModelFuture, ModelRef};
#[derive(Clone, Copy, Debug, Default, Deserialize, Eq, PartialEq, Serialize)]
#[serde(rename_all = "snake_case")]
#[non_exhaustive]
pub enum ImageFormat {
#[default]
Png,
Webp,
Jpeg,
}
impl ImageFormat {
#[must_use]
pub const fn media_type(self) -> &'static str {
match self {
Self::Png => "image/png",
Self::Webp => "image/webp",
Self::Jpeg => "image/jpeg",
}
}
}
#[derive(Clone, Debug, Deserialize, PartialEq, Serialize)]
pub struct ImageGenerationRequest {
pub model: ModelRef,
pub prompt: String,
pub count: u8,
pub size: Option<String>,
pub quality: Option<String>,
pub format: ImageFormat,
pub transparent: bool,
}
#[derive(Clone, Debug, Deserialize, Eq, PartialEq, Serialize)]
pub struct GeneratedImage {
pub source: MediaSource,
pub revised_prompt: Option<String>,
}
#[derive(Clone, Debug, Deserialize, Eq, PartialEq, Serialize)]
pub struct ImageGenerationResponse {
pub images: Vec<GeneratedImage>,
}
pub trait ImageGenerationModel: Send + Sync {
fn generate_image(
&self,
request: ImageGenerationRequest,
context: ModelCallContext,
) -> ModelFuture<'_, Result<ImageGenerationResponse, ModelError>>;
}
#[derive(Clone, Copy, Debug, Default, Deserialize, Eq, PartialEq, Serialize)]
#[serde(rename_all = "snake_case")]
#[non_exhaustive]
pub enum SpeechFormat {
#[default]
Mp3,
Opus,
Aac,
Flac,
Wav,
Pcm,
}
impl SpeechFormat {
#[must_use]
pub const fn media_type(self) -> &'static str {
match self {
Self::Mp3 => "audio/mpeg",
Self::Opus => "audio/opus",
Self::Aac => "audio/aac",
Self::Flac => "audio/flac",
Self::Wav => "audio/wav",
Self::Pcm => "audio/L16",
}
}
}
#[derive(Clone, Debug, Deserialize, PartialEq, Serialize)]
pub struct SpeechRequest {
pub model: ModelRef,
pub input: String,
pub voice: String,
pub instructions: Option<String>,
pub format: SpeechFormat,
pub speed: Option<f32>,
}
#[derive(Clone, Debug, Deserialize, Eq, PartialEq, Serialize)]
pub struct SpeechResponse {
pub media_type: String,
#[serde(with = "byte_serde")]
pub bytes: Vec<u8>,
}
pub trait SpeechModel: Send + Sync {
fn synthesize_speech(
&self,
request: SpeechRequest,
context: ModelCallContext,
) -> ModelFuture<'_, Result<SpeechResponse, ModelError>>;
}
#[derive(Clone, Debug, Deserialize, PartialEq, Serialize)]
pub struct TranscriptionRequest {
pub model: ModelRef,
pub file_name: String,
pub media_type: String,
#[serde(with = "byte_serde")]
pub bytes: Vec<u8>,
pub language: Option<String>,
pub prompt: Option<String>,
}
#[derive(Clone, Debug, Deserialize, PartialEq, Serialize)]
pub struct TranscriptionResponse {
pub text: String,
pub language: Option<String>,
pub duration_seconds: Option<f64>,
}
pub trait TranscriptionModel: Send + Sync {
fn transcribe(
&self,
request: TranscriptionRequest,
context: ModelCallContext,
) -> ModelFuture<'_, Result<TranscriptionResponse, ModelError>>;
}
mod byte_serde {
use serde::{Deserialize, Deserializer, Serializer};
pub fn serialize<S>(value: &[u8], serializer: S) -> Result<S::Ok, S::Error>
where
S: Serializer,
{
serializer.serialize_bytes(value)
}
pub fn deserialize<'de, D>(deserializer: D) -> Result<Vec<u8>, D::Error>
where
D: Deserializer<'de>,
{
Vec::<u8>::deserialize(deserializer)
}
}