合并 master 并解决鉴权与生成任务冲突
Project CI / Repository checks (pull_request) Successful in 3m0s
Project CI / Frontend tests (pull_request) Successful in 3m50s
Project CI / Native shell tests (pull_request) Successful in 16m31s
Project CI / Backend tests (pull_request) Successful in 19m35s

合并 master 的外部生成历史清理、任务内存上限和备份 OOM 兜底
保留退款 outbox 与短期认证 typed projection/CAS 及容量清理
合并 SpacetimeDB 生成 bindings、架构文档和决策记录
补齐认证投影容量测试字段并通过定向门禁
This commit is contained in:
2026-08-27 21:59:51 +08:00
27 changed files with 1772 additions and 153 deletions
@@ -14,6 +14,7 @@ use spacetime_client::{
ExternalGenerationJobRenewLeaseRecordInput, ExternalGenerationQueueWakeSubscription,
};
use tokio::{
sync::{OwnedSemaphorePermit, Semaphore},
task::{JoinHandle, JoinSet},
time::sleep,
};
@@ -92,6 +93,10 @@ pub(crate) async fn run_external_generation_worker(state: AppState) -> Result<()
let concurrency = state.config.external_generation_worker_concurrency.max(1);
let poll_interval = state.config.external_generation_worker_poll_interval;
let lease = state.config.external_generation_worker_lease;
// 超时任务不能立即取消(在途 procedure 仍可能写回),因此执行容量必须同时
// 约束 active 与 detached work;否则每次超时都会释放 tasks 槽位,实际内存占用
// 会超过配置并发。
let work_slots = std::sync::Arc::new(Semaphore::new(concurrency));
let mut tasks = JoinSet::new();
let mut shutdown = external_generation_worker_shutdown_signal();
let mut queue_wake = None;
@@ -113,16 +118,29 @@ pub(crate) async fn run_external_generation_worker(state: AppState) -> Result<()
);
loop {
// 持续有队列任务时不会进入等待分支,因此必须在每轮主动回收已完成的
// JoinHandle;否则 permit 虽已归还,JoinSet 仍会保留每个历史任务的句柄。
reap_finished_external_generation_worker_tasks(&mut tasks);
ensure_external_generation_queue_wake_subscription(&state, &mut queue_wake).await;
while tasks.len() >= concurrency {
if await_worker_task_or_shutdown(&mut tasks, &mut shutdown).await {
drain_external_generation_worker_tasks(&mut tasks).await;
return Ok(());
while work_slots.available_permits() == 0 {
tokio::select! {
_ = shutdown.as_mut() => {
drain_external_generation_worker_tasks(&mut tasks).await;
return Ok(());
}
permit = work_slots.clone().acquire_owned() => {
if permit.is_err() {
drain_external_generation_worker_tasks(&mut tasks).await;
return Ok(());
}
// 只用 acquire 作为容量变化唤醒信号,许可立即归还;真正领取任务
// 时在下方按返回的 job 数量逐个 try_acquire。
}
}
}
let available = concurrency.saturating_sub(tasks.len()).max(1);
let available = work_slots.available_permits().max(1);
let now_micros = current_utc_micros();
let lease_expires_at_micros = now_micros.saturating_add(duration_micros_i64(lease));
@@ -178,9 +196,13 @@ pub(crate) async fn run_external_generation_worker(state: AppState) -> Result<()
for job in jobs {
let state = state.clone();
let worker_id = worker_id.clone();
let permit = work_slots
.clone()
.try_acquire_owned()
.expect("claimed job must have an execution capacity permit");
tasks.spawn(async move {
if let Err(error) =
process_external_generation_job(state, worker_id, lease, job).await
process_external_generation_job(state, worker_id, lease, job, permit).await
{
error!(error = %error, "external generation worker 执行任务失败");
}
@@ -255,13 +277,11 @@ async fn await_worker_task(tasks: &mut JoinSet<()>) {
}
}
async fn await_worker_task_or_shutdown(
tasks: &mut JoinSet<()>,
shutdown: &mut ExternalGenerationShutdownSignal,
) -> bool {
tokio::select! {
_ = shutdown.as_mut() => true,
_ = await_worker_task(tasks) => false,
fn reap_finished_external_generation_worker_tasks(tasks: &mut JoinSet<()>) {
while let Some(result) = tasks.try_join_next() {
if let Err(error) = result {
error!(error = %error, "external generation worker 子任务 panic");
}
}
}
@@ -326,6 +346,7 @@ async fn process_external_generation_job(
worker_id: String,
lease: Duration,
job: ExternalGenerationJobRecord,
permit: OwnedSemaphorePermit,
) -> Result<(), String> {
let heartbeat_interval = external_generation_worker_heartbeat_interval(lease);
let job_timeout = external_generation_worker_job_timeout(&state.config, job.job_kind.as_str());
@@ -377,13 +398,14 @@ async fn process_external_generation_job(
job_id = %job.job_id,
job_kind = %job.job_kind,
timeout_seconds = job_timeout.as_secs(),
"external generation worker 任务超过执行预算,停止续租并释放 worker 槽位,在途执行交由租约仲裁"
"external generation worker 任务超过执行预算,停止续租并保留 worker 槽位,在途执行交由租约仲裁"
);
detach_external_generation_work_until_lease_expiry(
work_handle,
&job,
lease,
"任务超过执行预算",
Some(permit),
);
Err(message)
}
@@ -393,6 +415,7 @@ async fn process_external_generation_job(
&job,
lease,
"任务租约续期失败",
Some(permit),
);
Err(error)
}
@@ -417,6 +440,7 @@ fn detach_external_generation_work_until_lease_expiry(
job: &ExternalGenerationJobRecord,
lease: Duration,
reason: &'static str,
permit: Option<OwnedSemaphorePermit>,
) {
let job_id = job.job_id.clone();
let job_kind = job.job_kind.clone();
@@ -445,6 +469,9 @@ fn detach_external_generation_work_until_lease_expiry(
),
Err(_) => {
work_handle.abort();
// 仅调用 abort 不会从 JoinHandle/JoinSet 中消费完成结果;等待被取消
// 的 handle,确保 permit 与任务句柄在同一生命周期内一起释放。
let _ = work_handle.await;
warn!(
job_id = %job_id,
job_kind = %job_kind,
@@ -454,6 +481,9 @@ fn detach_external_generation_work_until_lease_expiry(
);
}
}
// 保持执行许可直到 work 真正结束或被取消,避免超时任务脱管后继续
// 累积图片/音频响应占用。
drop(permit);
});
}
@@ -1972,6 +2002,7 @@ mod tests {
&job,
Duration::from_millis(200),
"任务超过执行预算",
None,
);
tokio::time::sleep(Duration::from_millis(100)).await;
@@ -1981,6 +2012,61 @@ mod tests {
);
}
#[tokio::test]
async fn worker_detached_work_keeps_execution_slot_until_finished() {
let slots = std::sync::Arc::new(tokio::sync::Semaphore::new(1));
let permit = slots
.clone()
.acquire_owned()
.await
.expect("the only execution slot should be available");
let work_handle = tokio::spawn(async {
tokio::time::sleep(Duration::from_millis(20)).await;
Ok(())
});
let job = external_generation_job_record_fixture(Some("lease-1"));
detach_external_generation_work_until_lease_expiry(
work_handle,
&job,
Duration::from_millis(200),
"任务超过执行预算",
Some(permit),
);
assert!(
slots.try_acquire().is_err(),
"脱管 work 完成前不得重新领取执行容量"
);
tokio::time::sleep(Duration::from_millis(100)).await;
assert!(
slots.try_acquire().is_ok(),
"脱管 work 完成后应归还执行容量"
);
}
#[tokio::test]
async fn worker_reaps_completed_tasks_while_queue_remains_busy() {
let mut tasks = JoinSet::new();
let completed = std::sync::Arc::new(std::sync::atomic::AtomicUsize::new(0));
const TASK_COUNT: usize = 128;
for _ in 0..TASK_COUNT {
let completed = completed.clone();
tasks.spawn(async move {
completed.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
});
}
while completed.load(std::sync::atomic::Ordering::SeqCst) < TASK_COUNT {
tokio::task::yield_now().await;
}
assert_eq!(tasks.len(), TASK_COUNT);
reap_finished_external_generation_worker_tasks(&mut tasks);
assert!(tasks.is_empty(), "已完成任务的 JoinHandle 应在每轮被回收");
}
#[tokio::test]
async fn worker_detached_work_is_aborted_after_lease_arbitration_window() {
let connection = std::sync::Arc::new(tokio::sync::Semaphore::new(1));
@@ -1999,6 +2085,7 @@ mod tests {
&job,
Duration::from_millis(10),
"任务超过执行预算",
None,
);
let reacquired =
@@ -28,7 +28,7 @@ impl AiTaskService {
validate_task_create_input(&input).map_err(AiTaskServiceError::Field)?;
let snapshot = AiTaskSnapshot {
task_id: input.task_id.clone(),
task_id: normalize_required_string(input.task_id).unwrap_or_default(),
task_kind: input.task_kind,
owner_user_id: normalize_required_string(input.owner_user_id).unwrap_or_default(),
request_label: normalize_required_string(input.request_label).unwrap_or_default(),
@@ -1,10 +1,12 @@
use std::{
collections::HashMap,
collections::{BTreeMap, HashMap},
sync::{Arc, Mutex},
};
use crate::{
AiTaskServiceError, AiTaskSnapshot, AiTaskStageStatus, AiTaskStatus, AiTextChunkSnapshot,
MAX_AI_TASK_RETAINED_OUTPUT_BYTES, MAX_AI_TASK_RETAINED_TASKS, MAX_AI_TASK_TEXT_OUTPUT_BYTES,
validate_ai_task_snapshot_memory_limits,
};
use super::ensure_task_is_not_terminal;
@@ -17,7 +19,9 @@ pub struct InMemoryAiTaskStore {
#[derive(Debug, Default)]
struct InMemoryAiTaskStoreState {
tasks: HashMap<String, AiTaskSnapshot>,
text_chunks: HashMap<String, Vec<AiTextChunkSnapshot>>,
// Keep only the ordered deltas needed to handle an out-of-order chunk.
// Completed tasks drop this map immediately; it is not a second durable log.
text_chunks: HashMap<String, HashMap<crate::AiTaskStageKind, BTreeMap<u32, String>>>,
}
impl InMemoryAiTaskStore {
@@ -34,7 +38,48 @@ impl InMemoryAiTaskStore {
return Err(AiTaskServiceError::TaskAlreadyExists);
}
state.text_chunks.insert(task.task_id.clone(), Vec::new());
validate_task_memory_limits(&task)?;
let oldest_terminal = if state.tasks.len() >= MAX_AI_TASK_RETAINED_TASKS {
let oldest_terminal = state
.tasks
.values()
.filter(|value| value.status.is_terminal())
.min_by_key(|value| value.completed_at_micros.or(Some(value.updated_at_micros)))
.map(|value| value.task_id.clone());
if oldest_terminal.is_none() {
return Err(AiTaskServiceError::Store(
"AI 任务仓储已达到内存容量上限".to_string(),
));
}
oldest_terminal
} else {
None
};
let retained_output_bytes = retained_output_bytes(&state)
.saturating_sub(
oldest_terminal
.as_deref()
.and_then(|task_id| state.tasks.get(task_id))
.map(task_output_bytes)
.unwrap_or_default(),
)
.saturating_add(task_output_bytes(&task));
if retained_output_bytes > MAX_AI_TASK_RETAINED_OUTPUT_BYTES {
return Err(AiTaskServiceError::Store(
"AI 任务仓储输出工作集超过内存上限".to_string(),
));
}
if let Some(task_id) = oldest_terminal {
state.tasks.remove(&task_id);
state.text_chunks.remove(&task_id);
}
state
.text_chunks
.insert(task.task_id.clone(), HashMap::new());
state.tasks.insert(task.task_id.clone(), task.clone());
Ok(task)
}
@@ -51,12 +96,37 @@ impl InMemoryAiTaskStore {
.inner
.lock()
.map_err(|_| AiTaskServiceError::Store("AI 任务仓储锁已中毒".to_string()))?;
let task = state
.tasks
.get_mut(task_id.trim())
.ok_or(AiTaskServiceError::TaskNotFound)?;
apply(task)?;
Ok(task.clone())
let (previous_task, snapshot) = {
let task = state
.tasks
.get_mut(task_id.trim())
.ok_or(AiTaskServiceError::TaskNotFound)?;
let previous_task = task.clone();
if let Err(error) = apply(task) {
*task = previous_task;
return Err(error);
}
(previous_task, task.clone())
};
if let Err(error) = validate_task_memory_limits(&snapshot) {
state
.tasks
.insert(task_id.trim().to_string(), previous_task);
return Err(error);
}
let retained_output_bytes = retained_output_bytes(&state);
if retained_output_bytes > MAX_AI_TASK_RETAINED_OUTPUT_BYTES {
state
.tasks
.insert(task_id.trim().to_string(), previous_task);
return Err(AiTaskServiceError::Store(
"AI 任务仓储输出工作集超过内存上限".to_string(),
));
}
if snapshot.status.is_terminal() {
state.text_chunks.remove(task_id.trim());
}
Ok(snapshot)
}
pub(super) fn append_text_chunk(
@@ -67,13 +137,75 @@ impl InMemoryAiTaskStore {
.inner
.lock()
.map_err(|_| AiTaskServiceError::Store("AI 任务仓储锁已中毒".to_string()))?;
{
if chunk.delta_text.len() > MAX_AI_TASK_TEXT_OUTPUT_BYTES {
return Err(AiTaskServiceError::Store(
"AI 任务文本输出超过内存上限".to_string(),
));
}
let (previous_stage_output_bytes, previous_latest_output_bytes, previous_task) = {
let task = state
.tasks
.get(&chunk.task_id)
.ok_or(AiTaskServiceError::TaskNotFound)?;
ensure_task_is_not_terminal(task.status)?;
let stage = task
.stages
.iter()
.find(|stage| stage.stage_kind == chunk.stage_kind)
.ok_or(AiTaskServiceError::StageNotFound)?;
(
stage.text_output.as_ref().map_or(0, String::len),
task.latest_text_output.as_ref().map_or(0, String::len),
task.clone(),
)
};
let (previous_chunk, aggregated_bytes, aggregated_text) = {
let chunks = state
.text_chunks
.get_mut(&chunk.task_id)
.ok_or(AiTaskServiceError::TaskNotFound)?;
let stage_chunks = chunks.entry(chunk.stage_kind).or_default();
let previous_chunk = stage_chunks.insert(chunk.sequence, chunk.delta_text.clone());
let aggregated_bytes = stage_chunks
.values()
.fold(0_usize, |total, delta| total.saturating_add(delta.len()));
let mut aggregated_text = String::with_capacity(aggregated_bytes);
for delta in stage_chunks.values() {
aggregated_text.push_str(delta);
}
(previous_chunk, aggregated_bytes, aggregated_text)
};
if aggregated_bytes > MAX_AI_TASK_TEXT_OUTPUT_BYTES {
rollback_text_chunk(&mut state, &chunk, previous_chunk);
return Err(AiTaskServiceError::Store(
"AI 任务文本输出超过内存上限".to_string(),
));
}
let projected_retained_output_bytes = retained_output_bytes(&state)
.saturating_sub(previous_stage_output_bytes)
.saturating_sub(previous_latest_output_bytes)
.saturating_add(aggregated_bytes.saturating_mul(2));
if projected_retained_output_bytes > MAX_AI_TASK_RETAINED_OUTPUT_BYTES {
rollback_text_chunk(&mut state, &chunk, previous_chunk);
return Err(AiTaskServiceError::Store(
"AI 任务仓储输出工作集超过内存上限".to_string(),
));
}
let normalized_output = if aggregated_text.trim().is_empty() {
None
} else {
Some(aggregated_text)
};
let snapshot = {
let task = state
.tasks
.get_mut(&chunk.task_id)
.ok_or(AiTaskServiceError::TaskNotFound)?;
ensure_task_is_not_terminal(task.status)?;
let stage = task
.stages
.iter_mut()
@@ -83,45 +215,23 @@ impl InMemoryAiTaskStore {
stage.status = AiTaskStageStatus::Running;
stage.started_at_micros = Some(chunk.created_at_micros);
}
task.status = AiTaskStatus::Running;
task.started_at_micros
.get_or_insert(chunk.created_at_micros);
}
let chunks = state
.text_chunks
.get_mut(&chunk.task_id)
.ok_or(AiTaskServiceError::TaskNotFound)?;
chunks.push(chunk.clone());
chunks.sort_by_key(|value| value.sequence);
let aggregated_text = chunks
.iter()
.filter(|value| value.stage_kind == chunk.stage_kind)
.map(|value| value.delta_text.as_str())
.collect::<Vec<_>>()
.join("");
let normalized_output = if aggregated_text.trim().is_empty() {
None
} else {
Some(aggregated_text)
stage.text_output = normalized_output.clone();
task.latest_text_output = normalized_output;
task.updated_at_micros = chunk.created_at_micros;
task.version += 1;
task.clone()
};
let task = state
.tasks
.get_mut(&chunk.task_id)
.ok_or(AiTaskServiceError::TaskNotFound)?;
let stage = task
.stages
.iter_mut()
.find(|stage| stage.stage_kind == chunk.stage_kind)
.ok_or(AiTaskServiceError::StageNotFound)?;
stage.text_output = normalized_output.clone();
task.latest_text_output = normalized_output;
task.updated_at_micros = chunk.created_at_micros;
task.version += 1;
Ok(task.clone())
if let Err(error) = validate_task_memory_limits(&snapshot)
.and_then(|_| validate_retained_output_bytes(&state))
{
state.tasks.insert(chunk.task_id.clone(), previous_task);
rollback_text_chunk(&mut state, &chunk, previous_chunk);
return Err(error);
}
Ok(snapshot)
}
pub(super) fn get_task(&self, task_id: &str) -> Result<AiTaskSnapshot, AiTaskServiceError> {
@@ -136,3 +246,109 @@ impl InMemoryAiTaskStore {
.ok_or(AiTaskServiceError::TaskNotFound)
}
}
fn rollback_text_chunk(
state: &mut InMemoryAiTaskStoreState,
chunk: &AiTextChunkSnapshot,
previous_chunk: Option<String>,
) {
if let Some(stage_chunks) = state
.text_chunks
.get_mut(&chunk.task_id)
.and_then(|chunks| chunks.get_mut(&chunk.stage_kind))
{
if let Some(previous_chunk) = previous_chunk {
stage_chunks.insert(chunk.sequence, previous_chunk);
} else {
stage_chunks.remove(&chunk.sequence);
}
}
}
fn retained_output_bytes(state: &InMemoryAiTaskStoreState) -> usize {
let snapshot_bytes = state.tasks.values().fold(0_usize, |total, task| {
total.saturating_add(task_output_bytes(task))
});
state
.text_chunks
.values()
.fold(snapshot_bytes, |total, stages| {
stages.values().fold(total, |stage_total, chunks| {
chunks.values().fold(stage_total, |chunk_total, delta| {
chunk_total.saturating_add(delta.len())
})
})
})
}
fn validate_retained_output_bytes(
state: &InMemoryAiTaskStoreState,
) -> Result<(), AiTaskServiceError> {
if retained_output_bytes(state) > MAX_AI_TASK_RETAINED_OUTPUT_BYTES {
return Err(AiTaskServiceError::Store(
"AI 任务仓储输出工作集超过内存上限".to_string(),
));
}
Ok(())
}
fn task_output_bytes(task: &AiTaskSnapshot) -> usize {
let task_metadata = task
.task_id
.len()
.saturating_add(task.owner_user_id.len())
.saturating_add(task.request_label.len())
.saturating_add(task.source_module.len())
.saturating_add(task.source_entity_id.as_ref().map_or(0, String::len))
.saturating_add(task.stages.iter().fold(0_usize, |total, stage| {
total
.saturating_add(stage.label.len())
.saturating_add(stage.detail.len())
}));
let request_payload = task.request_payload_json.as_ref().map_or(0, String::len);
let failure_message = task.failure_message.as_ref().map_or(0, String::len);
let result_references = task
.result_references
.iter()
.fold(0_usize, |total, reference| {
total
.saturating_add(reference.result_ref_id.len())
.saturating_add(reference.task_id.len())
.saturating_add(reference.reference_id.len())
.saturating_add(reference.label.as_ref().map_or(0, String::len))
});
let latest_text = task.latest_text_output.as_ref().map_or(0, String::len);
let latest_structured = task
.latest_structured_payload_json
.as_ref()
.map_or(0, String::len);
let stage_bytes = task.stages.iter().fold(0_usize, |total, stage| {
let text = stage.text_output.as_ref().map_or(0, String::len);
let structured = stage
.structured_payload_json
.as_ref()
.map_or(0, String::len);
let warnings = stage
.warning_messages
.iter()
.fold(0_usize, |warning_total, warning| {
warning_total.saturating_add(warning.len())
});
total
.saturating_add(text)
.saturating_add(structured)
.saturating_add(warnings)
});
task_metadata
.saturating_add(request_payload)
.saturating_add(failure_message)
.saturating_add(result_references)
.saturating_add(latest_text)
.saturating_add(latest_structured)
.saturating_add(stage_bytes)
}
fn validate_task_memory_limits(task: &AiTaskSnapshot) -> Result<(), AiTaskServiceError> {
validate_ai_task_snapshot_memory_limits(task)
.map_err(|message| AiTaskServiceError::Store(message.to_string()))
}
+11
View File
@@ -1,4 +1,5 @@
mod ids;
mod limits;
mod stages;
mod types;
@@ -8,6 +9,16 @@ pub use ids::{
generate_ai_task_stage_id, generate_ai_text_chunk_id, normalize_optional_text,
normalize_string_list,
};
pub use limits::{
MAX_AI_TASK_FAILURE_MESSAGE_BYTES, MAX_AI_TASK_ID_BYTES, MAX_AI_TASK_OWNER_USER_ID_BYTES,
MAX_AI_TASK_REFERENCE_ID_BYTES, MAX_AI_TASK_REFERENCE_LABEL_BYTES,
MAX_AI_TASK_REQUEST_LABEL_BYTES, MAX_AI_TASK_REQUEST_PAYLOAD_BYTES,
MAX_AI_TASK_RESULT_REFERENCES, MAX_AI_TASK_RETAINED_OUTPUT_BYTES, MAX_AI_TASK_RETAINED_TASKS,
MAX_AI_TASK_SOURCE_ENTITY_ID_BYTES, MAX_AI_TASK_SOURCE_MODULE_BYTES,
MAX_AI_TASK_STAGE_DETAIL_BYTES, MAX_AI_TASK_STAGE_LABEL_BYTES,
MAX_AI_TASK_STRUCTURED_OUTPUT_BYTES, MAX_AI_TASK_TEXT_OUTPUT_BYTES, MAX_AI_TASK_WARNING_BYTES,
validate_ai_task_snapshot_memory_limits,
};
pub use types::{
AiResultReferenceKind, AiResultReferenceSnapshot, AiTaskKind, AiTaskSnapshot,
AiTaskStageBlueprint, AiTaskStageKind, AiTaskStageSnapshot, AiTaskStageStatus, AiTaskStatus,
@@ -0,0 +1,115 @@
use super::types::AiTaskSnapshot;
pub const MAX_AI_TASK_RETAINED_TASKS: usize = 1024;
pub const MAX_AI_TASK_ID_BYTES: usize = 256;
pub const MAX_AI_TASK_OWNER_USER_ID_BYTES: usize = 256;
pub const MAX_AI_TASK_REQUEST_LABEL_BYTES: usize = 4 * 1024;
pub const MAX_AI_TASK_SOURCE_MODULE_BYTES: usize = 256;
pub const MAX_AI_TASK_SOURCE_ENTITY_ID_BYTES: usize = 512;
pub const MAX_AI_TASK_STAGE_LABEL_BYTES: usize = 4 * 1024;
pub const MAX_AI_TASK_STAGE_DETAIL_BYTES: usize = 8 * 1024;
pub const MAX_AI_TASK_TEXT_OUTPUT_BYTES: usize = 512 * 1024;
pub const MAX_AI_TASK_STRUCTURED_OUTPUT_BYTES: usize = 512 * 1024;
pub const MAX_AI_TASK_WARNING_BYTES: usize = 64 * 1024;
pub const MAX_AI_TASK_REQUEST_PAYLOAD_BYTES: usize = 512 * 1024;
pub const MAX_AI_TASK_FAILURE_MESSAGE_BYTES: usize = 64 * 1024;
pub const MAX_AI_TASK_RESULT_REFERENCES: usize = 64;
pub const MAX_AI_TASK_REFERENCE_ID_BYTES: usize = 512;
pub const MAX_AI_TASK_REFERENCE_LABEL_BYTES: usize = 2 * 1024;
pub const MAX_AI_TASK_RETAINED_OUTPUT_BYTES: usize = 64 * 1024 * 1024;
pub fn validate_ai_task_snapshot_memory_limits(task: &AiTaskSnapshot) -> Result<(), &'static str> {
if task.task_id.len() > MAX_AI_TASK_ID_BYTES {
return Err("AI 任务 ID 超过内存上限");
}
if task.owner_user_id.len() > MAX_AI_TASK_OWNER_USER_ID_BYTES {
return Err("AI 任务用户 ID 超过内存上限");
}
if task.request_label.len() > MAX_AI_TASK_REQUEST_LABEL_BYTES {
return Err("AI 任务请求标签超过内存上限");
}
if task.source_module.len() > MAX_AI_TASK_SOURCE_MODULE_BYTES {
return Err("AI 任务来源模块超过内存上限");
}
if task
.source_entity_id
.as_ref()
.is_some_and(|entity_id| entity_id.len() > MAX_AI_TASK_SOURCE_ENTITY_ID_BYTES)
{
return Err("AI 任务来源实体 ID 超过内存上限");
}
if task.stages.iter().any(|stage| {
stage.label.len() > MAX_AI_TASK_STAGE_LABEL_BYTES
|| stage.detail.len() > MAX_AI_TASK_STAGE_DETAIL_BYTES
}) {
return Err("AI 任务阶段元数据超过内存上限");
}
if task
.request_payload_json
.as_ref()
.is_some_and(|payload| payload.len() > MAX_AI_TASK_REQUEST_PAYLOAD_BYTES)
{
return Err("AI 任务请求 payload 超过内存上限");
}
if task
.failure_message
.as_ref()
.is_some_and(|message| message.len() > MAX_AI_TASK_FAILURE_MESSAGE_BYTES)
{
return Err("AI 任务失败消息超过内存上限");
}
if task.stages.iter().any(|stage| {
stage
.text_output
.as_ref()
.is_some_and(|text| text.len() > MAX_AI_TASK_TEXT_OUTPUT_BYTES)
}) || task
.latest_text_output
.as_ref()
.is_some_and(|text| text.len() > MAX_AI_TASK_TEXT_OUTPUT_BYTES)
{
return Err("AI 任务文本输出超过内存上限");
}
if task.stages.iter().any(|stage| {
stage
.structured_payload_json
.as_ref()
.is_some_and(|payload| payload.len() > MAX_AI_TASK_STRUCTURED_OUTPUT_BYTES)
}) || task
.latest_structured_payload_json
.as_ref()
.is_some_and(|payload| payload.len() > MAX_AI_TASK_STRUCTURED_OUTPUT_BYTES)
{
return Err("AI 任务结构化输出超过内存上限");
}
if task.stages.iter().any(|stage| {
stage
.warning_messages
.iter()
.fold(0_usize, |total, warning| {
total.saturating_add(warning.len())
})
> MAX_AI_TASK_WARNING_BYTES
}) {
return Err("AI 任务 warning 输出超过内存上限");
}
if task.result_references.len() > MAX_AI_TASK_RESULT_REFERENCES {
return Err("AI 任务结果引用数量超过内存上限");
}
if task
.result_references
.iter()
.any(|reference| reference.reference_id.len() > MAX_AI_TASK_REFERENCE_ID_BYTES)
{
return Err("AI 任务结果引用 ID 超过内存上限");
}
if task.result_references.iter().any(|reference| {
reference
.label
.as_ref()
.is_some_and(|label| label.len() > MAX_AI_TASK_REFERENCE_LABEL_BYTES)
}) {
return Err("AI 任务结果引用标签超过内存上限");
}
Ok(())
}
+11 -3
View File
@@ -14,9 +14,17 @@ pub use domain::{
AI_RESULT_REF_ID_PREFIX, AI_TASK_ID_PREFIX, AI_TASK_STAGE_ID_PREFIX, AI_TEXT_CHUNK_ID_PREFIX,
AiResultReferenceKind, AiResultReferenceSnapshot, AiTaskKind, AiTaskSnapshot,
AiTaskStageBlueprint, AiTaskStageKind, AiTaskStageSnapshot, AiTaskStageStatus, AiTaskStatus,
AiTextChunkSnapshot, INITIAL_AI_TASK_VERSION, generate_ai_result_ref_id, generate_ai_task_id,
generate_ai_task_stage_id, generate_ai_text_chunk_id, normalize_optional_text,
normalize_string_list,
AiTextChunkSnapshot, INITIAL_AI_TASK_VERSION, MAX_AI_TASK_FAILURE_MESSAGE_BYTES,
MAX_AI_TASK_ID_BYTES, MAX_AI_TASK_OWNER_USER_ID_BYTES, MAX_AI_TASK_REFERENCE_ID_BYTES,
MAX_AI_TASK_REFERENCE_LABEL_BYTES, MAX_AI_TASK_REQUEST_LABEL_BYTES,
MAX_AI_TASK_REQUEST_PAYLOAD_BYTES, MAX_AI_TASK_RESULT_REFERENCES,
MAX_AI_TASK_RETAINED_OUTPUT_BYTES, MAX_AI_TASK_RETAINED_TASKS,
MAX_AI_TASK_SOURCE_ENTITY_ID_BYTES, MAX_AI_TASK_SOURCE_MODULE_BYTES,
MAX_AI_TASK_STAGE_DETAIL_BYTES, MAX_AI_TASK_STAGE_LABEL_BYTES,
MAX_AI_TASK_STRUCTURED_OUTPUT_BYTES, MAX_AI_TASK_TEXT_OUTPUT_BYTES, MAX_AI_TASK_WARNING_BYTES,
generate_ai_result_ref_id, generate_ai_task_id, generate_ai_task_stage_id,
generate_ai_text_chunk_id, normalize_optional_text, normalize_string_list,
validate_ai_task_snapshot_memory_limits,
};
pub use errors::{AiTaskFieldError, AiTaskServiceError};
pub use events::AiTaskDomainEvent;
+252
View File
@@ -43,6 +43,32 @@ fn create_task_rejects_duplicate_stage_blueprints() {
assert_eq!(error, AiTaskFieldError::DuplicateStageBlueprint);
}
#[test]
fn create_task_rejects_oversized_request_payload() {
let service = build_service();
let mut input = build_create_input(AiTaskKind::StoryGeneration);
input.request_payload_json = Some("x".repeat(MAX_AI_TASK_REQUEST_PAYLOAD_BYTES + 1));
let error = service
.create_task(input)
.expect_err("request payload over the memory cap should fail");
assert!(
matches!(error, AiTaskServiceError::Store(message) if message.contains("请求 payload"))
);
}
#[test]
fn create_task_rejects_oversized_request_metadata() {
let service = build_service();
let mut input = build_create_input(AiTaskKind::StoryGeneration);
input.request_label = "x".repeat(MAX_AI_TASK_REQUEST_LABEL_BYTES + 1);
let error = service
.create_task(input)
.expect_err("request metadata over the memory cap should fail");
assert!(matches!(error, AiTaskServiceError::Store(message) if message.contains("请求标签")));
}
#[test]
fn generate_ai_task_stage_id_contains_task_and_stage_slug() {
let stage_id = generate_ai_task_stage_id("aitask_demo", AiTaskStageKind::NormalizeResult);
@@ -112,6 +138,47 @@ fn append_text_chunk_aggregates_stream_output_by_stage() {
assert_eq!(second_chunk.sequence, 2);
}
#[test]
fn append_text_chunk_rejects_output_over_stage_memory_limit_without_mutating_task() {
let service = build_service();
let task = service
.create_task(build_create_input(AiTaskKind::CharacterChat))
.expect("task should create");
let max_output = "a".repeat(512 * 1024);
let (updated, _) = service
.append_text_chunk(
&task.task_id,
AiTaskStageKind::RequestModel,
1,
max_output.clone(),
task.created_at_micros + 1,
)
.expect("the stage limit itself should be accepted");
assert_eq!(
updated.latest_text_output.as_deref().map(str::len),
Some(max_output.len())
);
let error = service
.append_text_chunk(
&task.task_id,
AiTaskStageKind::RequestModel,
2,
"b".to_string(),
task.created_at_micros + 2,
)
.expect_err("output beyond the stage limit should fail");
assert!(matches!(error, AiTaskServiceError::Store(_)));
let after_rejection = service
.get_task(&task.task_id)
.expect("task should remain readable");
assert_eq!(
after_rejection.latest_text_output.as_deref().map(str::len),
Some(max_output.len())
);
}
#[test]
fn complete_stage_updates_latest_outputs() {
let service = build_service();
@@ -147,6 +214,150 @@ fn complete_stage_updates_latest_outputs() {
assert_eq!(stage.warning_messages, vec!["使用了 fallback 选项池"]);
}
#[test]
fn complete_stage_rejects_oversized_text_output_without_mutating_task() {
let service = build_service();
let task = service
.create_task(build_create_input(AiTaskKind::StoryGeneration))
.expect("task should create");
let oversized = "x".repeat(512 * 1024 + 1);
let error = service
.complete_stage(AiStageCompletionInput {
task_id: task.task_id.clone(),
stage_kind: AiTaskStageKind::NormalizeResult,
text_output: Some(oversized),
structured_payload_json: None,
warning_messages: Vec::new(),
completed_at_micros: task.created_at_micros + 1,
})
.expect_err("text output over the per-stage cap should fail");
assert!(matches!(error, AiTaskServiceError::Store(message) if message.contains("文本输出")));
let unchanged = service
.get_task(&task.task_id)
.expect("task should remain readable");
let stage = unchanged
.stages
.iter()
.find(|stage| stage.stage_kind == AiTaskStageKind::NormalizeResult)
.expect("normalize stage should exist");
assert_eq!(stage.status, AiTaskStageStatus::Pending);
assert!(stage.text_output.is_none());
assert!(unchanged.latest_text_output.is_none());
}
#[test]
fn complete_stage_rejects_oversized_structured_output_without_mutating_task() {
let service = build_service();
let task = service
.create_task(build_create_input(AiTaskKind::StoryGeneration))
.expect("task should create");
let oversized = "x".repeat(512 * 1024 + 1);
let error = service
.complete_stage(AiStageCompletionInput {
task_id: task.task_id.clone(),
stage_kind: AiTaskStageKind::NormalizeResult,
text_output: None,
structured_payload_json: Some(oversized),
warning_messages: Vec::new(),
completed_at_micros: task.created_at_micros + 1,
})
.expect_err("structured output over the per-stage cap should fail");
assert!(matches!(error, AiTaskServiceError::Store(message) if message.contains("结构化输出")));
let unchanged = service
.get_task(&task.task_id)
.expect("task should remain readable");
let stage = unchanged
.stages
.iter()
.find(|stage| stage.stage_kind == AiTaskStageKind::NormalizeResult)
.expect("normalize stage should exist");
assert_eq!(stage.status, AiTaskStageStatus::Pending);
assert!(stage.structured_payload_json.is_none());
assert!(unchanged.latest_structured_payload_json.is_none());
}
#[test]
fn complete_stage_rejects_oversized_warning_output_without_mutating_task() {
let service = build_service();
let task = service
.create_task(build_create_input(AiTaskKind::StoryGeneration))
.expect("task should create");
let oversized_warning = "w".repeat(64 * 1024 + 1);
let error = service
.complete_stage(AiStageCompletionInput {
task_id: task.task_id.clone(),
stage_kind: AiTaskStageKind::NormalizeResult,
text_output: None,
structured_payload_json: None,
warning_messages: vec![oversized_warning],
completed_at_micros: task.created_at_micros + 1,
})
.expect_err("warning output over the per-stage cap should fail");
assert!(matches!(error, AiTaskServiceError::Store(message) if message.contains("warning")));
let unchanged = service
.get_task(&task.task_id)
.expect("task should remain readable");
let stage = unchanged
.stages
.iter()
.find(|stage| stage.stage_kind == AiTaskStageKind::NormalizeResult)
.expect("normalize stage should exist");
assert_eq!(stage.status, AiTaskStageStatus::Pending);
assert!(stage.warning_messages.is_empty());
}
#[test]
fn complete_stage_enforces_global_retained_output_cap() {
let service = build_service();
let structured_payload = "x".repeat(512 * 1024);
for index in 0..63 {
let task = service
.create_task(AiTaskCreateInput {
task_id: format!("task-structured-cap-{index}"),
..build_create_input(AiTaskKind::StoryGeneration)
})
.expect("task should create");
service
.complete_stage(AiStageCompletionInput {
task_id: task.task_id,
stage_kind: AiTaskStageKind::NormalizeResult,
text_output: None,
structured_payload_json: Some(structured_payload.clone()),
warning_messages: Vec::new(),
completed_at_micros: task.created_at_micros + 1,
})
.expect("63 MiB retained output should remain within the global cap");
}
let task = service
.create_task(AiTaskCreateInput {
task_id: "task-structured-cap-overflow".to_string(),
..build_create_input(AiTaskKind::StoryGeneration)
})
.expect("the overflow candidate task itself should create");
let error = service
.complete_stage(AiStageCompletionInput {
task_id: task.task_id.clone(),
stage_kind: AiTaskStageKind::NormalizeResult,
text_output: None,
structured_payload_json: Some(structured_payload),
warning_messages: Vec::new(),
completed_at_micros: task.created_at_micros + 1,
})
.expect_err("global retained output cap should reject the overflow");
assert!(matches!(error, AiTaskServiceError::Store(message) if message.contains("工作集")));
let unchanged = service
.get_task(&task.task_id)
.expect("overflow task should remain readable");
assert!(unchanged.latest_structured_payload_json.is_none());
}
#[test]
fn attach_result_reference_appends_binding() {
let service = build_service();
@@ -172,6 +383,47 @@ fn attach_result_reference_appends_binding() {
assert_eq!(updated.result_references[0].reference_id, "profile_001");
}
#[test]
fn attach_result_reference_rejects_unbounded_reference_growth() {
let service = build_service();
let task = service
.create_task(build_create_input(AiTaskKind::CustomWorldGeneration))
.expect("task should create");
for index in 0..MAX_AI_TASK_RESULT_REFERENCES {
service
.attach_result_reference(
&task.task_id,
AiResultReferenceKind::CustomWorldProfile,
format!("profile_{index}"),
None,
task.created_at_micros + index as i64 + 1,
)
.expect("references within the cap should attach");
}
let error = service
.attach_result_reference(
&task.task_id,
AiResultReferenceKind::CustomWorldProfile,
"profile_overflow".to_string(),
None,
task.created_at_micros + MAX_AI_TASK_RESULT_REFERENCES as i64 + 1,
)
.expect_err("references over the cap should fail");
assert!(
matches!(error, AiTaskServiceError::Store(message) if message.contains("结果引用数量"))
);
let unchanged = service
.get_task(&task.task_id)
.expect("task should remain readable");
assert_eq!(
unchanged.result_references.len(),
MAX_AI_TASK_RESULT_REFERENCES
);
}
#[test]
fn fail_and_cancel_task_move_into_terminal_states() {
let service = build_service();
+266
View File
@@ -34,6 +34,9 @@ use tracing::{info, warn};
const DEFAULT_PHONE_VERIFY_CODE_SALT: &str = "genarrative-phone-verify-code-v1";
const PHONE_CODE_RESERVATION_MARKER: &str = "__genarrative_phone_code_reservation__";
const MAX_ACTIVE_WECHAT_AUTH_STATES: usize = 1024;
const REFRESH_SESSION_STALE_RETENTION: Duration = Duration::days(1);
const MAX_REFRESH_SESSIONS: usize = 8_192;
const MAX_PHONE_CODES: usize = 4_096;
#[derive(Clone, Debug)]
pub struct InMemoryAuthStore {
@@ -390,6 +393,7 @@ impl RefreshSessionService {
input: CreateRefreshSessionInput,
now: OffsetDateTime,
) -> Result<CreateRefreshSessionResult, RefreshSessionError> {
self.store.prune_stale_sessions(now)?;
self.store
.find_by_user_id(&input.user_id)
.map_err(map_password_store_error)?
@@ -426,6 +430,7 @@ impl RefreshSessionService {
input: RotateRefreshSessionInput,
now: OffsetDateTime,
) -> Result<RotateRefreshSessionResult, RefreshSessionError> {
self.store.prune_stale_sessions(now)?;
let Some(refresh_token_hash) = normalize_required_string(&input.refresh_token_hash) else {
return Err(RefreshSessionError::MissingToken);
};
@@ -480,6 +485,7 @@ impl RefreshSessionService {
user_id: &str,
now: OffsetDateTime,
) -> Result<ListActiveRefreshSessionsResult, RefreshSessionError> {
self.store.prune_stale_sessions(now)?;
self.store
.find_by_user_id(user_id)
.map_err(map_password_store_error)?
@@ -494,6 +500,7 @@ impl RefreshSessionService {
input: RevokeRefreshSessionByUserInput,
now: OffsetDateTime,
) -> Result<RevokeRefreshSessionResult, RefreshSessionError> {
self.store.prune_stale_sessions(now)?;
self.store
.find_by_user_id(&input.user_id)
.map_err(map_password_store_error)?
@@ -518,6 +525,7 @@ impl RefreshSessionService {
session_id: &str,
now: OffsetDateTime,
) -> Result<bool, RefreshSessionError> {
self.store.prune_stale_sessions(now)?;
self.store
.is_session_active_for_user(user_id, session_id.trim(), now)
}
@@ -556,6 +564,7 @@ impl PhoneAuthService {
input: &SendPhoneCodeInput,
now: OffsetDateTime,
) -> Result<(), PhoneAuthError> {
self.store.prune_expired_phone_codes(now)?;
let scene = input.scene.clone();
validate_mainland_china_country_code(input.country_code.as_deref())?;
let normalized_phone = normalize_mainland_china_phone_number(&input.pure_phone_number)?;
@@ -605,6 +614,7 @@ impl PhoneAuthService {
now: OffsetDateTime,
check_local_cooldown: bool,
) -> Result<SendPhoneCodeResult, PhoneAuthError> {
self.store.prune_expired_phone_codes(now)?;
let scene = input.scene.clone();
validate_mainland_china_country_code(input.country_code.as_deref())?;
let normalized_phone = normalize_mainland_china_phone_number(&input.pure_phone_number)?;
@@ -621,6 +631,8 @@ impl PhoneAuthService {
self.store
.ensure_phone_code_not_cooling_down(&normalized_phone.e164, &scene, now)?;
}
self.store
.ensure_phone_code_capacity(&normalized_phone.e164, &scene)?;
let expires_at = now
.checked_add(Duration::minutes(SMS_CODE_TTL_MINUTES))
.ok_or_else(|| PhoneAuthError::Store("短信验证码过期时间计算溢出".to_string()))?;
@@ -883,6 +895,7 @@ impl WechatAuthStateService {
input: CreateWechatAuthStateInput,
now: OffsetDateTime,
) -> Result<CreateWechatAuthStateResult, WechatAuthError> {
self.store.prune_wechat_states(now)?;
let created_at = format_rfc3339(now).map_err(|message| {
WechatAuthError::Store(format!("微信 state 时间格式化失败:{message}"))
})?;
@@ -1162,10 +1175,25 @@ impl InMemoryAuthStoreState {
}
}
let now = OffsetDateTime::now_utc();
let mut retained_refresh_session_count = 0_usize;
for session in view.refresh_sessions {
if !existing_user_ids.contains(&session.user_id) {
continue;
}
if should_prune_refresh_session_fields(
&session.expires_at,
session.revoked_at.as_deref(),
now,
) {
continue;
}
retained_refresh_session_count += 1;
if retained_refresh_session_count > MAX_REFRESH_SESSIONS {
return Err(format!(
"认证投影中的 refresh session 数量超过内存上限(最多 {MAX_REFRESH_SESSIONS} 条)"
));
}
let client_info =
serde_json::from_str::<RefreshSessionClientInfo>(&session.client_info_json)
.map_err(|error| format!("解析 refresh session 客户端信息失败:{error}"))?;
@@ -1371,6 +1399,8 @@ impl InMemoryAuthStore {
&self,
updated_at_micros: i64,
) -> Result<AuthStoreProjectionView, String> {
self.prune_stale_sessions(OffsetDateTime::now_utc())
.map_err(|error| error.to_string())?;
let mut state = self
.inner
.lock()
@@ -1499,6 +1529,38 @@ impl InMemoryAuthStore {
Err("认证工作集在导出期间持续发生变化".to_string())
}
fn prune_stale_sessions(&self, now: OffsetDateTime) -> Result<(), RefreshSessionError> {
let mut state = self
.inner
.lock()
.map_err(|_| RefreshSessionError::Store("会话仓储锁已中毒".to_string()))?;
let stale_session_ids = state
.sessions_by_id
.iter()
.filter(|(_, stored)| should_prune_refresh_session(&stored.session, now))
.map(|(session_id, _)| session_id.clone())
.collect::<Vec<_>>();
if stale_session_ids.is_empty() {
return Ok(());
}
for session_id in stale_session_ids {
let Some(stored) = state.sessions_by_id.remove(&session_id) else {
continue;
};
if state
.session_id_by_refresh_token_hash
.get(&stored.session.refresh_token_hash)
.is_some_and(|mapped_id| mapped_id == &session_id)
{
state
.session_id_by_refresh_token_hash
.remove(&stored.session.refresh_token_hash);
}
}
self.persist_refresh_state(&state)
}
fn persist_state(&self, state: &InMemoryAuthStoreState) -> Result<(), String> {
let _ = state;
self.revision.fetch_add(1, Ordering::Release);
@@ -2132,6 +2194,11 @@ impl InMemoryAuthStore {
"refresh token hash 已存在,无法重复创建会话".to_string(),
));
}
if state.sessions_by_id.len() >= MAX_REFRESH_SESSIONS {
return Err(RefreshSessionError::Store(
"refresh session 内存容量已达到上限".to_string(),
));
}
state.session_id_by_refresh_token_hash.insert(
session.refresh_token_hash.clone(),
@@ -2156,11 +2223,42 @@ impl InMemoryAuthStore {
.map_err(|_| PhoneAuthError::Store("短信验证码仓储锁已中毒".to_string()))?;
// 手机号和业务场景共同决定同一份验证码快照,重复发送时直接覆盖旧值。
let key = build_phone_code_key(&code.phone_number, &code.scene);
if !state.phone_codes_by_key.contains_key(&key)
&& state.phone_codes_by_key.len() >= MAX_PHONE_CODES
{
return Err(PhoneAuthError::Store(
"短信验证码内存容量已达到上限,请稍后重试".to_string(),
));
}
state.phone_codes_by_key.insert(key, code);
self.persist_phone_state(&state)?;
Ok(())
}
fn prune_expired_phone_codes(&self, now: OffsetDateTime) -> Result<(), PhoneAuthError> {
let mut state = self
.inner
.lock()
.map_err(|_| PhoneAuthError::Store("短信验证码仓储锁已中毒".to_string()))?;
let expired_keys = state
.phone_codes_by_key
.iter()
.filter_map(|(key, stored)| {
OffsetDateTime::parse(
&stored.expires_at,
&time::format_description::well_known::Rfc3339,
)
.ok()
.filter(|expires_at| *expires_at <= now)
.map(|_| key.clone())
})
.collect::<Vec<_>>();
for key in expired_keys {
state.phone_codes_by_key.remove(&key);
}
Ok(())
}
fn ensure_phone_code_not_cooling_down(
&self,
phone_number: &str,
@@ -2200,6 +2298,26 @@ impl InMemoryAuthStore {
})
}
fn ensure_phone_code_capacity(
&self,
phone_number: &str,
scene: &PhoneAuthScene,
) -> Result<(), PhoneAuthError> {
let state = self
.inner
.lock()
.map_err(|_| PhoneAuthError::Store("短信验证码仓储锁已中毒".to_string()))?;
let key = build_phone_code_key(phone_number, scene);
if state.phone_codes_by_key.contains_key(&key)
|| state.phone_codes_by_key.len() < MAX_PHONE_CODES
{
return Ok(());
}
Err(PhoneAuthError::Store(
"短信验证码内存容量已达到上限,请稍后重试".to_string(),
))
}
fn get_active_phone_code(
&self,
phone_number: &str,
@@ -2658,6 +2776,33 @@ impl InMemoryAuthStore {
Ok(())
}
fn prune_wechat_states(&self, now: OffsetDateTime) -> Result<(), WechatAuthError> {
let mut state = self
.inner
.lock()
.map_err(|_| WechatAuthError::Store("微信 state 仓储锁已中毒".to_string()))?;
let stale_tokens = state
.wechat_states_by_token
.iter()
.filter_map(|(token, stored)| {
if stored.state.consumed_at.is_some() {
return Some(token.clone());
}
OffsetDateTime::parse(
&stored.state.expires_at,
&time::format_description::well_known::Rfc3339,
)
.ok()
.filter(|expires_at| *expires_at <= now)
.map(|_| token.clone())
})
.collect::<Vec<_>>();
for token in stale_tokens {
state.wechat_states_by_token.remove(&token);
}
Ok(())
}
fn revoke_session_by_user_and_session_id(
&self,
user_id: &str,
@@ -2821,6 +2966,25 @@ impl InMemoryAuthStore {
}
}
fn should_prune_refresh_session(session: &RefreshSessionRecord, now: OffsetDateTime) -> bool {
should_prune_refresh_session_fields(&session.expires_at, session.revoked_at.as_deref(), now)
}
fn should_prune_refresh_session_fields(
expires_at: &str,
revoked_at: Option<&str>,
now: OffsetDateTime,
) -> bool {
let stale_before = now.saturating_sub(REFRESH_SESSION_STALE_RETENTION);
if let Some(revoked_at) = revoked_at {
return OffsetDateTime::parse(revoked_at, &time::format_description::well_known::Rfc3339)
.is_ok_and(|timestamp| timestamp <= stale_before);
}
OffsetDateTime::parse(expires_at, &time::format_description::well_known::Rfc3339)
.is_ok_and(|timestamp| timestamp <= stale_before)
}
fn map_sms_provider_error_to_phone_error(error: SmsProviderError) -> PhoneAuthError {
match error {
SmsProviderError::InvalidVerifyCode => PhoneAuthError::InvalidVerifyCode,
@@ -4413,6 +4577,108 @@ mod tests {
);
}
#[tokio::test]
async fn stale_refresh_sessions_are_pruned_from_both_indexes() {
let store = build_store();
let refresh_service = build_refresh_service(store.clone());
let user = create_phone_login_user(store.clone(), "13800138008").await;
let now = OffsetDateTime::now_utc();
refresh_service
.create_session(
CreateRefreshSessionInput {
user_id: user.id.clone(),
refresh_token_hash: hash_refresh_session_token("stale-revoked"),
issued_by_provider: AuthLoginMethod::Password,
client_info: build_client_info(),
},
now - Duration::days(2),
)
.expect("stale session should create");
store
.revoke_session_by_refresh_token_hash(
&hash_refresh_session_token("stale-revoked"),
now - Duration::days(2),
)
.expect("stale session should revoke");
refresh_service
.create_session(
CreateRefreshSessionInput {
user_id: user.id.clone(),
refresh_token_hash: hash_refresh_session_token("recent-revoked"),
issued_by_provider: AuthLoginMethod::Password,
client_info: build_client_info(),
},
now,
)
.expect("recent session should create");
store
.revoke_session_by_refresh_token_hash(
&hash_refresh_session_token("recent-revoked"),
now,
)
.expect("recent session should revoke");
let projection = store
.export_projection_view(now.unix_timestamp())
.expect("projection export should prune stale sessions");
assert_eq!(projection.refresh_sessions.len(), 1);
assert_eq!(
projection.refresh_sessions[0].refresh_token_hash,
hash_refresh_session_token("recent-revoked")
);
let stale_error = refresh_service
.rotate_session(
RotateRefreshSessionInput {
refresh_token_hash: hash_refresh_session_token("stale-revoked"),
next_refresh_token_hash: hash_refresh_session_token("stale-next"),
},
now,
)
.expect_err("pruned session should no longer be indexed");
assert_eq!(stale_error, RefreshSessionError::SessionNotFound);
}
#[test]
fn projection_restore_rejects_too_many_retained_refresh_sessions() {
let client_info_json =
serde_json::to_string(&build_client_info()).expect("client info should serialize");
let refresh_sessions = (0..=MAX_REFRESH_SESSIONS)
.map(|index| AuthStoreProjectionRefreshSession {
session_id: format!("session-{index}"),
user_id: "user_projection_cap".to_string(),
refresh_token_hash: format!("hash-{index}"),
issued_by_provider: "password".to_string(),
client_info_json: client_info_json.clone(),
expires_at: "2999-01-01T00:00:00Z".to_string(),
revoked_at: None,
created_at: "2026-01-01T00:00:00Z".to_string(),
updated_at: "2026-01-01T00:00:00Z".to_string(),
last_seen_at: "2026-01-01T00:00:00Z".to_string(),
})
.collect();
let error = InMemoryAuthStore::from_projection_view(AuthStoreProjectionView {
base_updated_at_micros: 0,
updated_at_micros: 1,
users: vec![projection_user(
"user_projection_cap",
"projection_cap",
None,
)],
identities: vec![],
refresh_sessions,
phone_codes: vec![],
wechat_states: vec![],
})
.expect_err("projection restore must enforce the refresh session cap");
assert!(error.contains("refresh session"));
assert!(error.contains(&MAX_REFRESH_SESSIONS.to_string()));
}
#[tokio::test]
async fn wechat_login_hits_existing_user_by_union_id_before_openid() {
let store = build_store();
@@ -414,6 +414,8 @@ pub mod external_generation_job_procedure_result_type;
pub mod external_generation_job_renew_lease_input_type;
pub mod external_generation_job_result_procedure_result_type;
pub mod external_generation_job_result_snapshot_type;
pub mod external_generation_job_retention_input_type;
pub mod external_generation_job_retention_procedure_result_type;
pub mod external_generation_job_snapshot_type;
pub mod external_generation_job_summary_backfill_input_type;
pub mod external_generation_job_summary_backfill_procedure_result_type;
@@ -578,6 +580,7 @@ pub mod profile_wallet_manual_restriction_table;
pub mod profile_wallet_manual_restriction_type;
pub mod profile_wallet_refund_outbox_table;
pub mod profile_wallet_refund_outbox_type;
pub mod prune_external_generation_job_history_and_return_procedure;
pub mod public_work_like_table;
pub mod public_work_like_type;
pub mod public_work_play_daily_stat_table;
@@ -1285,6 +1288,8 @@ pub use external_generation_job_procedure_result_type::ExternalGenerationJobProc
pub use external_generation_job_renew_lease_input_type::ExternalGenerationJobRenewLeaseInput;
pub use external_generation_job_result_procedure_result_type::ExternalGenerationJobResultProcedureResult;
pub use external_generation_job_result_snapshot_type::ExternalGenerationJobResultSnapshot;
pub use external_generation_job_retention_input_type::ExternalGenerationJobRetentionInput;
pub use external_generation_job_retention_procedure_result_type::ExternalGenerationJobRetentionProcedureResult;
pub use external_generation_job_snapshot_type::ExternalGenerationJobSnapshot;
pub use external_generation_job_summary_backfill_input_type::ExternalGenerationJobSummaryBackfillInput;
pub use external_generation_job_summary_backfill_procedure_result_type::ExternalGenerationJobSummaryBackfillProcedureResult;
@@ -1449,6 +1454,7 @@ pub use profile_wallet_manual_restriction_table::*;
pub use profile_wallet_manual_restriction_type::ProfileWalletManualRestriction;
pub use profile_wallet_refund_outbox_table::*;
pub use profile_wallet_refund_outbox_type::ProfileWalletRefundOutbox;
pub use prune_external_generation_job_history_and_return_procedure::prune_external_generation_job_history_and_return;
pub use public_work_like_table::*;
pub use public_work_like_type::PublicWorkLike;
pub use public_work_play_daily_stat_table::*;
@@ -56,6 +56,7 @@ impl __sdk::__query_builder::HasCols for ExternalGenerationJobEvent {
/// Provides typed access to indexed columns for query building.
pub struct ExternalGenerationJobEventIxCols {
pub event_id: __sdk::__query_builder::IxCol<ExternalGenerationJobEvent, String>,
pub job_id: __sdk::__query_builder::IxCol<ExternalGenerationJobEvent, String>,
}
impl __sdk::__query_builder::HasIxCols for ExternalGenerationJobEvent {
@@ -63,6 +64,7 @@ impl __sdk::__query_builder::HasIxCols for ExternalGenerationJobEvent {
fn ix_cols(table_name: &'static str) -> Self::IxCols {
ExternalGenerationJobEventIxCols {
event_id: __sdk::__query_builder::IxCol::new(table_name, "event_id"),
job_id: __sdk::__query_builder::IxCol::new(table_name, "job_id"),
}
}
}
@@ -0,0 +1,19 @@
// THIS FILE IS AUTOMATICALLY GENERATED BY SPACETIMEDB. EDITS TO THIS FILE
// WILL NOT BE SAVED. MODIFY TABLES IN YOUR MODULE SOURCE CODE INSTEAD.
#![allow(unused, clippy::all)]
use spacetimedb_sdk::__codegen::{self as __sdk, __lib, __sats, __ws};
#[derive(__lib::ser::Serialize, __lib::de::Deserialize, Clone, PartialEq, Debug)]
#[sats(crate = __lib)]
pub struct ExternalGenerationJobRetentionInput {
pub source_module: String,
pub limit: u32,
pub cursor_job_id: Option<String>,
pub completed_before_micros: i64,
pub dry_run: bool,
}
impl __sdk::InModule for ExternalGenerationJobRetentionInput {
type Module = super::RemoteModule;
}
@@ -0,0 +1,24 @@
// THIS FILE IS AUTOMATICALLY GENERATED BY SPACETIMEDB. EDITS TO THIS FILE
// WILL NOT BE SAVED. MODIFY TABLES IN YOUR MODULE SOURCE CODE INSTEAD.
#![allow(unused, clippy::all)]
use spacetimedb_sdk::__codegen::{self as __sdk, __lib, __sats, __ws};
#[derive(__lib::ser::Serialize, __lib::de::Deserialize, Clone, PartialEq, Debug)]
#[sats(crate = __lib)]
pub struct ExternalGenerationJobRetentionProcedureResult {
pub ok: bool,
pub dry_run: bool,
pub scanned_count: u64,
pub selected_count: u32,
pub deleted_job_count: u32,
pub deleted_summary_count: u32,
pub deleted_event_count: u32,
pub next_cursor_job_id: Option<String>,
pub has_more: bool,
pub error_message: Option<String>,
}
impl __sdk::InModule for ExternalGenerationJobRetentionProcedureResult {
type Module = super::RemoteModule;
}
@@ -0,0 +1,62 @@
// THIS FILE IS AUTOMATICALLY GENERATED BY SPACETIMEDB. EDITS TO THIS FILE
// WILL NOT BE SAVED. MODIFY TABLES IN YOUR MODULE SOURCE CODE INSTEAD.
#![allow(unused, clippy::all)]
use spacetimedb_sdk::__codegen::{self as __sdk, __lib, __sats, __ws};
use super::external_generation_job_retention_input_type::ExternalGenerationJobRetentionInput;
use super::external_generation_job_retention_procedure_result_type::ExternalGenerationJobRetentionProcedureResult;
#[derive(__lib::ser::Serialize, __lib::de::Deserialize, Clone, PartialEq, Debug)]
#[sats(crate = __lib)]
struct PruneExternalGenerationJobHistoryAndReturnArgs {
pub input: ExternalGenerationJobRetentionInput,
}
impl __sdk::InModule for PruneExternalGenerationJobHistoryAndReturnArgs {
type Module = super::RemoteModule;
}
#[allow(non_camel_case_types)]
/// Extension trait for access to the procedure `prune_external_generation_job_history_and_return`.
///
/// Implemented for [`super::RemoteProcedures`].
pub trait prune_external_generation_job_history_and_return {
fn prune_external_generation_job_history_and_return(
&self,
input: ExternalGenerationJobRetentionInput,
) {
self.prune_external_generation_job_history_and_return_then(input, |_, _| {});
}
fn prune_external_generation_job_history_and_return_then(
&self,
input: ExternalGenerationJobRetentionInput,
__callback: impl FnOnce(
&super::ProcedureEventContext,
Result<ExternalGenerationJobRetentionProcedureResult, __sdk::InternalError>,
) + Send
+ 'static,
);
}
impl prune_external_generation_job_history_and_return for super::RemoteProcedures {
fn prune_external_generation_job_history_and_return_then(
&self,
input: ExternalGenerationJobRetentionInput,
__callback: impl FnOnce(
&super::ProcedureEventContext,
Result<ExternalGenerationJobRetentionProcedureResult, __sdk::InternalError>,
) + Send
+ 'static,
) {
self.imp
.invoke_procedure_with_callback::<_, ExternalGenerationJobRetentionProcedureResult>(
"prune_external_generation_job_history_and_return",
PruneExternalGenerationJobHistoryAndReturnArgs { input },
__callback,
);
}
}
@@ -139,17 +139,6 @@ pub(crate) fn build_ai_text_chunk_row_id(snapshot: &AiTextChunkSnapshot) -> Stri
)
}
pub(crate) fn build_ai_text_chunk_snapshot_from_row(row: &AiTextChunk) -> AiTextChunkSnapshot {
AiTextChunkSnapshot {
chunk_id: row.chunk_id.clone(),
task_id: row.task_id.clone(),
stage_kind: row.stage_kind,
sequence: row.sequence,
delta_text: row.delta_text.clone(),
created_at_micros: row.created_at.to_micros_since_unix_epoch(),
}
}
pub(crate) fn build_ai_result_reference_row(
snapshot: &AiResultReferenceSnapshot,
) -> AiResultReference {
@@ -1,7 +1,7 @@
use crate::*;
use module_ai::{
generate_ai_result_ref_id, generate_ai_text_chunk_id, normalize_optional_text,
normalize_string_list,
MAX_AI_TASK_TEXT_OUTPUT_BYTES, generate_ai_result_ref_id, generate_ai_text_chunk_id,
normalize_optional_text, normalize_string_list, validate_ai_task_snapshot_memory_limits,
};
#[spacetimedb::table(
@@ -178,6 +178,9 @@ pub(crate) fn append_ai_text_chunk_tx(
if input.sequence == 0 {
return Err("ai_text_chunk.sequence 必须大于 0".to_string());
}
if input.delta_text.trim().len() > MAX_AI_TASK_TEXT_OUTPUT_BYTES {
return Err("AI 任务文本输出超过内存上限".to_string());
}
let mut snapshot = get_ai_task_snapshot_tx(ctx, &input.task_id)?;
ensure_ai_task_can_transition(snapshot.status)?;
@@ -200,7 +203,7 @@ pub(crate) fn append_ai_text_chunk_tx(
.ai_text_chunk()
.insert(build_ai_text_chunk_row(&chunk));
let aggregated_text = collect_ai_stage_text_output(ctx, &chunk.task_id, chunk.stage_kind);
let aggregated_text = collect_ai_stage_text_output(ctx, &chunk.task_id, chunk.stage_kind)?;
snapshot.status = AiTaskStatus::Running;
if snapshot.started_at_micros.is_none() {
@@ -215,6 +218,7 @@ pub(crate) fn append_ai_text_chunk_tx(
snapshot.updated_at_micros = input.created_at_micros;
snapshot.version += 1;
validate_ai_task_snapshot_memory_limits(&snapshot).map_err(str::to_string)?;
persist_ai_task_snapshot(ctx, &snapshot)?;
emit_ai_task_event(
ctx,
@@ -252,6 +256,7 @@ pub(crate) fn complete_ai_stage_tx(
snapshot.updated_at_micros = input.completed_at_micros;
snapshot.version += 1;
validate_ai_task_snapshot_memory_limits(&snapshot).map_err(str::to_string)?;
persist_ai_task_snapshot(ctx, &snapshot)?;
emit_ai_task_event(
ctx,
@@ -285,26 +290,27 @@ pub(crate) fn attach_ai_result_reference_tx(
label: normalize_optional_text(input.label),
created_at_micros: input.created_at_micros,
};
ctx.db
.ai_result_reference()
.insert(build_ai_result_reference_row(&reference));
snapshot.result_references.push(reference);
snapshot.updated_at_micros = input.created_at_micros;
snapshot.version += 1;
persist_ai_task_snapshot(ctx, &snapshot)?;
validate_ai_task_snapshot_memory_limits(&snapshot).map_err(str::to_string)?;
let reference = snapshot
.result_references
.last()
.cloned()
.ok_or_else(|| "ai_result_reference 写入后缺少快照".to_string())?;
ctx.db
.ai_result_reference()
.insert(build_ai_result_reference_row(&reference));
persist_ai_task_snapshot(ctx, &snapshot)?;
emit_ai_task_event(
ctx,
&snapshot,
AiTaskEventKind::ResultReferenceAttached,
None,
None,
Some(build_ai_result_reference_row_id(reference)),
Some(build_ai_result_reference_row_id(&reference)),
input.created_at_micros,
);
Ok(snapshot)
@@ -333,29 +339,48 @@ pub(crate) fn replace_ai_task_stages(
}
}
pub(crate) fn delete_ai_text_chunks_for_task(ctx: &ReducerContext, task_id: &str) {
let chunk_row_ids = ctx
.db
.ai_text_chunk()
.by_ai_text_chunk_task_id()
.filter(task_id)
.map(|row| row.text_chunk_row_id.clone())
.collect::<Vec<_>>();
for row_id in chunk_row_ids {
ctx.db.ai_text_chunk().text_chunk_row_id().delete(&row_id);
}
}
pub(crate) fn collect_ai_stage_text_output(
ctx: &ReducerContext,
task_id: &str,
stage_kind: AiTaskStageKind,
) -> Option<String> {
let mut chunks = ctx
) -> Result<Option<String>, String> {
let mut chunks = Vec::new();
let mut aggregated_bytes = 0_usize;
for row in ctx
.db
.ai_text_chunk()
.by_ai_text_chunk_task_id()
.filter(task_id)
.filter(|row| row.task_id == task_id && row.stage_kind == stage_kind)
.map(|row| build_ai_text_chunk_snapshot_from_row(&row))
.collect::<Vec<_>>();
chunks.sort_by_key(|chunk| chunk.sequence);
{
aggregated_bytes = aggregated_bytes.saturating_add(row.delta_text.len());
if aggregated_bytes > MAX_AI_TASK_TEXT_OUTPUT_BYTES {
return Err("AI 任务文本输出超过内存上限".to_string());
}
chunks.push((row.sequence, row.delta_text.clone()));
}
chunks.sort_by_key(|(sequence, _)| *sequence);
let aggregated = chunks
.into_iter()
.map(|chunk| chunk.delta_text)
.collect::<Vec<_>>()
.join("");
let mut aggregated = String::with_capacity(aggregated_bytes);
for (_, delta) in chunks {
aggregated.push_str(&delta);
}
if aggregated.trim().is_empty() {
None
Ok(None)
} else {
Some(aggregated)
Ok(Some(aggregated))
}
}
@@ -1,5 +1,8 @@
use crate::*;
use module_ai::{INITIAL_AI_TASK_VERSION, normalize_optional_text, validate_task_create_input};
use module_ai::{
INITIAL_AI_TASK_VERSION, normalize_optional_text, validate_ai_task_snapshot_memory_limits,
validate_task_create_input,
};
#[spacetimedb::table(
accessor = ai_task,
@@ -133,6 +136,7 @@ fn create_ai_task_tx(
}
let task_snapshot = build_ai_task_snapshot_from_create_input(&input);
validate_ai_task_snapshot_memory_limits(&task_snapshot).map_err(str::to_string)?;
ctx.db.ai_task().insert(build_ai_task_row(&task_snapshot));
replace_ai_task_stages(ctx, &task_snapshot.task_id, &task_snapshot.stages);
emit_ai_task_event(
@@ -187,7 +191,9 @@ fn complete_ai_task_tx(
snapshot.updated_at_micros = input.completed_at_micros;
snapshot.version += 1;
validate_ai_task_snapshot_memory_limits(&snapshot).map_err(str::to_string)?;
persist_ai_task_snapshot(ctx, &snapshot)?;
delete_ai_text_chunks_for_task(ctx, &snapshot.task_id);
emit_ai_task_event(
ctx,
&snapshot,
@@ -218,7 +224,9 @@ fn fail_ai_task_tx(
snapshot.updated_at_micros = input.completed_at_micros;
snapshot.version += 1;
validate_ai_task_snapshot_memory_limits(&snapshot).map_err(str::to_string)?;
persist_ai_task_snapshot(ctx, &snapshot)?;
delete_ai_text_chunks_for_task(ctx, &snapshot.task_id);
emit_ai_task_event(
ctx,
&snapshot,
@@ -243,7 +251,9 @@ fn cancel_ai_task_tx(
snapshot.updated_at_micros = input.completed_at_micros;
snapshot.version += 1;
validate_ai_task_snapshot_memory_limits(&snapshot).map_err(str::to_string)?;
persist_ai_task_snapshot(ctx, &snapshot)?;
delete_ai_text_chunks_for_task(ctx, &snapshot.task_id);
emit_ai_task_event(
ctx,
&snapshot,
@@ -97,6 +97,10 @@ pub struct ExternalGenerationJob {
accessor = by_external_generation_job_event_job_id,
btree(columns = [job_id, created_at])
),
index(
accessor = by_external_generation_job_event_job_id_only,
btree(columns = [job_id])
),
index(
accessor = by_external_generation_job_event_owner,
btree(columns = [owner_user_id, created_at])
@@ -378,6 +382,29 @@ pub struct ExternalGenerationJobPayloadCompactionProcedureResult {
pub error_message: Option<String>,
}
#[derive(Clone, Debug, PartialEq, Eq, SpacetimeType)]
pub struct ExternalGenerationJobRetentionInput {
pub source_module: String,
pub limit: u32,
pub cursor_job_id: Option<String>,
pub completed_before_micros: i64,
pub dry_run: bool,
}
#[derive(Clone, Debug, PartialEq, Eq, SpacetimeType)]
pub struct ExternalGenerationJobRetentionProcedureResult {
pub ok: bool,
pub dry_run: bool,
pub scanned_count: u64,
pub selected_count: u32,
pub deleted_job_count: u32,
pub deleted_summary_count: u32,
pub deleted_event_count: u32,
pub next_cursor_job_id: Option<String>,
pub has_more: bool,
pub error_message: Option<String>,
}
#[derive(Clone, Debug, PartialEq, Eq, SpacetimeType)]
pub struct ExternalGenerationQueueStatsSnapshot {
pub pending_count: u32,
@@ -676,6 +703,21 @@ pub fn compact_external_generation_job_payloads_and_return(
}
}
#[spacetimedb::procedure]
pub fn prune_external_generation_job_history_and_return(
ctx: &mut ProcedureContext,
input: ExternalGenerationJobRetentionInput,
) -> ExternalGenerationJobRetentionProcedureResult {
let caller = ctx.sender();
match ctx.try_with_tx(|tx| {
crate::migration::require_migration_operator(tx, caller)?;
prune_external_generation_job_history_tx(tx, input.clone())
}) {
Ok(result) => result,
Err(message) => failed_external_generation_job_retention_result(input.dry_run, message),
}
}
#[spacetimedb::procedure]
pub fn get_external_generation_queue_stats_and_return(
ctx: &mut ProcedureContext,
@@ -1369,6 +1411,105 @@ fn compact_external_generation_job_payloads_tx(
})
}
fn prune_external_generation_job_history_tx(
ctx: &ReducerContext,
input: ExternalGenerationJobRetentionInput,
) -> Result<ExternalGenerationJobRetentionProcedureResult, String> {
let source_module = input.source_module.trim().to_string();
validate_required("external_generation_job.source_module", &source_module)?;
let now_micros = ctx.timestamp.to_micros_since_unix_epoch();
if input.completed_before_micros > now_micros {
return Err(
"external_generation_job.completed_before_micros 不能晚于数据库当前时间".to_string(),
);
}
let cursor_job_id = input
.cursor_job_id
.as_deref()
.and_then(normalize_optional_text);
let limit = input
.limit
.clamp(1, MAX_EXTERNAL_GENERATION_MAINTENANCE_BATCH_SIZE) as usize;
let cursor_range = external_generation_job_maintenance_cursor_range(cursor_job_id.as_deref());
let cursor_to_skip = cursor_job_id.clone();
let rows = ctx
.db
.external_generation_job()
.by_external_generation_job_source_cursor()
.filter((source_module.as_str(), cursor_range))
.filter(move |row| {
cursor_to_skip
.as_deref()
.is_none_or(|cursor| row.job_id != cursor)
});
let (job_ids, next_cursor_job_id, has_more, scanned_count) =
select_external_generation_job_ids_for_maintenance(rows, limit, |row| {
ctx.db
.external_generation_job_summary()
.job_id()
.find(&row.job_id)
.is_some_and(|summary| {
is_external_generation_job_retention_candidate(
row,
&summary,
&source_module,
input.completed_before_micros,
)
})
});
let mut deleted_job_count = 0u32;
let mut deleted_summary_count = 0u32;
let mut deleted_event_count = 0u32;
if !input.dry_run {
for job_id in &job_ids {
let Some(row) = ctx.db.external_generation_job().job_id().find(job_id) else {
continue;
};
let Some(summary) = ctx
.db
.external_generation_job_summary()
.job_id()
.find(job_id)
else {
continue;
};
if !is_external_generation_job_retention_candidate(
&row,
&summary,
&source_module,
input.completed_before_micros,
) {
continue;
}
deleted_event_count = deleted_event_count
.saturating_add(delete_external_generation_job_events_for_job(ctx, job_id));
ctx.db
.external_generation_job_summary()
.job_id()
.delete(job_id);
deleted_summary_count = deleted_summary_count.saturating_add(1);
ctx.db.external_generation_job().job_id().delete(job_id);
deleted_job_count = deleted_job_count.saturating_add(1);
}
}
Ok(ExternalGenerationJobRetentionProcedureResult {
ok: true,
dry_run: input.dry_run,
scanned_count,
selected_count: job_ids.len() as u32,
deleted_job_count,
deleted_summary_count,
deleted_event_count,
next_cursor_job_id,
has_more,
error_message: None,
})
}
fn renew_external_generation_job_lease_tx(
ctx: &ReducerContext,
input: ExternalGenerationJobRenewLeaseInput,
@@ -1801,6 +1942,25 @@ fn should_compact_external_generation_job_payloads(
})
}
fn is_external_generation_job_retention_candidate(
row: &ExternalGenerationJob,
summary: &ExternalGenerationJobSummary,
source_module: &str,
completed_before_micros: i64,
) -> bool {
row.source_module.trim() == source_module.trim()
&& summary.job_id == row.job_id
&& summary.status == row.status
&& is_external_generation_job_terminal(row)
&& is_external_generation_job_summary_terminal(summary)
&& summary.notification_acknowledged_at.is_some()
&& row
.completed_at
.unwrap_or(row.updated_at)
.to_micros_since_unix_epoch()
<= completed_before_micros
}
fn external_generation_job_maintenance_cursor_range(
cursor_job_id: Option<&str>,
) -> RangeFrom<&str> {
@@ -1836,6 +1996,24 @@ fn select_external_generation_job_ids_for_maintenance(
)
}
fn delete_external_generation_job_events_for_job(ctx: &ReducerContext, job_id: &str) -> u32 {
let event_ids = ctx
.db
.external_generation_job_event()
.by_external_generation_job_event_job_id_only()
.filter(job_id)
.map(|event| event.event_id.clone())
.collect::<Vec<_>>();
let deleted_count = event_ids.len() as u32;
for event_id in event_ids {
ctx.db
.external_generation_job_event()
.event_id()
.delete(&event_id);
}
deleted_count
}
fn count_external_generation_job_summaries_for_owner(
ctx: &ReducerContext,
owner_user_id: &str,
@@ -2548,6 +2726,24 @@ fn failed_external_generation_job_payload_compaction_result(
}
}
fn failed_external_generation_job_retention_result(
dry_run: bool,
message: String,
) -> ExternalGenerationJobRetentionProcedureResult {
ExternalGenerationJobRetentionProcedureResult {
ok: false,
dry_run,
scanned_count: 0,
selected_count: 0,
deleted_job_count: 0,
deleted_summary_count: 0,
deleted_event_count: 0,
next_cursor_job_id: None,
has_more: false,
error_message: Some(message),
}
}
fn validate_required(field: &str, value: &str) -> Result<(), String> {
if value.trim().is_empty() {
return Err(format!("{field} 不能为空"));
@@ -3438,6 +3634,86 @@ mod tests {
));
}
#[test]
fn retention_only_selects_acknowledged_terminal_rows_before_cutoff() {
let mut row = external_generation_job_fixture(EXTERNAL_GENERATION_STATUS_COMPLETED);
row.source_module = EXTERNAL_GENERATION_EDITOR_SOURCE_MODULE.to_string();
row.completed_at = Some(micros(1_000));
row.updated_at = micros(1_000);
let mut summary = build_external_generation_job_summary_row(&row, None);
summary.notification_acknowledged_at = Some(micros(2_000));
assert!(is_external_generation_job_retention_candidate(
&row,
&summary,
EXTERNAL_GENERATION_EDITOR_SOURCE_MODULE,
1_000,
));
summary.notification_acknowledged_at = None;
assert!(!is_external_generation_job_retention_candidate(
&row,
&summary,
EXTERNAL_GENERATION_EDITOR_SOURCE_MODULE,
1_000,
));
summary.notification_acknowledged_at = Some(micros(2_000));
row.status = EXTERNAL_GENERATION_STATUS_RUNNING.to_string();
summary.status = EXTERNAL_GENERATION_STATUS_RUNNING.to_string();
assert!(!is_external_generation_job_retention_candidate(
&row,
&summary,
EXTERNAL_GENERATION_EDITOR_SOURCE_MODULE,
1_000,
));
row.status = EXTERNAL_GENERATION_STATUS_COMPLETED.to_string();
summary.status = EXTERNAL_GENERATION_STATUS_COMPLETED.to_string();
row.completed_at = Some(micros(1_001));
assert!(!is_external_generation_job_retention_candidate(
&row,
&summary,
EXTERNAL_GENERATION_EDITOR_SOURCE_MODULE,
1_000,
));
row.completed_at = Some(micros(1_000));
row.source_module = "puzzle".to_string();
assert!(!is_external_generation_job_retention_candidate(
&row,
&summary,
EXTERNAL_GENERATION_EDITOR_SOURCE_MODULE,
1_000,
));
}
#[test]
fn retention_rejects_mismatched_summary_identity_or_status() {
let mut row = external_generation_job_fixture(EXTERNAL_GENERATION_STATUS_FAILED);
row.source_module = EXTERNAL_GENERATION_EDITOR_SOURCE_MODULE.to_string();
row.completed_at = Some(micros(1_000));
let mut summary = build_external_generation_job_summary_row(&row, None);
summary.notification_acknowledged_at = Some(micros(2_000));
summary.job_id = "different-job".to_string();
assert!(!is_external_generation_job_retention_candidate(
&row,
&summary,
EXTERNAL_GENERATION_EDITOR_SOURCE_MODULE,
1_000,
));
summary.job_id = row.job_id.clone();
summary.status = EXTERNAL_GENERATION_STATUS_CANCELLED.to_string();
assert!(!is_external_generation_job_retention_candidate(
&row,
&summary,
EXTERNAL_GENERATION_EDITOR_SOURCE_MODULE,
1_000,
));
}
#[test]
fn maintenance_selector_bounds_scanned_rows_and_advances_by_last_scanned_job() {
let rows = (1..=4).map(|index| {