合并 master 并解决鉴权与生成任务冲突
合并 master 的外部生成历史清理、任务内存上限和备份 OOM 兜底 保留退款 outbox 与短期认证 typed projection/CAS 及容量清理 合并 SpacetimeDB 生成 bindings、架构文档和决策记录 补齐认证投影容量测试字段并通过定向门禁
This commit is contained in:
@@ -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()))
|
||||
}
|
||||
|
||||
@@ -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(())
|
||||
}
|
||||
@@ -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;
|
||||
|
||||
@@ -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();
|
||||
|
||||
@@ -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::*;
|
||||
|
||||
+2
@@ -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"),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
+19
@@ -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;
|
||||
}
|
||||
+24
@@ -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;
|
||||
}
|
||||
+62
@@ -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| {
|
||||
|
||||
Reference in New Issue
Block a user