Files
ystp/src/worker/mod.rs

2361 lines
75 KiB
Rust

use crate::error::{AppError, ErrorCode};
use crate::services::billing;
use crate::services::compress;
use crate::services::metrics;
use crate::services::object_lifecycle;
use crate::services::quota;
use crate::services::storage;
use crate::state::AppState;
use redis::streams::{StreamClaimReply, StreamId, StreamPendingCountReply, StreamReadOptions};
use redis::AsyncCommands;
use sqlx::FromRow;
use std::net::IpAddr;
use std::sync::Arc;
use std::time::Duration;
use tokio::sync::Semaphore;
use tokio::task::JoinSet;
use uuid::Uuid;
const STREAM_KEY: &str = metrics::QUEUE_STREAM_KEY;
const GROUP_NAME: &str = metrics::QUEUE_GROUP_NAME;
const DEAD_STREAM_KEY: &str = metrics::DEAD_STREAM_KEY;
const MAX_DELIVERIES: usize = 3;
const STALE_MESSAGE_IDLE_MS: usize = 5 * 60 * 1000;
const MESSAGE_HEARTBEAT_SECONDS: u64 = 30;
const PROCESSING_LEASE_SECONDS: i64 = 120;
const LEASE_BUSY_RETRY_SECONDS: u64 = 5;
const QUEUE_BLOCK_MS: usize = 1_000;
const MAINTENANCE_INTERVAL_SECONDS: u64 = 300;
const MAINTENANCE_BATCH_SIZE: i64 = 1_000;
const MAX_MAINTENANCE_BATCHES: usize = 20;
const INITIAL_RETRY_SECONDS: u64 = 2;
const MAX_RETRY_SECONDS: u64 = 30;
pub async fn run(state: AppState) -> Result<(), AppError> {
tracing::info!("Worker started");
crate::services::bootstrap::ensure_schema(&state).await?;
let worker_id = Uuid::new_v4();
let consumer = format!("worker_{worker_id}");
ensure_group(&state).await?;
tokio::spawn(maintenance_loop(state.clone()));
tokio::spawn(crate::services::task_queue::dispatch_loop(state.clone()));
tokio::spawn(crate::services::object_lifecycle::maintenance_loop(
state.clone(),
));
let task_concurrency = state.config.worker_task_concurrency.max(1) as usize;
let mut inflight = JoinSet::new();
let mut poll_backoff = Duration::from_secs(INITIAL_RETRY_SECONDS);
tracing::info!(task_concurrency, "Worker task scheduler ready");
loop {
while let Some(result) = inflight.try_join_next() {
log_message_task_result(result);
}
let available = task_concurrency.saturating_sub(inflight.len());
if available > 0 {
match read_messages(&state, &consumer, available).await {
Ok(messages) => {
poll_backoff = Duration::from_secs(INITIAL_RETRY_SECONDS);
for message in messages {
let state = state.clone();
let consumer = consumer.clone();
// A process can receive duplicate messages for the same task. Keep their
// database leases distinct even though they share one Redis consumer.
let lease_owner = Uuid::new_v4();
inflight.spawn(async move {
process_message_with_retries(state, lease_owner, consumer, message)
.await
});
}
}
Err(err) => {
tracing::error!(
error = ?err,
retry_in_seconds = poll_backoff.as_secs(),
"worker poll error"
);
tokio::time::sleep(poll_backoff).await;
poll_backoff = next_backoff(poll_backoff);
continue;
}
}
}
if inflight.len() >= task_concurrency {
if let Some(result) = inflight.join_next().await {
log_message_task_result(result);
}
} else {
tokio::select! {
result = inflight.join_next(), if !inflight.is_empty() => {
if let Some(result) = result {
log_message_task_result(result);
}
}
_ = tokio::time::sleep(Duration::from_millis(25)) => {}
}
}
}
}
async fn maintenance_loop(state: AppState) {
let mut interval = tokio::time::interval_at(
tokio::time::Instant::now() + Duration::from_secs(MAINTENANCE_INTERVAL_SECONDS),
Duration::from_secs(MAINTENANCE_INTERVAL_SECONDS),
);
interval.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Skip);
loop {
interval.tick().await;
if let Err(err) = maintenance(&state).await {
tracing::error!(error = ?err, "maintenance failed");
}
}
}
fn log_message_task_result(result: Result<Result<(), AppError>, tokio::task::JoinError>) {
match result {
Ok(Ok(())) => {}
Ok(Err(err)) => tracing::error!(error = ?err, "worker message task stopped"),
Err(err) => tracing::error!(error = ?err, "worker message task panicked"),
}
}
async fn ensure_group(state: &AppState) -> Result<(), AppError> {
let mut conn = state.redis.clone();
let res: Result<redis::Value, redis::RedisError> = redis::cmd("XGROUP")
.arg("CREATE")
.arg(STREAM_KEY)
.arg(GROUP_NAME)
.arg("0")
.arg("MKSTREAM")
.query_async(&mut conn)
.await;
match res {
Ok(_) => Ok(()),
Err(err) => {
let msg = err.to_string();
if msg.contains("BUSYGROUP") {
return Ok(());
}
Err(AppError::new(ErrorCode::Internal, "初始化队列失败").with_source(err))
}
}
}
async fn read_messages(
state: &AppState,
consumer: &str,
count: usize,
) -> Result<Vec<StreamId>, AppError> {
let mut conn = state.redis.clone();
if let Some(msg) = claim_stale_message(&mut conn, consumer).await? {
return Ok(vec![msg]);
}
let opts = StreamReadOptions::default()
.group(GROUP_NAME, consumer)
.count(count)
.block(QUEUE_BLOCK_MS);
let reply: redis::streams::StreamReadReply = conn
.xread_options(&[STREAM_KEY], &[">"], &opts)
.await
.map_err(|err| AppError::new(ErrorCode::Internal, "读取队列失败").with_source(err))?;
Ok(reply
.keys
.into_iter()
.flat_map(|key| key.ids.into_iter())
.collect())
}
async fn process_message_with_retries(
state: AppState,
worker_id: Uuid,
consumer: String,
mut message: StreamId,
) -> Result<(), AppError> {
loop {
match handle_message_with_heartbeat(&state, worker_id, &consumer, &message).await {
Ok(MessageOutcome::Done | MessageOutcome::OwnershipLost) => return Ok(()),
Ok(MessageOutcome::LeaseBusy) => {
tokio::time::sleep(Duration::from_secs(LEASE_BUSY_RETRY_SECONDS)).await;
continue;
}
Err(err) => {
let mut conn = state.redis.clone();
let deliveries = match pending_delivery_count(&mut conn, &message.id).await {
Ok(deliveries) => deliveries,
Err(count_err) => {
tracing::warn!(
message_id = %message.id,
error = ?count_err,
"failed to read delivery count; using initial retry delay"
);
1
}
};
let delay = delivery_backoff(deliveries);
tracing::warn!(
message_id = %message.id,
deliveries,
retry_in_seconds = delay.as_secs(),
error = ?err,
"worker message handling failed; retry scheduled"
);
tokio::time::sleep(delay).await;
let mut reclaim_backoff = Duration::from_secs(INITIAL_RETRY_SECONDS);
loop {
match redeliver_message(&mut conn, &consumer, &message.id).await {
Ok(Some(redelivered)) => {
message = redelivered;
break;
}
Ok(None) => return Ok(()),
Err(reclaim_err) => {
tracing::error!(
message_id = %message.id,
retry_in_seconds = reclaim_backoff.as_secs(),
error = ?reclaim_err,
"failed to reclaim pending worker message"
);
tokio::time::sleep(reclaim_backoff).await;
reclaim_backoff = next_backoff(reclaim_backoff);
}
}
}
}
}
}
}
async fn handle_message_with_heartbeat(
state: &AppState,
worker_id: Uuid,
consumer: &str,
message: &StreamId,
) -> Result<MessageOutcome, AppError> {
let mut conn = state.redis.clone();
let handling = handle_message(state, worker_id, &mut conn, message);
tokio::pin!(handling);
let task_id = message
.get::<String>("task_id")
.and_then(|value| Uuid::parse_str(&value).ok());
let mut heartbeat = tokio::time::interval_at(
tokio::time::Instant::now() + Duration::from_secs(MESSAGE_HEARTBEAT_SECONDS),
Duration::from_secs(MESSAGE_HEARTBEAT_SECONDS),
);
heartbeat.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Skip);
loop {
tokio::select! {
result = &mut handling => return result,
_ = heartbeat.tick() => {
match touch_pending_message(state, consumer, &message.id).await {
Ok(true) => {}
Ok(false) => {
if let Some(task_id) = task_id {
release_processing_lease(state, task_id, worker_id).await;
}
return Ok(MessageOutcome::OwnershipLost);
}
Err(err) => {
tracing::warn!(
message_id = %message.id,
error = ?err,
"failed to refresh worker message heartbeat"
);
continue;
}
}
if let Some(task_id) = task_id {
if !renew_processing_lease(state, task_id, worker_id).await? {
return Ok(MessageOutcome::OwnershipLost);
}
}
}
}
}
}
async fn touch_pending_message(
state: &AppState,
consumer: &str,
message_id: &str,
) -> Result<bool, AppError> {
let mut conn = state.redis.clone();
let script = redis::Script::new(
r#"
local pending = redis.call('XPENDING', KEYS[1], ARGV[1], ARGV[2], ARGV[2], 1)
if #pending == 0 then return 0 end
if pending[1][2] ~= ARGV[3] then return -1 end
redis.call('XCLAIM', KEYS[1], ARGV[1], ARGV[3], 0, ARGV[2], 'JUSTID')
return 1
"#,
);
let owned: i64 = script
.key(STREAM_KEY)
.arg(GROUP_NAME)
.arg(message_id)
.arg(consumer)
.invoke_async(&mut conn)
.await
.map_err(|err| {
AppError::new(ErrorCode::Internal, "刷新队列任务心跳失败").with_source(err)
})?;
Ok(owned == 1)
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum MessageOutcome {
Done,
LeaseBusy,
OwnershipLost,
}
async fn handle_message(
state: &AppState,
worker_id: Uuid,
conn: &mut redis::aio::ConnectionManager,
msg: &StreamId,
) -> Result<MessageOutcome, AppError> {
let Some(task_id_str) = msg.get::<String>("task_id") else {
ack_message(conn, &msg.id).await?;
return Ok(MessageOutcome::Done);
};
let task_id = match Uuid::parse_str(&task_id_str) {
Ok(value) => value,
Err(_) => {
ack_message(conn, &msg.id).await?;
return Ok(MessageOutcome::Done);
}
};
match process_task(state, task_id, worker_id).await {
Ok(TaskProcessOutcome::Done) => {
ack_message(conn, &msg.id).await?;
Ok(MessageOutcome::Done)
}
Ok(TaskProcessOutcome::LeaseBusy) => Ok(MessageOutcome::LeaseBusy),
Err(err) => {
let deliveries = pending_delivery_count(conn, &msg.id).await?;
if should_dead_letter(deliveries) {
if !mark_task_dead_letter(state, task_id, worker_id, &err.message).await? {
return Ok(MessageOutcome::OwnershipLost);
}
write_dead_letter(conn, &msg.id, task_id, deliveries, &err).await?;
ack_message(conn, &msg.id).await?;
metrics::record_dead_letter(state);
tracing::error!(
task_id = %task_id,
deliveries,
error = %err,
"task moved to dead-letter stream"
);
return Ok(MessageOutcome::Done);
}
Err(err)
}
}
}
async fn redeliver_message(
conn: &mut redis::aio::ConnectionManager,
consumer: &str,
message_id: &str,
) -> Result<Option<StreamId>, AppError> {
let claimed: StreamClaimReply = conn
.xclaim(STREAM_KEY, GROUP_NAME, consumer, 0, &[message_id])
.await
.map_err(|err| {
AppError::new(ErrorCode::Internal, "重新认领待重试任务失败").with_source(err)
})?;
Ok(claimed.ids.into_iter().next())
}
async fn claim_stale_message(
conn: &mut redis::aio::ConnectionManager,
consumer: &str,
) -> Result<Option<StreamId>, AppError> {
let pending: StreamPendingCountReply = conn
.xpending_count(STREAM_KEY, GROUP_NAME, "-", "+", 100)
.await
.map_err(|err| {
AppError::new(ErrorCode::Internal, "查询停滞队列消息失败").with_source(err)
})?;
let Some(stale) = pending
.ids
.into_iter()
.find(|item| item.consumer != consumer && item.last_delivered_ms >= STALE_MESSAGE_IDLE_MS)
else {
return Ok(None);
};
let claimed: StreamClaimReply = conn
.xclaim(
STREAM_KEY,
GROUP_NAME,
consumer,
STALE_MESSAGE_IDLE_MS,
&[stale.id],
)
.await
.map_err(|err| {
AppError::new(ErrorCode::Internal, "认领停滞队列消息失败").with_source(err)
})?;
Ok(claimed.ids.into_iter().next())
}
async fn pending_delivery_count(
conn: &mut redis::aio::ConnectionManager,
msg_id: &str,
) -> Result<usize, AppError> {
let pending: StreamPendingCountReply = conn
.xpending_count(STREAM_KEY, GROUP_NAME, msg_id, msg_id, 1)
.await
.map_err(|err| {
AppError::new(ErrorCode::Internal, "查询队列重试次数失败").with_source(err)
})?;
Ok(pending
.ids
.first()
.map(|item| item.times_delivered)
.unwrap_or(1))
}
fn should_dead_letter(deliveries: usize) -> bool {
deliveries >= MAX_DELIVERIES
}
fn delivery_backoff(deliveries: usize) -> Duration {
let exponent = deliveries.saturating_sub(1).min(4) as u32;
Duration::from_secs(
INITIAL_RETRY_SECONDS
.saturating_mul(2_u64.saturating_pow(exponent))
.min(MAX_RETRY_SECONDS),
)
}
fn next_backoff(current: Duration) -> Duration {
current
.saturating_mul(2)
.min(Duration::from_secs(MAX_RETRY_SECONDS))
}
async fn write_dead_letter(
conn: &mut redis::aio::ConnectionManager,
message_id: &str,
task_id: Uuid,
deliveries: usize,
error: &AppError,
) -> Result<(), AppError> {
redis::cmd("XADD")
.arg(DEAD_STREAM_KEY)
.arg("MAXLEN")
.arg("~")
.arg(10_000)
.arg("*")
.arg("task_id")
.arg(task_id.to_string())
.arg("source_message_id")
.arg(message_id)
.arg("deliveries")
.arg(deliveries)
.arg("error_code")
.arg(error.code.as_str())
.arg("error_message")
.arg(&error.message)
.arg("failed_at")
.arg(chrono::Utc::now().to_rfc3339())
.query_async::<_, redis::Value>(conn)
.await
.map_err(|err| AppError::new(ErrorCode::Internal, "写入死信队列失败").with_source(err))?;
Ok(())
}
async fn mark_task_dead_letter(
state: &AppState,
task_id: Uuid,
worker_id: Uuid,
message: &str,
) -> Result<bool, AppError> {
let mut tx =
state.db.begin().await.map_err(|err| {
AppError::new(ErrorCode::Internal, "开启死信事务失败").with_source(err)
})?;
let owned_task: Option<Uuid> = sqlx::query_scalar(
r#"
SELECT id
FROM tasks
WHERE id = $1
AND status = 'processing'
AND lease_owner = $2
AND lease_until > NOW()
FOR UPDATE
"#,
)
.bind(task_id)
.bind(worker_id)
.fetch_optional(&mut *tx)
.await
.map_err(|err| AppError::new(ErrorCode::Internal, "锁定死信任务失败").with_source(err))?;
if owned_task.is_none() {
tx.rollback().await.ok();
return Ok(false);
}
sqlx::query(
r#"
UPDATE task_files
SET status = 'failed',
error_message = $2,
completed_at = NOW(),
lease_owner = NULL,
lease_until = NULL
WHERE task_id = $1
AND status IN ('pending', 'processing')
"#,
)
.bind(task_id)
.bind(message)
.execute(&mut *tx)
.await
.map_err(|err| AppError::new(ErrorCode::Internal, "标记死信文件失败").with_source(err))?;
sqlx::query(
r#"
UPDATE tasks
SET status = 'failed',
completed_files = (
SELECT COUNT(*) FROM task_files WHERE task_id = $1 AND status = 'completed'
),
failed_files = (
SELECT COUNT(*) FROM task_files WHERE task_id = $1 AND status = 'failed'
),
total_compressed_size = COALESCE((
SELECT SUM(compressed_size)
FROM task_files
WHERE task_id = $1 AND status = 'completed'
), 0)::bigint,
error_message = $2,
completed_at = NOW(),
lease_owner = NULL,
lease_until = NULL
WHERE id = $1 AND status = 'processing' AND lease_owner = $3
"#,
)
.bind(task_id)
.bind(message)
.bind(worker_id)
.execute(&mut *tx)
.await
.map_err(|err| AppError::new(ErrorCode::Internal, "标记死信任务失败").with_source(err))?;
tx.commit()
.await
.map_err(|err| AppError::new(ErrorCode::Internal, "提交死信事务失败").with_source(err))?;
quota::settle_anonymous_task_reservation(state, task_id).await?;
Ok(true)
}
async fn renew_processing_lease(
state: &AppState,
task_id: Uuid,
worker_id: Uuid,
) -> Result<bool, AppError> {
let mut tx =
state.db.begin().await.map_err(|err| {
AppError::new(ErrorCode::Internal, "开启续租事务失败").with_source(err)
})?;
let renewed = sqlx::query(
r#"
UPDATE tasks
SET lease_until = NOW() + $3 * INTERVAL '1 second'
WHERE id = $1
AND status = 'processing'
AND lease_owner = $2
AND lease_until > NOW()
"#,
)
.bind(task_id)
.bind(worker_id)
.bind(PROCESSING_LEASE_SECONDS)
.execute(&mut *tx)
.await
.map_err(|err| AppError::new(ErrorCode::Internal, "续租任务失败").with_source(err))?;
if renewed.rows_affected() == 0 {
tx.rollback().await.ok();
return Ok(false);
}
sqlx::query(
r#"
UPDATE task_files
SET lease_until = NOW() + $3 * INTERVAL '1 second'
WHERE task_id = $1
AND status = 'processing'
AND lease_owner = $2
AND lease_until > NOW()
"#,
)
.bind(task_id)
.bind(worker_id)
.bind(PROCESSING_LEASE_SECONDS)
.execute(&mut *tx)
.await
.map_err(|err| AppError::new(ErrorCode::Internal, "续租任务文件失败").with_source(err))?;
tx.commit()
.await
.map_err(|err| AppError::new(ErrorCode::Internal, "提交续租失败").with_source(err))?;
Ok(true)
}
async fn release_processing_lease(state: &AppState, task_id: Uuid, worker_id: Uuid) {
let result = async {
let mut tx = state.db.begin().await?;
sqlx::query(
"UPDATE tasks SET lease_until = NOW() WHERE id = $1 AND status = 'processing' AND lease_owner = $2",
)
.bind(task_id)
.bind(worker_id)
.execute(&mut *tx)
.await?;
sqlx::query(
"UPDATE task_files SET lease_until = NOW() WHERE task_id = $1 AND status = 'processing' AND lease_owner = $2",
)
.bind(task_id)
.bind(worker_id)
.execute(&mut *tx)
.await?;
tx.commit().await
}
.await;
if let Err(err) = result {
tracing::warn!(task_id = %task_id, worker_id = %worker_id, error = ?err, "failed to release processing lease");
}
}
async fn ack_message(
conn: &mut redis::aio::ConnectionManager,
msg_id: &str,
) -> Result<(), AppError> {
let _: i64 = redis::cmd("XACK")
.arg(STREAM_KEY)
.arg(GROUP_NAME)
.arg(msg_id)
.query_async(conn)
.await
.map_err(|err| AppError::new(ErrorCode::Internal, "确认队列消息失败").with_source(err))?;
Ok(())
}
#[derive(Debug, FromRow)]
struct TaskProcRow {
compression_level: String,
compression_rate: Option<i16>,
max_width: Option<i32>,
max_height: Option<i32>,
preserve_metadata: bool,
user_id: Option<Uuid>,
session_id: Option<String>,
api_key_id: Option<Uuid>,
source: String,
client_ip: Option<String>,
retention_hours: i32,
anonymous_units_reserved: i32,
processing_attempt: i64,
}
#[derive(Debug, FromRow)]
struct TaskFileProcRow {
id: Uuid,
input_path: Option<String>,
original_format: String,
output_format: String,
}
#[derive(Clone)]
struct TaskContext {
api_key_id: Option<Uuid>,
source: String,
preserve_metadata: bool,
session_id: Option<String>,
anon_ip: Option<IpAddr>,
is_anonymous: bool,
retention_hours: i32,
anonymous_quota_reserved: bool,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(crate) enum TaskProcessOutcome {
Done,
LeaseBusy,
}
pub(crate) async fn process_task(
state: &AppState,
task_id: Uuid,
worker_id: Uuid,
) -> Result<TaskProcessOutcome, AppError> {
let task: Option<TaskProcRow> = sqlx::query_as(
r#"
UPDATE tasks
SET status = 'processing',
started_at = COALESCE(started_at, NOW()),
processing_attempt = CASE
WHEN status = 'pending'
OR lease_owner IS DISTINCT FROM $2
OR lease_until IS NULL
OR lease_until <= NOW()
THEN processing_attempt + 1
ELSE processing_attempt
END,
lease_owner = $2,
lease_until = NOW() + $3 * INTERVAL '1 second'
WHERE id = $1
AND deletion_started_at IS NULL
AND (
status = 'pending'
OR (
status = 'processing'
AND (
lease_owner = $2
OR lease_owner IS NULL
OR lease_until IS NULL
OR lease_until <= NOW()
)
)
)
RETURNING
compression_level::text AS compression_level,
compression_rate,
max_width,
max_height,
preserve_metadata,
user_id,
session_id,
api_key_id,
source::text AS source,
host(client_ip) AS client_ip,
retention_hours,
anonymous_units_reserved,
processing_attempt
"#,
)
.bind(task_id)
.bind(worker_id)
.bind(PROCESSING_LEASE_SECONDS)
.fetch_optional(&state.db)
.await
.map_err(|err| AppError::new(ErrorCode::Internal, "领取任务失败").with_source(err))?;
let Some(task) = task else {
let status: Option<String> =
sqlx::query_scalar("SELECT status::text FROM tasks WHERE id = $1")
.bind(task_id)
.fetch_optional(&state.db)
.await
.map_err(|err| {
AppError::new(ErrorCode::Internal, "查询任务状态失败").with_source(err)
})?;
match status.as_deref() {
None => {
tracing::info!(task_id = %task_id, "task was deleted before processing; acknowledging message");
return Ok(TaskProcessOutcome::Done);
}
Some("cancelled") => {
finalize_cancelled_task(state, task_id).await?;
return Ok(TaskProcessOutcome::Done);
}
Some("completed" | "failed") => {
quota::settle_anonymous_task_reservation(state, task_id).await?;
return Ok(TaskProcessOutcome::Done);
}
Some("processing") => return Ok(TaskProcessOutcome::LeaseBusy),
Some(_) => return Ok(TaskProcessOutcome::LeaseBusy),
}
};
let compression_rate = task.compression_rate.and_then(|v| u8::try_from(v).ok());
let level = compression_rate
.map(compress::rate_to_level)
.unwrap_or(compress::parse_level(&task.compression_level)?);
let max_width = task.max_width.and_then(|v| u32::try_from(v).ok());
let max_height = task.max_height.and_then(|v| u32::try_from(v).ok());
let mut files: Vec<TaskFileProcRow> = sqlx::query_as(
r#"
SELECT
id,
input_path,
original_format,
output_format
FROM task_files
WHERE task_id = $1
AND status IN ('pending', 'processing')
ORDER BY created_at ASC
"#,
)
.bind(task_id)
.fetch_all(&state.db)
.await
.map_err(|err| AppError::new(ErrorCode::Internal, "查询任务文件失败").with_source(err))?;
let billing_ctx = if let Some(user_id) = task.user_id {
Some(billing::get_user_billing(state, user_id).await?)
} else {
None
};
let anon_ip: Option<IpAddr> = task
.client_ip
.as_deref()
.and_then(|s| s.parse::<IpAddr>().ok());
let ctx = TaskContext {
api_key_id: task.api_key_id,
source: task.source.clone(),
preserve_metadata: task.preserve_metadata,
session_id: task.session_id.clone(),
anon_ip,
is_anonymous: task.user_id.is_none(),
retention_hours: task.retention_hours,
anonymous_quota_reserved: task.anonymous_units_reserved > 0,
};
let concurrency = state.config.worker_concurrency.max(1) as usize;
let semaphore = Arc::new(Semaphore::new(concurrency));
let mut join_set = JoinSet::new();
for file in files.drain(..) {
let permit = semaphore.clone().acquire_owned().await.unwrap();
let state = state.clone();
let ctx = ctx.clone();
let billing_ctx = billing_ctx.clone();
let file_id = file.id;
join_set.spawn(async move {
let _permit = permit;
let result = process_task_file(
state,
task_id,
task.processing_attempt,
worker_id,
file,
level,
compression_rate,
max_width,
max_height,
ctx,
billing_ctx,
)
.await;
if let Err(err) = &result {
tracing::error!(task_id = %task_id, file_id = %file_id, error = %err, "file processing failed");
}
result
});
}
let mut first_error = None;
while let Some(result) = join_set.join_next().await {
match result {
Ok(Ok(())) => {}
Ok(Err(err)) => {
if first_error.is_none() {
first_error = Some(err);
}
}
Err(err) => {
return Err(
AppError::new(ErrorCode::Internal, "文件处理线程异常退出").with_source(err)
);
}
}
}
if let Some(err) = first_error {
return Err(err);
}
if finalize_task_status(state, task_id, task.processing_attempt, worker_id).await? {
Ok(TaskProcessOutcome::Done)
} else {
Ok(TaskProcessOutcome::LeaseBusy)
}
}
fn parse_image_fmt(value: &str) -> Result<compress::ImageFmt, AppError> {
match value.trim().to_ascii_lowercase().as_str() {
"png" => Ok(compress::ImageFmt::Png),
"jpeg" | "jpg" => Ok(compress::ImageFmt::Jpeg),
"webp" => Ok(compress::ImageFmt::Webp),
"avif" => Ok(compress::ImageFmt::Avif),
"gif" => Ok(compress::ImageFmt::Gif),
"bmp" => Ok(compress::ImageFmt::Bmp),
"tif" | "tiff" => Ok(compress::ImageFmt::Tiff),
"ico" => Ok(compress::ImageFmt::Ico),
_ => Err(AppError::new(ErrorCode::InvalidRequest, "未知图片格式")),
}
}
async fn is_task_cancelled(state: &AppState, task_id: Uuid) -> Result<bool, AppError> {
let status: Option<String> = sqlx::query_scalar("SELECT status::text FROM tasks WHERE id = $1")
.bind(task_id)
.fetch_optional(&state.db)
.await
.map_err(|err| AppError::new(ErrorCode::Internal, "查询任务状态失败").with_source(err))?;
Ok(matches!(status.as_deref(), Some("cancelled")))
}
#[derive(Debug, Clone, Copy)]
struct FileFence {
task_id: Uuid,
task_attempt: i64,
file_id: Uuid,
file_attempt: i64,
worker_id: Uuid,
}
async fn file_attempt_is_current(state: &AppState, fence: &FileFence) -> Result<bool, AppError> {
sqlx::query_scalar(
r#"
SELECT EXISTS(
SELECT 1
FROM tasks t
JOIN task_files f ON f.task_id = t.id
WHERE t.id = $1
AND t.status = 'processing'
AND t.deletion_started_at IS NULL
AND t.processing_attempt = $2
AND t.lease_owner = $5
AND t.lease_until > NOW()
AND f.id = $3
AND f.status = 'processing'
AND f.processing_attempt = $4
AND f.lease_owner = $5
AND f.lease_until > NOW()
)
"#,
)
.bind(fence.task_id)
.bind(fence.task_attempt)
.bind(fence.file_id)
.bind(fence.file_attempt)
.bind(fence.worker_id)
.fetch_one(&state.db)
.await
.map_err(|err| AppError::new(ErrorCode::Internal, "检查文件处理租约失败").with_source(err))
}
#[allow(clippy::too_many_arguments)]
async fn process_task_file(
state: AppState,
task_id: Uuid,
task_attempt: i64,
worker_id: Uuid,
file: TaskFileProcRow,
level: compress::CompressionLevel,
compression_rate: Option<u8>,
max_width: Option<u32>,
max_height: Option<u32>,
ctx: TaskContext,
billing_ctx: Option<billing::BillingContext>,
) -> Result<(), AppError> {
let file_attempt: Option<i64> = sqlx::query_scalar(
r#"
UPDATE task_files AS f
SET status = 'processing',
processing_attempt = processing_attempt + 1,
lease_owner = $4,
lease_until = NOW() + $5 * INTERVAL '1 second',
error_message = NULL
FROM tasks AS t
WHERE f.id = $1
AND f.task_id = t.id
AND t.id = $2
AND t.status = 'processing'
AND t.processing_attempt = $3
AND t.lease_owner = $4
AND t.lease_until > NOW()
AND (
f.status = 'pending'
OR (
f.status = 'processing'
AND (
f.lease_owner = $4
OR f.lease_owner IS NULL
OR f.lease_until IS NULL
OR f.lease_until <= NOW()
)
)
)
RETURNING f.processing_attempt
"#,
)
.bind(file.id)
.bind(task_id)
.bind(task_attempt)
.bind(worker_id)
.bind(PROCESSING_LEASE_SECONDS)
.fetch_optional(&state.db)
.await
.map_err(|err| AppError::new(ErrorCode::Internal, "更新文件处理状态失败").with_source(err))?;
let Some(file_attempt) = file_attempt else {
return Ok(());
};
let fence = FileFence {
task_id,
task_attempt,
file_id: file.id,
file_attempt,
worker_id,
};
let Some(input_path) = file.input_path.clone() else {
mark_file_failed(&state, &fence, "原文件不存在").await?;
return Ok(());
};
let input_bytes = match tokio::fs::read(&input_path).await {
Ok(v) => v,
Err(_) => {
mark_file_failed_and_cleanup(&state, &fence, "读取原文件失败", &input_path).await?;
return Ok(());
}
};
let format_in = match parse_image_fmt(&file.original_format) {
Ok(format) => format,
Err(err) => {
mark_file_failed_and_cleanup(&state, &fence, &err.message, &input_path).await?;
return Ok(());
}
};
let format_out = match parse_image_fmt(&file.output_format) {
Ok(format) => format,
Err(err) => {
mark_file_failed_and_cleanup(&state, &fence, &err.message, &input_path).await?;
return Ok(());
}
};
let original_size = input_bytes.len() as u64;
let compressed = match compress::compress_image_bytes(
&state,
input_bytes,
format_in,
format_out,
level,
compression_rate,
None, // target_size_bytes: worker 批量任务不支持精确大小
max_width,
max_height,
ctx.preserve_metadata,
)
.await
{
Ok(v) => v,
Err(err) => {
mark_file_failed_and_cleanup(&state, &fence, &err.message, &input_path).await?;
return Ok(());
}
};
if !file_attempt_is_current(&state, &fence).await? || is_task_cancelled(&state, task_id).await?
{
return Ok(());
}
let compressed_size = compressed.len() as u64;
let saved_percent = if original_size == 0 {
0.0
} else {
(original_size.saturating_sub(compressed_size) as f64) * 100.0 / (original_size as f64)
};
let charge_units = quota::output_consumes_unit(
compression_rate,
format_in == format_out,
max_width.is_some() || max_height.is_some(),
false,
original_size,
compressed_size,
);
let object_key = storage::result_attempt_key(
ctx.retention_hours as i64,
task_id,
file.id,
task_attempt,
file_attempt,
format_out.extension(),
);
let tracked = match object_lifecycle::store_tracked_bytes(
&state,
task_id,
Some(file.id),
"result",
&object_key,
compressed.into(),
format_out.content_type(),
)
.await
{
Ok(value) => value,
Err(err) => {
if reset_file_for_retry(&state, &fence, "对象存储暂时不可用").await? {
return Err(err);
}
return Ok(());
}
};
if ctx.is_anonymous && charge_units && !ctx.anonymous_quota_reserved {
let Some(session_id) = ctx.session_id.as_deref() else {
discard_tracked_result(&state, &tracked, None).await;
mark_file_failed_and_cleanup(&state, &fence, "匿名任务缺少 session_id", &input_path)
.await?;
return Ok(());
};
let Some(ip) = ctx.anon_ip else {
discard_tracked_result(&state, &tracked, None).await;
mark_file_failed_and_cleanup(&state, &fence, "匿名任务缺少 client_ip", &input_path)
.await?;
return Ok(());
};
if let Err(err) = quota::consume_anonymous_units(&state, session_id, ip, 1).await {
discard_tracked_result(&state, &tracked, Some(&err)).await;
mark_file_failed_and_cleanup(&state, &fence, &err.message, &input_path).await?;
return Ok(());
}
}
match finalize_file(
&state,
&billing_ctx,
ctx.api_key_id,
&ctx.source,
&fence,
&tracked,
original_size as i64,
compressed_size as i64,
saved_percent,
format_in,
format_out,
charge_units,
)
.await
{
Ok(FinalizeFileOutcome::Committed) => {
let _ = tokio::fs::remove_file(&input_path).await;
}
Ok(FinalizeFileOutcome::LeaseLost) => {
discard_tracked_result(&state, &tracked, None).await;
}
Err(err) => match worker_result_was_committed(&state, &fence, &tracked).await {
Ok(true) => {
tracing::warn!(task_id = %task_id, file_id = %fence.file_id, error = %err, "worker result commit response was lost; recovered committed publication");
let _ = tokio::fs::remove_file(&input_path).await;
}
Ok(false) => {
discard_tracked_result(&state, &tracked, Some(&err)).await;
mark_file_failed_and_cleanup(&state, &fence, &err.message, &input_path).await?;
}
Err(probe_err) => {
tracing::error!(task_id = %task_id, file_id = %fence.file_id, error = %probe_err, original_error = %err, "worker result commit state is unknown; staging lease will reconcile object");
return Err(err);
}
},
}
Ok(())
}
async fn discard_tracked_result(
state: &AppState,
tracked: &object_lifecycle::TrackedStoredObject,
error: Option<&AppError>,
) {
if let Err(schedule_err) =
object_lifecycle::schedule_tracked_delete(state, tracked, error).await
{
tracing::error!(storage_object_id = %tracked.lifecycle_id, error = %schedule_err, "failed to persist discarded worker object cleanup");
return;
}
if let Err(cleanup_err) =
object_lifecycle::cleanup_ready_objects(state, 1, Some(tracked.task_id)).await
{
tracing::warn!(storage_object_id = %tracked.lifecycle_id, error = %cleanup_err, "discarded worker object cleanup deferred");
}
}
async fn worker_result_was_committed(
state: &AppState,
fence: &FileFence,
tracked: &object_lifecycle::TrackedStoredObject,
) -> Result<bool, AppError> {
sqlx::query_scalar(
r#"
SELECT EXISTS(
SELECT 1
FROM tasks AS task
JOIN task_files AS file ON file.task_id = task.id
JOIN storage_objects AS object ON object.id = $3
WHERE task.id = $1
AND file.id = $2
AND file.status = 'completed'
AND file.storage_backend = $4
AND file.storage_endpoint_id IS NOT DISTINCT FROM $5
AND COALESCE(file.storage_key, file.storage_path) = $6
AND object.state = 'published'
AND object.task_id = task.id
AND object.task_file_id = file.id
)
"#,
)
.bind(fence.task_id)
.bind(fence.file_id)
.bind(tracked.lifecycle_id)
.bind(&tracked.stored.backend)
.bind(tracked.stored.endpoint_id)
.bind(&tracked.stored.key)
.fetch_one(&state.db)
.await
.map_err(|err| {
AppError::new(
ErrorCode::StorageUnavailable,
"核验 Worker 结果提交状态失败",
)
.with_source(err)
})
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum FinalizeFileOutcome {
Committed,
LeaseLost,
}
#[allow(clippy::too_many_arguments)]
async fn finalize_file(
state: &AppState,
billing_ctx: &Option<billing::BillingContext>,
api_key_id: Option<Uuid>,
source: &str,
fence: &FileFence,
tracked: &object_lifecycle::TrackedStoredObject,
bytes_in: i64,
bytes_out: i64,
saved_percent: f64,
format_in: compress::ImageFmt,
format_out: compress::ImageFmt,
charge_units: bool,
) -> Result<FinalizeFileOutcome, AppError> {
let stored = &tracked.stored;
let mut tx = state
.db
.begin()
.await
.map_err(|err| AppError::new(ErrorCode::Internal, "开启事务失败").with_source(err))?;
let task_fence: Option<(String, i64, Option<Uuid>, bool)> = sqlx::query_as(
r#"
SELECT status::text, processing_attempt, lease_owner,
COALESCE(lease_until > NOW(), false)
FROM tasks
WHERE id = $1
FOR UPDATE
"#,
)
.bind(fence.task_id)
.fetch_optional(&mut *tx)
.await
.map_err(|err| AppError::new(ErrorCode::Internal, "锁定任务状态失败").with_source(err))?;
let Some((task_status, task_attempt, task_owner, task_lease_valid)) = task_fence else {
return Ok(FinalizeFileOutcome::LeaseLost);
};
if task_status != "processing"
|| task_attempt != fence.task_attempt
|| task_owner != Some(fence.worker_id)
|| !task_lease_valid
{
return Ok(FinalizeFileOutcome::LeaseLost);
}
let file_fence: Option<(String, i64, Option<Uuid>, bool)> = sqlx::query_as(
r#"
SELECT status::text, processing_attempt, lease_owner,
COALESCE(lease_until > NOW(), false)
FROM task_files
WHERE id = $1 AND task_id = $2
FOR UPDATE
"#,
)
.bind(fence.file_id)
.bind(fence.task_id)
.fetch_optional(&mut *tx)
.await
.map_err(|err| AppError::new(ErrorCode::Internal, "锁定文件状态失败").with_source(err))?;
let Some((file_status, file_attempt, file_owner, file_lease_valid)) = file_fence else {
return Ok(FinalizeFileOutcome::LeaseLost);
};
if file_status != "processing"
|| file_attempt != fence.file_attempt
|| file_owner != Some(fence.worker_id)
|| !file_lease_valid
{
return Ok(FinalizeFileOutcome::LeaseLost);
}
// Paid users: charge before marking file completed (atomic w/ status update).
if charge_units {
if let Some(billing) = billing_ctx {
charge_one_unit(
&mut tx,
billing,
api_key_id,
source,
fence.task_id,
fence.file_id,
format_in,
format_out,
bytes_in as u64,
bytes_out as u64,
)
.await?;
}
}
let file_updated = sqlx::query(
r#"
UPDATE task_files
SET storage_path = $2,
storage_backend = $3,
storage_endpoint_id = $4,
storage_key = $5,
storage_etag = $6,
input_path = NULL,
compressed_size = $7,
saved_percent = $8,
status = 'completed',
completed_at = NOW(),
lease_owner = NULL,
lease_until = NULL
WHERE id = $1
AND task_id = $9
AND status = 'processing'
AND processing_attempt = $10
AND lease_owner = $11
"#,
)
.bind(fence.file_id)
.bind(if stored.backend == "local" {
Some(stored.key.as_str())
} else {
None
})
.bind(&stored.backend)
.bind(stored.endpoint_id)
.bind(&stored.key)
.bind(&stored.etag)
.bind(bytes_out)
.bind(saved_percent)
.bind(fence.task_id)
.bind(fence.file_attempt)
.bind(fence.worker_id)
.execute(&mut *tx)
.await
.map_err(|err| AppError::new(ErrorCode::Internal, "更新文件失败").with_source(err))?;
if file_updated.rows_affected() == 0 {
return Ok(FinalizeFileOutcome::LeaseLost);
}
let task_updated = sqlx::query(
r#"
UPDATE tasks
SET completed_files = LEAST(total_files, completed_files + 1),
total_compressed_size = total_compressed_size + $2
WHERE id = $1
AND status = 'processing'
AND processing_attempt = $3
AND lease_owner = $4
"#,
)
.bind(fence.task_id)
.bind(bytes_out)
.bind(fence.task_attempt)
.bind(fence.worker_id)
.execute(&mut *tx)
.await
.map_err(|err| AppError::new(ErrorCode::Internal, "更新任务统计失败").with_source(err))?;
if task_updated.rows_affected() == 0 {
return Ok(FinalizeFileOutcome::LeaseLost);
}
object_lifecycle::publish_in_tx(&mut tx, tracked).await?;
tx.commit()
.await
.map_err(|err| AppError::new(ErrorCode::Internal, "提交事务失败").with_source(err))?;
Ok(FinalizeFileOutcome::Committed)
}
async fn mark_file_failed(
state: &AppState,
fence: &FileFence,
message: &str,
) -> Result<bool, AppError> {
let mut tx = state
.db
.begin()
.await
.map_err(|err| AppError::new(ErrorCode::Internal, "开启事务失败").with_source(err))?;
let task_owned: Option<Uuid> = sqlx::query_scalar(
r#"
SELECT id
FROM tasks
WHERE id = $1
AND status = 'processing'
AND processing_attempt = $2
AND lease_owner = $3
AND lease_until > NOW()
FOR UPDATE
"#,
)
.bind(fence.task_id)
.bind(fence.task_attempt)
.bind(fence.worker_id)
.fetch_optional(&mut *tx)
.await
.map_err(|err| AppError::new(ErrorCode::Internal, "锁定任务状态失败").with_source(err))?;
if task_owned.is_none() {
tx.rollback().await.ok();
return Ok(false);
}
let updated = sqlx::query(
r#"
UPDATE task_files
SET status = 'failed',
error_message = $2,
storage_path = NULL,
input_path = NULL,
completed_at = NOW(),
lease_owner = NULL,
lease_until = NULL
WHERE id = $1
AND task_id = $3
AND status = 'processing'
AND processing_attempt = $4
AND lease_owner = $5
"#,
)
.bind(fence.file_id)
.bind(message)
.bind(fence.task_id)
.bind(fence.file_attempt)
.bind(fence.worker_id)
.execute(&mut *tx)
.await
.map_err(|err| AppError::new(ErrorCode::Internal, "更新文件失败").with_source(err))?;
if updated.rows_affected() > 0 {
sqlx::query(
r#"
UPDATE tasks
SET failed_files = LEAST(total_files, failed_files + 1)
WHERE id = $1
AND status = 'processing'
AND processing_attempt = $2
AND lease_owner = $3
"#,
)
.bind(fence.task_id)
.bind(fence.task_attempt)
.bind(fence.worker_id)
.execute(&mut *tx)
.await
.map_err(|err| AppError::new(ErrorCode::Internal, "更新任务统计失败").with_source(err))?;
}
tx.commit()
.await
.map_err(|err| AppError::new(ErrorCode::Internal, "提交事务失败").with_source(err))?;
Ok(updated.rows_affected() > 0)
}
async fn mark_file_failed_and_cleanup(
state: &AppState,
fence: &FileFence,
message: &str,
input_path: &str,
) -> Result<(), AppError> {
if mark_file_failed(state, fence, message).await? {
let _ = tokio::fs::remove_file(input_path).await;
}
Ok(())
}
async fn reset_file_for_retry(
state: &AppState,
fence: &FileFence,
message: &str,
) -> Result<bool, AppError> {
let mut tx =
state.db.begin().await.map_err(|err| {
AppError::new(ErrorCode::Internal, "开启重试事务失败").with_source(err)
})?;
let task_owned: Option<Uuid> = sqlx::query_scalar(
r#"
SELECT id FROM tasks
WHERE id = $1
AND status = 'processing'
AND processing_attempt = $2
AND lease_owner = $3
AND lease_until > NOW()
FOR UPDATE
"#,
)
.bind(fence.task_id)
.bind(fence.task_attempt)
.bind(fence.worker_id)
.fetch_optional(&mut *tx)
.await
.map_err(|err| AppError::new(ErrorCode::Internal, "锁定重试任务失败").with_source(err))?;
if task_owned.is_none() {
tx.rollback().await.ok();
return Ok(false);
}
let updated = sqlx::query(
r#"
UPDATE task_files
SET status = 'pending',
error_message = $2,
lease_owner = NULL,
lease_until = NULL
WHERE id = $1
AND task_id = $3
AND status = 'processing'
AND processing_attempt = $4
AND lease_owner = $5
"#,
)
.bind(fence.file_id)
.bind(message)
.bind(fence.task_id)
.bind(fence.file_attempt)
.bind(fence.worker_id)
.execute(&mut *tx)
.await
.map_err(|err| AppError::new(ErrorCode::Internal, "恢复待重试文件失败").with_source(err))?;
tx.commit()
.await
.map_err(|err| AppError::new(ErrorCode::Internal, "提交重试状态失败").with_source(err))?;
Ok(updated.rows_affected() > 0)
}
async fn finalize_task_status(
state: &AppState,
task_id: Uuid,
task_attempt: i64,
worker_id: Uuid,
) -> Result<bool, AppError> {
let mut tx = state.db.begin().await.map_err(|err| {
AppError::new(ErrorCode::Internal, "开启任务终结事务失败").with_source(err)
})?;
let task: Option<(String, i64, Option<Uuid>, bool, i32)> = sqlx::query_as(
r#"
SELECT status::text, processing_attempt, lease_owner,
COALESCE(lease_until > NOW(), false), total_files
FROM tasks
WHERE id = $1
FOR UPDATE
"#,
)
.bind(task_id)
.fetch_optional(&mut *tx)
.await
.map_err(|err| AppError::new(ErrorCode::Internal, "锁定任务失败").with_source(err))?;
let Some((status, current_attempt, current_owner, lease_valid, total)) = task else {
return Ok(true);
};
if status == "cancelled" {
tx.rollback().await.ok();
finalize_cancelled_task(state, task_id).await?;
return Ok(true);
}
if status != "processing"
|| current_attempt != task_attempt
|| current_owner != Some(worker_id)
|| !lease_valid
{
return Ok(false);
}
let (completed, failed, total_compressed_size): (i64, i64, i64) = sqlx::query_as(
r#"
SELECT
COUNT(*) FILTER (WHERE status = 'completed'),
COUNT(*) FILTER (WHERE status = 'failed'),
COALESCE(SUM(compressed_size) FILTER (WHERE status = 'completed'), 0)::bigint
FROM task_files
WHERE task_id = $1
"#,
)
.bind(task_id)
.fetch_one(&mut *tx)
.await
.map_err(|err| AppError::new(ErrorCode::Internal, "统计任务文件失败").with_source(err))?;
let finished = completed + failed >= i64::from(total) && total > 0;
if !finished {
sqlx::query(
r#"
UPDATE tasks
SET completed_files = $4, failed_files = $5, total_compressed_size = $6
WHERE id = $1 AND status = 'processing'
AND processing_attempt = $2 AND lease_owner = $3
"#,
)
.bind(task_id)
.bind(task_attempt)
.bind(worker_id)
.bind(completed as i32)
.bind(failed as i32)
.bind(total_compressed_size)
.execute(&mut *tx)
.await
.map_err(|err| AppError::new(ErrorCode::Internal, "刷新任务统计失败").with_source(err))?;
tx.commit().await.map_err(|err| {
AppError::new(ErrorCode::Internal, "提交任务统计失败").with_source(err)
})?;
return Ok(false);
}
let final_status = if completed == 0 && failed == i64::from(total) {
"failed"
} else {
"completed"
};
let updated = sqlx::query(
r#"
UPDATE tasks
SET status = $4::task_status,
completed_files = $5,
failed_files = $6,
total_compressed_size = $7,
completed_at = NOW(),
lease_owner = NULL,
lease_until = NULL
WHERE id = $1 AND status = 'processing'
AND processing_attempt = $2 AND lease_owner = $3
"#,
)
.bind(task_id)
.bind(task_attempt)
.bind(worker_id)
.bind(final_status)
.bind(completed as i32)
.bind(failed as i32)
.bind(total_compressed_size)
.execute(&mut *tx)
.await
.map_err(|err| AppError::new(ErrorCode::Internal, "更新任务状态失败").with_source(err))?;
if updated.rows_affected() == 0 {
return Ok(false);
}
tx.commit()
.await
.map_err(|err| AppError::new(ErrorCode::Internal, "提交任务状态失败").with_source(err))?;
quota::settle_anonymous_task_reservation(state, task_id).await?;
Ok(true)
}
async fn finalize_cancelled_task(state: &AppState, task_id: Uuid) -> Result<(), AppError> {
let mut tx = state.db.begin().await.map_err(|err| {
AppError::new(ErrorCode::Internal, "开启取消清理事务失败").with_source(err)
})?;
let status: Option<String> =
sqlx::query_scalar("SELECT status::text FROM tasks WHERE id = $1 FOR UPDATE")
.bind(task_id)
.fetch_optional(&mut *tx)
.await
.map_err(|err| {
AppError::new(ErrorCode::Internal, "锁定取消任务失败").with_source(err)
})?;
if status.as_deref() != Some("cancelled") {
tx.rollback().await.ok();
return Ok(());
}
let paths: Vec<Option<String>> = sqlx::query_scalar(
"SELECT input_path FROM task_files WHERE task_id = $1 AND status IN ('pending','processing') FOR UPDATE",
)
.bind(task_id)
.fetch_all(&mut *tx)
.await
.map_err(|err| AppError::new(ErrorCode::Internal, "锁定取消文件失败").with_source(err))?;
sqlx::query(
r#"
UPDATE task_files
SET status = 'failed', error_message = '已取消', storage_path = NULL,
input_path = NULL, completed_at = NOW(), lease_owner = NULL, lease_until = NULL
WHERE task_id = $1 AND status IN ('pending','processing')
"#,
)
.bind(task_id)
.execute(&mut *tx)
.await
.map_err(|err| AppError::new(ErrorCode::Internal, "取消任务文件失败").with_source(err))?;
sqlx::query(
r#"
UPDATE tasks
SET completed_files = (SELECT COUNT(*) FROM task_files WHERE task_id = $1 AND status = 'completed'),
failed_files = (SELECT COUNT(*) FROM task_files WHERE task_id = $1 AND status = 'failed'),
total_compressed_size = COALESCE((SELECT SUM(compressed_size) FROM task_files WHERE task_id = $1 AND status = 'completed'), 0)::bigint,
completed_at = COALESCE(completed_at, NOW()), lease_owner = NULL, lease_until = NULL
WHERE id = $1 AND status = 'cancelled'
"#,
)
.bind(task_id)
.execute(&mut *tx)
.await
.map_err(|err| AppError::new(ErrorCode::Internal, "完成取消任务失败").with_source(err))?;
tx.commit()
.await
.map_err(|err| AppError::new(ErrorCode::Internal, "提交取消清理失败").with_source(err))?;
for path in paths.into_iter().flatten() {
let _ = tokio::fs::remove_file(path).await;
}
quota::settle_anonymous_task_reservation(state, task_id).await?;
Ok(())
}
#[allow(clippy::too_many_arguments)]
async fn charge_one_unit(
tx: &mut sqlx::Transaction<'_, sqlx::Postgres>,
billing: &billing::BillingContext,
api_key_id: Option<Uuid>,
source: &str,
task_id: Uuid,
task_file_id: Uuid,
format_in: compress::ImageFmt,
format_out: compress::ImageFmt,
bytes_in: u64,
bytes_out: u64,
) -> Result<(), AppError> {
quota::consume_user_unit(tx, billing, bytes_in, bytes_out).await?;
sqlx::query(
r#"
INSERT INTO usage_events (
user_id, api_key_id, source,
task_id, task_file_id,
units, bytes_in, bytes_out, format_in, format_out
) VALUES (
$1, $2, $3::task_source,
$4, $5,
1, $6, $7, $8, $9
)
"#,
)
.bind(billing.user_id)
.bind(api_key_id)
.bind(source)
.bind(task_id)
.bind(task_file_id)
.bind(bytes_in as i64)
.bind(bytes_out as i64)
.bind(format_in.as_str())
.bind(format_out.as_str())
.execute(&mut **tx)
.await
.map_err(|err| AppError::new(ErrorCode::Internal, "写入用量明细失败").with_source(err))?;
Ok(())
}
async fn maintenance(state: &AppState) -> Result<(), AppError> {
settle_stale_anonymous_single_reservations(state).await?;
settle_finished_anonymous_reservations(state).await?;
cleanup_expired_tasks(state).await?;
cleanup_stale_zip_temp(state).await?;
cleanup_expired_records(state).await?;
Ok(())
}
async fn settle_stale_anonymous_single_reservations(state: &AppState) -> Result<(), AppError> {
for _ in 0..MAX_MAINTENANCE_BATCHES {
let settled =
quota::settle_stale_anonymous_single_reservations(state, MAINTENANCE_BATCH_SIZE)
.await?;
if settled < MAINTENANCE_BATCH_SIZE as usize {
break;
}
tokio::task::yield_now().await;
}
Ok(())
}
async fn settle_finished_anonymous_reservations(state: &AppState) -> Result<(), AppError> {
for _ in 0..MAX_MAINTENANCE_BATCHES {
let task_ids: Vec<Uuid> = sqlx::query_scalar(
r#"
SELECT id
FROM tasks
WHERE anonymous_units_reserved > 0
AND status IN ('completed', 'failed', 'cancelled')
ORDER BY completed_at ASC NULLS FIRST
LIMIT $1
"#,
)
.bind(MAINTENANCE_BATCH_SIZE)
.fetch_all(&state.db)
.await
.map_err(|err| {
AppError::new(ErrorCode::Internal, "查询待结算匿名任务失败").with_source(err)
})?;
let batch_len = task_ids.len();
let mut settled = 0usize;
for task_id in task_ids {
match quota::settle_anonymous_task_reservation(state, task_id).await {
Ok(_) => settled += 1,
Err(err) => {
tracing::warn!(task_id = %task_id, error = %err, "anonymous quota settlement deferred")
}
}
}
if settled == 0 || batch_len < MAINTENANCE_BATCH_SIZE as usize {
break;
}
tokio::task::yield_now().await;
}
Ok(())
}
async fn cleanup_stale_zip_temp(state: &AppState) -> Result<(), AppError> {
let root = std::path::Path::new(&state.config.storage_path).join("tmp/zips");
let mut entries = match tokio::fs::read_dir(&root).await {
Ok(entries) => entries,
Err(err) if err.kind() == std::io::ErrorKind::NotFound => return Ok(()),
Err(err) => {
return Err(
AppError::new(ErrorCode::StorageUnavailable, "读取 ZIP 临时目录失败")
.with_source(err),
)
}
};
let cutoff = std::time::SystemTime::now()
.checked_sub(std::time::Duration::from_secs(6 * 60 * 60))
.unwrap_or(std::time::UNIX_EPOCH);
while let Some(entry) = entries.next_entry().await.map_err(|err| {
AppError::new(ErrorCode::StorageUnavailable, "遍历 ZIP 临时目录失败").with_source(err)
})? {
let metadata = match entry.metadata().await {
Ok(metadata) => metadata,
Err(_) => continue,
};
if !metadata.is_dir()
|| metadata
.modified()
.map(|time| time >= cutoff)
.unwrap_or(true)
{
continue;
}
let _ = tokio::fs::remove_dir_all(entry.path()).await;
}
Ok(())
}
async fn cleanup_expired_records(state: &AppState) -> Result<(), AppError> {
let _ = sqlx::query("DELETE FROM idempotency_keys WHERE expires_at < NOW()")
.execute(&state.db)
.await;
let _ = sqlx::query(
"DELETE FROM email_verifications WHERE expires_at < NOW() AND verified_at IS NULL",
)
.execute(&state.db)
.await;
let _ = sqlx::query("DELETE FROM password_resets WHERE expires_at < NOW() - INTERVAL '7 days'")
.execute(&state.db)
.await;
let _ = sqlx::query(
"DELETE FROM email_change_requests WHERE expires_at < NOW() - INTERVAL '7 days'",
)
.execute(&state.db)
.await;
let _ = sqlx::query(
r#"
DELETE FROM anonymous_single_reservations
WHERE status IN ('charged', 'refunded')
AND settled_at < NOW() - INTERVAL '7 days'
"#,
)
.execute(&state.db)
.await;
let _ = sqlx::query(
"DELETE FROM storage_objects WHERE state = 'deleted' AND deleted_at < NOW() - INTERVAL '7 days'",
)
.execute(&state.db)
.await;
let _ =
sqlx::query("DELETE FROM webhook_events WHERE received_at < NOW() - INTERVAL '90 days'")
.execute(&state.db)
.await;
let _ = sqlx::query(
r#"
DELETE FROM storage_endpoints e
WHERE e.deleted_at < NOW() - INTERVAL '30 days'
AND NOT EXISTS (SELECT 1 FROM task_files f WHERE f.storage_endpoint_id = e.id)
AND NOT EXISTS (SELECT 1 FROM tasks t WHERE t.zip_storage_endpoint_id = e.id)
AND NOT EXISTS (SELECT 1 FROM storage_objects o WHERE o.storage_endpoint_id = e.id)
"#,
)
.execute(&state.db)
.await;
Ok(())
}
async fn cleanup_expired_tasks(state: &AppState) -> Result<(), AppError> {
for _ in 0..MAX_MAINTENANCE_BATCHES {
let task_ids: Vec<Uuid> = sqlx::query_scalar(
"SELECT id FROM tasks WHERE expires_at < NOW() ORDER BY expires_at ASC LIMIT $1",
)
.bind(MAINTENANCE_BATCH_SIZE)
.fetch_all(&state.db)
.await
.map_err(|err| AppError::new(ErrorCode::Internal, "查询过期任务失败").with_source(err))?;
if task_ids.is_empty() {
break;
}
let batch_len = task_ids.len();
let mut cleaned = 0usize;
for task_id in task_ids {
match cleanup_expired_task(state, task_id).await {
Ok(()) => cleaned += 1,
Err(err) => {
tracing::warn!(task_id = %task_id, error = %err, "expired task cleanup deferred")
}
}
}
if cleaned == 0 || batch_len < MAINTENANCE_BATCH_SIZE as usize {
break;
}
tokio::task::yield_now().await;
}
Ok(())
}
async fn cleanup_expired_task(state: &AppState, task_id: Uuid) -> Result<(), AppError> {
if object_lifecycle::mark_expired_task(state, task_id).await? {
object_lifecycle::finalize_task_deletion(state, task_id).await?;
}
Ok(())
}
#[cfg(test)]
fn stored_locator(stored: &storage::StoredObject) -> storage::ObjectLocator {
storage::ObjectLocator {
backend: stored.backend.clone(),
endpoint_id: stored.endpoint_id,
key: stored.key.clone(),
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::config::Config;
use crate::services::mail::Mailer;
use bytes::Bytes;
use chrono::Utc;
use sqlx::postgres::PgPoolOptions;
use tokio::sync::Barrier;
#[test]
fn third_delivery_moves_message_to_dead_letter() {
assert!(!should_dead_letter(1));
assert!(!should_dead_letter(2));
assert!(should_dead_letter(3));
assert!(should_dead_letter(10));
}
#[test]
fn retry_backoff_is_exponential_and_capped() {
assert_eq!(delivery_backoff(1), Duration::from_secs(2));
assert_eq!(delivery_backoff(2), Duration::from_secs(4));
assert_eq!(delivery_backoff(3), Duration::from_secs(8));
assert_eq!(delivery_backoff(10), Duration::from_secs(30));
assert_eq!(
next_backoff(Duration::from_secs(30)),
Duration::from_secs(30)
);
}
#[test]
fn message_heartbeat_precedes_stale_claim_threshold() {
assert!(MESSAGE_HEARTBEAT_SECONDS * 1_000 < STALE_MESSAGE_IDLE_MS as u64);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
#[ignore = "requires isolated IMAGEFORGE_TEST_DATABASE_URL and IMAGEFORGE_TEST_REDIS_URL"]
async fn concurrent_attempts_finalize_once_and_keep_winning_object() {
let database_url = std::env::var("IMAGEFORGE_TEST_DATABASE_URL")
.expect("IMAGEFORGE_TEST_DATABASE_URL must be set");
assert!(
database_url.to_ascii_lowercase().contains("test"),
"refusing to run destructive integration test outside a test database"
);
let redis_url = std::env::var("IMAGEFORGE_TEST_REDIS_URL")
.expect("IMAGEFORGE_TEST_REDIS_URL must be set");
let pool = PgPoolOptions::new()
.max_connections(16)
.connect(&database_url)
.await
.expect("connect test database");
sqlx::migrate!().run(&pool).await.expect("run migrations");
let mut config = Config::from_env().expect("load test config");
config.database_url = database_url;
config.redis_url = redis_url;
config.storage_path = format!(
"{}/imageforge-worker-fence-{}",
std::env::temp_dir().display(),
Uuid::new_v4()
);
tokio::fs::create_dir_all(&config.storage_path)
.await
.expect("create test storage directory");
let redis = redis::Client::open(config.redis_url.clone())
.expect("create test redis client")
.get_connection_manager()
.await
.expect("connect test redis");
let state = AppState {
mailer: Arc::new(Mailer::new(&config).expect("create disabled test mailer")),
image_processing_semaphore: Arc::new(Semaphore::new(2)),
zip_build_semaphore: Arc::new(Semaphore::new(2)),
runtime_policy_cache: crate::services::settings::RuntimePolicyCache::new(),
storage_cache: storage::StorageCache::new(),
config,
db: pool.clone(),
redis,
};
ensure_group(&state)
.await
.expect("create test stream group");
let stream_task_id = Uuid::new_v4();
let mut redis_conn = state.redis.clone();
let stream_message_id: String = redis::cmd("XADD")
.arg(STREAM_KEY)
.arg("*")
.arg("task_id")
.arg(stream_task_id.to_string())
.query_async(&mut redis_conn)
.await
.expect("append test stream message");
let stream_owner = format!("test-owner-{}", Uuid::new_v4());
let options = StreamReadOptions::default()
.group(GROUP_NAME, &stream_owner)
.count(100);
let _: redis::streams::StreamReadReply = redis_conn
.xread_options(&[STREAM_KEY], &[">"], &options)
.await
.expect("claim test stream message");
assert!(
!touch_pending_message(&state, "test-intruder", &stream_message_id)
.await
.expect("reject intruder heartbeat")
);
assert!(
touch_pending_message(&state, &stream_owner, &stream_message_id)
.await
.expect("accept owner heartbeat")
);
ack_message(&mut redis_conn, &stream_message_id)
.await
.expect("ack test stream message");
let _: i64 = redis::cmd("XDEL")
.arg(STREAM_KEY)
.arg(&stream_message_id)
.query_async(&mut redis_conn)
.await
.expect("delete test stream message");
let marker = Uuid::new_v4().simple().to_string();
let user_id = Uuid::new_v4();
let task_id = Uuid::new_v4();
let file_id = Uuid::new_v4();
let stale_owner = Uuid::new_v4();
let winning_owner = Uuid::new_v4();
sqlx::query(
"INSERT INTO users (id, email, username, password_hash) VALUES ($1, $2, $3, 'test-only')",
)
.bind(user_id)
.bind(format!("worker-{marker}@example.test"))
.bind(format!("worker_{marker}"))
.execute(&pool)
.await
.expect("insert test user");
sqlx::query(
r#"
INSERT INTO tasks (
id, user_id, status, total_files, total_original_size,
processing_attempt, lease_owner, lease_until
) VALUES ($1, $2, 'processing', 1, 100, 2, $3, NOW() + INTERVAL '5 minutes')
"#,
)
.bind(task_id)
.bind(user_id)
.bind(winning_owner)
.execute(&pool)
.await
.expect("insert test task");
sqlx::query(
r#"
INSERT INTO task_files (
id, task_id, original_name, original_format, output_format,
original_size, status, processing_attempt, lease_owner, lease_until
) VALUES (
$1, $2, 'fence.png', 'png', 'png',
100, 'processing', 2, $3, NOW() + INTERVAL '5 minutes'
)
"#,
)
.bind(file_id)
.bind(task_id)
.bind(winning_owner)
.execute(&pool)
.await
.expect("insert test task file");
let stale_key = storage::result_attempt_key(24, task_id, file_id, 1, 1, "png");
let winning_key = storage::result_attempt_key(24, task_id, file_id, 2, 2, "png");
let stale_object = object_lifecycle::store_tracked_bytes(
&state,
task_id,
Some(file_id),
"result",
&stale_key,
Bytes::from_static(b"stale-attempt"),
"image/png",
)
.await
.expect("store stale attempt object");
let winning_object = object_lifecycle::store_tracked_bytes(
&state,
task_id,
Some(file_id),
"result",
&winning_key,
Bytes::from_static(b"winning-attempt"),
"image/png",
)
.await
.expect("store winning attempt object");
if let Ok(expected_backend) = std::env::var("IMAGEFORGE_TEST_EXPECT_STORAGE_BACKEND") {
assert_eq!(stale_object.stored.backend, expected_backend);
assert_eq!(winning_object.stored.backend, expected_backend);
}
let period_start = Utc::now() - chrono::Duration::hours(1);
let period_end = Utc::now() + chrono::Duration::days(30);
let billing = billing::BillingContext {
user_id,
subscription_id: None,
plan: billing::Plan {
included_units_per_period: 100,
max_file_size_mb: 10,
max_files_per_batch: 10,
retention_days: 1,
feature_api_enabled: false,
},
period_start,
period_end,
};
let stale_fence = FileFence {
task_id,
task_attempt: 1,
file_id,
file_attempt: 1,
worker_id: stale_owner,
};
let winning_fence = FileFence {
task_id,
task_attempt: 2,
file_id,
file_attempt: 2,
worker_id: winning_owner,
};
let barrier = Arc::new(Barrier::new(2));
let stale_join = {
let state = state.clone();
let billing = billing.clone();
let stored = stale_object.clone();
let barrier = barrier.clone();
tokio::spawn(async move {
barrier.wait().await;
finalize_file(
&state,
&Some(billing),
None,
"web",
&stale_fence,
&stored,
100,
50,
50.0,
compress::ImageFmt::Png,
compress::ImageFmt::Png,
true,
)
.await
})
};
let winning_join = {
let state = state.clone();
let billing = billing.clone();
let stored = winning_object.clone();
let barrier = barrier.clone();
tokio::spawn(async move {
barrier.wait().await;
finalize_file(
&state,
&Some(billing),
None,
"web",
&winning_fence,
&stored,
100,
40,
60.0,
compress::ImageFmt::Png,
compress::ImageFmt::Png,
true,
)
.await
})
};
let stale_result = stale_join
.await
.expect("stale finalize task")
.expect("stale finalize");
let winning_result = winning_join
.await
.expect("winning finalize task")
.expect("winning finalize");
assert_eq!(stale_result, FinalizeFileOutcome::LeaseLost);
assert_eq!(winning_result, FinalizeFileOutcome::Committed);
discard_tracked_result(&state, &stale_object, None).await;
assert!(
storage::read_bytes(&state, &stored_locator(&stale_object.stored))
.await
.is_err()
);
assert_eq!(
storage::read_bytes(&state, &stored_locator(&winning_object.stored))
.await
.expect("read winning object"),
b"winning-attempt"
);
assert!(finalize_task_status(&state, task_id, 2, winning_owner)
.await
.expect("finalize test task"));
let task: (String, i32, i32, i64) = sqlx::query_as(
"SELECT status::text, completed_files, failed_files, total_compressed_size FROM tasks WHERE id = $1",
)
.bind(task_id)
.fetch_one(&pool)
.await
.expect("query test task");
assert_eq!(task, ("completed".to_string(), 1, 0, 40));
let file: (String, String, i64) = sqlx::query_as(
"SELECT status::text, storage_key, compressed_size FROM task_files WHERE id = $1",
)
.bind(file_id)
.fetch_one(&pool)
.await
.expect("query test file");
assert_eq!(
file,
(
"completed".to_string(),
winning_object.stored.key.clone(),
40
)
);
let usage_event_count: i64 =
sqlx::query_scalar("SELECT COUNT(*) FROM usage_events WHERE task_file_id = $1")
.bind(file_id)
.fetch_one(&pool)
.await
.expect("count usage events");
assert_eq!(usage_event_count, 1);
let used_units: i32 = sqlx::query_scalar(
"SELECT used_units FROM usage_periods WHERE user_id = $1 AND period_start = $2 AND period_end = $3",
)
.bind(user_id)
.bind(period_start)
.bind(period_end)
.fetch_one(&pool)
.await
.expect("query used units");
assert_eq!(used_units, 1);
discard_tracked_result(&state, &winning_object, None).await;
sqlx::query("DELETE FROM usage_events WHERE task_id = $1")
.bind(task_id)
.execute(&pool)
.await
.expect("delete test usage events");
sqlx::query("DELETE FROM tasks WHERE id = $1")
.bind(task_id)
.execute(&pool)
.await
.expect("delete test task");
sqlx::query("DELETE FROM storage_objects WHERE task_id = $1")
.bind(task_id)
.execute(&pool)
.await
.expect("delete test storage lifecycle rows");
sqlx::query("DELETE FROM usage_periods WHERE user_id = $1")
.bind(user_id)
.execute(&pool)
.await
.expect("delete test usage period");
sqlx::query("DELETE FROM users WHERE id = $1")
.bind(user_id)
.execute(&pool)
.await
.expect("delete test user");
tokio::fs::remove_dir_all(&state.config.storage_path)
.await
.expect("delete test storage directory");
}
}