diff --git a/docs/architecture.md b/docs/architecture.md index 15be66f..393e062 100644 --- a/docs/architecture.md +++ b/docs/architecture.md @@ -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) diff --git a/docs/database.md b/docs/database.md index 82d5b41..7266148 100644 --- a/docs/database.md +++ b/docs/database.md @@ -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} diff --git a/migrations/007_queue_and_anonymous_quota.sql b/migrations/007_queue_and_anonymous_quota.sql new file mode 100644 index 0000000..503230e --- /dev/null +++ b/migrations/007_queue_and_anonymous_quota.sql @@ -0,0 +1,2 @@ +ALTER TABLE tasks + ADD COLUMN IF NOT EXISTS anonymous_units_reserved INTEGER NOT NULL DEFAULT 0; diff --git a/src/api/admin.rs b/src/api/admin.rs index 9192578..8d9b96d 100644 --- a/src/api/admin.rs +++ b/src/api/admin.rs @@ -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::()) } #[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"), "🔑🔑🔑🔑🔑🔑🔑🔑..."); + } +} diff --git a/src/api/compress.rs b/src/api/compress.rs index ed7747d..59644ae 100644 --- a/src/api/compress.rs +++ b/src/api/compress.rs @@ -263,11 +263,15 @@ async fn compress_json( } } + let mut anonymous_reserved = false; let op: Result = (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(), diff --git a/src/api/tasks.rs b/src/api/tasks.rs index a6a3d28..b17b33c 100644 --- a/src/api/tasks.rs +++ b/src/api/tasks.rs @@ -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 = (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 { - 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 = redis::cmd("GET") - .arg(session_key) - .query_async(&mut conn) - .await - .unwrap_or(None); - let v2: Option = 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, diff --git a/src/api/webhooks.rs b/src/api/webhooks.rs index 3f5980d..1ce69b1 100644 --- a/src/api/webhooks.rs +++ b/src/api/webhooks.rs @@ -52,11 +52,23 @@ async fn stripe_webhook( AppError::new(ErrorCode::InvalidRequest, "Webhook JSON 解析失败").with_source(err) })?; - let inserted: Option = sqlx::query_scalar( + let claimed: Option = 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 = 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"); + } +} diff --git a/src/services/quota.rs b/src/services/quota.rs index a7e71fd..cc57381 100644 --- a/src/services/quota.rs +++ b/src/services/quota.rs @@ -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() diff --git a/src/worker/mod.rs b/src/worker/mod.rs index 80dc407..2f39421 100644 --- a/src/worker/mod.rs +++ b/src/worker/mod.rs @@ -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::("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::("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, 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 { + 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, retention_hours: i32, + anonymous_units_reserved: i32, } #[derive(Debug, FromRow)] @@ -185,10 +349,11 @@ struct TaskContext { anon_ip: Option, 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 = 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)); + } +}