use std::fmt;
use serde_json::{Value, json};
use crate::{DEFAULT_TRANSCRIPTION_PROMPT, Error, Result};
pub(crate) const MAX_AUDIO_BYTES: usize = 25 * 1024 * 1024;
const MAX_AUDIO_FILENAME_CHARACTERS: usize = 120;
const MAX_TRANSCRIPTION_PROMPT_CHARACTERS: usize = 32_000;
const MAX_IMAGE_PROMPT_CHARACTERS: usize = 32_000;
const MIN_IMAGE_PIXELS: u64 = 655_360;
const MAX_IMAGE_PIXELS: u64 = 8_294_400;
const MAX_IMAGE_EDGE: u16 = 3_840;
#[derive(Clone, Eq, PartialEq)]
pub struct AudioInput {
file_name: String,
mime_type: String,
data: Vec<u8>,
}
impl AudioInput {
pub fn new(
file_name: impl Into<String>,
mime_type: impl Into<String>,
data: Vec<u8>,
) -> Result<Self> {
let value = Self {
file_name: file_name.into(),
mime_type: mime_type.into().to_ascii_lowercase(),
data,
};
value.validate()?;
Ok(value)
}
pub fn file_name(&self) -> &str {
&self.file_name
}
pub fn mime_type(&self) -> &str {
&self.mime_type
}
pub fn data(&self) -> &[u8] {
&self.data
}
pub fn len(&self) -> usize {
self.data.len()
}
pub fn is_empty(&self) -> bool {
self.data.is_empty()
}
pub(crate) fn into_parts(self) -> (String, String, Vec<u8>) {
(self.file_name, self.mime_type, self.data)
}
fn validate(&self) -> Result<()> {
if self.file_name.is_empty()
|| self.file_name.chars().count() > MAX_AUDIO_FILENAME_CHARACTERS
|| !self.file_name.chars().all(|character| {
character.is_ascii_alphanumeric() || matches!(character, '.' | '-' | '_')
})
|| matches!(self.file_name.as_str(), "." | "..")
{
return Err(Error::InvalidInput(format!(
"audio filename must contain 1 through {MAX_AUDIO_FILENAME_CHARACTERS} ASCII letters, digits, dots, hyphens, or underscores"
)));
}
if !matches!(
self.mime_type.as_str(),
"audio/flac"
| "audio/x-flac"
| "audio/m4a"
| "audio/mp3"
| "audio/mp4"
| "audio/mpeg"
| "audio/mpga"
| "audio/ogg"
| "audio/opus"
| "audio/wav"
| "audio/x-wav"
| "audio/webm"
| "application/ogg"
| "video/mp4"
| "video/webm"
) {
return Err(Error::InvalidInput(
"audio MIME type must describe a supported FLAC, MP3, MP4, M4A, OGG, WAV, or WebM recording".into(),
));
}
let extension = self
.file_name
.rsplit_once('.')
.map(|(_, extension)| extension.to_ascii_lowercase());
if !matches!(
extension.as_deref(),
Some("flac" | "mp3" | "mp4" | "mpeg" | "mpga" | "m4a" | "ogg" | "wav" | "webm")
) {
return Err(Error::InvalidInput(
"audio filename must use a supported flac, mp3, mp4, mpeg, mpga, m4a, ogg, wav, or webm extension".into(),
));
}
if self.data.is_empty() || self.data.len() > MAX_AUDIO_BYTES {
return Err(Error::InvalidInput(format!(
"audio must contain between 1 and {MAX_AUDIO_BYTES} bytes"
)));
}
Ok(())
}
}
impl fmt::Debug for AudioInput {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("AudioInput")
.field("file_name", &self.file_name)
.field("mime_type", &self.mime_type)
.field("bytes", &self.data.len())
.finish()
}
}
#[derive(Clone, Debug, Eq, PartialEq)]
pub struct TranscriptionRequest {
pub audio: AudioInput,
pub prompt: Option<String>,
pub language: Option<String>,
}
impl TranscriptionRequest {
pub fn new(audio: AudioInput) -> Self {
Self {
audio,
prompt: Some(DEFAULT_TRANSCRIPTION_PROMPT.into()),
language: None,
}
}
pub(crate) fn validate(&self) -> Result<()> {
self.audio.validate()?;
if let Some(prompt) = &self.prompt
&& (prompt.trim().is_empty()
|| prompt.chars().count() > MAX_TRANSCRIPTION_PROMPT_CHARACTERS)
{
return Err(Error::InvalidInput(format!(
"transcription prompt must contain 1 through {MAX_TRANSCRIPTION_PROMPT_CHARACTERS} characters when supplied"
)));
}
if let Some(language) = &self.language
&& (language.len() != 2 || !language.bytes().all(|value| value.is_ascii_lowercase()))
{
return Err(Error::InvalidInput(
"transcription language must be a two-letter lowercase ISO-639-1 code".into(),
));
}
Ok(())
}
}
#[derive(Clone, Debug, Default, Eq, PartialEq)]
pub struct TranscriptionTokenDetails {
pub audio_tokens: Option<u64>,
pub text_tokens: Option<u64>,
}
#[derive(Clone, Debug, Eq, PartialEq)]
pub struct TranscriptionTokenUsage {
pub input_tokens: u64,
pub output_tokens: u64,
pub total_tokens: u64,
pub input_details: Option<TranscriptionTokenDetails>,
}
#[derive(Clone, Debug, PartialEq)]
pub enum TranscriptionUsage {
Tokens(TranscriptionTokenUsage),
DurationSeconds(f64),
}
#[derive(Clone, Debug, PartialEq)]
pub struct Transcription {
pub text: String,
pub usage: Option<TranscriptionUsage>,
pub request_id: Option<String>,
}
#[derive(Clone, Copy, Debug, Default, Eq, Hash, PartialEq)]
pub enum ImageSize {
#[default]
Auto,
Dimensions {
width: u16,
height: u16,
},
}
impl ImageSize {
pub fn dimensions(width: u16, height: u16) -> Result<Self> {
let value = Self::Dimensions { width, height };
value.validate()?;
Ok(value)
}
pub(crate) fn as_api_value(self) -> String {
match self {
Self::Auto => "auto".into(),
Self::Dimensions { width, height } => format!("{width}x{height}"),
}
}
pub(crate) fn validate(self) -> Result<()> {
let Self::Dimensions { width, height } = self else {
return Ok(());
};
let pixels = u64::from(width).saturating_mul(u64::from(height));
let short = width.min(height);
let long = width.max(height);
if width % 16 != 0
|| height % 16 != 0
|| width > MAX_IMAGE_EDGE
|| height > MAX_IMAGE_EDGE
|| short == 0
|| u32::from(long) > u32::from(short).saturating_mul(3)
|| !(MIN_IMAGE_PIXELS..=MAX_IMAGE_PIXELS).contains(&pixels)
{
return Err(Error::InvalidInput(format!(
"GPT Image 2 dimensions must be multiples of 16, no edge may exceed {MAX_IMAGE_EDGE}, the aspect ratio must be at most 3:1, and total pixels must be between {MIN_IMAGE_PIXELS} and {MAX_IMAGE_PIXELS}"
)));
}
Ok(())
}
}
#[derive(Clone, Copy, Debug, Default, Eq, Hash, PartialEq)]
pub enum ImageQuality {
#[default]
Auto,
Low,
Medium,
High,
}
impl ImageQuality {
pub(crate) const fn as_str(self) -> &'static str {
match self {
Self::Auto => "auto",
Self::Low => "low",
Self::Medium => "medium",
Self::High => "high",
}
}
pub(crate) fn parse(value: &str) -> Option<Self> {
match value {
"auto" => Some(Self::Auto),
"low" => Some(Self::Low),
"medium" => Some(Self::Medium),
"high" => Some(Self::High),
_ => None,
}
}
}
#[derive(Clone, Copy, Debug, Default, Eq, Hash, PartialEq)]
pub enum ImageFormat {
#[default]
Png,
Jpeg,
WebP,
}
impl ImageFormat {
pub const fn mime_type(self) -> &'static str {
match self {
Self::Png => "image/png",
Self::Jpeg => "image/jpeg",
Self::WebP => "image/webp",
}
}
pub(crate) const fn as_str(self) -> &'static str {
match self {
Self::Png => "png",
Self::Jpeg => "jpeg",
Self::WebP => "webp",
}
}
pub(crate) fn parse(value: &str) -> Option<Self> {
match value {
"png" => Some(Self::Png),
"jpeg" => Some(Self::Jpeg),
"webp" => Some(Self::WebP),
_ => None,
}
}
}
#[derive(Clone, Copy, Debug, Default, Eq, Hash, PartialEq)]
pub enum ImageBackground {
#[default]
Auto,
Opaque,
}
impl ImageBackground {
pub(crate) const fn as_str(self) -> &'static str {
match self {
Self::Auto => "auto",
Self::Opaque => "opaque",
}
}
}
#[derive(Clone, Copy, Debug, Default, Eq, Hash, PartialEq)]
pub enum Moderation {
#[default]
Auto,
Low,
}
impl Moderation {
pub(crate) const fn as_str(self) -> &'static str {
match self {
Self::Auto => "auto",
Self::Low => "low",
}
}
}
#[derive(Clone, Debug, Eq, PartialEq)]
pub struct ImageGenerationRequest {
pub prompt: String,
pub size: ImageSize,
pub quality: ImageQuality,
pub output_format: ImageFormat,
pub output_compression: Option<u8>,
pub background: ImageBackground,
pub moderation: Moderation,
pub user: Option<String>,
}
impl ImageGenerationRequest {
pub fn new(prompt: impl Into<String>) -> Self {
Self {
prompt: prompt.into(),
size: ImageSize::Auto,
quality: ImageQuality::Auto,
output_format: ImageFormat::Png,
output_compression: None,
background: ImageBackground::Auto,
moderation: Moderation::Auto,
user: None,
}
}
pub(crate) fn validate(&self) -> Result<()> {
if self.prompt.trim().is_empty()
|| self.prompt.chars().count() > MAX_IMAGE_PROMPT_CHARACTERS
{
return Err(Error::InvalidInput(format!(
"image prompt must contain 1 through {MAX_IMAGE_PROMPT_CHARACTERS} characters"
)));
}
self.size.validate()?;
if self.output_compression.is_some_and(|value| value > 100) {
return Err(Error::InvalidInput(
"output compression must be between 0 and 100".into(),
));
}
if self.output_compression.is_some() && self.output_format == ImageFormat::Png {
return Err(Error::InvalidInput(
"output compression is supported only for JPEG and WebP images".into(),
));
}
if let Some(user) = &self.user
&& (user.trim().is_empty()
|| user.chars().count() > 512
|| user.chars().any(char::is_control))
{
return Err(Error::InvalidInput(
"image user identifier must contain 1 through 512 non-control characters when supplied".into(),
));
}
Ok(())
}
pub(crate) fn payload(&self) -> Value {
let mut payload = json!({
"model": crate::GPT_IMAGE_2,
"prompt": self.prompt,
"n": 1,
"size": self.size.as_api_value(),
"quality": self.quality.as_str(),
"output_format": self.output_format.as_str(),
"background": self.background.as_str(),
"moderation": self.moderation.as_str(),
"stream": false
});
let object = payload.as_object_mut().expect("image payload is an object");
if let Some(compression) = self.output_compression {
object.insert("output_compression".into(), json!(compression));
}
if let Some(user) = &self.user {
object.insert("user".into(), json!(user));
}
payload
}
}
#[derive(Clone, Eq, PartialEq)]
pub struct GeneratedImage {
pub data: Vec<u8>,
pub format: ImageFormat,
}
impl fmt::Debug for GeneratedImage {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("GeneratedImage")
.field("format", &self.format)
.field("bytes", &self.data.len())
.finish()
}
}
#[derive(Clone, Debug, Default, Eq, PartialEq)]
pub struct ImageTokenDetails {
pub text_tokens: u64,
pub image_tokens: u64,
}
#[derive(Clone, Debug, Eq, PartialEq)]
pub struct ImageUsage {
pub input_tokens: u64,
pub output_tokens: u64,
pub total_tokens: u64,
pub input_details: ImageTokenDetails,
pub output_details: Option<ImageTokenDetails>,
}
#[derive(Clone, Debug, PartialEq)]
pub struct ImageGeneration {
pub created: u64,
pub image: GeneratedImage,
pub size: Option<String>,
pub quality: Option<ImageQuality>,
pub usage: Option<ImageUsage>,
pub request_id: Option<String>,
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn audio_debug_omits_bytes_and_rejects_unsafe_names() {
let audio = AudioInput::new("note.webm", "audio/webm", vec![7, 8, 9]).unwrap();
let debug = format!("{audio:?}");
assert!(debug.contains("bytes: 3"));
assert!(!debug.contains("7, 8, 9"));
assert!(AudioInput::new("../note.webm", "audio/webm", vec![1]).is_err());
}
#[test]
fn image_dimensions_enforce_current_gpt_image_2_constraints() {
assert_eq!(
ImageSize::dimensions(2048, 2048).unwrap(),
ImageSize::Dimensions {
width: 2048,
height: 2048
}
);
assert!(ImageSize::dimensions(1000, 1000).is_err());
assert!(ImageSize::dimensions(3840, 3840).is_err());
assert!(ImageSize::dimensions(3072, 1024).is_ok());
assert!(ImageSize::dimensions(3088, 1024).is_err());
}
#[test]
fn png_rejects_compression_but_jpeg_accepts_it() {
let mut request = ImageGenerationRequest::new("draw a lighthouse");
request.output_compression = Some(80);
assert!(request.validate().is_err());
request.output_format = ImageFormat::Jpeg;
assert!(request.validate().is_ok());
request.output_compression = Some(101);
assert!(request.validate().is_err());
}
#[test]
fn image_payload_is_single_shot_gpt_image_2() {
let request = ImageGenerationRequest::new("draw a lighthouse");
let payload = request.payload();
assert_eq!(payload["model"], "gpt-image-2");
assert_eq!(payload["n"], 1);
assert_eq!(payload["stream"], false);
assert_eq!(payload["background"], "auto");
assert_eq!(payload["output_format"], "png");
}
}