Files
Genarrative/server-rs/crates/api-server/src/raw_image.rs
T
k88936 04a0d21351 避免原始 PNG 的重复解析
使用解码后的颜色类型校验位深,将读取限制统一留在 image 解码器。
2026-09-11 19:03:03 +08:00

621 lines
22 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::{GenericImageView, ImageFormat, ImageReader};
use platform_image::{
GPT_IMAGE_2_2K_LONG_EDGE_THRESHOLD, RAW_IMAGE_MAX_EDGE, RAW_IMAGE_MAX_PIXELS,
RawImageEditOptions, ReferenceImage, create_vector_engine_raw_image_edit,
validate_raw_image_edit_dimensions,
};
use serde::Serialize;
use serde_json::json;
use std::io::Cursor;
use crate::{
asset_billing::{
execute_billable_asset_operation_with_cost, with_editor_generation_durable_billing_boundary,
},
auth::AuthenticatedAccessToken,
http_error::AppError,
openai_image_generation::{
map_platform_image_error, record_openai_image_failure_if_configured,
require_openai_image_settings,
},
request_context::RequestContext,
state::AppState,
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;
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 payload = parse_multipart_request(multipart).await?;
let prepared = tokio::task::spawn_blocking(move || prepare_request(payload))
.await
.map_err(|error| {
AppError::from_status(StatusCode::INTERNAL_SERVER_ERROR).with_message(error.to_string())
})??;
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 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(
&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;
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,
Some("raw-image-edit".to_string()),
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))
}
struct PreparedRawImageEdit {
image: ReferenceImage,
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;
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").await?);
}
"mask" => {
if mask.is_some() {
return Err(bad_request("mask 字段不能重复"));
}
mask = Some(read_multipart_image(field, "mask").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(
field: axum::extract::multipart::Field<'_>,
name: &str,
) -> Result<RawImageData, AppError> {
let mime_type = field.content_type().unwrap_or_default().to_string();
if !mime_type.eq_ignore_ascii_case("image/png") {
return Err(bad_request(format!("{name} 必须为 image/png")));
}
let bytes = field.bytes().await.map_err(|error| {
tracing::warn!(field = name, error = %error, "raw image multipart 图片读取失败");
bad_request(format!("{name} 字段读取失败"))
})?;
if bytes.is_empty() {
return Err(bad_request(format!("{name} 文件不能为空")));
}
Ok(RawImageData {
bytes,
mime_type: "image/png".to_string(),
file_name: format!("{name}.png"),
})
}
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<(ReferenceImage, 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 decoded = reader
.decode()
.map_err(|error| map_decode_image_error(field, error))?;
if matches!(
decoded.color(),
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) = decoded.dimensions();
Ok((
ReferenceImage {
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(),
}))
}
#[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 字节"));
}
}