Files
Genarrative/server-rs/crates/platform-matting/src/lib.rs
T
lhk229 0c04bbbea3 内存优化 (#91)
BGfilter服务调用改为传url,减少内存占用

Reviewed-on: https://git.genarrative.world/git/GenarrativeAI/Genarrative/pulls/91
Reviewed-by: 段舒康 <kdletters@qq.com>
Co-authored-by: Linghong <ink29535@proton.me>
Co-committed-by: Linghong <ink29535@proton.me>
2026-07-18 21:17:35 +08:00

1535 lines
57 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.
//! 阿里云 VIAPI 通用抠图(SegmentCommonImage)客户端。
//!
//! 官方没有 Rust SDK,这里按 ACS3-HMAC-SHA256 签名协议手搓 HTTP 调用,
//! 签名实现与 platform-auth 的阿里云短信调用保持一致。
//! 文档:https://help.aliyun.com/zh/viapi/use-cases/general-image-segmentation
use std::collections::BTreeMap;
use hmac::{Hmac, Mac};
use reqwest::{Client, StatusCode};
use sha2::{Digest, Sha256};
use time::OffsetDateTime;
use tracing::{info, warn};
type HmacSha256 = Hmac<Sha256>;
pub const DEFAULT_IMAGESEG_ENDPOINT: &str = "imageseg.cn-shanghai.aliyuncs.com";
const IMAGESEG_API_VERSION: &str = "2019-12-30";
const SEGMENT_COMMON_IMAGE_ACTION: &str = "SegmentCommonImage";
// VIAPI 新版官方 SDK 的 AdvanceRequest 上传通道:先向 Open Platform 申请单对象
// Policy,再 multipart POST 到授权返回的上海地域临时 OSS。
const OPEN_PLATFORM_ENDPOINT: &str = "openplatform.aliyuncs.com";
const AUTHORIZE_FILE_UPLOAD_ACTION: &str = "AuthorizeFileUpload";
const AUTHORIZE_FILE_UPLOAD_VERSION: &str = "2019-12-19";
const IMAGESEG_PRODUCT: &str = "imageseg";
pub const DEFAULT_MATTING_REQUEST_TIMEOUT_MS: u64 = 30_000;
/// SegmentCommonImage 要求每条边小于 2000。
const MAX_INPUT_EDGE: u32 = 1_999;
/// SegmentCommonImage 要求输入体积不超过 3 MB。
const MAX_INPUT_BYTES: usize = 3 * 1024 * 1024;
/// SegmentCommonImage 要求每条边大于 32 像素。
const MIN_INPUT_EDGE: u32 = 32;
/// 抠图结果下载的最大字节数。上游是阿里云返回的原尺寸 RGBA PNG,正常远小于此值;
/// 设上限是为了防止异常 / 恶意上游用超大响应撑爆 API 进程内存(逐帧动画可并发多次调用)。
/// 与 api-server BgFilter 降级路径的 `EDITOR_BACKGROUND_REMOVAL_MAX_RESPONSE_BYTES` 对齐。
const MAX_RESULT_RESPONSE_BYTES: usize = 32 * 1024 * 1024;
/// 图片解码的最大边长像素(源图与抠图结果共用)。防止解码阶段按声明的巨幅尺寸分配
/// 像素缓冲。与 api-server BgFilter 路径的 `EDITOR_BACKGROUND_REMOVAL_MAX_IMAGE_DIMENSION` 对齐。
const MAX_DECODE_IMAGE_DIMENSION: u32 = 8192;
#[derive(Clone, Debug)]
pub struct MattingConfig {
pub endpoint: String,
pub access_key_id: String,
pub access_key_secret: String,
pub request_timeout_ms: u64,
}
impl MattingConfig {
pub fn new(
endpoint: String,
access_key_id: String,
access_key_secret: String,
) -> Result<Self, MattingError> {
Self::with_timeout(
endpoint,
access_key_id,
access_key_secret,
DEFAULT_MATTING_REQUEST_TIMEOUT_MS,
)
}
pub fn with_timeout(
endpoint: String,
access_key_id: String,
access_key_secret: String,
request_timeout_ms: u64,
) -> Result<Self, MattingError> {
let endpoint = endpoint
.trim()
.trim_start_matches("https://")
.trim_start_matches("http://")
.trim_matches('/')
.to_string();
if endpoint.is_empty() {
return Err(MattingError::InvalidConfig(
"抠图服务 endpoint 不能为空".to_string(),
));
}
let access_key_id = access_key_id.trim().to_string();
let access_key_secret = access_key_secret.trim().to_string();
if access_key_id.is_empty() || access_key_secret.is_empty() {
return Err(MattingError::InvalidConfig(
"抠图服务 AccessKeyId/AccessKeySecret 不能为空".to_string(),
));
}
Ok(Self {
endpoint,
access_key_id,
access_key_secret,
request_timeout_ms: request_timeout_ms.max(1),
})
}
}
/// SegmentCommonImage 的 ReturnForm 取值。
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum SegmentReturnForm {
/// 透明背景 PNG(默认,抠像主用途)。
Crop,
/// 黑白 mask 图。
Mask,
/// 白底图。
WhiteBackground,
}
impl SegmentReturnForm {
fn as_str(&self) -> &'static str {
match self {
Self::Crop => "crop",
Self::Mask => "mask",
Self::WhiteBackground => "whiteBK",
}
}
}
#[derive(Clone, Debug)]
pub struct SegmentCommonImageRequest {
/// 待抠图图片的公网可访问 URL(推荐上海地域 OSS 签名 URL)。
pub image_url: String,
/// None 表示不传 ReturnForm,使用服务端默认行为。
pub return_form: Option<SegmentReturnForm>,
}
#[derive(Clone, Debug)]
pub struct SegmentCommonImageResult {
/// 抠图结果图片 URL(阿里云临时地址,30 分钟内有效,需及时转存)。
pub image_url: String,
pub request_id: Option<String>,
}
#[derive(Debug)]
pub enum MattingError {
InvalidConfig(String),
InvalidRequest(String),
LocalProcessing(LocalProcessingFailure),
Sign(String),
Upstream(UpstreamFailure),
}
/// 外部调用已经开始,但本地处理阶段失败。
///
/// 这类错误仍然需要进入外部失败审计,但不能被误标成上游 HTTP / 传输故障。
#[derive(Debug)]
pub struct LocalProcessingFailure {
message: String,
failure_stage: &'static str,
}
/// 结构化的上游调用失败分类。协议层归属(是否传输层故障、是否超时、上游 HTTP 状态)由
/// platform-matting 在错误发生处直接捕获,调用方(api-server BFF)不再从中文 message 反推。
#[derive(Debug)]
pub struct UpstreamFailure {
message: String,
/// 传输层超时(reqwest `is_timeout`)。
timeout: bool,
/// 传输层故障:请求已发出但未拿到有效 HTTP 响应(连接失败、读体中断、超时)。
transport: bool,
/// 上游返回的 HTTP 状态码;`None` 表示没拿到状态(传输层故障或响应体不可用)。
upstream_status: Option<u16>,
failure_stage: &'static str,
}
impl MattingError {
pub fn message(&self) -> &str {
match self {
Self::InvalidConfig(message) | Self::InvalidRequest(message) | Self::Sign(message) => {
message
}
Self::LocalProcessing(failure) => &failure.message,
Self::Upstream(failure) => &failure.message,
}
}
/// 是否已经开始外部调用链路。
///
/// `LocalProcessing` 表示外部调用已经开始,但解码、尺寸校验或其它本地处理失败;它不能
/// 与真正发请求前的 `InvalidConfig` / `InvalidRequest` / `Sign` 混为一谈。
pub fn external_call_attempted(&self) -> bool {
matches!(self, Self::LocalProcessing(_) | Self::Upstream(_))
}
/// 上游调用是否为传输层超时。
pub fn is_timeout(&self) -> bool {
matches!(self, Self::Upstream(failure) if failure.timeout)
}
/// 上游调用是否为传输层故障(未拿到有效 HTTP 响应)。
pub fn is_transport(&self) -> bool {
matches!(self, Self::Upstream(failure) if failure.transport)
}
/// 上游返回的 HTTP 状态码(若有)。
pub fn upstream_status(&self) -> Option<u16> {
match self {
Self::Upstream(failure) => failure.upstream_status,
_ => None,
}
}
/// 失败发生阶段,供 api-server 生成准确的外部失败审计 metadata。
pub fn failure_stage(&self) -> &'static str {
match self {
Self::InvalidConfig(_) | Self::InvalidRequest(_) | Self::Sign(_) => "preflight",
Self::LocalProcessing(failure) => failure.failure_stage,
Self::Upstream(failure) => failure.failure_stage,
}
}
/// 将一个已经发生外部调用后的本地错误转换为可审计的处理失败。
pub fn with_failure_stage(self, failure_stage: &'static str) -> Self {
match self {
Self::InvalidConfig(message) | Self::InvalidRequest(message) | Self::Sign(message) => {
Self::LocalProcessing(LocalProcessingFailure {
message,
failure_stage,
})
}
Self::LocalProcessing(mut failure) => {
failure.failure_stage = failure_stage;
Self::LocalProcessing(failure)
}
Self::Upstream(mut failure) => {
failure.failure_stage = failure_stage;
Self::Upstream(failure)
}
}
}
/// 传输层故障构造器:请求已发出但没拿到有效 HTTP 响应(连接、读体、超时)。
pub fn upstream_transport_error(message: String, timeout: bool) -> Self {
Self::Upstream(UpstreamFailure {
message,
timeout,
transport: true,
upstream_status: None,
failure_stage: "aliyun_segment",
})
}
/// 上游返回了 HTTP 响应但状态非成功。
pub fn upstream_http_error(message: String, status: u16) -> Self {
Self::Upstream(UpstreamFailure {
message,
timeout: false,
transport: false,
upstream_status: Some(status),
failure_stage: "aliyun_segment",
})
}
/// 已拿到 HTTP 响应,但响应体 / 结果不可用(JSON 非法、缺字段、结果图不可解码 / 尺寸不符、
/// 结果编码失败等)。不是传输层故障,也不归因某个错误状态码。
pub fn upstream_response_error(message: String, upstream_status: Option<u16>) -> Self {
Self::Upstream(UpstreamFailure {
message,
timeout: false,
transport: false,
upstream_status,
failure_stage: "aliyun_segment",
})
}
}
impl std::fmt::Display for MattingError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str(self.message())
}
}
impl std::error::Error for MattingError {}
#[derive(Clone)]
pub struct MattingClient {
config: MattingConfig,
client: Client,
}
impl std::fmt::Debug for MattingClient {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("MattingClient")
.field("endpoint", &self.config.endpoint)
.finish_non_exhaustive()
}
}
impl MattingClient {
pub fn new(config: MattingConfig) -> Result<Self, MattingError> {
let client = Client::builder()
.timeout(std::time::Duration::from_millis(config.request_timeout_ms))
.build()
.map_err(|error| {
MattingError::InvalidConfig(format!("构建 reqwest client 失败:{error}"))
})?;
Ok(Self { config, client })
}
pub async fn segment_common_image(
&self,
request: SegmentCommonImageRequest,
) -> Result<SegmentCommonImageResult, MattingError> {
let image_url = request.image_url.trim().to_string();
if image_url.is_empty() {
return Err(MattingError::InvalidRequest(
"imageUrl 不能为空".to_string(),
));
}
let mut form = BTreeMap::new();
form.insert(
"Action".to_string(),
SEGMENT_COMMON_IMAGE_ACTION.to_string(),
);
form.insert("Format".to_string(), "json".to_string());
form.insert("Version".to_string(), IMAGESEG_API_VERSION.to_string());
form.insert("ImageURL".to_string(), image_url);
if let Some(return_form) = request.return_form {
form.insert("ReturnForm".to_string(), return_form.as_str().to_string());
}
let payload = build_aliyun_form_body(&form);
let headers = self.build_signature_headers(SEGMENT_COMMON_IMAGE_ACTION, &payload)?;
info!(
provider = "aliyun-imageseg",
endpoint = self.config.endpoint.as_str(),
action = SEGMENT_COMMON_IMAGE_ACTION,
return_form = request
.return_form
.map_or("default", |return_form| return_form.as_str()),
"准备调用阿里云通用抠图接口"
);
let response = self
.client
.post(format!("https://{}/", self.config.endpoint))
.headers(headers)
.header(
reqwest::header::CONTENT_TYPE,
"application/x-www-form-urlencoded",
)
.body(payload)
.send()
.await
.map_err(|error| {
MattingError::upstream_transport_error(
format!("通用抠图请求失败:{error}"),
error.is_timeout(),
)
})?;
let http_status = response.status();
let body_text = response.text().await.map_err(|error| {
MattingError::upstream_transport_error(
format!("通用抠图响应读取失败:{error}"),
error.is_timeout(),
)
})?;
let body: serde_json::Value = serde_json::from_str(&body_text).map_err(|error| {
MattingError::upstream_response_error(
format!(
"通用抠图响应不是合法 JSON:{error};原始响应:{}",
truncate_for_log(&body_text)
),
Some(http_status.as_u16()),
)
})?;
let request_id = body
.get("RequestId")
.and_then(|value| value.as_str())
.map(|value| value.to_string());
if http_status != StatusCode::OK {
let code = body.get("Code").and_then(|value| value.as_str());
let message = body.get("Message").and_then(|value| value.as_str());
warn!(
provider = "aliyun-imageseg",
http_status = http_status.as_u16(),
provider_code = code.unwrap_or("unknown"),
provider_message = message.unwrap_or("unknown"),
provider_request_id = request_id.as_deref().unwrap_or("unknown"),
"阿里云通用抠图接口返回失败"
);
return Err(MattingError::upstream_http_error(
format!(
"通用抠图接口返回失败(HTTP {},Code={}):{}",
http_status.as_u16(),
code.unwrap_or("unknown"),
message.unwrap_or("unknown")
),
http_status.as_u16(),
));
}
let result_image_url = body
.get("Data")
.and_then(|data| data.get("ImageURL"))
.and_then(|value| value.as_str())
.map(|value| value.to_string())
.ok_or_else(|| {
MattingError::upstream_response_error(
format!(
"通用抠图响应缺少 Data.ImageURL;原始响应:{}",
truncate_for_log(&body_text)
),
Some(http_status.as_u16()),
)
})?;
info!(
provider = "aliyun-imageseg",
provider_request_id = request_id.as_deref().unwrap_or("unknown"),
"阿里云通用抠图接口调用成功"
);
Ok(SegmentCommonImageResult {
image_url: result_image_url,
request_id,
})
}
/// 私有 OSS 签名 URL → 下载 → AuthorizeFileUpload 临时对象 → 通用抠图 → 透明 PNG。
///
/// 源图下载缓冲和解码图只活到临时对象上传完成;若输入发生过降尺寸,结果返回后再单独下载
/// 一次源图完成 Alpha 回贴,避免在整个阿里云推理期间常驻原图内存。
pub async fn segment_image_url_to_transparent_png(
&self,
source_url: &str,
file_name: &str,
) -> Result<Vec<u8>, MattingError> {
let source_url = source_url.trim();
if source_url.is_empty() {
return Err(MattingError::InvalidRequest(
"待抠图源图 URL 不能为空".to_string(),
));
}
let (source_dims, upload_dims, image_url) = {
let source_bytes = self
.download_source_image(source_url)
.await
.map_err(|error| error.with_failure_stage("source_download"))?;
let source = decode_source_image(&source_bytes)
.map_err(|error| error.with_failure_stage("source_decode"))?;
let source_dims = (source.width(), source.height());
validate_source_dimensions(source_dims)
.map_err(|error| error.with_failure_stage("source_validate"))?;
let (upload_bytes, upload_dims) = normalize_matting_input_png(&source)
.map_err(|error| error.with_failure_stage("source_normalize"))?;
let image_url = self
.upload_temp_image(upload_bytes, file_name, "image/png")
.await
.map_err(|error| error.with_failure_stage("temp_upload"))?;
(source_dims, upload_dims, image_url)
};
let result_image = self
.segment_uploaded_input_to_rgba(image_url, source_dims, upload_dims)
.await?;
if upload_dims == source_dims {
return encode_rgba_png(&result_image)
.map_err(|error| error.with_failure_stage("result_encode"));
}
// 只有降尺寸送抠时才重新下载源图恢复原始 RGB;第一次下载缓冲已在上传临时对象后释放。
let source_rgba = {
let source_bytes = self
.download_source_image(source_url)
.await
.map_err(|error| error.with_failure_stage("source_download"))?;
let source = decode_source_image(&source_bytes)
.map_err(|error| error.with_failure_stage("source_decode"))?;
if (source.width(), source.height()) != source_dims {
return Err(MattingError::upstream_response_error(
"待抠图源图在处理期间尺寸发生变化".to_string(),
None,
)
.with_failure_stage("source_validate"));
}
source.to_rgba8()
};
compose_source_with_result_alpha(source_rgba, result_image)
.map_err(|error| error.with_failure_stage("result_encode"))
}
async fn segment_uploaded_input_to_rgba(
&self,
image_url: String,
source_dims: (u32, u32),
upload_dims: (u32, u32),
) -> Result<image::RgbaImage, MattingError> {
let result = self
.segment_common_image(SegmentCommonImageRequest {
image_url,
// 默认 ReturnForm:原尺寸 + 透明背景,无需本地合成。
return_form: None,
})
.await
.map_err(|error| error.with_failure_stage("aliyun_segment"))?;
let result_image = {
let result_bytes = self
.download_result_image(&result.image_url)
.await
.map_err(|error| error.with_failure_stage("result_download"))?;
decode_result_image(&result_bytes)
.map_err(|error| error.with_failure_stage("result_decode"))?
.to_rgba8()
};
if upload_dims == source_dims && result_image.dimensions() != source_dims {
return Err(MattingError::upstream_response_error(
format!(
"抠图结果尺寸 {}x{} 与输入 {}x{} 不一致",
result_image.width(),
result_image.height(),
source_dims.0,
source_dims.1
),
None,
)
.with_failure_stage("result_validate"));
}
Ok(result_image)
}
async fn download_source_image(&self, url: &str) -> Result<Vec<u8>, MattingError> {
let mut response = self.client.get(url).send().await.map_err(|error| {
MattingError::upstream_transport_error(
format!(
"下载待抠图源图失败(transport={}, timeout={}, connect={})",
classify_reqwest_error(&error),
error.is_timeout(),
error.is_connect()
),
error.is_timeout(),
)
})?;
let status = response.status();
if !status.is_success() {
return Err(MattingError::upstream_http_error(
format!("下载待抠图源图失败(HTTP {})", status.as_u16()),
status.as_u16(),
));
}
if response
.content_length()
.is_some_and(|length| length > MAX_RESULT_RESPONSE_BYTES as u64)
{
return Err(source_response_too_large_error());
}
let mut bytes = Vec::new();
while let Some(chunk) = response.chunk().await.map_err(|error| {
MattingError::upstream_transport_error(
format!(
"读取待抠图源图失败(transport={}, timeout={}, connect={})",
classify_reqwest_error(&error),
error.is_timeout(),
error.is_connect()
),
error.is_timeout(),
)
})? {
if bytes.len().saturating_add(chunk.len()) > MAX_RESULT_RESPONSE_BYTES {
return Err(source_response_too_large_error());
}
bytes.extend_from_slice(chunk.as_ref());
}
if bytes.is_empty() {
return Err(MattingError::upstream_response_error(
"待抠图源图为空".to_string(),
Some(status.as_u16()),
));
}
Ok(bytes)
}
async fn download_result_image(&self, url: &str) -> Result<Vec<u8>, MattingError> {
let mut response = self.client.get(url).send().await.map_err(|error| {
MattingError::upstream_transport_error(
describe_result_download_transport_error(&error),
error.is_timeout(),
)
})?;
let status = response.status();
if !status.is_success() {
return Err(MattingError::upstream_http_error(
format!("下载抠图结果失败(HTTP {})", status.as_u16()),
status.as_u16(),
));
}
// 先按 Content-Length 快速拒绝,再流式累加做兜底:不信任上游声明的长度,
// 逐块累计超阈值立即中断,避免 response.bytes() 一次性分配任意大小响应撑爆内存。
if response
.content_length()
.is_some_and(|length| length > MAX_RESULT_RESPONSE_BYTES as u64)
{
return Err(result_response_too_large_error());
}
let mut bytes = Vec::new();
while let Some(chunk) = response.chunk().await.map_err(|error| {
MattingError::upstream_transport_error(
describe_result_download_body_error(&error),
error.is_timeout(),
)
})? {
if bytes.len().saturating_add(chunk.len()) > MAX_RESULT_RESPONSE_BYTES {
return Err(result_response_too_large_error());
}
bytes.extend_from_slice(chunk.as_ref());
}
Ok(bytes)
}
/// 按 VIAPI 新版官方 SDK 的 AdvanceRequest 协议,把本地图片字节上传到授权的
/// 上海地域临时 OSS,返回可直接作为 ImageURL 的公网地址。
pub async fn upload_temp_image(
&self,
bytes: Vec<u8>,
file_name: &str,
content_type: &str,
) -> Result<String, MattingError> {
if bytes.is_empty() {
return Err(MattingError::InvalidRequest("上传内容不能为空".to_string()));
}
let file_name = file_name.trim().trim_matches('/');
if file_name.is_empty() {
return Err(MattingError::InvalidRequest(
"file_name 不能为空".to_string(),
));
}
let authorized = self.authorize_file_upload().await?;
let upload_host = format!("{}.{}", authorized.bucket, authorized.endpoint);
let target_url = format!("https://{upload_host}");
let file_part = reqwest::multipart::Part::bytes(bytes)
.file_name(file_name.to_string())
.mime_str(content_type)
.map_err(|error| {
MattingError::InvalidRequest(format!("上传内容类型不合法:{error}"))
})?;
let form = reqwest::multipart::Form::new()
.text("OSSAccessKeyId", authorized.access_key_id.clone())
.text("policy", authorized.encoded_policy.clone())
.text("Signature", authorized.signature.clone())
.text("key", authorized.object_key.clone())
.text("success_action_status", "201")
.part("file", file_part);
let response = self
.client
.post(&target_url)
.multipart(form)
.send()
.await
.map_err(|error| {
MattingError::upstream_transport_error(
format!("上传 AuthorizeFileUpload 临时对象请求失败:{error}"),
error.is_timeout(),
)
})?;
let status = response.status();
if !status.is_success() {
let body = response.text().await.unwrap_or_default();
return Err(MattingError::upstream_http_error(
format!(
"上传 AuthorizeFileUpload 临时对象失败(HTTP {}):{}",
status.as_u16(),
sanitize_policy_upload_error_body(&body, &authorized)
),
status.as_u16(),
));
}
Ok(format!(
"https://{upload_host}/{}",
encode_oss_object_key(&authorized.object_key)
))
}
async fn authorize_file_upload(&self) -> Result<AuthorizedFileUpload, MattingError> {
let canonical_query = format!("Product={IMAGESEG_PRODUCT}");
let headers = self.build_acs3_headers_for_request(
OPEN_PLATFORM_ENDPOINT,
AUTHORIZE_FILE_UPLOAD_ACTION,
AUTHORIZE_FILE_UPLOAD_VERSION,
"GET",
&canonical_query,
"",
)?;
let response = self
.client
.get(format!(
"https://{OPEN_PLATFORM_ENDPOINT}/?{canonical_query}"
))
.headers(headers)
.send()
.await
.map_err(|error| {
MattingError::upstream_transport_error(
format!("AuthorizeFileUpload 请求失败:{error}"),
error.is_timeout(),
)
})?;
let http_status = response.status();
let body_text = response.text().await.map_err(|error| {
MattingError::upstream_transport_error(
format!("AuthorizeFileUpload 响应读取失败:{error}"),
error.is_timeout(),
)
})?;
parse_authorized_file_upload_response(&body_text, http_status)
}
fn build_signature_headers(
&self,
action: &str,
payload: &str,
) -> Result<reqwest::header::HeaderMap, MattingError> {
self.build_acs3_headers(&self.config.endpoint, action, IMAGESEG_API_VERSION, payload)
}
fn build_acs3_headers(
&self,
endpoint: &str,
action: &str,
version: &str,
payload: &str,
) -> Result<reqwest::header::HeaderMap, MattingError> {
self.build_acs3_headers_for_request(endpoint, action, version, "POST", "", payload)
}
fn build_acs3_headers_for_request(
&self,
endpoint: &str,
action: &str,
version: &str,
method: &str,
canonical_query: &str,
payload: &str,
) -> Result<reqwest::header::HeaderMap, MattingError> {
let date = current_aliyun_timestamp();
let nonce = uuid::Uuid::new_v4().simple().to_string();
let payload_hash = sha256_hex(payload.as_bytes());
let canonical_headers = format!(
"host:{}\nx-acs-action:{}\nx-acs-content-sha256:{}\nx-acs-date:{}\nx-acs-signature-nonce:{}\nx-acs-version:{}\n",
endpoint, action, payload_hash, date, nonce, version
);
let signed_headers =
"host;x-acs-action;x-acs-content-sha256;x-acs-date;x-acs-signature-nonce;x-acs-version";
let canonical_request = format!(
"{method}\n/\n{canonical_query}\n{canonical_headers}\n{signed_headers}\n{payload_hash}"
);
let string_to_sign = format!(
"ACS3-HMAC-SHA256\n{}",
sha256_hex(canonical_request.as_bytes())
);
let signature = hmac_sha256_hex(
self.config.access_key_secret.as_bytes(),
string_to_sign.as_bytes(),
)?;
let authorization = format!(
"ACS3-HMAC-SHA256 Credential={},SignedHeaders={signed_headers},Signature={signature}",
self.config.access_key_id
);
let mut headers = reqwest::header::HeaderMap::new();
insert_header(&mut headers, "x-acs-action", action)?;
insert_header(&mut headers, "x-acs-version", version)?;
insert_header(&mut headers, "x-acs-date", &date)?;
insert_header(&mut headers, "x-acs-signature-nonce", &nonce)?;
insert_header(&mut headers, "x-acs-content-sha256", &payload_hash)?;
insert_header(&mut headers, "authorization", &authorization)?;
Ok(headers)
}
}
/// 把源图归一化成阿里云 SegmentCommonImage 接受的输入:统一 8 位 RGBA PNG、最长边 ≤1999、体积 ≤3MB。
/// 返回编码后的 PNG 字节及其实际尺寸;尺寸若与源图不同说明发生了缩放,调用方需 alpha 上采样回原尺寸。
fn normalize_matting_input_png(
source: &image::DynamicImage,
) -> Result<(Vec<u8>, (u32, u32)), MattingError> {
normalize_matting_input_png_within(source, MAX_INPUT_BYTES)
}
fn validate_source_dimensions(source_dims: (u32, u32)) -> Result<(), MattingError> {
if source_dims.0 <= MIN_INPUT_EDGE || source_dims.1 <= MIN_INPUT_EDGE {
return Err(MattingError::InvalidRequest(format!(
"待抠图图片尺寸 {}x{} 过小,阿里云通用抠图要求每条边大于 {MIN_INPUT_EDGE} 像素",
source_dims.0, source_dims.1
)));
}
Ok(())
}
fn compose_source_with_result_alpha(
mut source_rgba: image::RgbaImage,
result_image: image::RgbaImage,
) -> Result<Vec<u8>, MattingError> {
let (source_width, source_height) = source_rgba.dimensions();
let alpha_mask = image::DynamicImage::ImageRgba8(result_image).resize_exact(
source_width,
source_height,
image::imageops::FilterType::Triangle,
);
let alpha_mask = alpha_mask.to_rgba8();
for (target, mask) in source_rgba.pixels_mut().zip(alpha_mask.pixels()) {
target.0[3] = target.0[3].min(mask.0[3]);
}
encode_rgba_png(&source_rgba)
}
/// `normalize_matting_input_png` 的可注入体积上限版本,便于单测用小阈值触发降尺寸循环。
fn normalize_matting_input_png_within(
source: &image::DynamicImage,
max_bytes: usize,
) -> Result<(Vec<u8>, (u32, u32)), MattingError> {
let mut max_edge = MAX_INPUT_EDGE;
loop {
// to_rgba8 统一成 8 位/通道 RGBA(32 位 PNG),规避阿里云不支持的 8/16/64 位 PNG 与非 PNG 格式。
let rgba = if source.width() > max_edge || source.height() > max_edge {
source
.resize(max_edge, max_edge, image::imageops::FilterType::CatmullRom)
.to_rgba8()
} else {
source.to_rgba8()
};
let encoded = encode_rgba_png(&rgba)?;
let (width, height) = rgba.dimensions();
if encoded.len() <= max_bytes {
return Ok((encoded, (width, height)));
}
// 仍超体积:等比缩到约 0.8 再试;一旦再缩会让边跌破阿里云 >32 下限就尽力返回当前
// (体积仍超限的图交给上游 4xx → 本地兜底,现实内容不会走到这一步)。
let longest = width.max(height);
let next = (longest as f32 * 0.8) as u32;
if next <= MIN_INPUT_EDGE || next >= longest {
return Ok((encoded, (width, height)));
}
max_edge = next;
}
}
fn encode_rgba_png(image: &image::RgbaImage) -> Result<Vec<u8>, MattingError> {
use image::ImageEncoder as _;
let mut encoded = Vec::new();
image::codecs::png::PngEncoder::new(&mut encoded)
.write_image(
image.as_raw(),
image.width(),
image.height(),
image::ExtendedColorType::Rgba8,
)
.map_err(|error| {
MattingError::upstream_response_error(format!("编码抠图结果 PNG 失败:{error}"), None)
})?;
Ok(encoded)
}
fn result_response_too_large_error() -> MattingError {
MattingError::upstream_response_error(
format!("抠图结果响应过大,超过 {MAX_RESULT_RESPONSE_BYTES} 字节上限"),
None,
)
}
fn source_response_too_large_error() -> MattingError {
MattingError::upstream_response_error(
format!("待抠图源图响应过大,超过 {MAX_RESULT_RESPONSE_BYTES} 字节上限"),
None,
)
}
/// 解码待抠图源图时套上尺寸 / 分配上限。源图直接来自请求体 / 生成产物,缺省
/// `image::Limits` 无尺寸上限、仅 512MiB alloc 兜底,12MB 请求体即可构造巨幅压缩图
/// (解压炸弹)在缩图前撑爆内存;逐帧动画并发调用时风险叠加。
fn decode_source_image(bytes: &[u8]) -> Result<image::DynamicImage, MattingError> {
decode_image_within_limits(bytes)
.map_err(|error| MattingError::InvalidRequest(format!("解析待抠图图片失败:{error}")))
}
/// 解码抠图结果时套上尺寸 / 分配上限,防止上游用巨幅尺寸声明在解码阶段撑爆内存。
fn decode_result_image(bytes: &[u8]) -> Result<image::DynamicImage, MattingError> {
decode_image_within_limits(bytes).map_err(|error| {
MattingError::upstream_response_error(format!("解析抠图结果图片失败:{error}"), None)
})
}
/// 统一的有界解码:尺寸上限 `MAX_DECODE_IMAGE_DIMENSION`、分配上限
/// `MAX_RESULT_RESPONSE_BYTES * 4`,覆盖源图与抠图结果两条解码路径。
fn decode_image_within_limits(bytes: &[u8]) -> Result<image::DynamicImage, String> {
use std::io::Cursor;
let mut reader = image::ImageReader::new(Cursor::new(bytes))
.with_guessed_format()
.map_err(|error| format!("识别图片格式失败:{error}"))?;
let mut limits = image::Limits::default();
limits.max_image_width = Some(MAX_DECODE_IMAGE_DIMENSION);
limits.max_image_height = Some(MAX_DECODE_IMAGE_DIMENSION);
limits.max_alloc = Some(MAX_RESULT_RESPONSE_BYTES as u64 * 4);
reader.limits(limits);
reader.decode().map_err(|error| error.to_string())
}
#[derive(Debug)]
struct AuthorizedFileUpload {
bucket: String,
endpoint: String,
access_key_id: String,
encoded_policy: String,
signature: String,
object_key: String,
}
fn parse_authorized_file_upload_response(
body_text: &str,
http_status: StatusCode,
) -> Result<AuthorizedFileUpload, MattingError> {
// 成功响应包含 Policy 与 Signature;解析错误中不能回显原始正文。
let body: serde_json::Value = serde_json::from_str(body_text).map_err(|error| {
MattingError::upstream_response_error(
format!("AuthorizeFileUpload 响应不是合法 JSON:{error}"),
Some(http_status.as_u16()),
)
})?;
if http_status != StatusCode::OK {
return Err(MattingError::upstream_http_error(
format!(
"AuthorizeFileUpload 返回失败(HTTP {},Code={}):{}",
http_status.as_u16(),
body.get("Code")
.and_then(|value| value.as_str())
.unwrap_or("unknown"),
body.get("Message")
.and_then(|value| value.as_str())
.unwrap_or("unknown")
),
http_status.as_u16(),
));
}
let read_field = |name: &str| -> Result<String, MattingError> {
body.get(name)
.and_then(|value| value.as_str())
.map(str::trim)
.filter(|value| !value.is_empty())
.map(str::to_string)
.ok_or_else(|| {
MattingError::upstream_response_error(
format!("AuthorizeFileUpload 响应缺少 {name}"),
Some(http_status.as_u16()),
)
})
};
let bucket = read_field("Bucket")?.trim_matches('/').to_string();
let endpoint = read_field("Endpoint")?
.trim_start_matches("https://")
.trim_start_matches("http://")
.trim_matches('/')
.to_string();
if bucket.is_empty() || endpoint.is_empty() {
return Err(MattingError::upstream_response_error(
"AuthorizeFileUpload 响应中的 Bucket/Endpoint 不合法".to_string(),
Some(http_status.as_u16()),
));
}
Ok(AuthorizedFileUpload {
bucket,
endpoint,
access_key_id: read_field("AccessKeyId")?,
encoded_policy: read_field("EncodedPolicy")?,
signature: read_field("Signature")?,
object_key: read_field("ObjectKey")?,
})
}
fn describe_result_download_transport_error(error: &reqwest::Error) -> String {
sanitize_result_download_error_message(format!(
"下载抠图结果失败(transport={}, timeout={}, connect={})",
classify_reqwest_error(error),
error.is_timeout(),
error.is_connect()
))
}
fn describe_result_download_body_error(error: &reqwest::Error) -> String {
sanitize_result_download_error_message(format!(
"读取抠图结果字节失败(transport={}, timeout={}, connect={})",
classify_reqwest_error(error),
error.is_timeout(),
error.is_connect()
))
}
fn classify_reqwest_error(error: &reqwest::Error) -> &'static str {
if error.is_timeout() {
"timeout"
} else if error.is_connect() {
"connect"
} else if error.is_body() || error.is_decode() {
"body"
} else if error.is_request() {
"request"
} else {
"unknown"
}
}
fn sanitize_result_download_error_message(message: String) -> String {
if message.contains("OSSAccessKeyId=")
|| message.contains("Signature=")
|| message.contains("Expires=")
{
return "下载抠图结果失败(临时签名 URL 已脱敏)".to_string();
}
message
}
/// 脱敏 Policy POST 上传错误体:OSS 签名类错误可能回显 StringToSign、Policy 或签名串。
/// 剥掉已知签名元素,并替换本次授权材料,保留 `<Code>` 等可诊断信息。
fn sanitize_policy_upload_error_body(body: &str, authorized: &AuthorizedFileUpload) -> String {
let mut sanitized = body.to_string();
for tag in [
"StringToSign",
"StringToSignBytes",
"SignatureProvided",
"Policy",
] {
sanitized = redact_xml_element(&sanitized, tag);
}
for secret in [
authorized.access_key_id.as_str(),
authorized.encoded_policy.as_str(),
authorized.signature.as_str(),
] {
if !secret.is_empty() {
sanitized = sanitized.replace(secret, "***");
}
}
truncate_for_log(&sanitized)
}
fn encode_oss_object_key(object_key: &str) -> String {
object_key
.split('/')
.map(|segment| urlencoding_encode(segment).replace('+', "%20"))
.collect::<Vec<_>>()
.join("/")
}
/// 把 `<tag>…</tag>` 的内容替换成 `[redacted]`(tag 内容可跨行);只有开标签无闭标签时丢弃其后全部内容。
fn redact_xml_element(text: &str, tag: &str) -> String {
let open = format!("<{tag}>");
let close = format!("</{tag}>");
let mut result = String::with_capacity(text.len());
let mut rest = text;
while let Some(start) = rest.find(&open) {
result.push_str(&rest[..start]);
result.push_str(&open);
result.push_str("[redacted]");
let after_open = start + open.len();
match rest[after_open..].find(&close) {
Some(rel_end) => {
result.push_str(&close);
rest = &rest[after_open + rel_end + close.len()..];
}
None => {
rest = "";
break;
}
}
}
result.push_str(rest);
result
}
fn insert_header(
headers: &mut reqwest::header::HeaderMap,
name: &'static str,
value: &str,
) -> Result<(), MattingError> {
let value = reqwest::header::HeaderValue::from_str(value)
.map_err(|error| MattingError::Sign(format!("构造请求头 {name} 失败:{error}")))?;
headers.insert(name, value);
Ok(())
}
fn current_aliyun_timestamp() -> String {
// 阿里云 OpenAPI ACS3 签名头 x-acs-date 要求不带小数秒的 UTC ISO 8601 格式
// (yyyy-MM-dd'T'HH:mm:ss'Z'),Rfc3339 默认会保留纳秒,网关会判定时间格式非法。
let now = OffsetDateTime::now_utc();
format!(
"{:04}-{:02}-{:02}T{:02}:{:02}:{:02}Z",
now.year(),
u8::from(now.month()),
now.day(),
now.hour(),
now.minute(),
now.second()
)
}
fn canonicalize_aliyun_form_params(params: &BTreeMap<String, String>) -> String {
params
.iter()
.map(|(key, value)| format!("{}={}", urlencoding_encode(key), urlencoding_encode(value)))
.collect::<Vec<_>>()
.join("&")
}
fn urlencoding_encode(value: &str) -> String {
serde_urlencoded::to_string([("k", value)])
.map(|encoded| encoded.trim_start_matches("k=").to_string())
.unwrap_or_else(|_| value.to_string())
}
fn build_aliyun_form_body(params: &BTreeMap<String, String>) -> String {
serde_urlencoded::to_string(params).unwrap_or_else(|_| canonicalize_aliyun_form_params(params))
}
fn hmac_sha256_hex(key: &[u8], content: &[u8]) -> Result<String, MattingError> {
let mut signer = HmacSha256::new_from_slice(key)
.map_err(|error| MattingError::Sign(format!("初始化抠图签名器失败:{error}")))?;
signer.update(content);
Ok(hex::encode(signer.finalize().into_bytes()))
}
fn sha256_hex(content: &[u8]) -> String {
let mut hasher = Sha256::new();
hasher.update(content);
hex::encode(hasher.finalize())
}
fn truncate_for_log(text: &str) -> String {
const LIMIT: usize = 512;
if text.len() <= LIMIT {
text.to_string()
} else {
let mut end = LIMIT;
while !text.is_char_boundary(end) {
end -= 1;
}
format!("{}...(截断)", &text[..end])
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn local_preflight_variants_report_no_external_call() {
for error in [
MattingError::InvalidConfig("endpoint 为空".to_string()),
MattingError::InvalidRequest("待抠图图片尺寸过小".to_string()),
MattingError::Sign("初始化签名器失败".to_string()),
] {
assert!(!error.external_call_attempted(), "{}", error.message());
assert!(!error.is_timeout());
assert!(!error.is_transport());
assert_eq!(error.upstream_status(), None);
}
}
#[test]
fn local_processing_after_external_call_is_auditable_with_stage() {
let error = MattingError::InvalidRequest("解析待抠图图片失败:invalid png".to_string())
.with_failure_stage("source_decode");
assert!(error.external_call_attempted());
assert_eq!(error.failure_stage(), "source_decode");
assert!(!error.is_timeout());
assert!(!error.is_transport());
assert_eq!(error.upstream_status(), None);
}
#[test]
fn upstream_transport_error_classifies_as_transport() {
let error = MattingError::upstream_transport_error(
"通用抠图请求失败:dns error".to_string(),
false,
);
assert!(error.external_call_attempted());
assert!(error.is_transport());
assert!(!error.is_timeout());
assert_eq!(error.upstream_status(), None);
}
#[test]
fn upstream_transport_error_preserves_timeout_flag() {
let error =
MattingError::upstream_transport_error("通用抠图请求失败:timed out".to_string(), true);
assert!(error.external_call_attempted());
assert!(error.is_transport());
assert!(error.is_timeout());
}
#[test]
fn upstream_http_error_carries_status_and_is_not_transport() {
let error = MattingError::upstream_http_error(
"通用抠图接口返回失败(HTTP 429,Code=Throttled)".to_string(),
429,
);
assert!(error.external_call_attempted());
assert!(!error.is_transport());
assert!(!error.is_timeout());
assert_eq!(error.upstream_status(), Some(429));
}
#[test]
fn upstream_response_error_is_external_but_not_transport() {
let error = MattingError::upstream_response_error("抠图结果尺寸不一致".to_string(), None);
assert!(error.external_call_attempted());
assert!(!error.is_transport());
assert!(!error.is_timeout());
assert_eq!(error.upstream_status(), None);
}
#[test]
fn form_body_is_sorted_and_url_encoded() {
let mut params = BTreeMap::new();
params.insert("ReturnForm".to_string(), "crop".to_string());
params.insert(
"ImageURL".to_string(),
"https://example.com/a.png?x=1&y=2".to_string(),
);
params.insert("Action".to_string(), "SegmentCommonImage".to_string());
let body = build_aliyun_form_body(&params);
assert_eq!(
body,
"Action=SegmentCommonImage&ImageURL=https%3A%2F%2Fexample.com%2Fa.png%3Fx%3D1%26y%3D2&ReturnForm=crop"
);
}
#[test]
fn signature_headers_use_acs3_sha256() {
let config = MattingConfig::new(
DEFAULT_IMAGESEG_ENDPOINT.to_string(),
"test-key-id".to_string(),
"test-key-secret".to_string(),
)
.expect("config should build");
let client = MattingClient::new(config).expect("client should build");
let headers = client
.build_signature_headers("SegmentCommonImage", "ImageURL=x")
.expect("headers should build");
let authorization = headers
.get("authorization")
.expect("authorization header should exist")
.to_str()
.expect("authorization header should be ascii");
assert!(authorization.starts_with("ACS3-HMAC-SHA256 Credential=test-key-id,"));
assert!(authorization.contains("SignedHeaders=host;x-acs-action;x-acs-content-sha256;"));
assert_eq!(
headers
.get("x-acs-version")
.and_then(|value| value.to_str().ok()),
Some("2019-12-30")
);
let date = headers
.get("x-acs-date")
.and_then(|value| value.to_str().ok())
.expect("x-acs-date should exist");
assert!(!date.contains('.'), "x-acs-date 不能带小数秒:{date}");
}
#[test]
fn authorize_file_upload_response_parses_policy_fields() {
let response = r#"{
"RequestId":"request-id",
"Bucket":"viapi-customer-pop",
"Endpoint":"oss-cn-shanghai.aliyuncs.com",
"AccessKeyId":"temporary-key-id",
"EncodedPolicy":"encoded-policy",
"Signature":"policy-signature",
"ObjectKey":"imageseg/2026/07/input image.png"
}"#;
let authorized = parse_authorized_file_upload_response(response, StatusCode::OK)
.expect("AuthorizeFileUpload response should parse");
assert_eq!(authorized.bucket, "viapi-customer-pop");
assert_eq!(authorized.endpoint, "oss-cn-shanghai.aliyuncs.com");
assert_eq!(authorized.access_key_id, "temporary-key-id");
assert_eq!(authorized.encoded_policy, "encoded-policy");
assert_eq!(authorized.signature, "policy-signature");
assert_eq!(authorized.object_key, "imageseg/2026/07/input image.png");
assert_eq!(
encode_oss_object_key(&authorized.object_key),
"imageseg/2026/07/input%20image.png"
);
}
#[test]
fn authorize_file_upload_response_does_not_echo_sensitive_body_on_parse_error() {
let sensitive_body = "not-json encoded-policy policy-signature temporary-key-id";
let error = parse_authorized_file_upload_response(sensitive_body, StatusCode::OK)
.expect_err("invalid response should fail");
assert!(error.message().contains("响应不是合法 JSON"));
assert!(!error.message().contains("encoded-policy"));
assert!(!error.message().contains("policy-signature"));
assert!(!error.message().contains("temporary-key-id"));
}
#[test]
fn authorize_file_upload_headers_sign_get_query() {
let config = MattingConfig::new(
DEFAULT_IMAGESEG_ENDPOINT.to_string(),
"test-key-id".to_string(),
"test-key-secret".to_string(),
)
.expect("config should build");
let client = MattingClient::new(config).expect("client should build");
let headers = client
.build_acs3_headers_for_request(
OPEN_PLATFORM_ENDPOINT,
AUTHORIZE_FILE_UPLOAD_ACTION,
AUTHORIZE_FILE_UPLOAD_VERSION,
"GET",
"Product=imageseg",
"",
)
.expect("headers should build");
assert_eq!(
headers
.get("x-acs-action")
.and_then(|value| value.to_str().ok()),
Some(AUTHORIZE_FILE_UPLOAD_ACTION)
);
assert_eq!(
headers
.get("x-acs-version")
.and_then(|value| value.to_str().ok()),
Some(AUTHORIZE_FILE_UPLOAD_VERSION)
);
assert_eq!(
headers
.get("x-acs-content-sha256")
.and_then(|value| value.to_str().ok()),
Some(sha256_hex(b"").as_str())
);
}
#[test]
fn result_download_error_message_redacts_signed_oss_url() {
let message = sanitize_result_download_error_message(
"下载抠图结果失败:error sending request for url (https://viapi-customer-temp.oss-cn-shanghai.aliyuncs.com/a.png?OSSAccessKeyId=ak&Expires=123&Signature=secret)".to_string(),
);
assert!(!message.contains("OSSAccessKeyId"));
assert!(!message.contains("Signature"));
assert!(!message.contains("Expires=123"));
assert_eq!(message, "下载抠图结果失败(临时签名 URL 已脱敏)");
}
#[test]
fn result_download_error_message_keeps_safe_text() {
let message = sanitize_result_download_error_message(
"下载抠图结果失败(transport=connect, timeout=false, connect=true)".to_string(),
);
assert_eq!(
message,
"下载抠图结果失败(transport=connect, timeout=false, connect=true)"
);
}
#[test]
fn policy_upload_error_body_redacts_authorization_material() {
let authorized = AuthorizedFileUpload {
bucket: "viapi-customer-pop".to_string(),
endpoint: "oss-cn-shanghai.aliyuncs.com".to_string(),
access_key_id: "temporary-key-id".to_string(),
encoded_policy: "encoded-policy".to_string(),
signature: "policy-signature".to_string(),
object_key: "imageseg/input.png".to_string(),
};
let body = format!(
"<Error><Code>SignatureDoesNotMatch</Code>\
<Message>mismatch for {} and {}</Message>\
<Policy>{}</Policy>\
<SignatureProvided>{}</SignatureProvided>\
<StringToSignBytes>50 55 54 0a</StringToSignBytes>\
<StringToSign>POST\npolicy={}</StringToSign></Error>",
authorized.access_key_id,
authorized.signature,
authorized.encoded_policy,
authorized.signature,
authorized.encoded_policy
);
let sanitized = sanitize_policy_upload_error_body(&body, &authorized);
assert!(!sanitized.contains(&authorized.access_key_id));
assert!(!sanitized.contains(&authorized.encoded_policy));
assert!(!sanitized.contains(&authorized.signature));
assert!(
!sanitized.contains("50 55 54 0a"),
"StringToSignBytes 不能残留"
);
assert!(
sanitized.contains("SignatureDoesNotMatch"),
"OSS Code 应保留供诊断"
);
assert!(
sanitized.contains("[redacted]"),
"签名材料元素应被脱敏为 [redacted]"
);
}
fn noise_image(width: u32, height: u32) -> image::DynamicImage {
// 不易压缩的伪随机噪声,保证 PNG 体积随像素数量增长,可触发降尺寸循环。
let mut img = image::RgbaImage::new(width, height);
for (x, y, px) in img.enumerate_pixels_mut() {
let v = ((x * 71 + y * 131 + x * y) % 256) as u8;
*px = image::Rgba([v, v.wrapping_mul(3), v.wrapping_add(97), 255]);
}
image::DynamicImage::ImageRgba8(img)
}
#[test]
fn normalize_keeps_small_image_and_outputs_png() {
let source = image::DynamicImage::ImageRgba8(image::RgbaImage::from_pixel(
200,
150,
image::Rgba([10, 20, 30, 255]),
));
let (bytes, dims) = normalize_matting_input_png(&source).expect("normalize should succeed");
assert_eq!(dims, (200, 150), "小图不缩放,尺寸原样");
assert_eq!(
image::guess_format(&bytes).expect("guess format"),
image::ImageFormat::Png,
"输出必须是 PNG"
);
let decoded = image::load_from_memory(&bytes).expect("output should decode");
assert_eq!((decoded.width(), decoded.height()), (200, 150));
}
#[test]
fn normalize_caps_oversized_edge_to_1999() {
let source = image::DynamicImage::ImageRgba8(image::RgbaImage::from_pixel(
2400,
1200,
image::Rgba([0, 0, 0, 255]),
));
let (_bytes, (width, height)) =
normalize_matting_input_png(&source).expect("normalize should succeed");
assert!(
width <= MAX_INPUT_EDGE && height <= MAX_INPUT_EDGE,
"两边都 ≤1999,实得 {width}x{height}"
);
assert_eq!(width.max(height), MAX_INPUT_EDGE, "最长边压到 1999");
}
#[test]
fn normalize_downscales_until_under_byte_limit() {
let source = noise_image(300, 300);
let full = normalize_matting_input_png(&source)
.expect("full-size normalize should succeed")
.0;
// 阈值取全尺寸 PNG 的一半,必须触发降尺寸。
let limit = full.len() / 2;
let (bytes, (width, height)) = normalize_matting_input_png_within(&source, limit)
.expect("limited normalize should succeed");
assert!(
width < 300 && height < 300,
"应从 300x300 降尺寸,实得 {width}x{height}"
);
assert!(
bytes.len() <= limit || width.min(height) <= MIN_INPUT_EDGE + 1,
"编码 {} 字节应落在 {limit} 内(或已触最小边下限)",
bytes.len()
);
assert_eq!(
image::guess_format(&bytes).expect("guess format"),
image::ImageFormat::Png,
"降尺寸后仍是 PNG"
);
}
#[test]
fn url_input_upload_scope_ends_before_aliyun_inference() {
let source = include_str!("lib.rs");
let start = source
.find("pub async fn segment_image_url_to_transparent_png")
.expect("URL input method should exist");
let tail = &source[start..];
let end = tail
.find("async fn segment_uploaded_input_to_rgba")
.expect("uploaded input helper should follow URL method");
let body = &tail[..end];
let download = body
.find("let source_bytes = self\n .download_source_image(source_url)")
.expect("source should download lazily");
let upload = body
.find(".upload_temp_image(upload_bytes, file_name, \"image/png\")")
.expect("source should upload to temporary OSS");
let scope_end = body
.find("let result_image = self")
.expect("Aliyun inference should start after upload scope");
assert!(download < upload && upload < scope_end);
assert!(body.contains("let (source_dims, upload_dims, image_url) = {"));
assert_eq!(body.matches("download_source_image(source_url)").count(), 2);
}
#[test]
fn compose_source_with_result_alpha_preserves_rgb() {
let source = image::RgbaImage::from_pixel(2, 2, image::Rgba([10, 20, 30, 180]));
let mask = image::RgbaImage::from_pixel(1, 1, image::Rgba([200, 210, 220, 90]));
let encoded = compose_source_with_result_alpha(source, mask)
.expect("alpha composition should encode");
let output = image::load_from_memory(&encoded)
.expect("output should decode")
.to_rgba8();
assert_eq!(output.dimensions(), (2, 2));
assert!(output.pixels().all(|pixel| pixel.0 == [10, 20, 30, 90]));
}
}