fix: harden queue retries and quota reservation
This commit is contained in:
@@ -220,7 +220,7 @@ pub enum Permission {
|
||||
|
||||
计量口径见 `docs/billing.md`,架构上建议:
|
||||
- Worker 在每个文件成功输出后写入 `usage_events`(账本明细),并更新 `usage_periods`(按订阅周期聚合)。
|
||||
- API 在创建任务/接收同步压缩请求时做**配额预检**(快速失败),Worker 做**最终扣减**(账本落地,保证一致性)。
|
||||
- 登录用户与 API Key 在创建任务/接收同步压缩请求时做**配额预检**,Worker 成功输出后做最终账本扣减。匿名请求在压缩前按会话和 IP 原子预占每日试用次数;同步处理或任务创建失败时退还,避免超额请求先消耗 CPU、批量并发穿透配额。
|
||||
- 对外 API 强烈建议支持 `Idempotency-Key`;DB 侧存储“幂等记录 + 响应摘要”,避免重复扣减与重复任务。
|
||||
|
||||
### 7. 支付回调(Webhooks)
|
||||
|
||||
@@ -533,9 +533,15 @@ TTL: 48 hours
|
||||
```
|
||||
Stream: stream:compress_jobs
|
||||
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 任务进度(可选)
|
||||
```
|
||||
PubSub: pubsub:task:{task_id}
|
||||
|
||||
2
migrations/007_queue_and_anonymous_quota.sql
Normal file
2
migrations/007_queue_and_anonymous_quota.sql
Normal file
@@ -0,0 +1,2 @@
|
||||
ALTER TABLE tasks
|
||||
ADD COLUMN IF NOT EXISTS anonymous_units_reserved INTEGER NOT NULL DEFAULT 0;
|
||||
@@ -1585,10 +1585,10 @@ async fn test_mail(
|
||||
|
||||
fn mask_secret(secret: &str) -> String {
|
||||
let trimmed = secret.trim();
|
||||
if trimmed.len() <= 8 {
|
||||
if trimmed.chars().count() <= 8 {
|
||||
return trimmed.to_string();
|
||||
}
|
||||
format!("{}...", &trimmed[..8])
|
||||
format!("{}...", trimmed.chars().take(8).collect::<String>())
|
||||
}
|
||||
|
||||
#[derive(Debug, FromRow, Serialize)]
|
||||
@@ -1678,3 +1678,15 @@ async fn update_config(
|
||||
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"), "🔑🔑🔑🔑🔑🔑🔑🔑...");
|
||||
}
|
||||
}
|
||||
|
||||
@@ -263,11 +263,15 @@ async fn compress_json(
|
||||
}
|
||||
}
|
||||
|
||||
let mut anonymous_reserved = false;
|
||||
let op: Result<CompressResponse, AppError> = (async {
|
||||
match "a_ctx {
|
||||
QuotaContext::User(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;
|
||||
@@ -298,7 +302,7 @@ async fn compress_json(
|
||||
&& format_in == format_out
|
||||
&& req.max_width.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 file_id = Uuid::new_v4();
|
||||
@@ -309,23 +313,6 @@ async fn compress_json(
|
||||
storage::store_bytes(&state, &object_key, compressed, format_out.content_type())
|
||||
.await?;
|
||||
|
||||
if charge_units {
|
||||
if let QuotaContext::Anonymous { session_id, ip } = "a_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;
|
||||
|
||||
if let Err(err) = record_task_and_metering(
|
||||
@@ -410,6 +397,11 @@ async fn compress_json(
|
||||
))
|
||||
}
|
||||
Err(err) => {
|
||||
if anonymous_reserved {
|
||||
if let QuotaContext::Anonymous { session_id, ip } = "a_ctx {
|
||||
let _ = quota::refund_anonymous_units(&state, session_id, *ip, 1).await;
|
||||
}
|
||||
}
|
||||
if let (Some(scope), Some(idem_key), Some(request_hash)) = (
|
||||
idempotency_scope,
|
||||
idempotency_key.as_deref(),
|
||||
|
||||
@@ -18,7 +18,7 @@ use chrono::{DateTime, Duration, Utc};
|
||||
use serde::{Deserialize, Serialize};
|
||||
use sha2::{Digest, Sha256};
|
||||
use sqlx::FromRow;
|
||||
use std::net::{IpAddr, SocketAddr};
|
||||
use std::net::SocketAddr;
|
||||
use tokio::io::AsyncWriteExt;
|
||||
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 (retention, task_owner, source) = match &principal {
|
||||
context::Principal::Anonymous { session_id } => {
|
||||
enforce_batch_limits_anonymous(&state, &files)?;
|
||||
let remaining = anonymous_remaining_units(&state, session_id, ip).await?;
|
||||
if remaining < files.len() as i64 {
|
||||
return Err(AppError::new(
|
||||
ErrorCode::QuotaExceeded,
|
||||
"匿名试用次数已用完(每日 10 次)",
|
||||
));
|
||||
}
|
||||
let units = u32::try_from(files.len()).map_err(|_| {
|
||||
AppError::new(ErrorCode::InvalidRequest, "批量文件数量超出限制")
|
||||
})?;
|
||||
quota::consume_anonymous_units(&state, session_id, ip, units).await?;
|
||||
anonymous_reserved_units = units;
|
||||
Ok((
|
||||
Duration::hours(state.config.anon_retention_hours as i64),
|
||||
TaskOwner::Anonymous {
|
||||
@@ -238,13 +237,13 @@ async fn create_batch_task(
|
||||
compression_rate,
|
||||
total_files, completed_files, failed_files,
|
||||
total_original_size, total_compressed_size,
|
||||
expires_at, retention_hours
|
||||
expires_at, retention_hours, anonymous_units_reserved
|
||||
) VALUES (
|
||||
$1, $2, $3, $4, $5::inet, $6::task_source, 'pending',
|
||||
$7::compression_level, $8, $9, $10, $11, $12,
|
||||
$13, 0, 0,
|
||||
$14, 0,
|
||||
$15, $16
|
||||
$15, $16, $17
|
||||
)
|
||||
"#,
|
||||
)
|
||||
@@ -264,6 +263,7 @@ async fn create_batch_task(
|
||||
.bind(total_original_size)
|
||||
.bind(expires_at)
|
||||
.bind(retention_hours as i32)
|
||||
.bind(anonymous_reserved_units as i32)
|
||||
.execute(&mut *tx)
|
||||
.await
|
||||
.map_err(|err| AppError::new(ErrorCode::Internal, "创建任务失败").with_source(err))?;
|
||||
@@ -345,6 +345,17 @@ async fn create_batch_task(
|
||||
))
|
||||
}
|
||||
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 idem_acquired {
|
||||
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();
|
||||
redis::cmd("XADD")
|
||||
.arg("stream:compress_jobs")
|
||||
.arg("MAXLEN")
|
||||
.arg("~")
|
||||
.arg(100_000)
|
||||
.arg("*")
|
||||
.arg("task_id")
|
||||
.arg(task_id.to_string())
|
||||
@@ -651,39 +665,6 @@ async fn ensure_quota_available(
|
||||
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)]
|
||||
struct TaskRow {
|
||||
status: String,
|
||||
|
||||
@@ -52,11 +52,23 @@ async fn stripe_webhook(
|
||||
AppError::new(ErrorCode::InvalidRequest, "Webhook JSON 解析失败").with_source(err)
|
||||
})?;
|
||||
|
||||
let inserted: Option<String> = sqlx::query_scalar(
|
||||
let claimed: Option<String> = sqlx::query_scalar(
|
||||
r#"
|
||||
INSERT INTO webhook_events (provider, provider_event_id, event_type, payload)
|
||||
VALUES ('stripe', $1, $2, $3)
|
||||
ON CONFLICT (provider, provider_event_id) DO NOTHING
|
||||
INSERT INTO webhook_events (
|
||||
provider, provider_event_id, event_type, payload, status
|
||||
) 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
|
||||
"#,
|
||||
)
|
||||
@@ -67,16 +79,29 @@ async fn stripe_webhook(
|
||||
.await
|
||||
.map_err(|err| AppError::new(ErrorCode::Internal, "Webhook 入库失败").with_source(err))?;
|
||||
|
||||
if inserted.is_none() {
|
||||
return Ok(Json(Envelope {
|
||||
success: true,
|
||||
data: serde_json::json!({ "status": "duplicate" }),
|
||||
}));
|
||||
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 {
|
||||
success: true,
|
||||
data: serde_json::json!({ "status": "duplicate" }),
|
||||
}));
|
||||
}
|
||||
return Err(AppError::new(
|
||||
ErrorCode::StorageUnavailable,
|
||||
"Webhook 事件正在处理,请稍后重试",
|
||||
));
|
||||
}
|
||||
|
||||
if let Err(err) = process_stripe_event(&state, &event).await {
|
||||
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(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 {
|
||||
if s.len() > max {
|
||||
s.truncate(max);
|
||||
let mut end = max;
|
||||
while !s.is_char_boundary(end) {
|
||||
end -= 1;
|
||||
}
|
||||
s.truncate(end);
|
||||
}
|
||||
s
|
||||
}
|
||||
@@ -552,3 +581,15 @@ fn map_invoice_status(status: &str) -> &'static str {
|
||||
_ => "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");
|
||||
}
|
||||
}
|
||||
|
||||
@@ -306,6 +306,46 @@ pub async fn consume_anonymous_units(
|
||||
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 {
|
||||
let now = Utc::now() + Duration::hours(8);
|
||||
now.format("%Y-%m-%d").to_string()
|
||||
|
||||
@@ -5,7 +5,7 @@ use crate::services::quota;
|
||||
use crate::services::storage;
|
||||
use crate::state::AppState;
|
||||
|
||||
use redis::streams::StreamReadOptions;
|
||||
use redis::streams::{StreamClaimReply, StreamId, StreamPendingCountReply, StreamReadOptions};
|
||||
use redis::AsyncCommands;
|
||||
use sqlx::FromRow;
|
||||
use std::net::IpAddr;
|
||||
@@ -17,6 +17,9 @@ use uuid::Uuid;
|
||||
|
||||
const STREAM_KEY: &str = "stream:compress_jobs";
|
||||
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> {
|
||||
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> {
|
||||
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.
|
||||
let pending_opts = StreamReadOptions::default()
|
||||
.group(GROUP_NAME, consumer)
|
||||
@@ -97,31 +104,187 @@ async fn poll_once(state: &AppState, consumer: &str) -> Result<(), AppError> {
|
||||
|
||||
for key in reply.keys {
|
||||
for msg in key.ids {
|
||||
let Some(task_id_str) = msg.get::<String>("task_id") else {
|
||||
ack_message(&mut conn, &msg.id).await?;
|
||||
continue;
|
||||
};
|
||||
|
||||
let task_id = match Uuid::parse_str(&task_id_str) {
|
||||
Ok(v) => v,
|
||||
Err(_) => {
|
||||
ack_message(&mut conn, &msg.id).await?;
|
||||
continue;
|
||||
}
|
||||
};
|
||||
|
||||
process_task(state, task_id).await.map_err(|err| {
|
||||
tracing::error!(task_id = %task_id, error = %err, "task processing failed; message left pending for retry");
|
||||
err
|
||||
})?;
|
||||
|
||||
ack_message(&mut conn, &msg.id).await?;
|
||||
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 {
|
||||
ack_message(conn, &msg.id).await?;
|
||||
return Ok(());
|
||||
};
|
||||
let task_id = match Uuid::parse_str(&task_id_str) {
|
||||
Ok(value) => value,
|
||||
Err(_) => {
|
||||
ack_message(conn, &msg.id).await?;
|
||||
return Ok(());
|
||||
}
|
||||
};
|
||||
|
||||
match process_task(state, task_id).await {
|
||||
Ok(()) => ack_message(conn, &msg.id).await,
|
||||
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);
|
||||
};
|
||||
|
||||
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(())
|
||||
}
|
||||
|
||||
async fn ack_message(
|
||||
conn: &mut redis::aio::ConnectionManager,
|
||||
msg_id: &str,
|
||||
@@ -150,6 +313,7 @@ struct TaskProcRow {
|
||||
source: String,
|
||||
client_ip: Option<String>,
|
||||
retention_hours: i32,
|
||||
anonymous_units_reserved: i32,
|
||||
}
|
||||
|
||||
#[derive(Debug, FromRow)]
|
||||
@@ -185,10 +349,11 @@ struct TaskContext {
|
||||
anon_ip: Option<IpAddr>,
|
||||
is_anonymous: bool,
|
||||
retention_hours: i32,
|
||||
anonymous_quota_reserved: bool,
|
||||
}
|
||||
|
||||
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#"
|
||||
SELECT
|
||||
status::text AS status,
|
||||
@@ -202,7 +367,8 @@ async fn process_task(state: &AppState, task_id: Uuid) -> Result<(), AppError> {
|
||||
api_key_id,
|
||||
source::text AS source,
|
||||
host(client_ip) AS client_ip,
|
||||
retention_hours
|
||||
retention_hours,
|
||||
anonymous_units_reserved
|
||||
FROM tasks
|
||||
WHERE id = $1
|
||||
"#,
|
||||
@@ -211,7 +377,10 @@ async fn process_task(state: &AppState, task_id: Uuid) -> Result<(), AppError> {
|
||||
.fetch_optional(&state.db)
|
||||
.await
|
||||
.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") {
|
||||
return Ok(());
|
||||
@@ -293,6 +462,7 @@ async fn process_task(state: &AppState, task_id: Uuid) -> Result<(), AppError> {
|
||||
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;
|
||||
@@ -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 _ = storage::delete_object(&state, &stored_locator(&stored)).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(),
|
||||
}
|
||||
}
|
||||
|
||||
#[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));
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user