diff --git a/src/agent/provider/openai/config.rs b/src/agent/provider/openai/config.rs index 2a8b820..8e24cae 100644 --- a/src/agent/provider/openai/config.rs +++ b/src/agent/provider/openai/config.rs @@ -149,13 +149,13 @@ pub struct ImageGenerationConfig { pub model_id: String, #[serde(default = "default_image_style")] - pub style: async_openai::types::ImageStyle, + pub style: Option, #[serde(default = "default_image_size")] pub size: async_openai::types::ImageSize, #[serde(default = "default_image_quality")] - pub quality: async_openai::types::ImageQuality, + pub quality: Option, } impl Default for ImageGenerationConfig { @@ -181,14 +181,14 @@ impl ImageGenerationConfig { } } -fn default_image_style() -> async_openai::types::ImageStyle { - async_openai::types::ImageStyle::Vivid +fn default_image_style() -> Option { + None } fn default_image_size() -> async_openai::types::ImageSize { async_openai::types::ImageSize::S1024x1024 } -fn default_image_quality() -> async_openai::types::ImageQuality { - async_openai::types::ImageQuality::Standard +fn default_image_quality() -> Option { + None } diff --git a/src/agent/provider/openai/controller.rs b/src/agent/provider/openai/controller.rs index f2c68cf..efed833 100644 --- a/src/agent/provider/openai/controller.rs +++ b/src/agent/provider/openai/controller.rs @@ -265,12 +265,15 @@ impl ControllerTrait for Controller { let quality = if params.cheaper_quality_switching_allowed { // Switch to a cheaper quality match &image_generation_config.quality { - async_openai::types::ImageQuality::Standard => { - async_openai::types::ImageQuality::Standard - } - async_openai::types::ImageQuality::HD => { - async_openai::types::ImageQuality::Standard + Some(quality) => match quality { + async_openai::types::ImageQuality::Standard => { + Some(async_openai::types::ImageQuality::Standard) + } + async_openai::types::ImageQuality::HD => { + Some(async_openai::types::ImageQuality::Standard) + } } + None => None } } else { image_generation_config.quality.clone() @@ -284,14 +287,34 @@ impl ControllerTrait for Controller { }) .unwrap_or(image_generation_config.size); - let request = CreateImageRequestArgs::default() + let response_format = match model.clone() { + async_openai::types::ImageModel::DallE2 => Some(async_openai::types::ImageResponseFormat::B64Json), + async_openai::types::ImageModel::DallE3 => Some(async_openai::types::ImageResponseFormat::B64Json), + async_openai::types::ImageModel::Other(model_str) => match model_str.as_str() { + _ => Some(async_openai::types::ImageResponseFormat::B64Json), + }, + }; + + let mut request_builder = CreateImageRequestArgs::default(); + + request_builder .model(model) .prompt(prompt.to_owned()) - .response_format(async_openai::types::ImageResponseFormat::B64Json) - .size(size) - .style(image_generation_config.style.clone()) - .quality(quality) - .build()?; + .size(size); + + if let Some(response_format) = response_format { + request_builder.response_format(response_format); + } + + if let Some(style) = &image_generation_config.style { + request_builder.style(style.clone()); + } + + if let Some(quality) = quality { + request_builder.quality(quality.clone()); + } + + let request = request_builder.build()?; tracing::trace!( ?prompt, diff --git a/src/agent/provider/openai_compat/config.rs b/src/agent/provider/openai_compat/config.rs index 18fcf8e..58f1795 100644 --- a/src/agent/provider/openai_compat/config.rs +++ b/src/agent/provider/openai_compat/config.rs @@ -230,15 +230,15 @@ impl TryInto for ImageGenerationConfig { }; let style = if let Some(style) = &self.style { - convert_string_to_enum::(style)? + Some(convert_string_to_enum::(style)?) } else { - async_openai::types::ImageStyle::Vivid + None }; let quality = if let Some(quality) = &self.quality { - convert_string_to_enum::(quality)? + Some(convert_string_to_enum::(quality)?) } else { - async_openai::types::ImageQuality::Standard + None }; Ok(OpenAIImageGenerationConfig {