fix: harden queue retries and quota reservation
This commit is contained in:
@@ -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)
|
||||||
|
|||||||
@@ -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}
|
||||||
|
|||||||
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 {
|
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"), "🔑🔑🔑🔑🔑🔑🔑🔑...");
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
@@ -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 "a_ctx {
|
match "a_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 } = "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;
|
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 } = "a_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(),
|
||||||
|
|||||||
@@ -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,
|
||||||
|
|||||||
@@ -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() {
|
||||||
return Ok(Json(Envelope {
|
let status: Option<String> = sqlx::query_scalar(
|
||||||
success: true,
|
"SELECT status FROM webhook_events WHERE provider = 'stripe' AND provider_event_id = $1",
|
||||||
data: serde_json::json!({ "status": "duplicate" }),
|
)
|
||||||
}));
|
.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 {
|
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");
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
@@ -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()
|
||||||
|
|||||||
@@ -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,31 +104,187 @@ 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 {
|
||||||
let Some(task_id_str) = msg.get::<String>("task_id") else {
|
handle_message(state, &mut conn, msg).await?;
|
||||||
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?;
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
Ok(())
|
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(
|
async fn ack_message(
|
||||||
conn: &mut redis::aio::ConnectionManager,
|
conn: &mut redis::aio::ConnectionManager,
|
||||||
msg_id: &str,
|
msg_id: &str,
|
||||||
@@ -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));
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user