0c04bbbea3
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>
1535 lines
57 KiB
Rust
1535 lines
57 KiB
Rust
//! 阿里云 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(¶ms);
|
||
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]));
|
||
}
|
||
}
|