Files
Genarrative/server-rs/crates/api-server/src/raw_image.rs
T
k88936 462fc856aa 为原始图片上传解析增加超时
从请求入口开始计算准备截止时间,并限制 multipart 上传解析阶段的最长等待。
2026-09-12 21:33:48 +08:00

843 lines
29 KiB
Rust
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
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);
}
}