fix: harden submit queue durability

This commit is contained in:
2026-09-06 01:35:26 +09:00
parent bb4c1dfe4f
commit b038f022d3
14 changed files with 1370 additions and 161 deletions
+259 -43
View File
@@ -229,6 +229,43 @@ enum PendingRun {
Resume,
}
fn stage_pending_notification<St: Store + Clone>(
pending_submissions: &crate::worker::PendingSubmissionHandle<St>,
notify_buffer: &NotifyBuffer,
source_namespace: &str,
notification_request_id: &str,
) -> bool {
let Some(notification) =
pending_submissions.prepare_notification(source_namespace, notification_request_id)
else {
return false;
};
let extension = pending_submissions.notification_activation_extension();
notify_buffer.push_durable_notify(
notification.message,
notification.auto_run,
notification.provenance,
extension,
);
true
}
fn stage_oldest_passive_notification<St: Store + Clone>(
pending_submissions: &crate::worker::PendingSubmissionHandle<St>,
notify_buffer: &NotifyBuffer,
) -> bool {
pending_submissions
.next_passive_notification_identity()
.is_some_and(|(source_namespace, request_id)| {
stage_pending_notification(
pending_submissions,
notify_buffer,
&source_namespace,
&request_id,
)
})
}
fn prepare_pending_run<St: Store + Clone>(
pending_submissions: &crate::worker::PendingSubmissionHandle<St>,
notify_buffer: &NotifyBuffer,
@@ -241,7 +278,12 @@ fn prepare_pending_run<St: Store + Clone>(
Some(crate::worker::PendingActivation::Notification(notification)) => {
let extension = pending_submissions.notification_activation_extension();
let notification_request_id = notification.notification_request_id.clone();
notify_buffer.push_durable_notify(notification.message, extension);
notify_buffer.push_durable_notify(
notification.message,
notification.auto_run,
notification.provenance,
extension,
);
Some(PendingRun::RunForNotification {
invoke_kind: protocol::InvokeKind::Notify,
notification_request_id: Some(notification_request_id),
@@ -1314,6 +1356,7 @@ async fn controller_loop<C, St>(
);
let mut pending: Option<PendingRun> = None;
let pending_submissions = worker.pending_submission_handle();
stage_oldest_passive_notification(&pending_submissions, &notify_buffer);
loop {
// Top-of-iteration: if an event handler staged a run, fire it
@@ -1347,6 +1390,8 @@ async fn controller_loop<C, St>(
} => notification_request_id.clone(),
_ => None,
};
let passive_notification_request_id =
pending_submissions.activating_passive_notification_id();
let (mut new_status, shutdown, may_drain_pending) = match run {
PendingRun::Submit(submission) => {
let (input_commit_tx, input_commit_rx) = oneshot::channel();
@@ -1356,6 +1401,7 @@ async fn controller_loop<C, St>(
worker.run_with_input_extensions_and_commit_hook(
submission.input,
vec![extension],
submission.provenance,
move || {
let _ = input_commit_tx.send(());
},
@@ -1415,8 +1461,11 @@ async fn controller_loop<C, St>(
.await
}
};
if let Some(notification_request_id) = notification_request_id {
if let Some(notification_request_id) =
notification_request_id.or(passive_notification_request_id)
{
pending_submissions.finish_notification_activation(&notification_request_id);
stage_oldest_passive_notification(&pending_submissions, &notify_buffer);
}
if !shutdown && may_drain_pending && new_status == WorkerStatus::Idle {
@@ -1470,13 +1519,47 @@ async fn controller_loop<C, St>(
Method::Submit {
submission_request_id,
input,
}
| Method::SubmitTracked {
submission_request_id,
input,
} => {
let request_id = submission_request_id.clone();
match pending_submissions.accept(submission_request_id, input, true) {
match pending_submissions.accept_from_source(
submission_request_id,
input,
pending_submissions.direct_client_namespace(),
session_store::LoggedSessionHistoryOrigin::LegacyUnknown,
true,
) {
Ok(acceptance) => {
if let Some(activation) = acceptance.activation {
pending = Some(PendingRun::Submit(activation));
} else {
let _ = working_event_tx.send(Event::SubmissionAccepted {
submission_request_id: acceptance.submission_request_id,
submission_id: acceptance.submission_id,
disposition: acceptance.disposition,
});
}
}
Err(error) => {
let _ = working_event_tx.send(Event::SubmissionRejected {
submission_request_id: request_id,
message: error.to_string(),
});
}
}
}
Method::SubmitTracked {
submission_request_id,
input,
source,
} => {
let request_id = submission_request_id.clone();
match pending_submissions.accept_from_source(
submission_request_id,
input,
source.namespace(),
crate::worker::authenticated_input_provenance(&source),
true,
) {
Ok(acceptance) => {
if let Some(activation) = acceptance.activation {
pending = Some(PendingRun::Submit(activation));
@@ -1502,31 +1585,85 @@ async fn controller_loop<C, St>(
message,
auto_run,
} => {
if auto_run {
match pending_submissions.accept_notification(notification_request_id, message)
{
Ok(true) => {
match prepare_pending_run(&pending_submissions, &notify_buffer, None) {
Ok(Some(next)) => pending = Some(next),
Ok(None) => {}
Err(error) => {
let _ = working_event_tx.send(Event::Error {
code: ErrorCode::Internal,
message: error.to_string(),
});
}
let request_id = notification_request_id.clone();
let source_namespace = pending_submissions.direct_client_namespace();
match pending_submissions.accept_notification_from_source(
notification_request_id,
message,
source_namespace.clone(),
session_store::LoggedSessionHistoryOrigin::LegacyUnknown,
auto_run,
) {
Ok(_) if auto_run => {
match prepare_pending_run(&pending_submissions, &notify_buffer, None) {
Ok(Some(next)) => pending = Some(next),
Ok(None) => {}
Err(error) => {
let _ = working_event_tx.send(Event::Error {
code: ErrorCode::Internal,
message: error.to_string(),
});
}
}
Ok(false) => {}
Err(error) => {
let _ = working_event_tx.send(Event::Error {
code: ErrorCode::InvalidRequest,
message: error.to_string(),
});
}
Ok(_) => {
stage_pending_notification(
&pending_submissions,
&notify_buffer,
&source_namespace,
&request_id,
);
}
Err(error) => {
let _ = working_event_tx.send(Event::Error {
code: ErrorCode::InvalidRequest,
message: error.to_string(),
});
}
}
}
Method::NotifyTracked {
notification_request_id,
message,
auto_run,
source,
} => {
let request_id = notification_request_id.clone();
let source_namespace = source.namespace();
match pending_submissions.accept_notification_from_source(
notification_request_id,
message,
source_namespace.clone(),
crate::worker::authenticated_input_provenance(&source),
auto_run,
) {
Ok(_) if auto_run => {
match prepare_pending_run(&pending_submissions, &notify_buffer, None) {
Ok(Some(next)) => pending = Some(next),
Ok(None) => {}
Err(error) => {
let _ = working_event_tx.send(Event::Error {
code: ErrorCode::Internal,
message: error.to_string(),
});
}
}
}
} else {
worker.push_notify(message, false);
Ok(_) => {
stage_pending_notification(
&pending_submissions,
&notify_buffer,
&source_namespace,
&request_id,
);
}
Err(error) => {
let _ = working_event_tx.send(Event::Error {
code: ErrorCode::InvalidRequest,
message: error.to_string(),
});
}
}
}
@@ -2070,13 +2207,46 @@ where
Some(Method::Submit {
submission_request_id,
input,
}
| Method::SubmitTracked {
submission_request_id,
input,
}) => {
let request_id = submission_request_id.clone();
match pending_submissions.accept(submission_request_id, input, false) {
match pending_submissions.accept_from_source(
submission_request_id,
input,
pending_submissions.direct_client_namespace(),
session_store::LoggedSessionHistoryOrigin::LegacyUnknown,
false,
) {
Ok(acceptance) => {
let _ = working_event_tx.send(Event::SubmissionAccepted {
submission_request_id: acceptance.submission_request_id,
submission_id: acceptance.submission_id,
disposition: acceptance.disposition,
});
let _ = working_event_tx.send(Event::PendingSubmissionsChanged {
pending: pending_submissions.snapshot(),
});
}
Err(error) => {
let _ = working_event_tx.send(Event::SubmissionRejected {
submission_request_id: request_id,
message: error.to_string(),
});
}
}
}
Some(Method::SubmitTracked {
submission_request_id,
input,
source,
}) => {
let request_id = submission_request_id.clone();
match pending_submissions.accept_from_source(
submission_request_id,
input,
source.namespace(),
crate::worker::authenticated_input_provenance(&source),
false,
) {
Ok(acceptance) => {
let _ = working_event_tx.send(Event::SubmissionAccepted {
submission_request_id: acceptance.submission_request_id,
@@ -2147,23 +2317,69 @@ where
message,
auto_run,
}) => {
if auto_run {
if let Err(error) = pending_submissions.accept_notification(
notification_request_id,
message,
) {
let request_id = notification_request_id.clone();
let source_namespace = pending_submissions.direct_client_namespace();
match pending_submissions.accept_notification_from_source(
notification_request_id,
message,
source_namespace.clone(),
session_store::LoggedSessionHistoryOrigin::LegacyUnknown,
auto_run,
) {
Ok(_) if !auto_run => {
stage_pending_notification(
&pending_submissions,
notify_buffer,
&source_namespace,
&request_id,
);
}
Ok(_) => {}
Err(error) => {
let _ = working_event_tx.send(Event::Error {
code: ErrorCode::InvalidRequest,
message: error.to_string(),
});
} else {
let _ = working_event_tx.send(Event::PendingSubmissionsChanged {
pending: pending_submissions.snapshot(),
}
}
let _ = working_event_tx.send(Event::PendingSubmissionsChanged {
pending: pending_submissions.snapshot(),
});
}
Some(Method::NotifyTracked {
notification_request_id,
message,
auto_run,
source,
}) => {
let request_id = notification_request_id.clone();
let source_namespace = source.namespace();
match pending_submissions.accept_notification_from_source(
notification_request_id,
message,
source_namespace.clone(),
crate::worker::authenticated_input_provenance(&source),
auto_run,
) {
Ok(_) if !auto_run => {
stage_pending_notification(
&pending_submissions,
notify_buffer,
&source_namespace,
&request_id,
);
}
Ok(_) => {}
Err(error) => {
let _ = working_event_tx.send(Event::Error {
code: ErrorCode::InvalidRequest,
message: error.to_string(),
});
}
} else {
notify_buffer.push_notify(message, false);
}
let _ = working_event_tx.send(Event::PendingSubmissionsChanged {
pending: pending_submissions.snapshot(),
});
}
Some(Method::ListCompletions { .. }) => {}
Some(Method::ListWorkers | Method::RestoreWorker { .. } | Method::RegisterPeer { .. }) => {
+18 -8
View File
@@ -178,14 +178,21 @@ impl WorkerInterceptor {
/// matches worker-history order.
fn commit_system_items_with_extensions(
&self,
items: &[(SystemItem, Vec<session_store::SessionExtension>)],
items: &[(
SystemItem,
Vec<session_store::SessionExtension>,
Option<session_store::LoggedSessionHistoryOrigin>,
)],
) -> Result<(), session_store::StoreError> {
let Some(writer) = self.log_writer.as_ref() else {
return Ok(());
};
for (item, extensions) in items {
let entry =
writer.commit_system_item_with_extensions(item.clone(), extensions.clone())?;
for (item, extensions, history_provenance) in items {
let entry = writer.commit_system_item_with_extensions(
item.clone(),
extensions.clone(),
history_provenance.clone(),
)?;
self.pending_committed_history
.lock()
.expect("pending committed history poisoned")
@@ -199,7 +206,7 @@ impl WorkerInterceptor {
&items
.iter()
.cloned()
.map(|item| (item, Vec::new()))
.map(|item| (item, Vec::new(), None))
.collect::<Vec<_>>(),
)
}
@@ -341,8 +348,11 @@ impl Interceptor<SessionHistoryMetadata> for WorkerInterceptor {
projection_digest: projection.catalog_digest.clone(),
logical_name: "internal.notify_wrapper".to_string(),
};
let mut system_items: Vec<(SystemItem, Vec<session_store::SessionExtension>)> =
Vec::with_capacity(drained.len());
let mut system_items: Vec<(
SystemItem,
Vec<session_store::SessionExtension>,
Option<session_store::LoggedSessionHistoryOrigin>,
)> = Vec::with_capacity(drained.len());
let mut items: Vec<Item> = Vec::with_capacity(drained.len());
for entry in &drained {
let system_item = match build_system_item_with_provenance(
@@ -360,7 +370,7 @@ impl Interceptor<SessionHistoryMetadata> for WorkerInterceptor {
}
};
items.push(system_item.to_history_item());
system_items.push((system_item, entry.extensions()));
system_items.push((system_item, entry.extensions(), entry.history_provenance()));
}
if let Err(error) = self.commit_system_items_with_extensions(&system_items) {
self.pending_notifies.requeue_front(drained);
+22 -3
View File
@@ -25,7 +25,7 @@ use std::collections::VecDeque;
use std::sync::{Arc, Mutex};
use protocol::WorkerEvent;
use session_store::{SessionExtension, SystemItem};
use session_store::{LoggedSessionHistoryOrigin, SessionExtension, SystemItem};
use tracing::warn;
use crate::prompt::catalog::{CatalogError, PromptCatalog};
@@ -45,6 +45,7 @@ pub enum PendingNotify {
message: String,
auto_run: bool,
extensions: Vec<SessionExtension>,
history_provenance: Option<LoggedSessionHistoryOrigin>,
},
WorkerEvent {
event: WorkerEvent,
@@ -58,6 +59,15 @@ impl PendingNotify {
PendingNotify::WorkerEvent { .. } => Vec::new(),
}
}
pub(crate) fn history_provenance(&self) -> Option<LoggedSessionHistoryOrigin> {
match self {
PendingNotify::Notify {
history_provenance, ..
} => history_provenance.clone(),
PendingNotify::WorkerEvent { .. } => None,
}
}
}
/// Shared, mutex-guarded buffer of pending entries.
@@ -81,14 +91,22 @@ impl NotifyBuffer {
message,
auto_run,
extensions: Vec::new(),
history_provenance: None,
});
}
pub fn push_durable_notify(&self, message: String, extension: SessionExtension) {
pub fn push_durable_notify(
&self,
message: String,
auto_run: bool,
history_provenance: LoggedSessionHistoryOrigin,
extension: SessionExtension,
) {
self.push_entry(PendingNotify::Notify {
message,
auto_run: true,
auto_run,
extensions: vec![extension],
history_provenance: Some(history_provenance),
});
}
@@ -230,6 +248,7 @@ mod tests {
message: "hello".into(),
auto_run: false,
extensions: Vec::new(),
history_provenance: None,
};
let catalog = PromptCatalog::builtins_only().unwrap();
let item = build_system_item(&entry, &catalog).unwrap();
+557 -48
View File
@@ -79,11 +79,12 @@ const MAX_SUBMISSION_RECEIPTS: usize = 128;
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
pub(crate) struct PendingSubmission {
pub(crate) submission_request_id: String,
source_namespace: String,
pub(crate) submission_id: String,
payload_digest: String,
accepted_at_ms: u64,
activation_sequence: u64,
provenance: WorkerHistoryProvenance,
pub(crate) provenance: WorkerHistoryProvenance,
#[serde(default)]
was_queued: bool,
pub(crate) input: Vec<Segment>,
@@ -92,6 +93,7 @@ pub(crate) struct PendingSubmission {
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
struct SubmissionReceipt {
submission_request_id: String,
source_namespace: String,
submission_id: String,
payload_digest: String,
disposition: protocol::SubmissionDisposition,
@@ -100,17 +102,21 @@ struct SubmissionReceipt {
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
pub(crate) struct PendingNotification {
pub(crate) notification_request_id: String,
source_namespace: String,
pub(crate) message: String,
payload_digest: String,
pub(crate) auto_run: bool,
accepted_at_ms: u64,
activation_sequence: u64,
provenance: WorkerHistoryProvenance,
pub(crate) provenance: WorkerHistoryProvenance,
}
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
struct NotificationReceipt {
notification_request_id: String,
source_namespace: String,
payload_digest: String,
auto_run: bool,
}
#[derive(Debug, Clone, Default, serde::Serialize, serde::Deserialize)]
@@ -130,14 +136,24 @@ pub(crate) struct PendingActivationState {
impl PendingActivationState {
pub(crate) fn snapshot(&self) -> protocol::PendingSubmissionsSnapshot {
let head_id = match (self.pending.front(), self.pending_notifications.front()) {
let pending_notification = self
.pending_notifications
.iter()
.find(|notification| notification.auto_run);
let head_id = match (self.pending.front(), pending_notification) {
(Some(submission), Some(notification))
if notification.activation_sequence < submission.activation_sequence =>
{
Some(notification.notification_request_id.clone())
Some(notification_head_id(
&notification.source_namespace,
&notification.notification_request_id,
))
}
(Some(submission), _) => Some(submission.submission_id.clone()),
(None, Some(notification)) => Some(notification.notification_request_id.clone()),
(None, Some(notification)) => Some(notification_head_id(
&notification.source_namespace,
&notification.notification_request_id,
)),
(None, None) => None,
};
protocol::PendingSubmissionsSnapshot {
@@ -172,6 +188,65 @@ impl PendingActivationState {
}
}
fn notification_payload_digest(message: &str, auto_run: bool) -> String {
use sha2::Digest as _;
let mut hasher = sha2::Sha256::new();
hasher.update(if auto_run {
&b"auto\0"[..]
} else {
&b"deferred\0"[..]
});
hasher.update(message.as_bytes());
hasher
.finalize()
.iter()
.map(|byte| format!("{byte:02x}"))
.collect()
}
pub(crate) fn authenticated_input_provenance(
source: &protocol::AuthenticatedInputSource,
) -> WorkerHistoryProvenance {
match source {
protocol::AuthenticatedInputSource::Account { account_id } => {
WorkerHistoryProvenance::HumanInput {
account_id: account_id.clone(),
}
}
protocol::AuthenticatedInputSource::Worker {
runtime_id,
worker_id,
} => WorkerHistoryProvenance::WorkerInput {
actor: session_store::LoggedWorkerSubject {
workspace_id: None,
runtime_id: Some(runtime_id.clone()),
worker_id: worker_id.clone(),
},
},
protocol::AuthenticatedInputSource::Backend { operation_id } => {
WorkerHistoryProvenance::BackendInstruction {
operation_id: Some(operation_id.clone()),
}
}
}
}
fn notification_head_id(source_namespace: &str, request_id: &str) -> String {
use sha2::Digest as _;
let mut hasher = sha2::Sha256::new();
hasher.update(source_namespace.as_bytes());
hasher.update(b"\0");
hasher.update(request_id.as_bytes());
format!(
"notification:{}",
hasher
.finalize()
.iter()
.map(|byte| format!("{byte:02x}"))
.collect::<String>()
)
}
fn submission_payload_len(input: &[Segment]) -> u64 {
serde_json::to_vec(input)
.map(|bytes| u64::try_from(bytes.len()).unwrap_or(u64::MAX))
@@ -186,6 +261,15 @@ fn submission_payload_digest(input: &[Segment]) -> String {
.collect()
}
fn submission_uploaded_file_refs(
input: &[Segment],
) -> impl Iterator<Item = &protocol::UploadedFileRef> {
input.iter().filter_map(|segment| match segment {
Segment::UploadedFile { file } => Some(file),
_ => None,
})
}
fn submission_artifact_ref_count(input: &[Segment]) -> usize {
input
.iter()
@@ -1191,11 +1275,57 @@ where
Ok(())
}
fn pin_submission_files(
&self,
pending: &PendingSubmission,
) -> Result<(), PendingSubmissionError> {
let session_id = self.writer.state.location().session_id;
for reference in submission_uploaded_file_refs(&pending.input) {
self.writer
.store
.pin_uploaded_file(session_id, reference, &pending.submission_id)?;
}
Ok(())
}
fn release_submission_files(
&self,
pending: &PendingSubmission,
) -> Result<(), PendingSubmissionError> {
let session_id = self.writer.state.location().session_id;
for reference in submission_uploaded_file_refs(&pending.input) {
self.writer.store.release_uploaded_file_pin(
session_id,
&reference.artifact_id,
&pending.submission_id,
)?;
}
Ok(())
}
#[cfg(test)]
pub(crate) fn accept(
&self,
submission_request_id: String,
input: Vec<Segment>,
activate_now: bool,
) -> Result<SubmissionAcceptance, PendingSubmissionError> {
self.accept_from_source(
submission_request_id,
input,
self.direct_client_namespace(),
WorkerHistoryProvenance::LegacyUnknown,
activate_now,
)
}
pub(crate) fn accept_from_source(
&self,
submission_request_id: String,
input: Vec<Segment>,
source_namespace: String,
provenance: WorkerHistoryProvenance,
activate_now: bool,
) -> Result<SubmissionAcceptance, PendingSubmissionError> {
if submission_request_id.trim().is_empty() {
return Err(PendingSubmissionError::EmptyRequestId);
@@ -1218,11 +1348,10 @@ where
.lock()
.expect("pending activation state poisoned");
let original = current.clone();
if let Some(receipt) = current
.receipts
.iter()
.find(|receipt| receipt.submission_request_id == submission_request_id)
{
if let Some(receipt) = current.receipts.iter().find(|receipt| {
receipt.submission_request_id == submission_request_id
&& receipt.source_namespace == source_namespace
}) {
if receipt.payload_digest != payload_digest {
return Err(PendingSubmissionError::IdempotencyConflict);
}
@@ -1237,11 +1366,12 @@ where
let submission_id = uuid::Uuid::now_v7().to_string();
let pending = PendingSubmission {
submission_request_id: submission_request_id.clone(),
source_namespace: source_namespace.clone(),
submission_id: submission_id.clone(),
payload_digest: payload_digest.clone(),
accepted_at_ms: segment_log::now_millis(),
activation_sequence: current.next_activation_sequence,
provenance: WorkerHistoryProvenance::LegacyUnknown,
provenance,
was_queued: !activate_now,
input,
};
@@ -1253,6 +1383,7 @@ where
};
current.remember_receipt(SubmissionReceipt {
submission_request_id: submission_request_id.clone(),
source_namespace,
submission_id: submission_id.clone(),
payload_digest,
disposition,
@@ -1299,7 +1430,14 @@ where
return Err(PendingSubmissionError::ArtifactLimit);
}
current.pending.push_back(pending.clone());
}
if !activate_now {
if let Err(error) = self.pin_submission_files(&pending) {
*current = original;
return Err(error);
}
if let Err(error) = self.persist_locked(&current) {
let _ = self.release_submission_files(&pending);
*current = original;
return Err(error);
}
@@ -1312,10 +1450,29 @@ where
})
}
#[cfg(test)]
pub(crate) fn accept_notification(
&self,
notification_request_id: String,
message: String,
auto_run: bool,
) -> Result<bool, PendingSubmissionError> {
self.accept_notification_from_source(
notification_request_id,
message,
self.direct_client_namespace(),
WorkerHistoryProvenance::LegacyUnknown,
auto_run,
)
}
pub(crate) fn accept_notification_from_source(
&self,
notification_request_id: String,
message: String,
source_namespace: String,
provenance: WorkerHistoryProvenance,
auto_run: bool,
) -> Result<bool, PendingSubmissionError> {
if notification_request_id.trim().is_empty() {
return Err(PendingSubmissionError::EmptyRequestId);
@@ -1323,7 +1480,7 @@ where
if notification_request_id.len() > MAX_ACTIVATION_REQUEST_ID_BYTES {
return Err(PendingSubmissionError::RequestIdLimit);
}
let payload_digest = submission_payload_digest(&[Segment::text(message.clone())]);
let payload_digest = notification_payload_digest(&message, auto_run);
let _append_guard = self
.writer
.state
@@ -1334,12 +1491,11 @@ where
.state
.lock()
.expect("pending activation state poisoned");
if let Some(receipt) = state
.notification_receipts
.iter()
.find(|receipt| receipt.notification_request_id == notification_request_id)
{
if receipt.payload_digest != payload_digest {
if let Some(receipt) = state.notification_receipts.iter().find(|receipt| {
receipt.notification_request_id == notification_request_id
&& receipt.source_namespace == source_namespace
}) {
if receipt.payload_digest != payload_digest || receipt.auto_run != auto_run {
return Err(PendingSubmissionError::IdempotencyConflict);
}
return Ok(false);
@@ -1373,17 +1529,19 @@ where
state.next_activation_sequence = state.next_activation_sequence.saturating_add(1);
state.pending_notifications.push_back(PendingNotification {
notification_request_id: notification_request_id.clone(),
source_namespace: source_namespace.clone(),
message,
payload_digest: payload_digest.clone(),
auto_run,
accepted_at_ms: segment_log::now_millis(),
activation_sequence,
provenance: WorkerHistoryProvenance::BackendInstruction {
operation_id: Some(notification_request_id.clone()),
},
provenance,
});
state.remember_notification_receipt(NotificationReceipt {
notification_request_id,
source_namespace,
payload_digest,
auto_run,
});
state.revision = state.revision.saturating_add(1);
if let Err(error) = self.persist_locked(&state) {
@@ -1393,6 +1551,59 @@ where
Ok(true)
}
pub(crate) fn activating_passive_notification_id(&self) -> Option<String> {
self.state
.lock()
.expect("pending activation state poisoned")
.activating_notification
.as_ref()
.filter(|notification| !notification.auto_run)
.map(|notification| notification.notification_request_id.clone())
}
pub(crate) fn next_passive_notification_identity(&self) -> Option<(String, String)> {
self.state
.lock()
.expect("pending activation state poisoned")
.pending_notifications
.iter()
.find(|notification| !notification.auto_run)
.map(|notification| {
(
notification.source_namespace.clone(),
notification.notification_request_id.clone(),
)
})
}
pub(crate) fn prepare_notification(
&self,
source_namespace: &str,
notification_request_id: &str,
) -> Option<PendingNotification> {
let mut state = self
.state
.lock()
.expect("pending activation state poisoned");
if state.activating_notification.is_some() {
return None;
}
let index = state
.pending_notifications
.iter()
.position(|notification| {
notification.notification_request_id == notification_request_id
&& notification.source_namespace == source_namespace
})?;
let notification = state
.pending_notifications
.remove(index)
.expect("located pending notification must exist");
state.activating_notification = Some(notification.clone());
state.revision = state.revision.saturating_add(1);
Some(notification)
}
pub(crate) fn prepare_next_activation(
&self,
fence: Option<(u64, &str)>,
@@ -1414,17 +1625,20 @@ where
return Ok(None);
}
let submission_sequence = state.pending.front().map(|item| item.activation_sequence);
let notification_sequence = state
let notification_index = state
.pending_notifications
.front()
.iter()
.position(|item| item.auto_run);
let notification_sequence = notification_index
.and_then(|index| state.pending_notifications.get(index))
.map(|item| item.activation_sequence);
if notification_sequence.is_some()
&& (submission_sequence.is_none() || notification_sequence < submission_sequence)
{
let notification = state
.pending_notifications
.pop_front()
.expect("notification sequence came from queue head");
.remove(notification_index.expect("notification sequence came from an item"))
.expect("notification sequence came from an existing item");
state.activating_notification = Some(notification.clone());
state.revision = state.revision.saturating_add(1);
return Ok(Some(PendingActivation::Notification(notification)));
@@ -1442,6 +1656,12 @@ where
}
pub(crate) fn abort_activation(&self, pending: PendingSubmission) {
let _append_guard = self
.writer
.state
.append_lock
.lock()
.expect("segment append lock poisoned");
let mut state = self
.state
.lock()
@@ -1452,13 +1672,16 @@ where
}
state.activating = None;
if pending.was_queued {
state.pending.push_front(pending);
state.pending.push_front(pending.clone());
} else {
state
.receipts
.retain(|receipt| receipt.submission_id != pending.submission_id);
}
state.revision = state.revision.saturating_add(1);
if let Err(error) = self.persist_locked(&state) {
tracing::error!(error = %error, "failed to persist aborted pending activation");
}
}
pub(crate) fn activation_extension(&self) -> SessionExtension {
@@ -1530,6 +1753,10 @@ where
}
}
pub(crate) fn direct_client_namespace(&self) -> String {
format!("direct:{}", self.writer.state.location().session_id)
}
pub(crate) fn snapshot(&self) -> protocol::PendingSubmissionsSnapshot {
self.state
.lock()
@@ -1561,12 +1788,16 @@ where
else {
return Err(PendingSubmissionError::NotFound(submission_id.to_owned()));
};
state.pending.remove(index);
let removed = state
.pending
.remove(index)
.expect("located pending submission must exist");
state.revision = state.revision.saturating_add(1);
if let Err(error) = self.persist_locked(&state) {
*state = original;
return Err(error);
}
self.release_submission_files(&removed)?;
Ok(state.snapshot())
}
@@ -1586,13 +1817,16 @@ where
.expect("pending activation state poisoned");
Self::validate_fence(&state, expected_revision, None)?;
let original = state.clone();
state.pending.clear();
let removed = state.pending.drain(..).collect::<Vec<_>>();
state.pending_notifications.clear();
state.revision = state.revision.saturating_add(1);
if let Err(error) = self.persist_locked(&state) {
*state = original;
return Err(error);
}
for pending in &removed {
self.release_submission_files(pending)?;
}
Ok(state.snapshot())
}
}
@@ -1627,9 +1861,11 @@ pub trait SystemItemCommitter: Send + Sync {
&self,
item: SystemItem,
extensions: Vec<SessionExtension>,
history_provenance: Option<WorkerHistoryProvenance>,
) -> Result<HistoryEntry<SessionHistoryMetadata>, StoreError> {
let metadata = new_history_metadata(
WorkerHistoryProvenance::BackendInstruction { operation_id: None },
history_provenance
.unwrap_or(WorkerHistoryProvenance::BackendInstruction { operation_id: None }),
None,
);
let history_item = item.to_history_item();
@@ -2666,9 +2902,11 @@ impl<C: LlmClient + 'static, St: Store> Worker<C, St> {
.expect("pending activation state poisoned")
.clone();
if !pending_state.pending.is_empty()
|| !pending_state.pending_notifications.is_empty()
|| pending_state.activating.is_some()
|| pending_state.activating_notification.is_some()
|| !pending_state.receipts.is_empty()
|| !pending_state.notification_receipts.is_empty()
{
let checkpoint = LogEntry::Extension {
ts: segment_log::now_millis(),
@@ -3501,8 +3739,13 @@ impl<C: LlmClient + 'static, St: Store> Worker<C, St> {
where
St: Clone + 'static,
{
self.run_with_input_extensions_and_commit_hook(input, input_extensions, || {})
.await
self.run_with_input_extensions_and_commit_hook(
input,
input_extensions,
WorkerHistoryProvenance::LegacyUnknown,
|| {},
)
.await
}
/// Run user input and invoke `on_input_committed` only after the annotated
@@ -3513,6 +3756,7 @@ impl<C: LlmClient + 'static, St: Store> Worker<C, St> {
&mut self,
input: Vec<Segment>,
mut input_extensions: Vec<SessionExtension>,
input_provenance: WorkerHistoryProvenance,
on_input_committed: F,
) -> Result<WorkerRunResult, WorkerError>
where
@@ -3560,8 +3804,12 @@ impl<C: LlmClient + 'static, St: Store> Worker<C, St> {
trigger: protocol::InvokeKind::UserSend,
})?;
let projected_input =
self.projected_input_history(&input, flow_projection.as_ref(), &projected_entry_ids);
let projected_input = self.projected_input_history(
&input,
flow_projection.as_ref(),
&projected_entry_ids,
&input_provenance,
);
// Persist original typed segments together with the exact ordered
// model-visible item+origin projection before any entry becomes live.
@@ -3867,6 +4115,7 @@ impl<C: LlmClient + 'static, St: Store> Worker<C, St> {
input: &[Segment],
flow_projection: Option<&PreparedFlowProjection>,
entry_ids: &[SessionHistoryEntryId],
provenance: &WorkerHistoryProvenance,
) -> Vec<HistoryEntry<SessionHistoryMetadata>> {
if let Some(flow) = flow_projection {
return input
@@ -3887,10 +4136,7 @@ impl<C: LlmClient + 'static, St: Store> Worker<C, St> {
other => history_entry_with_id(
Item::user_message(Segment::flatten_to_text(std::slice::from_ref(other))),
entry_id.clone(),
// Current public submit transport does not carry a
// trusted account/Worker subject envelope. Fail closed
// instead of promoting role=user to HumanInput.
WorkerHistoryProvenance::LegacyUnknown,
provenance.clone(),
),
})
.collect();
@@ -3902,7 +4148,7 @@ impl<C: LlmClient + 'static, St: Store> Worker<C, St> {
.first()
.expect("projected Worker input always has one entry id")
.clone(),
WorkerHistoryProvenance::LegacyUnknown,
provenance.clone(),
)]
}
@@ -7938,8 +8184,15 @@ mod build_summary_prompt_tests {
serde_json::to_value(&state).unwrap(),
);
let projected_ids = vec![SessionHistoryEntryId::new(), SessionHistoryEntryId::new()];
let projected =
worker.projected_input_history(&segments, projection.as_ref(), &projected_ids);
let input_provenance = WorkerHistoryProvenance::HumanInput {
account_id: "account-1".into(),
};
let projected = worker.projected_input_history(
&segments,
projection.as_ref(),
&projected_ids,
&input_provenance,
);
worker
.commit_entry(LogEntry::AnnotatedUserInput {
ts: segment_log::now_millis(),
@@ -7966,6 +8219,7 @@ mod build_summary_prompt_tests {
projected[0].annotation.origin,
WorkerHistoryProvenance::FlowInstruction { .. }
));
assert_eq!(projected[1].annotation.origin, input_provenance);
assert_eq!(state.instance.definition_revision, 3);
assert_eq!(state.instance.current_state.as_str(), "implement");
assert_eq!(workspace_client.requests.lock().unwrap().len(), 1);
@@ -8048,7 +8302,12 @@ mod build_summary_prompt_tests {
.delete_uploaded_file(worker.session_id(), &file.artifact_id),
Err(StoreError::ArtifactAlreadyCommitted)
));
let projected = worker.projected_input_history(&input, None, &[entry_id]);
let projected = worker.projected_input_history(
&input,
None,
&[entry_id],
&WorkerHistoryProvenance::LegacyUnknown,
);
let text = projected[0].item.as_text().unwrap();
assert!(text.contains("notes.md"));
assert!(text.contains(&file.artifact_id));
@@ -8130,7 +8389,12 @@ mod build_summary_prompt_tests {
if retained.source_entry_id == artifact.source_entry_id
));
let history = worker.projected_input_history(&input, None, &[entry_id]);
let history = worker.projected_input_history(
&input,
None,
&[entry_id],
&WorkerHistoryProvenance::LegacyUnknown,
);
assert!(!history[0].item.as_text().unwrap().contains("終端"));
append_test_entry(
&worker,
@@ -8406,6 +8670,47 @@ mod build_summary_prompt_tests {
assert_eq!(worker.history()[0].as_text().unwrap(), "first message");
}
#[tokio::test]
async fn rewind_preserves_notification_only_pending_activation_checkpoint() {
let (_dir, mut worker) = rewind_test_worker().await;
append_user_turn(&worker, 10, "first message");
append_user_turn(&worker, 20, "second message");
worker
.pending_submission_handle()
.accept_notification("notification-1".into(), "keep me".into(), true)
.unwrap();
let (head_entries, targets) = worker.list_rewind_targets().unwrap();
worker
.rewind_to(targets.last().unwrap().id.clone(), head_entries)
.await
.unwrap();
let location = worker.segment_state.location();
let entries = worker
.store
.read_all(location.session_id, location.segment_id)
.unwrap();
let restored: PendingActivationState = entries
.iter()
.rev()
.find_map(|entry| match entry {
LogEntry::Extension {
domain, payload, ..
} if domain == SESSION_PENDING_ACTIVATIONS_EXTENSION_DOMAIN => {
serde_json::from_value(payload.clone()).ok()
}
_ => None,
})
.unwrap();
assert_eq!(restored.pending_notifications.len(), 1);
assert_eq!(restored.notification_receipts.len(), 1);
assert_eq!(
restored.pending_notifications[0].notification_request_id,
"notification-1"
);
}
#[tokio::test]
async fn annotated_history_rewind_commits_authoritative_prefix() {
let (_dir, mut worker) = rewind_test_worker().await;
@@ -9251,6 +9556,110 @@ mod build_summary_prompt_tests {
);
}
#[test]
fn submission_retry_identity_is_scoped_to_authenticated_source_and_keeps_provenance() {
let temp = tempfile::tempdir().unwrap();
let handle = PendingSubmissionHandle::for_test(temp.path());
let input = vec![Segment::text("same request")];
let account_a = WorkerHistoryProvenance::HumanInput {
account_id: "account-a".into(),
};
let account_b = WorkerHistoryProvenance::HumanInput {
account_id: "account-b".into(),
};
let first = handle
.accept_from_source(
"request-1".into(),
input.clone(),
"account:account-a".into(),
account_a.clone(),
false,
)
.unwrap();
let replay = handle
.accept_from_source(
"request-1".into(),
input.clone(),
"account:account-a".into(),
account_a.clone(),
false,
)
.unwrap();
let other_source = handle
.accept_from_source(
"request-1".into(),
input,
"account:account-b".into(),
account_b.clone(),
false,
)
.unwrap();
assert_eq!(replay.submission_id, first.submission_id);
assert_ne!(other_source.submission_id, first.submission_id);
let state = handle.state.lock().unwrap();
assert_eq!(state.pending[0].provenance, account_a);
assert_eq!(state.pending[1].provenance, account_b);
}
#[test]
fn queued_submission_pins_uploaded_file_until_cancelled() {
let temp = tempfile::tempdir().unwrap();
let handle = PendingSubmissionHandle::for_test(temp.path());
let session_id = handle.writer.state.session_id();
let reference = handle
.writer
.store
.write_uploaded_file(
session_id,
"queued.txt",
"text/plain",
b"queued artifact",
session_store::UploadedFileLimits {
max_file_bytes: 1024,
max_session_bytes: 2048,
},
)
.unwrap();
let accepted = handle
.accept(
"artifact-request".into(),
vec![Segment::UploadedFile {
file: reference.clone(),
}],
false,
)
.unwrap();
assert_eq!(
handle
.writer
.store
.delete_uncommitted_uploaded_files(session_id)
.unwrap(),
0
);
assert!(
handle
.writer
.store
.read_uploaded_file_by_id(session_id, &reference.artifact_id)
.is_ok()
);
handle
.cancel(&accepted.submission_id, handle.snapshot().revision)
.unwrap();
assert_eq!(
handle
.writer
.store
.delete_uncommitted_uploaded_files(session_id)
.unwrap(),
1
);
}
#[test]
fn pending_submission_queue_is_durable_idempotent_and_bounded() {
let temp = tempfile::tempdir().unwrap();
@@ -9339,22 +9748,96 @@ mod build_summary_prompt_tests {
assert_eq!(cleared.notification_count, 0);
}
#[test]
fn durable_notification_commits_authenticated_history_provenance() {
let temp = tempfile::tempdir().unwrap();
let handle = PendingSubmissionHandle::for_test(temp.path());
let provenance = WorkerHistoryProvenance::HumanInput {
account_id: "account-1".into(),
};
let committed = handle
.writer
.commit_system_item_with_extensions(
SystemItem::Notification {
message: "notice".into(),
body: "notice".into(),
prompt_provenance: None,
},
Vec::new(),
Some(provenance.clone()),
)
.unwrap();
assert_eq!(committed.annotation.origin, provenance);
}
#[test]
fn notification_retry_identity_is_scoped_to_authenticated_source() {
let temp = tempfile::tempdir().unwrap();
let handle = PendingSubmissionHandle::for_test(temp.path());
let account_a = WorkerHistoryProvenance::HumanInput {
account_id: "account-a".into(),
};
let account_b = WorkerHistoryProvenance::HumanInput {
account_id: "account-b".into(),
};
assert!(
handle
.accept_notification_from_source(
"request-1".into(),
"notice".into(),
"account:account-a".into(),
account_a.clone(),
false,
)
.unwrap()
);
assert!(
!handle
.accept_notification_from_source(
"request-1".into(),
"notice".into(),
"account:account-a".into(),
account_a.clone(),
false,
)
.unwrap()
);
assert!(
handle
.accept_notification_from_source(
"request-1".into(),
"notice".into(),
"account:account-b".into(),
account_b.clone(),
false,
)
.unwrap()
);
let state = handle.state.lock().unwrap();
assert_eq!(state.pending_notifications[0].provenance, account_a);
assert_eq!(state.pending_notifications[1].provenance, account_b);
}
#[test]
fn notification_and_submit_share_activation_order_and_notification_dedupes() {
let temp = tempfile::tempdir().unwrap();
let handle = PendingSubmissionHandle::for_test(temp.path());
assert!(
handle
.accept_notification("notification-1".into(), "notice".into())
.accept_notification("notification-1".into(), "notice".into(), true)
.unwrap()
);
assert!(
!handle
.accept_notification("notification-1".into(), "notice".into())
.accept_notification("notification-1".into(), "notice".into(), true)
.unwrap()
);
assert!(matches!(
handle.accept_notification("notification-1".into(), "different".into()),
handle.accept_notification("notification-1".into(), "different".into(), true),
Err(PendingSubmissionError::IdempotencyConflict)
));
assert!(matches!(
handle.accept_notification("notification-1".into(), "notice".into(), false),
Err(PendingSubmissionError::IdempotencyConflict)
));
handle
@@ -9382,9 +9865,10 @@ mod build_summary_prompt_tests {
let mut session = WorkerSession::new(session_store::new_session_id(), Vec::new());
let state = PendingActivationState {
revision: 4,
next_activation_sequence: 2,
next_activation_sequence: 3,
activating: Some(PendingSubmission {
submission_request_id: "request-1".into(),
source_namespace: "direct:test".into(),
submission_id: "submission-1".into(),
payload_digest: submission_payload_digest(&[Segment::text("first")]),
accepted_at_ms: 1,
@@ -9393,9 +9877,21 @@ mod build_summary_prompt_tests {
was_queued: false,
input: vec![Segment::text("first")],
}),
activating_notification: None,
activating_notification: Some(PendingNotification {
notification_request_id: "notification-1".into(),
source_namespace: "account:account-1".into(),
message: "deferred notice".into(),
payload_digest: notification_payload_digest("deferred notice", false),
auto_run: false,
accepted_at_ms: 3,
activation_sequence: 2,
provenance: WorkerHistoryProvenance::HumanInput {
account_id: "account-1".into(),
},
}),
pending: VecDeque::from([PendingSubmission {
submission_request_id: "request-2".into(),
source_namespace: "direct:test".into(),
submission_id: "submission-2".into(),
payload_digest: submission_payload_digest(&[Segment::text("second")]),
accepted_at_ms: 2,
@@ -9406,7 +9902,12 @@ mod build_summary_prompt_tests {
}]),
pending_notifications: VecDeque::new(),
receipts: VecDeque::new(),
notification_receipts: VecDeque::new(),
notification_receipts: VecDeque::from([NotificationReceipt {
notification_request_id: "notification-1".into(),
source_namespace: "account:account-1".into(),
payload_digest: notification_payload_digest("deferred notice", false),
auto_run: false,
}]),
};
session.restore_pending_activations(&[(
SESSION_PENDING_ACTIVATIONS_EXTENSION_DOMAIN.into(),
@@ -9420,6 +9921,14 @@ mod build_summary_prompt_tests {
assert_eq!(state.pending.len(), 2);
assert_eq!(state.pending[0].submission_id, "submission-1");
assert_eq!(state.pending[1].submission_id, "submission-2");
assert!(state.activating_notification.is_none());
assert_eq!(state.pending_notifications.len(), 1);
assert!(!state.pending_notifications[0].auto_run);
assert!(matches!(
state.pending_notifications[0].provenance,
WorkerHistoryProvenance::HumanInput { ref account_id } if account_id == "account-1"
));
assert_eq!(state.notification_receipts.len(), 1);
}
fn minimal_manifest() -> WorkerManifest {
+71 -8
View File
@@ -1804,15 +1804,18 @@ async fn notify_while_idle_with_auto_run_false_waits_for_explicit_run() {
let client_for_assert = client.clone();
let worker = make_worker(client).await;
let handle = spawn_controller(worker).await;
let notification_request_id = protocol::new_submission_request_id();
handle
.send(Method::Notify {
notification_request_id: protocol::new_submission_request_id(),
message: "progress snapshot".into(),
auto_run: false,
})
.await
.unwrap();
for _ in 0..2 {
handle
.send(Method::Notify {
notification_request_id: notification_request_id.clone(),
message: "progress snapshot".into(),
auto_run: false,
})
.await
.unwrap();
}
tokio::time::sleep(std::time::Duration::from_millis(100)).await;
assert_eq!(handle.shared_state.get_status(), WorkerStatus::Idle);
@@ -2049,6 +2052,66 @@ async fn notify_while_running_does_not_emit_already_running_error() {
wait_for_status(&handle, WorkerStatus::Idle).await;
}
#[tokio::test]
async fn weak_notify_while_running_is_deduped_and_survives_until_next_submit() {
let client = MockClient::sequential(vec![
MockResponse::Hang(Vec::new()),
MockResponse::Complete(simple_text_events()),
]);
let client_for_assert = client.clone();
let worker = make_worker(client).await;
let handle = spawn_controller(worker).await;
handle
.send(Method::submit_text(
protocol::new_submission_request_id(),
"first",
))
.await
.unwrap();
wait_for_status(&handle, WorkerStatus::Running).await;
let notification_request_id = protocol::new_submission_request_id();
for _ in 0..2 {
handle
.send(Method::Notify {
notification_request_id: notification_request_id.clone(),
message: "durable weak notice".into(),
auto_run: false,
})
.await
.unwrap();
}
handle.send(Method::Cancel).await.unwrap();
wait_for_status(&handle, WorkerStatus::Idle).await;
let mut rx = handle.subscribe();
handle
.send(Method::submit_text(
protocol::new_submission_request_id(),
"second",
))
.await
.unwrap();
tokio::time::timeout(std::time::Duration::from_secs(2), async {
loop {
if matches!(rx.recv().await, Ok(Event::TurnEnd { .. })) {
break;
}
}
})
.await
.expect("second submit completes");
let requests = client_for_assert.captured_requests();
let notice_count = requests[1]
.items
.iter()
.filter_map(|item| item.as_text())
.filter(|text| text.contains("durable weak notice"))
.count();
assert_eq!(notice_count, 1);
}
#[tokio::test]
async fn status_json_reflects_worker_name() {
let client = MockClient::new(simple_text_events());