fix: harden queue retries and quota reservation

This commit is contained in:
237899745
2026-07-25 18:15:39 +08:00
parent 091a0cee75
commit f389dfb567
9 changed files with 358 additions and 101 deletions

View File

@@ -220,7 +220,7 @@ pub enum Permission {
计量口径见 `docs/billing.md`,架构上建议: 计量口径见 `docs/billing.md`,架构上建议:
- Worker 在每个文件成功输出后写入 `usage_events`(账本明细),并更新 `usage_periods`(按订阅周期聚合)。 - Worker 在每个文件成功输出后写入 `usage_events`(账本明细),并更新 `usage_periods`(按订阅周期聚合)。
- API 在创建任务/接收同步压缩请求时做**配额预检**(快速失败)Worker 做**最终扣减**(账本落地,保证一致性) - 登录用户与 API Key 在创建任务/接收同步压缩请求时做**配额预检**Worker 成功输出后做最终账本扣减。匿名请求在压缩前按会话和 IP 原子预占每日试用次数;同步处理或任务创建失败时退还,避免超额请求先消耗 CPU、批量并发穿透配额
- 对外 API 强烈建议支持 `Idempotency-Key`DB 侧存储“幂等记录 + 响应摘要”,避免重复扣减与重复任务。 - 对外 API 强烈建议支持 `Idempotency-Key`DB 侧存储“幂等记录 + 响应摘要”,避免重复扣减与重复任务。
### 7. 支付回调Webhooks ### 7. 支付回调Webhooks

View File

@@ -533,9 +533,15 @@ TTL: 48 hours
``` ```
Stream: stream:compress_jobs Stream: stream:compress_jobs
Group: compress_workers Group: compress_workers
Message fields: { task_id, priority, created_at } Message fields: { task_id, created_at }
Approximate max length: 100000
Dead-letter stream: stream:compress_jobs:dead
Dead-letter max length: 10000
``` ```
Worker 每次优先处理本消费者 pending并认领空闲超过 5 分钟的其他消费者消息。消息第 3 次投递仍失败时写入死信流、将任务标记为失败并 ACK 原消息;已删除或已进入终态的任务直接 ACK。
### 5.4 任务进度(可选) ### 5.4 任务进度(可选)
``` ```
PubSub: pubsub:task:{task_id} PubSub: pubsub:task:{task_id}

View File

@@ -0,0 +1,2 @@
ALTER TABLE tasks
ADD COLUMN IF NOT EXISTS anonymous_units_reserved INTEGER NOT NULL DEFAULT 0;

View File

@@ -1585,10 +1585,10 @@ async fn test_mail(
fn mask_secret(secret: &str) -> String { fn mask_secret(secret: &str) -> String {
let trimmed = secret.trim(); let trimmed = secret.trim();
if trimmed.len() <= 8 { if trimmed.chars().count() <= 8 {
return trimmed.to_string(); return trimmed.to_string();
} }
format!("{}...", &trimmed[..8]) format!("{}...", trimmed.chars().take(8).collect::<String>())
} }
#[derive(Debug, FromRow, Serialize)] #[derive(Debug, FromRow, Serialize)]
@@ -1678,3 +1678,15 @@ async fn update_config(
data: row, data: row,
})) }))
} }
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn secret_masking_never_splits_utf8() {
assert_eq!(mask_secret("12345678abcdef"), "12345678...");
assert_eq!(mask_secret("中文密钥测试内容"), "中文密钥测试内容");
assert_eq!(mask_secret("🔑🔑🔑🔑🔑🔑🔑🔑more"), "🔑🔑🔑🔑🔑🔑🔑🔑...");
}
}

View File

@@ -263,11 +263,15 @@ async fn compress_json(
} }
} }
let mut anonymous_reserved = false;
let op: Result<CompressResponse, AppError> = (async { let op: Result<CompressResponse, AppError> = (async {
match &quota_ctx { match &quota_ctx {
QuotaContext::User(billing) => ensure_quota_available(&state, billing, 1).await?, QuotaContext::User(billing) => ensure_quota_available(&state, billing, 1).await?,
QuotaContext::ApiKey(billing, _) => ensure_quota_available(&state, billing, 1).await?, QuotaContext::ApiKey(billing, _) => ensure_quota_available(&state, billing, 1).await?,
QuotaContext::Anonymous { .. } => {} QuotaContext::Anonymous { session_id, ip } => {
quota::consume_anonymous_units(&state, session_id, *ip, 1).await?;
anonymous_reserved = true;
}
} }
let original_size = req.file_bytes.len() as u64; let original_size = req.file_bytes.len() as u64;
@@ -298,7 +302,7 @@ async fn compress_json(
&& format_in == format_out && format_in == format_out
&& req.max_width.is_none() && req.max_width.is_none()
&& req.max_height.is_none(); && req.max_height.is_none();
let charge_units = !skip_charge && compressed_size < original_size; let charge_units = anonymous_reserved || (!skip_charge && compressed_size < original_size);
let task_id = Uuid::new_v4(); let task_id = Uuid::new_v4();
let file_id = Uuid::new_v4(); let file_id = Uuid::new_v4();
@@ -309,23 +313,6 @@ async fn compress_json(
storage::store_bytes(&state, &object_key, compressed, format_out.content_type()) storage::store_bytes(&state, &object_key, compressed, format_out.content_type())
.await?; .await?;
if charge_units {
if let QuotaContext::Anonymous { session_id, ip } = &quota_ctx {
if let Err(err) = quota::consume_anonymous_units(&state, session_id, *ip, 1).await {
let _ = storage::delete_object(
&state,
&storage::ObjectLocator {
backend: stored.backend.clone(),
endpoint_id: stored.endpoint_id,
key: stored.key.clone(),
},
)
.await;
return Err(err);
}
}
}
let expires_at = Utc::now() + retention; let expires_at = Utc::now() + retention;
if let Err(err) = record_task_and_metering( if let Err(err) = record_task_and_metering(
@@ -410,6 +397,11 @@ async fn compress_json(
)) ))
} }
Err(err) => { Err(err) => {
if anonymous_reserved {
if let QuotaContext::Anonymous { session_id, ip } = &quota_ctx {
let _ = quota::refund_anonymous_units(&state, session_id, *ip, 1).await;
}
}
if let (Some(scope), Some(idem_key), Some(request_hash)) = ( if let (Some(scope), Some(idem_key), Some(request_hash)) = (
idempotency_scope, idempotency_scope,
idempotency_key.as_deref(), idempotency_key.as_deref(),

View File

@@ -18,7 +18,7 @@ use chrono::{DateTime, Duration, Utc};
use serde::{Deserialize, Serialize}; use serde::{Deserialize, Serialize};
use sha2::{Digest, Sha256}; use sha2::{Digest, Sha256};
use sqlx::FromRow; use sqlx::FromRow;
use std::net::{IpAddr, SocketAddr}; use std::net::SocketAddr;
use tokio::io::AsyncWriteExt; use tokio::io::AsyncWriteExt;
use uuid::Uuid; use uuid::Uuid;
@@ -150,17 +150,16 @@ async fn create_batch_task(
} }
} }
let mut anonymous_reserved_units = 0u32;
let create_result: Result<BatchCreateResponse, AppError> = (async { let create_result: Result<BatchCreateResponse, AppError> = (async {
let (retention, task_owner, source) = match &principal { let (retention, task_owner, source) = match &principal {
context::Principal::Anonymous { session_id } => { context::Principal::Anonymous { session_id } => {
enforce_batch_limits_anonymous(&state, &files)?; enforce_batch_limits_anonymous(&state, &files)?;
let remaining = anonymous_remaining_units(&state, session_id, ip).await?; let units = u32::try_from(files.len()).map_err(|_| {
if remaining < files.len() as i64 { AppError::new(ErrorCode::InvalidRequest, "批量文件数量超出限制")
return Err(AppError::new( })?;
ErrorCode::QuotaExceeded, quota::consume_anonymous_units(&state, session_id, ip, units).await?;
"匿名试用次数已用完(每日 10 次)", anonymous_reserved_units = units;
));
}
Ok(( Ok((
Duration::hours(state.config.anon_retention_hours as i64), Duration::hours(state.config.anon_retention_hours as i64),
TaskOwner::Anonymous { TaskOwner::Anonymous {
@@ -238,13 +237,13 @@ async fn create_batch_task(
compression_rate, compression_rate,
total_files, completed_files, failed_files, total_files, completed_files, failed_files,
total_original_size, total_compressed_size, total_original_size, total_compressed_size,
expires_at, retention_hours expires_at, retention_hours, anonymous_units_reserved
) VALUES ( ) VALUES (
$1, $2, $3, $4, $5::inet, $6::task_source, 'pending', $1, $2, $3, $4, $5::inet, $6::task_source, 'pending',
$7::compression_level, $8, $9, $10, $11, $12, $7::compression_level, $8, $9, $10, $11, $12,
$13, 0, 0, $13, 0, 0,
$14, 0, $14, 0,
$15, $16 $15, $16, $17
) )
"#, "#,
) )
@@ -264,6 +263,7 @@ async fn create_batch_task(
.bind(total_original_size) .bind(total_original_size)
.bind(expires_at) .bind(expires_at)
.bind(retention_hours as i32) .bind(retention_hours as i32)
.bind(anonymous_reserved_units as i32)
.execute(&mut *tx) .execute(&mut *tx)
.await .await
.map_err(|err| AppError::new(ErrorCode::Internal, "创建任务失败").with_source(err))?; .map_err(|err| AppError::new(ErrorCode::Internal, "创建任务失败").with_source(err))?;
@@ -345,6 +345,17 @@ async fn create_batch_task(
)) ))
} }
Err(err) => { Err(err) => {
if anonymous_reserved_units > 0 {
if let context::Principal::Anonymous { session_id } = &principal {
let _ = quota::refund_anonymous_units(
&state,
session_id,
ip,
anonymous_reserved_units,
)
.await;
}
}
if let (Some(scope), Some(idem_key)) = (idempotency_scope, idempotency_key.as_deref()) { if let (Some(scope), Some(idem_key)) = (idempotency_scope, idempotency_key.as_deref()) {
if idem_acquired { if idem_acquired {
let _ = idempotency::abort(&state, scope, idem_key, &request_hash).await; let _ = idempotency::abort(&state, scope, idem_key, &request_hash).await;
@@ -368,6 +379,9 @@ async fn enqueue_task(state: &AppState, task_id: Uuid) -> Result<(), AppError> {
let now = Utc::now().to_rfc3339(); let now = Utc::now().to_rfc3339();
redis::cmd("XADD") redis::cmd("XADD")
.arg("stream:compress_jobs") .arg("stream:compress_jobs")
.arg("MAXLEN")
.arg("~")
.arg(100_000)
.arg("*") .arg("*")
.arg("task_id") .arg("task_id")
.arg(task_id.to_string()) .arg(task_id.to_string())
@@ -651,39 +665,6 @@ async fn ensure_quota_available(
quota::ensure_user_units(state, ctx, needed_units).await quota::ensure_user_units(state, ctx, needed_units).await
} }
async fn anonymous_remaining_units(
state: &AppState,
session_id: &str,
ip: IpAddr,
) -> Result<i64, AppError> {
let date = utc8_date();
let session_key = format!("anon_quota:{session_id}:{date}");
let ip_key = format!("anon_quota_ip:{ip}:{date}");
let mut conn = state.redis.clone();
let v1: Option<i64> = redis::cmd("GET")
.arg(session_key)
.query_async(&mut conn)
.await
.unwrap_or(None);
let v2: Option<i64> = redis::cmd("GET")
.arg(ip_key)
.query_async(&mut conn)
.await
.unwrap_or(None);
let limit = state.config.anon_daily_units as i64;
Ok(std::cmp::min(
limit - v1.unwrap_or(0),
limit - v2.unwrap_or(0),
))
}
fn utc8_date() -> String {
let now = Utc::now() + Duration::hours(8);
now.format("%Y-%m-%d").to_string()
}
#[derive(Debug, FromRow)] #[derive(Debug, FromRow)]
struct TaskRow { struct TaskRow {
status: String, status: String,

View File

@@ -52,11 +52,23 @@ async fn stripe_webhook(
AppError::new(ErrorCode::InvalidRequest, "Webhook JSON 解析失败").with_source(err) AppError::new(ErrorCode::InvalidRequest, "Webhook JSON 解析失败").with_source(err)
})?; })?;
let inserted: Option<String> = sqlx::query_scalar( let claimed: Option<String> = sqlx::query_scalar(
r#" r#"
INSERT INTO webhook_events (provider, provider_event_id, event_type, payload) INSERT INTO webhook_events (
VALUES ('stripe', $1, $2, $3) provider, provider_event_id, event_type, payload, status
ON CONFLICT (provider, provider_event_id) DO NOTHING ) VALUES ('stripe', $1, $2, $3, 'processing')
ON CONFLICT (provider, provider_event_id) DO UPDATE
SET event_type = EXCLUDED.event_type,
payload = EXCLUDED.payload,
received_at = NOW(),
processed_at = NULL,
status = 'processing',
error_message = NULL
WHERE webhook_events.status IN ('received', 'failed')
OR (
webhook_events.status = 'processing'
AND webhook_events.received_at < NOW() - INTERVAL '5 minutes'
)
RETURNING provider_event_id RETURNING provider_event_id
"#, "#,
) )
@@ -67,16 +79,29 @@ async fn stripe_webhook(
.await .await
.map_err(|err| AppError::new(ErrorCode::Internal, "Webhook 入库失败").with_source(err))?; .map_err(|err| AppError::new(ErrorCode::Internal, "Webhook 入库失败").with_source(err))?;
if inserted.is_none() { if claimed.is_none() {
let status: Option<String> = sqlx::query_scalar(
"SELECT status FROM webhook_events WHERE provider = 'stripe' AND provider_event_id = $1",
)
.bind(&event.id)
.fetch_optional(&state.db)
.await
.map_err(|err| AppError::new(ErrorCode::Internal, "Webhook 状态查询失败").with_source(err))?;
if status.as_deref() == Some("processed") {
return Ok(Json(Envelope { return Ok(Json(Envelope {
success: true, success: true,
data: serde_json::json!({ "status": "duplicate" }), data: serde_json::json!({ "status": "duplicate" }),
})); }));
} }
return Err(AppError::new(
ErrorCode::StorageUnavailable,
"Webhook 事件正在处理,请稍后重试",
));
}
if let Err(err) = process_stripe_event(&state, &event).await { if let Err(err) = process_stripe_event(&state, &event).await {
let _ = sqlx::query( let _ = sqlx::query(
"UPDATE webhook_events SET status = 'failed', error_message = $2, processed_at = NOW() WHERE provider = 'stripe' AND provider_event_id = $1", "UPDATE webhook_events SET status = 'failed', error_message = $2, processed_at = NULL WHERE provider = 'stripe' AND provider_event_id = $1 AND status = 'processing'",
) )
.bind(&event.id) .bind(&event.id)
.bind(err.to_string()) .bind(err.to_string())
@@ -528,7 +553,11 @@ async fn upsert_invoice(state: &AppState, object: &serde_json::Value) -> Result<
fn truncate(mut s: String, max: usize) -> String { fn truncate(mut s: String, max: usize) -> String {
if s.len() > max { if s.len() > max {
s.truncate(max); let mut end = max;
while !s.is_char_boundary(end) {
end -= 1;
}
s.truncate(end);
} }
s s
} }
@@ -552,3 +581,15 @@ fn map_invoice_status(status: &str) -> &'static str {
_ => "open", _ => "open",
} }
} }
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn truncate_preserves_utf8_boundaries() {
assert_eq!(truncate("中文测试".to_string(), 5), "");
assert_eq!(truncate("abc中文".to_string(), 5), "abc");
assert_eq!(truncate("short".to_string(), 20), "short");
}
}

View File

@@ -306,6 +306,46 @@ pub async fn consume_anonymous_units(
Ok(()) Ok(())
} }
pub async fn refund_anonymous_units(
state: &AppState,
session_id: &str,
ip: IpAddr,
units: u32,
) -> Result<(), AppError> {
if units == 0 {
return Ok(());
}
let date = utc8_date();
let session_key = format!("anon_quota:{session_id}:{date}");
let ip_key = format!("anon_quota_ip:{ip}:{date}");
let mut conn = state.redis.clone();
let script = redis::Script::new(
r#"
local dec = tonumber(ARGV[1])
local function refund(key)
local current = tonumber(redis.call('GET', key) or '0')
if current <= 0 then return 0 end
return redis.call('DECRBY', key, math.min(current, dec))
end
refund(KEYS[1])
refund(KEYS[2])
return 1
"#,
);
let _: i64 = script
.key(session_key)
.key(ip_key)
.arg(units as i64)
.invoke_async(&mut conn)
.await
.map_err(|err| AppError::new(ErrorCode::Internal, "退还匿名配额失败").with_source(err))?;
Ok(())
}
fn utc8_date() -> String { fn utc8_date() -> String {
let now = Utc::now() + Duration::hours(8); let now = Utc::now() + Duration::hours(8);
now.format("%Y-%m-%d").to_string() now.format("%Y-%m-%d").to_string()

View File

@@ -5,7 +5,7 @@ use crate::services::quota;
use crate::services::storage; use crate::services::storage;
use crate::state::AppState; use crate::state::AppState;
use redis::streams::StreamReadOptions; use redis::streams::{StreamClaimReply, StreamId, StreamPendingCountReply, StreamReadOptions};
use redis::AsyncCommands; use redis::AsyncCommands;
use sqlx::FromRow; use sqlx::FromRow;
use std::net::IpAddr; use std::net::IpAddr;
@@ -17,6 +17,9 @@ use uuid::Uuid;
const STREAM_KEY: &str = "stream:compress_jobs"; const STREAM_KEY: &str = "stream:compress_jobs";
const GROUP_NAME: &str = "compress_workers"; const GROUP_NAME: &str = "compress_workers";
const DEAD_STREAM_KEY: &str = "stream:compress_jobs:dead";
const MAX_DELIVERIES: usize = 3;
const STALE_MESSAGE_IDLE_MS: usize = 5 * 60 * 1000;
pub async fn run(state: AppState) -> Result<(), AppError> { pub async fn run(state: AppState) -> Result<(), AppError> {
tracing::info!("Worker started"); tracing::info!("Worker started");
@@ -70,6 +73,10 @@ async fn ensure_group(state: &AppState, _consumer: &str) -> Result<(), AppError>
async fn poll_once(state: &AppState, consumer: &str) -> Result<(), AppError> { async fn poll_once(state: &AppState, consumer: &str) -> Result<(), AppError> {
let mut conn = state.redis.clone(); let mut conn = state.redis.clone();
if let Some(msg) = claim_stale_message(&mut conn, consumer).await? {
return handle_message(state, &mut conn, msg).await;
}
// Retry messages already delivered to this consumer before taking new work. // Retry messages already delivered to this consumer before taking new work.
let pending_opts = StreamReadOptions::default() let pending_opts = StreamReadOptions::default()
.group(GROUP_NAME, consumer) .group(GROUP_NAME, consumer)
@@ -97,28 +104,184 @@ async fn poll_once(state: &AppState, consumer: &str) -> Result<(), AppError> {
for key in reply.keys { for key in reply.keys {
for msg in key.ids { for msg in key.ids {
handle_message(state, &mut conn, msg).await?;
}
}
Ok(())
}
async fn handle_message(
state: &AppState,
conn: &mut redis::aio::ConnectionManager,
msg: StreamId,
) -> Result<(), AppError> {
let Some(task_id_str) = msg.get::<String>("task_id") else { let Some(task_id_str) = msg.get::<String>("task_id") else {
ack_message(&mut conn, &msg.id).await?; ack_message(conn, &msg.id).await?;
continue; return Ok(());
}; };
let task_id = match Uuid::parse_str(&task_id_str) { let task_id = match Uuid::parse_str(&task_id_str) {
Ok(v) => v, Ok(value) => value,
Err(_) => { Err(_) => {
ack_message(&mut conn, &msg.id).await?; ack_message(conn, &msg.id).await?;
continue; return Ok(());
} }
}; };
process_task(state, task_id).await.map_err(|err| { match process_task(state, task_id).await {
tracing::error!(task_id = %task_id, error = %err, "task processing failed; message left pending for retry"); Ok(()) => ack_message(conn, &msg.id).await,
err Err(err) => {
let deliveries = pending_delivery_count(conn, &msg.id).await?;
if should_dead_letter(deliveries) {
write_dead_letter(conn, &msg.id, task_id, deliveries, &err).await?;
mark_task_dead_letter(state, task_id, &err.message).await?;
ack_message(conn, &msg.id).await?;
tracing::error!(
task_id = %task_id,
deliveries,
error = %err,
"task moved to dead-letter stream"
);
return Ok(());
}
tracing::warn!(
task_id = %task_id,
deliveries,
error = %err,
"task processing failed; message left pending for retry"
);
Err(err)
}
}
}
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);
};
ack_message(&mut conn, &msg.id).await?; 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
}
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,
message: &str,
) -> Result<(), AppError> {
let mut tx =
state.db.begin().await.map_err(|err| {
AppError::new(ErrorCode::Internal, "开启死信事务失败").with_source(err)
})?;
sqlx::query(
r#"
UPDATE task_files
SET status = 'failed',
error_message = $2,
completed_at = NOW()
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',
failed_files = GREATEST(total_files - completed_files, failed_files),
error_message = $2,
completed_at = NOW()
WHERE id = $1 AND status NOT IN ('completed', 'failed', 'cancelled')
"#,
)
.bind(task_id)
.bind(message)
.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(()) Ok(())
} }
@@ -150,6 +313,7 @@ struct TaskProcRow {
source: String, source: String,
client_ip: Option<String>, client_ip: Option<String>,
retention_hours: i32, retention_hours: i32,
anonymous_units_reserved: i32,
} }
#[derive(Debug, FromRow)] #[derive(Debug, FromRow)]
@@ -185,10 +349,11 @@ struct TaskContext {
anon_ip: Option<IpAddr>, anon_ip: Option<IpAddr>,
is_anonymous: bool, is_anonymous: bool,
retention_hours: i32, retention_hours: i32,
anonymous_quota_reserved: bool,
} }
async fn process_task(state: &AppState, task_id: Uuid) -> Result<(), AppError> { async fn process_task(state: &AppState, task_id: Uuid) -> Result<(), AppError> {
let mut task: TaskProcRow = sqlx::query_as( let Some(mut task): Option<TaskProcRow> = sqlx::query_as(
r#" r#"
SELECT SELECT
status::text AS status, status::text AS status,
@@ -202,7 +367,8 @@ async fn process_task(state: &AppState, task_id: Uuid) -> Result<(), AppError> {
api_key_id, api_key_id,
source::text AS source, source::text AS source,
host(client_ip) AS client_ip, host(client_ip) AS client_ip,
retention_hours retention_hours,
anonymous_units_reserved
FROM tasks FROM tasks
WHERE id = $1 WHERE id = $1
"#, "#,
@@ -211,7 +377,10 @@ async fn process_task(state: &AppState, task_id: Uuid) -> Result<(), AppError> {
.fetch_optional(&state.db) .fetch_optional(&state.db)
.await .await
.map_err(|err| AppError::new(ErrorCode::Internal, "查询任务失败").with_source(err))? .map_err(|err| AppError::new(ErrorCode::Internal, "查询任务失败").with_source(err))?
.ok_or_else(|| AppError::new(ErrorCode::NotFound, "任务不存在"))?; else {
tracing::info!(task_id = %task_id, "task was deleted before processing; acknowledging message");
return Ok(());
};
if matches!(task.status.as_str(), "completed" | "failed" | "cancelled") { if matches!(task.status.as_str(), "completed" | "failed" | "cancelled") {
return Ok(()); return Ok(());
@@ -293,6 +462,7 @@ async fn process_task(state: &AppState, task_id: Uuid) -> Result<(), AppError> {
anon_ip, anon_ip,
is_anonymous: task.user_id.is_none(), is_anonymous: task.user_id.is_none(),
retention_hours: task.retention_hours, retention_hours: task.retention_hours,
anonymous_quota_reserved: task.anonymous_units_reserved > 0,
}; };
let concurrency = state.config.worker_concurrency.max(1) as usize; let concurrency = state.config.worker_concurrency.max(1) as usize;
@@ -503,7 +673,7 @@ async fn process_task_file(
} }
}; };
if ctx.is_anonymous && charge_units { if ctx.is_anonymous && charge_units && !ctx.anonymous_quota_reserved {
let Some(session_id) = ctx.session_id.as_deref() else { let Some(session_id) = ctx.session_id.as_deref() else {
let _ = storage::delete_object(&state, &stored_locator(&stored)).await; let _ = storage::delete_object(&state, &stored_locator(&stored)).await;
mark_file_failed(&state, task_id, file.id, "匿名任务缺少 session_id").await?; mark_file_failed(&state, task_id, file.id, "匿名任务缺少 session_id").await?;
@@ -994,3 +1164,16 @@ fn stored_locator(stored: &storage::StoredObject) -> storage::ObjectLocator {
key: stored.key.clone(), key: stored.key.clone(),
} }
} }
#[cfg(test)]
mod tests {
use super::*;
#[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));
}
}