462fc856aa
从请求入口开始计算准备截止时间,并限制 multipart 上传解析阶段的最长等待。
843 lines
29 KiB
Rust
843 lines
29 KiB
Rust
use axum::{
|
||
Json,
|
||
extract::{Extension, Multipart, State},
|
||
http::StatusCode,
|
||
};
|
||
use bytes::Bytes;
|
||
use image::{ImageDecoder, ImageFormat, ImageReader};
|
||
use platform_image::{
|
||
GPT_IMAGE_2_2K_LONG_EDGE_THRESHOLD, RAW_IMAGE_MAX_EDGE, RAW_IMAGE_MAX_PIXELS,
|
||
RawImageEditImage, RawImageEditOptions, create_vector_engine_raw_image_edit,
|
||
validate_raw_image_edit_dimensions,
|
||
};
|
||
use serde::Serialize;
|
||
use serde_json::json;
|
||
use std::io::Cursor;
|
||
use std::time::{Duration, Instant};
|
||
|
||
use crate::{
|
||
asset_billing::{
|
||
execute_billable_asset_operation_with_cost, with_editor_generation_durable_billing_boundary,
|
||
},
|
||
auth::AuthenticatedAccessToken,
|
||
http_error::AppError,
|
||
openai_image_generation::{
|
||
build_openai_image_http_client, map_platform_image_error,
|
||
record_openai_image_failure_if_configured, require_openai_image_settings,
|
||
},
|
||
request_context::RequestContext,
|
||
state::{AppState, RAW_IMAGE_DECODE_MAX_CONCURRENCY},
|
||
tracking::record_external_generation_run_after_success,
|
||
};
|
||
use time::OffsetDateTime;
|
||
|
||
#[derive(Debug)]
|
||
struct RawImageData {
|
||
pub(crate) bytes: Bytes,
|
||
pub(crate) mime_type: String,
|
||
pub(crate) file_name: String,
|
||
}
|
||
|
||
#[derive(Debug)]
|
||
struct RawImageEditRequest {
|
||
pub(crate) image: RawImageData,
|
||
pub(crate) mask: Option<RawImageData>,
|
||
pub(crate) prompt: String,
|
||
pub(crate) quality: Option<String>,
|
||
pub(crate) background: Option<String>,
|
||
pub(crate) output_format: Option<String>,
|
||
pub(crate) width: u32,
|
||
pub(crate) height: u32,
|
||
}
|
||
|
||
#[derive(Debug, Serialize)]
|
||
pub(crate) struct RawImageEditItem {
|
||
pub(crate) b64_json: String,
|
||
}
|
||
|
||
#[derive(Debug, Serialize)]
|
||
pub(crate) struct RawImageEditResponse {
|
||
pub(crate) data: Vec<RawImageEditItem>,
|
||
}
|
||
|
||
const RAW_IMAGE_MAX_TEXT_FIELD_BYTES: usize = 16 * 1024;
|
||
const RAW_IMAGE_MAX_FILE_BYTES: usize = 32 * 1024 * 1024;
|
||
const RAW_IMAGE_MAX_INPUT_BYTES: usize = 48 * 1024 * 1024;
|
||
const RAW_IMAGE_PREPARE_TIMEOUT: Duration = Duration::from_secs(30);
|
||
|
||
pub(crate) async fn edit_raw_image(
|
||
State(state): State<AppState>,
|
||
Extension(request_context): Extension<RequestContext>,
|
||
Extension(authenticated): Extension<AuthenticatedAccessToken>,
|
||
multipart: Multipart,
|
||
) -> Result<Json<RawImageEditResponse>, AppError> {
|
||
let local_deadline = Instant::now()
|
||
.checked_add(RAW_IMAGE_PREPARE_TIMEOUT)
|
||
.unwrap_or_else(Instant::now);
|
||
let parse_deadline = request_context
|
||
.external_call_deadline()
|
||
.map(|deadline| deadline.min(local_deadline))
|
||
.unwrap_or(local_deadline);
|
||
let payload = match tokio::time::timeout_at(
|
||
tokio::time::Instant::from_std(parse_deadline),
|
||
parse_multipart_request(multipart),
|
||
)
|
||
.await
|
||
{
|
||
Ok(result) => result?,
|
||
Err(_) => return Err(raw_image_prepare_timeout_error("raw 图片请求解析超时")),
|
||
};
|
||
let processing_deadline = request_context
|
||
.external_call_deadline()
|
||
.map(|deadline| deadline.min(local_deadline))
|
||
.unwrap_or(local_deadline);
|
||
let permit = match tokio::time::timeout_at(
|
||
tokio::time::Instant::from_std(processing_deadline),
|
||
state.raw_image_decode_limiter().acquire_owned(),
|
||
)
|
||
.await
|
||
{
|
||
Ok(Ok(permit)) => permit,
|
||
Ok(Err(error)) => {
|
||
return Err(
|
||
AppError::from_status(StatusCode::SERVICE_UNAVAILABLE).with_details(json!({
|
||
"provider": "raw-image-edit",
|
||
"code": "RAW_IMAGE_DECODE_LIMITER_UNAVAILABLE",
|
||
"message": format!("raw 图片解码并发控制器不可用:{error}"),
|
||
})),
|
||
);
|
||
}
|
||
Err(_) => return Err(raw_image_prepare_timeout_error("等待 raw 图片解码槽位超时")),
|
||
};
|
||
let worker = tokio::task::spawn_blocking(move || {
|
||
// 超时只能停止 async 等待,permit 必须由 blocking closure 持有到解码真正结束。
|
||
let _permit = permit;
|
||
prepare_request(payload)
|
||
});
|
||
let prepared =
|
||
match tokio::time::timeout_at(tokio::time::Instant::from_std(processing_deadline), worker)
|
||
.await
|
||
{
|
||
Ok(Ok(result)) => result?,
|
||
Ok(Err(error)) => {
|
||
return Err(AppError::from_status(StatusCode::INTERNAL_SERVER_ERROR)
|
||
.with_message(error.to_string()));
|
||
}
|
||
Err(_) => return Err(raw_image_prepare_timeout_error("raw 图片解码处理超时")),
|
||
};
|
||
let settings = require_openai_image_settings(&state)?.with_external_api_audit_context(
|
||
&request_context,
|
||
Some(authenticated.claims().user_id().to_string()),
|
||
None,
|
||
);
|
||
let provider_settings = settings.provider_settings();
|
||
let http_client = build_openai_image_http_client(&settings)?;
|
||
let user_id = authenticated.claims().user_id().to_string();
|
||
let request_id = request_context.request_id().to_string();
|
||
let points_cost = raw_image_edit_price(&state, prepared.width, prepared.height).await?;
|
||
let audit_settings = settings.clone();
|
||
let tracking_state = audit_settings.external_api_audit_state.clone();
|
||
let tracking_payload = json!({
|
||
"width": prepared.width,
|
||
"height": prepared.height,
|
||
"promptChars": prepared.prompt.chars().count(),
|
||
"hasMask": prepared.options.mask.is_some(),
|
||
"quality": prepared.options.quality.as_deref(),
|
||
"background": prepared.options.background.as_deref(),
|
||
"outputFormat": prepared.options.output_format.as_deref(),
|
||
});
|
||
let started_at_micros = (OffsetDateTime::now_utc().unix_timestamp_nanos() / 1_000) as i64;
|
||
let operation = async move {
|
||
let generated = match create_vector_engine_raw_image_edit(
|
||
&http_client,
|
||
&provider_settings,
|
||
prepared.prompt.as_str(),
|
||
prepared.image,
|
||
prepared.options,
|
||
"raw_image_edit",
|
||
)
|
||
.await
|
||
{
|
||
Ok(generated) => generated,
|
||
Err(error) => {
|
||
record_openai_image_failure_if_configured(&audit_settings, &error).await;
|
||
if let Some(state) = tracking_state.as_ref() {
|
||
record_external_generation_run_after_success(
|
||
state,
|
||
platform_image::VECTOR_ENGINE_PROVIDER,
|
||
"raw_image_edit",
|
||
"raw_image_edit",
|
||
tracking_payload,
|
||
started_at_micros,
|
||
false,
|
||
Some(error.message().to_string()),
|
||
None,
|
||
None,
|
||
)
|
||
.await;
|
||
}
|
||
return Err(map_platform_image_error(error));
|
||
}
|
||
};
|
||
let data: Vec<RawImageEditItem> = generated
|
||
.b64_images
|
||
.into_iter()
|
||
.map(|b64_json| RawImageEditItem { b64_json })
|
||
.collect();
|
||
if let Some(state) = tracking_state.as_ref() {
|
||
record_external_generation_run_after_success(
|
||
state,
|
||
platform_image::VECTOR_ENGINE_PROVIDER,
|
||
"raw_image_edit",
|
||
"raw_image_edit",
|
||
tracking_payload,
|
||
started_at_micros,
|
||
true,
|
||
None,
|
||
None,
|
||
Some(json!({ "imageCount": data.len() })),
|
||
)
|
||
.await;
|
||
}
|
||
Ok::<_, AppError>(RawImageEditResponse { data })
|
||
};
|
||
let result = with_editor_generation_durable_billing_boundary(
|
||
execute_billable_asset_operation_with_cost(
|
||
&state,
|
||
user_id.as_str(),
|
||
"raw-image-edit",
|
||
request_id.as_str(),
|
||
u64::from(points_cost),
|
||
operation,
|
||
),
|
||
)
|
||
.await?;
|
||
Ok(Json(result))
|
||
}
|
||
|
||
#[derive(Debug)]
|
||
struct PreparedRawImageEdit {
|
||
image: RawImageEditImage,
|
||
prompt: String,
|
||
options: RawImageEditOptions,
|
||
width: u32,
|
||
height: u32,
|
||
}
|
||
|
||
async fn parse_multipart_request(
|
||
mut multipart: Multipart,
|
||
) -> Result<RawImageEditRequest, AppError> {
|
||
let mut image = None;
|
||
let mut mask = None;
|
||
let mut prompt = None;
|
||
let mut quality = None;
|
||
let mut background = None;
|
||
let mut output_format = None;
|
||
let mut width = None;
|
||
let mut height = None;
|
||
let mut image_bytes_total = 0usize;
|
||
|
||
while let Some(field) = multipart.next_field().await.map_err(|error| {
|
||
tracing::warn!(error = %error, "raw image multipart 字段解析失败");
|
||
bad_request("multipart 请求无效")
|
||
})? {
|
||
let name = field
|
||
.name()
|
||
.ok_or_else(|| bad_request("multipart 字段缺少名称"))?
|
||
.to_string();
|
||
match name.as_str() {
|
||
"image" => {
|
||
if image.is_some() {
|
||
return Err(bad_request("image 字段不能重复"));
|
||
}
|
||
image = Some(read_multipart_image(field, "image", &mut image_bytes_total).await?);
|
||
}
|
||
"mask" => {
|
||
if mask.is_some() {
|
||
return Err(bad_request("mask 字段不能重复"));
|
||
}
|
||
mask = Some(read_multipart_image(field, "mask", &mut image_bytes_total).await?);
|
||
}
|
||
"prompt" => set_text_field(&mut prompt, field, "prompt").await?,
|
||
"quality" => set_text_field(&mut quality, field, "quality").await?,
|
||
"background" => set_text_field(&mut background, field, "background").await?,
|
||
"output_format" => set_text_field(&mut output_format, field, "output_format").await?,
|
||
"width" => set_text_field(&mut width, field, "width").await?,
|
||
"height" => set_text_field(&mut height, field, "height").await?,
|
||
_ => return Err(bad_request(format!("不支持的 multipart 字段:{name}"))),
|
||
}
|
||
}
|
||
|
||
let image = image.ok_or_else(|| bad_request("image 字段不能为空"))?;
|
||
let prompt = prompt.ok_or_else(|| bad_request("prompt 字段不能为空"))?;
|
||
let width = parse_multipart_u32(width, "width")?;
|
||
let height = parse_multipart_u32(height, "height")?;
|
||
|
||
Ok(RawImageEditRequest {
|
||
image,
|
||
mask,
|
||
prompt,
|
||
quality,
|
||
background,
|
||
output_format,
|
||
width,
|
||
height,
|
||
})
|
||
}
|
||
|
||
async fn read_multipart_image(
|
||
mut field: axum::extract::multipart::Field<'_>,
|
||
name: &str,
|
||
total_bytes: &mut usize,
|
||
) -> Result<RawImageData, AppError> {
|
||
let mime_type = field
|
||
.content_type()
|
||
.map(|value| {
|
||
value
|
||
.split(';')
|
||
.next()
|
||
.unwrap_or_default()
|
||
.trim()
|
||
.to_string()
|
||
})
|
||
.unwrap_or_default();
|
||
if !mime_type.eq_ignore_ascii_case("image/png") {
|
||
return Err(bad_request(format!("{name} 必须为 image/png")));
|
||
}
|
||
let mut bytes = Vec::new();
|
||
let mut field_bytes = 0usize;
|
||
while let Some(chunk) = field.chunk().await.map_err(|error| {
|
||
tracing::warn!(field = name, error = %error, "raw image multipart 图片读取失败");
|
||
bad_request(format!("{name} 字段读取失败"))
|
||
})? {
|
||
append_bounded_image_chunk(
|
||
&mut bytes,
|
||
&mut field_bytes,
|
||
total_bytes,
|
||
&chunk,
|
||
name,
|
||
RAW_IMAGE_MAX_FILE_BYTES,
|
||
RAW_IMAGE_MAX_INPUT_BYTES,
|
||
)?;
|
||
}
|
||
if bytes.is_empty() {
|
||
return Err(bad_request(format!("{name} 文件不能为空")));
|
||
}
|
||
Ok(RawImageData {
|
||
bytes: Bytes::from(bytes),
|
||
mime_type: "image/png".to_string(),
|
||
file_name: format!("{name}.png"),
|
||
})
|
||
}
|
||
|
||
fn append_bounded_image_chunk(
|
||
bytes: &mut Vec<u8>,
|
||
field_bytes: &mut usize,
|
||
total_bytes: &mut usize,
|
||
chunk: &[u8],
|
||
name: &str,
|
||
max_field_bytes: usize,
|
||
max_total_bytes: usize,
|
||
) -> Result<(), AppError> {
|
||
let next_field_bytes = field_bytes.saturating_add(chunk.len());
|
||
if next_field_bytes > max_field_bytes {
|
||
tracing::warn!(
|
||
field = name,
|
||
bytes = next_field_bytes,
|
||
max_bytes = max_field_bytes,
|
||
"raw image multipart 图片字段超过大小限制"
|
||
);
|
||
return Err(payload_too_large(format!(
|
||
"{name} 图片字段不能超过 {max_field_bytes} 字节"
|
||
)));
|
||
}
|
||
let next_total_bytes = total_bytes.saturating_add(chunk.len());
|
||
if next_total_bytes > max_total_bytes {
|
||
tracing::warn!(
|
||
field = name,
|
||
bytes = next_total_bytes,
|
||
max_bytes = max_total_bytes,
|
||
"raw image multipart 图片总输入超过大小限制"
|
||
);
|
||
return Err(payload_too_large(format!(
|
||
"image 和 mask 图片总大小不能超过 {max_total_bytes} 字节"
|
||
)));
|
||
}
|
||
bytes.extend_from_slice(chunk);
|
||
*field_bytes = next_field_bytes;
|
||
*total_bytes = next_total_bytes;
|
||
Ok(())
|
||
}
|
||
|
||
async fn set_text_field(
|
||
target: &mut Option<String>,
|
||
mut field: axum::extract::multipart::Field<'_>,
|
||
name: &str,
|
||
) -> Result<(), AppError> {
|
||
if target.is_some() {
|
||
return Err(bad_request(format!("{name} 字段不能重复")));
|
||
}
|
||
let mut bytes = Vec::new();
|
||
while let Some(chunk) = field.chunk().await.map_err(|error| {
|
||
tracing::warn!(field = name, error = %error, "raw image multipart 文本读取失败");
|
||
bad_request(format!("{name} 字段读取失败"))
|
||
})? {
|
||
if bytes.len().saturating_add(chunk.len()) > RAW_IMAGE_MAX_TEXT_FIELD_BYTES {
|
||
tracing::warn!(
|
||
field = name,
|
||
limit_bytes = RAW_IMAGE_MAX_TEXT_FIELD_BYTES,
|
||
"raw image multipart 文本字段超过大小限制"
|
||
);
|
||
return Err(bad_request(format!(
|
||
"{name} 字段不能超过 {RAW_IMAGE_MAX_TEXT_FIELD_BYTES} 字节"
|
||
)));
|
||
}
|
||
bytes.extend_from_slice(&chunk);
|
||
}
|
||
*target = Some(String::from_utf8(bytes).map_err(|error| {
|
||
tracing::warn!(field = name, error = %error, "raw image multipart 文本字段不是有效 UTF-8");
|
||
bad_request(format!("{name} 字段必须为有效 UTF-8 文本"))
|
||
})?);
|
||
Ok(())
|
||
}
|
||
|
||
fn parse_multipart_u32(value: Option<String>, field: &str) -> Result<u32, AppError> {
|
||
let value = value.ok_or_else(|| bad_request(format!("{field} 字段不能为空")))?;
|
||
value
|
||
.trim()
|
||
.parse::<u32>()
|
||
.map_err(|_| bad_request(format!("{field} 必须为有效整数")))
|
||
}
|
||
|
||
fn prepare_request(payload: RawImageEditRequest) -> Result<PreparedRawImageEdit, AppError> {
|
||
if payload.prompt.trim().is_empty() {
|
||
return Err(bad_request("prompt 不能为空"));
|
||
}
|
||
if payload.prompt.len() > RAW_IMAGE_MAX_TEXT_FIELD_BYTES {
|
||
return Err(bad_request(format!(
|
||
"prompt 不能超过 {RAW_IMAGE_MAX_TEXT_FIELD_BYTES} 字节"
|
||
)));
|
||
}
|
||
validate_raw_image_edit_dimensions(payload.width, payload.height)
|
||
.map_err(|error| bad_request(error.to_string()))?;
|
||
validate_optional_value(
|
||
payload.quality.as_deref(),
|
||
"quality",
|
||
["low", "medium", "high", "auto"],
|
||
)?;
|
||
validate_optional_value(
|
||
payload.background.as_deref(),
|
||
"background",
|
||
["transparent", "opaque", "auto"],
|
||
)?;
|
||
validate_optional_value(
|
||
payload.output_format.as_deref(),
|
||
"output_format",
|
||
["png", "webp", "jpeg"],
|
||
)?;
|
||
let quality = normalize_optional(payload.quality);
|
||
let background = normalize_optional(payload.background);
|
||
let output_format = normalize_optional(payload.output_format);
|
||
let (image, image_width, image_height) = decode_image(payload.image, "image")?;
|
||
let mask = payload
|
||
.mask
|
||
.map(|value| {
|
||
let (mask, mask_width, mask_height) = decode_image(value, "mask")?;
|
||
if mask_width != image_width || mask_height != image_height {
|
||
return Err(bad_request("mask 尺寸必须与 image 一致"));
|
||
}
|
||
Ok(mask)
|
||
})
|
||
.transpose()?;
|
||
Ok(PreparedRawImageEdit {
|
||
image,
|
||
prompt: payload.prompt,
|
||
options: RawImageEditOptions {
|
||
quality,
|
||
background,
|
||
output_format,
|
||
width: payload.width,
|
||
height: payload.height,
|
||
mask,
|
||
},
|
||
width: payload.width,
|
||
height: payload.height,
|
||
})
|
||
}
|
||
|
||
fn normalize_optional(value: Option<String>) -> Option<String> {
|
||
value
|
||
.map(|value| value.trim().to_string())
|
||
.filter(|value| !value.is_empty())
|
||
}
|
||
|
||
fn validate_optional_value<const N: usize>(
|
||
value: Option<&str>,
|
||
field: &str,
|
||
allowed: [&str; N],
|
||
) -> Result<(), AppError> {
|
||
let Some(value) = value.map(str::trim).filter(|value| !value.is_empty()) else {
|
||
return Ok(());
|
||
};
|
||
if allowed.contains(&value) {
|
||
return Ok(());
|
||
}
|
||
Err(bad_request(format!("{field} 值无效")))
|
||
}
|
||
|
||
fn decode_image(
|
||
value: RawImageData,
|
||
field: &str,
|
||
) -> Result<(RawImageEditImage, u32, u32), AppError> {
|
||
let mime_type = value.mime_type.trim().to_string();
|
||
if !mime_type.eq_ignore_ascii_case("image/png") {
|
||
return Err(bad_request(format!(
|
||
"{field} Content-Type 必须为 image/png"
|
||
)));
|
||
}
|
||
let bytes = value.bytes;
|
||
if bytes.is_empty() {
|
||
return Err(bad_request(format!("{field} 文件不能为空")));
|
||
}
|
||
let mut reader = ImageReader::new(Cursor::new(bytes.as_ref()))
|
||
.with_guessed_format()
|
||
.map_err(|_| bad_request(format!("{field} 文件必须是有效 PNG 文件")))?;
|
||
let mut limits = image::Limits::default();
|
||
limits.max_image_width = Some(RAW_IMAGE_MAX_EDGE);
|
||
limits.max_image_height = Some(RAW_IMAGE_MAX_EDGE);
|
||
limits.max_alloc = Some(RAW_IMAGE_MAX_PIXELS.saturating_mul(4));
|
||
reader.limits(limits);
|
||
if reader.format() != Some(ImageFormat::Png) {
|
||
return Err(bad_request(format!("{field} 文件必须是有效 PNG 文件")));
|
||
}
|
||
let decoder = reader
|
||
.into_decoder()
|
||
.map_err(|error| map_decode_image_error(field, error))?;
|
||
if matches!(
|
||
decoder.color_type(),
|
||
image::ColorType::Rgba16
|
||
| image::ColorType::Rgb16
|
||
| image::ColorType::L16
|
||
| image::ColorType::La16
|
||
) {
|
||
return Err(bad_request(format!(
|
||
"{field} 必须为 8-bit PNG(每通道 8 位)"
|
||
)));
|
||
}
|
||
let (width, height) = decoder.dimensions();
|
||
let decoded_bytes = usize::try_from(decoder.total_bytes()).map_err(|_| {
|
||
bad_request(format!(
|
||
"{field} 文件超出 PNG 尺寸或解码资源上限(单边不超过 {RAW_IMAGE_MAX_EDGE}px)"
|
||
))
|
||
})?;
|
||
let max_decoded_bytes =
|
||
usize::try_from(RAW_IMAGE_MAX_PIXELS.saturating_mul(4)).unwrap_or(usize::MAX);
|
||
if decoded_bytes > max_decoded_bytes {
|
||
return Err(bad_request(format!(
|
||
"{field} 文件超出 PNG 尺寸或解码资源上限(单边不超过 {RAW_IMAGE_MAX_EDGE}px)"
|
||
)));
|
||
}
|
||
let mut decoded = vec![0_u8; decoded_bytes];
|
||
decoder
|
||
.read_image(&mut decoded)
|
||
.map_err(|error| map_decode_image_error(field, error))?;
|
||
Ok((
|
||
RawImageEditImage {
|
||
bytes,
|
||
file_name: value.file_name,
|
||
mime_type,
|
||
},
|
||
width,
|
||
height,
|
||
))
|
||
}
|
||
|
||
fn map_decode_image_error(field: &str, error: image::ImageError) -> AppError {
|
||
let message = match error {
|
||
image::ImageError::Limits(_) => {
|
||
format!("{field} 文件超出 PNG 尺寸或解码资源上限(单边不超过 {RAW_IMAGE_MAX_EDGE}px)")
|
||
}
|
||
_ => format!("{field} 文件必须是有效 PNG 文件"),
|
||
};
|
||
bad_request(message)
|
||
}
|
||
|
||
async fn raw_image_edit_price(state: &AppState, width: u32, height: u32) -> Result<u32, AppError> {
|
||
let tier = if width.max(height) > GPT_IMAGE_2_2K_LONG_EDGE_THRESHOLD {
|
||
"2K"
|
||
} else {
|
||
"1K"
|
||
};
|
||
state
|
||
.editor_generation_pricing()
|
||
.await
|
||
.map(|pricing| {
|
||
pricing.image_generation_mud_points(Some("quick-edit"), Some("gpt-image-2"), Some(tier))
|
||
})
|
||
.map_err(|error| {
|
||
AppError::from_status(StatusCode::INTERNAL_SERVER_ERROR).with_details(json!({
|
||
"provider": "editor-generation-pricing",
|
||
"message": error.to_string(),
|
||
}))
|
||
})
|
||
}
|
||
|
||
fn bad_request(message: impl Into<String>) -> AppError {
|
||
AppError::from_status(StatusCode::BAD_REQUEST).with_details(json!({
|
||
"provider": "raw-image-edit",
|
||
"message": message.into(),
|
||
}))
|
||
}
|
||
|
||
fn payload_too_large(message: impl Into<String>) -> AppError {
|
||
AppError::from_status(StatusCode::PAYLOAD_TOO_LARGE).with_details(json!({
|
||
"provider": "raw-image-edit",
|
||
"message": message.into(),
|
||
"maxFileBytes": RAW_IMAGE_MAX_FILE_BYTES,
|
||
"maxInputBytes": RAW_IMAGE_MAX_INPUT_BYTES,
|
||
}))
|
||
}
|
||
|
||
fn raw_image_prepare_timeout_error(message: &str) -> AppError {
|
||
AppError::from_status(StatusCode::GATEWAY_TIMEOUT).with_details(json!({
|
||
"provider": "raw-image-edit",
|
||
"code": "RAW_IMAGE_PREPARE_TIMEOUT",
|
||
"message": message,
|
||
"maxConcurrency": RAW_IMAGE_DECODE_MAX_CONCURRENCY,
|
||
}))
|
||
}
|
||
|
||
#[cfg(test)]
|
||
mod tests {
|
||
use super::*;
|
||
use axum::{body::Body, extract::FromRequest, http::Request};
|
||
use image::{ImageFormat, Rgba, RgbaImage};
|
||
use std::io::Cursor;
|
||
|
||
fn png_bytes(width: u32, height: u32) -> Vec<u8> {
|
||
let image = RgbaImage::from_pixel(width, height, Rgba([255, 0, 0, 255]));
|
||
let mut bytes = Vec::new();
|
||
image
|
||
.write_to(&mut Cursor::new(&mut bytes), ImageFormat::Png)
|
||
.expect("test PNG should encode");
|
||
bytes
|
||
}
|
||
|
||
fn request(image: Vec<u8>, mask: Option<Vec<u8>>) -> RawImageEditRequest {
|
||
RawImageEditRequest {
|
||
image: RawImageData {
|
||
bytes: Bytes::from(image),
|
||
mime_type: "image/png".to_string(),
|
||
file_name: "image.png".to_string(),
|
||
},
|
||
mask: mask.map(|bytes| RawImageData {
|
||
bytes: Bytes::from(bytes),
|
||
mime_type: "image/png".to_string(),
|
||
file_name: "mask.png".to_string(),
|
||
}),
|
||
prompt: "edit".to_string(),
|
||
width: 1024,
|
||
height: 1024,
|
||
quality: None,
|
||
background: None,
|
||
output_format: None,
|
||
}
|
||
}
|
||
|
||
fn multipart_body(boundary: &str, image: &[u8]) -> Vec<u8> {
|
||
let mut body = Vec::new();
|
||
let add_text = |body: &mut Vec<u8>, name: &str, value: &str| {
|
||
body.extend_from_slice(format!(
|
||
"--{boundary}\r\nContent-Disposition: form-data; name=\"{name}\"\r\n\r\n{value}\r\n"
|
||
).as_bytes());
|
||
};
|
||
add_text(&mut body, "prompt", "edit");
|
||
add_text(&mut body, "width", "1024");
|
||
add_text(&mut body, "height", "1024");
|
||
body.extend_from_slice(
|
||
format!(
|
||
"--{boundary}\r\nContent-Disposition: form-data; name=\"image\"; filename=\"ignored.png\"\r\nContent-Type: image/png\r\n\r\n"
|
||
)
|
||
.as_bytes(),
|
||
);
|
||
body.extend_from_slice(image);
|
||
body.extend_from_slice(format!("\r\n--{boundary}--\r\n").as_bytes());
|
||
body
|
||
}
|
||
|
||
#[tokio::test]
|
||
async fn multipart_parser_accepts_binary_image_and_text_fields() {
|
||
let boundary = "raw-test-boundary";
|
||
let image = png_bytes(1, 1);
|
||
let body = multipart_body(boundary, &image);
|
||
let request = Request::builder()
|
||
.header(
|
||
"content-type",
|
||
format!("multipart/form-data; boundary={boundary}"),
|
||
)
|
||
.body(Body::from(body))
|
||
.expect("multipart request");
|
||
let multipart = Multipart::from_request(request, &())
|
||
.await
|
||
.expect("multipart");
|
||
let parsed = parse_multipart_request(multipart)
|
||
.await
|
||
.expect("multipart fields should parse");
|
||
|
||
assert_eq!(parsed.prompt, "edit");
|
||
assert_eq!(parsed.width, 1024);
|
||
assert_eq!(parsed.height, 1024);
|
||
assert!(parsed.image.bytes.starts_with(b"\x89PNG\r\n\x1a\n"));
|
||
}
|
||
|
||
#[test]
|
||
fn response_contains_only_data_b64_json() {
|
||
let response = serde_json::to_value(RawImageEditResponse {
|
||
data: vec![RawImageEditItem {
|
||
b64_json: "aGVsbG8=".to_string(),
|
||
}],
|
||
})
|
||
.expect("response should serialize");
|
||
assert_eq!(
|
||
response,
|
||
serde_json::json!({"data": [{"b64_json": "aGVsbG8="}]})
|
||
);
|
||
}
|
||
|
||
#[test]
|
||
fn dimensions_follow_strict_raw_image_contract() {
|
||
assert!(validate_raw_image_edit_dimensions(1024, 1024).is_ok());
|
||
assert!(validate_raw_image_edit_dimensions(3840, 1280).is_ok());
|
||
assert!(validate_raw_image_edit_dimensions(3839, 1280).is_err());
|
||
assert!(validate_raw_image_edit_dimensions(3840, 1264).is_err());
|
||
assert!(validate_raw_image_edit_dimensions(1024, 1000).is_err());
|
||
assert!(validate_raw_image_edit_dimensions(16, 16).is_err());
|
||
assert!(validate_raw_image_edit_dimensions(3840, 3840).is_err());
|
||
}
|
||
|
||
#[test]
|
||
fn input_requires_decodable_png_and_png_mime() {
|
||
assert!(prepare_request(request(png_bytes(1, 1), None)).is_ok());
|
||
|
||
let invalid_bytes = RawImageEditRequest {
|
||
image: RawImageData {
|
||
bytes: Bytes::from_static(b"hello"),
|
||
mime_type: "image/png".to_string(),
|
||
file_name: "image.png".to_string(),
|
||
},
|
||
..request(png_bytes(1, 1), None)
|
||
};
|
||
assert!(prepare_request(invalid_bytes).is_err());
|
||
|
||
let invalid_mime = RawImageEditRequest {
|
||
image: RawImageData {
|
||
bytes: Bytes::from(png_bytes(1, 1)),
|
||
mime_type: "image/jpeg".to_string(),
|
||
file_name: "image.jpg".to_string(),
|
||
},
|
||
..request(png_bytes(1, 1), None)
|
||
};
|
||
assert!(prepare_request(invalid_mime).is_err());
|
||
}
|
||
|
||
#[test]
|
||
fn input_rejects_non_eight_bit_png() {
|
||
let mut bytes = Vec::new();
|
||
{
|
||
let mut encoder = png::Encoder::new(&mut bytes, 1, 1);
|
||
encoder.set_color(png::ColorType::Rgba);
|
||
encoder.set_depth(png::BitDepth::Sixteen);
|
||
let mut writer = encoder.write_header().expect("PNG header");
|
||
writer
|
||
.write_image_data(&[0, 0, 0, 0, 0, 0, 0, 0])
|
||
.expect("PNG body");
|
||
}
|
||
let error = prepare_request(request(bytes, None)).expect_err("16-bit PNG should fail");
|
||
assert!(format!("{error:?}").contains("image 必须为 8-bit PNG(每通道 8 位)"));
|
||
}
|
||
|
||
#[test]
|
||
fn oversized_valid_png_reports_resource_limit() {
|
||
let error = match prepare_request(request(png_bytes(RAW_IMAGE_MAX_EDGE + 1, 1), None)) {
|
||
Ok(_) => panic!("oversized PNG should fail"),
|
||
Err(error) => error,
|
||
};
|
||
|
||
assert!(format!("{error:?}").contains("超出 PNG 尺寸或解码资源上限"));
|
||
}
|
||
|
||
#[test]
|
||
fn mask_must_match_source_image_dimensions() {
|
||
let error = match prepare_request(request(png_bytes(2, 1), Some(png_bytes(1, 1)))) {
|
||
Ok(_) => panic!("mismatched mask should fail before billing"),
|
||
Err(error) => error,
|
||
};
|
||
|
||
assert!(format!("{error:?}").contains("mask 尺寸必须与 image 一致"));
|
||
}
|
||
|
||
#[test]
|
||
fn prompt_uses_raw_utf8_byte_limit() {
|
||
let mut parsed = request(png_bytes(1, 1), None);
|
||
parsed.prompt = "a".repeat(RAW_IMAGE_MAX_TEXT_FIELD_BYTES + 1);
|
||
let error = match prepare_request(parsed) {
|
||
Ok(_) => panic!("oversized prompt should fail before image decode and billing"),
|
||
Err(error) => error,
|
||
};
|
||
|
||
assert!(format!("{error:?}").contains("prompt 不能超过 16384 字节"));
|
||
}
|
||
|
||
#[test]
|
||
fn image_chunks_enforce_field_and_total_limits() {
|
||
let mut bytes = Vec::new();
|
||
let mut field_bytes = 0;
|
||
let mut total_bytes = 0;
|
||
append_bounded_image_chunk(
|
||
&mut bytes,
|
||
&mut field_bytes,
|
||
&mut total_bytes,
|
||
b"abc",
|
||
"image",
|
||
3,
|
||
8,
|
||
)
|
||
.expect("chunk at field limit should pass");
|
||
assert_eq!(bytes, b"abc");
|
||
assert_eq!(field_bytes, 3);
|
||
assert_eq!(total_bytes, 3);
|
||
|
||
let error = append_bounded_image_chunk(
|
||
&mut bytes,
|
||
&mut field_bytes,
|
||
&mut total_bytes,
|
||
b"d",
|
||
"image",
|
||
3,
|
||
8,
|
||
)
|
||
.expect_err("chunk over field limit should fail");
|
||
assert_eq!(error.status_code(), StatusCode::PAYLOAD_TOO_LARGE);
|
||
assert_eq!(bytes, b"abc");
|
||
assert_eq!(field_bytes, 3);
|
||
assert_eq!(total_bytes, 3);
|
||
|
||
let mut other_field = Vec::new();
|
||
let mut other_field_bytes = 0;
|
||
let error = append_bounded_image_chunk(
|
||
&mut other_field,
|
||
&mut other_field_bytes,
|
||
&mut total_bytes,
|
||
b"123456",
|
||
"mask",
|
||
8,
|
||
8,
|
||
)
|
||
.expect_err("chunk over total limit should fail");
|
||
assert_eq!(error.status_code(), StatusCode::PAYLOAD_TOO_LARGE);
|
||
assert!(other_field.is_empty());
|
||
assert_eq!(other_field_bytes, 0);
|
||
assert_eq!(total_bytes, 3);
|
||
}
|
||
}
|