use std::future::Future;
use std::str::FromStr;
use bytes::Bytes;
use serde::Deserialize;
use serde::Serialize;
use tokio_util::sync::CancellationToken;
use crate::error::ProviderError;
use crate::language_model::ResponseMetadata;
use crate::shared::FileData;
use crate::shared::Headers;
use crate::shared::MediaType;
use crate::shared::ModelId;
use crate::shared::ProviderId;
use crate::shared::ProviderMetadata;
use crate::shared::ProviderOptions;
use crate::shared::Warning;
use crate::shared::base64_bytes;
pub trait ImageModel: Send + Sync + 'static {
fn provider(&self) -> &ProviderId;
fn model_id(&self) -> &ModelId;
fn max_images_per_call(&self) -> Option<usize>;
fn do_generate(
&self,
options: ImageOptions,
) -> impl Future<Output = Result<ImageResult, ProviderError>> + Send;
}
#[derive(Debug, Clone)]
pub struct ImageOptions {
pub prompt: Option<String>,
pub n: u32,
pub size: Option<ImageSize>,
pub aspect_ratio: Option<AspectRatio>,
pub seed: Option<u64>,
pub files: Vec<ImageFile>,
pub mask: Option<ImageFile>,
pub provider_options: ProviderOptions,
pub headers: Headers,
pub cancellation: CancellationToken,
}
impl ImageOptions {
#[must_use]
pub fn new(prompt: impl Into<String>) -> Self {
Self {
prompt: Some(prompt.into()),
..Self::default()
}
}
}
impl Default for ImageOptions {
fn default() -> Self {
Self {
prompt: None,
n: 1,
size: None,
aspect_ratio: None,
seed: None,
files: Vec::new(),
mask: None,
provider_options: ProviderOptions::new(),
headers: Headers::new(),
cancellation: CancellationToken::new(),
}
}
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct ImageFile {
pub data: FileData,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub media_type: Option<MediaType>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub provider_options: Option<ProviderOptions>,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct GeneratedImage {
#[serde(with = "base64_bytes")]
pub data: Bytes,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub media_type: Option<MediaType>,
}
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Serialize, Deserialize)]
pub struct ImageUsage {
#[serde(default, skip_serializing_if = "Option::is_none")]
pub input_tokens: Option<u64>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub output_tokens: Option<u64>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub total_tokens: Option<u64>,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct ImageResult {
pub images: Vec<GeneratedImage>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub is_retryable: Option<bool>,
#[serde(default)]
pub warnings: Vec<Warning>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub provider_metadata: Option<ProviderMetadata>,
#[serde(default)]
pub response: ResponseMetadata,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub usage: Option<ImageUsage>,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)]
#[serde(try_from = "String", into = "String")]
pub struct ImageSize {
pub width: u32,
pub height: u32,
}
impl ImageSize {
#[must_use]
pub fn new(width: u32, height: u32) -> Self {
Self { width, height }
}
}
impl std::fmt::Display for ImageSize {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(f, "{}x{}", self.width, self.height)
}
}
impl FromStr for ImageSize {
type Err = InvalidDimension;
fn from_str(text: &str) -> Result<Self, Self::Err> {
parse_pair(text, 'x')
.map(|(width, height)| Self { width, height })
.ok_or_else(|| InvalidDimension {
text: text.to_owned(),
expected: "WIDTHxHEIGHT",
})
}
}
impl TryFrom<String> for ImageSize {
type Error = InvalidDimension;
fn try_from(text: String) -> Result<Self, Self::Error> {
text.parse()
}
}
impl From<ImageSize> for String {
fn from(size: ImageSize) -> Self {
size.to_string()
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)]
#[serde(try_from = "String", into = "String")]
pub struct AspectRatio {
pub width: u32,
pub height: u32,
}
impl AspectRatio {
#[must_use]
pub fn new(width: u32, height: u32) -> Self {
Self { width, height }
}
}
impl std::fmt::Display for AspectRatio {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(f, "{}:{}", self.width, self.height)
}
}
impl FromStr for AspectRatio {
type Err = InvalidDimension;
fn from_str(text: &str) -> Result<Self, Self::Err> {
parse_pair(text, ':')
.map(|(width, height)| Self { width, height })
.ok_or_else(|| InvalidDimension {
text: text.to_owned(),
expected: "WIDTH:HEIGHT",
})
}
}
impl TryFrom<String> for AspectRatio {
type Error = InvalidDimension;
fn try_from(text: String) -> Result<Self, Self::Error> {
text.parse()
}
}
impl From<AspectRatio> for String {
fn from(ratio: AspectRatio) -> Self {
ratio.to_string()
}
}
fn parse_pair(text: &str, separator: char) -> Option<(u32, u32)> {
let (left, right) = text.trim().split_once(separator)?;
let left: u32 = left.trim().parse().ok()?;
let right: u32 = right.trim().parse().ok()?;
(left > 0 && right > 0).then_some((left, right))
}
#[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)]
#[error("invalid dimension `{text}`: expected `{expected}`")]
pub struct InvalidDimension {
pub text: String,
pub expected: &'static str,
}