From b038f022d395fa05f7ec472650f60fd96daacfa8 Mon Sep 17 00:00:00 2001 From: Hare Date: Sun, 6 Sep 2026 01:35:26 +0900 Subject: [PATCH] fix: harden submit queue durability --- crates/protocol/src/lib.rs | 64 +- crates/session-store/src/fs_store.rs | 117 +++- crates/session-store/src/store.rs | 20 + crates/session-store/src/uploaded_file.rs | 82 ++- crates/worker-runtime/src/worker_backend.rs | 6 + crates/worker/src/controller.rs | 302 +++++++-- crates/worker/src/ipc/interceptor.rs | 26 +- crates/worker/src/ipc/notify_buffer.rs | 25 +- crates/worker/src/worker.rs | 605 ++++++++++++++++-- crates/worker/tests/controller_test.rs | 79 ++- crates/workspace-server/src/hosts.rs | 10 +- .../src/runtime_subscription_tests.rs | 8 +- crates/workspace-server/src/server.rs | 170 ++++- .../src/workspace_subscription.rs | 17 +- 14 files changed, 1370 insertions(+), 161 deletions(-) diff --git a/crates/protocol/src/lib.rs b/crates/protocol/src/lib.rs index e3330ba0..e6d790db 100644 --- a/crates/protocol/src/lib.rs +++ b/crates/protocol/src/lib.rs @@ -32,6 +32,38 @@ fn is_false(value: &bool) -> bool { // Method (Client → Worker via Unix Socket) // --------------------------------------------------------------------------- +/// Trusted source identity attached by an authenticated transport boundary. +/// +/// Public clients cannot select this value directly. Runtime/Backend adapters +/// stamp it before forwarding an accepted Submit or Notify to a Worker. +#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)] +#[serde(tag = "kind", rename_all = "snake_case", deny_unknown_fields)] +pub enum AuthenticatedInputSource { + Account { + account_id: String, + }, + Worker { + runtime_id: String, + worker_id: String, + }, + Backend { + operation_id: String, + }, +} + +impl AuthenticatedInputSource { + pub fn namespace(&self) -> String { + match self { + Self::Account { account_id } => format!("account:{account_id}"), + Self::Worker { + runtime_id, + worker_id, + } => format!("worker:{runtime_id}:{worker_id}"), + Self::Backend { operation_id } => format!("backend:{operation_id}"), + } + } +} + #[derive(Debug, Clone, Serialize, Deserialize)] #[cfg_attr(feature = "typescript", derive(ts_rs::TS))] #[serde(tag = "method", content = "params", rename_all = "snake_case")] @@ -45,13 +77,13 @@ pub enum Method { submission_request_id: String, input: Vec, }, - /// Runtime-internal Submit with the same request identity contract. This - /// variant is not serializable on the public Client → Worker protocol. - #[serde(skip)] + /// Authenticated transport form of Submit. Trusted adapters replace + /// public Submit before forwarding it to the Worker. #[cfg_attr(feature = "typescript", ts(skip))] SubmitTracked { submission_request_id: String, input: Vec, + source: AuthenticatedInputSource, }, /// Human-readable text injected into the target Worker's LLM context /// as a non-blocking system message. `auto_run` controls whether an @@ -65,6 +97,15 @@ pub enum Method { #[serde(default = "default_true", skip_serializing_if = "is_true")] auto_run: bool, }, + /// Authenticated transport form of Notify. + #[cfg_attr(feature = "typescript", ts(skip))] + NotifyTracked { + notification_request_id: String, + message: String, + #[serde(default = "default_true", skip_serializing_if = "is_true")] + auto_run: bool, + source: AuthenticatedInputSource, + }, /// Typed lifecycle report from a child Worker to its direct parent. WorkerEvent(WorkerEvent), /// Return the authoritative FIFO summary without exposing queued payloads. @@ -1497,15 +1538,26 @@ mod tests { } #[test] - fn runtime_tracked_submit_is_not_public_protocol_json() { + fn authenticated_submit_round_trips_trusted_source() { let method = Method::SubmitTracked { input: vec![Segment::text("private")], submission_request_id: "request-1".to_string(), + source: AuthenticatedInputSource::Account { + account_id: "account-1".into(), + }, }; - assert!(serde_json::to_string(&method).is_err()); + let json = serde_json::to_string(&method).unwrap(); + let decoded = serde_json::from_str::(&json).unwrap(); + assert!(matches!( + decoded, + Method::SubmitTracked { + source: AuthenticatedInputSource::Account { account_id }, + .. + } if account_id == "account-1" + )); assert!( serde_json::from_str::( - r#"{"method":"submit_tracked","input":[],"submission_id":"forged"}"#, + r#"{"method":"submit_tracked","input":[],"submission_request_id":"forged"}"#, ) .is_err() ); diff --git a/crates/session-store/src/fs_store.rs b/crates/session-store/src/fs_store.rs index a0e80a75..abdf14e4 100644 --- a/crates/session-store/src/fs_store.rs +++ b/crates/session-store/src/fs_store.rs @@ -22,7 +22,8 @@ use crate::store::{Store, StoreError}; use crate::uploaded_file::{ bind_uploaded_file, clear_uploaded_file_binding, copy_committed_uploaded_files, delete_uncommitted_uploaded_files, delete_uploaded_file, list_uploaded_file_refs, - read_uploaded_file, read_uploaded_file_by_id, write_uploaded_file, + pin_uploaded_file, read_uploaded_file, read_uploaded_file_by_id, release_uploaded_file_pin, + write_uploaded_file, }; use crate::{ PasteArtifactLimits, SegmentId, SessionId, UploadedFileLimits, UploadedFileUploadContext, @@ -518,6 +519,32 @@ impl Store for FsStore { } } + fn pin_uploaded_file( + &self, + session_id: SessionId, + reference: &UploadedFileRef, + owner_id: &str, + ) -> Result<(), StoreError> { + let _guard = self + .append_lock + .lock() + .map_err(|_| std::io::Error::other("session store append lock was poisoned"))?; + pin_uploaded_file(&self.paste_artifact_dir(session_id), reference, owner_id) + } + + fn release_uploaded_file_pin( + &self, + session_id: SessionId, + artifact_id: &str, + owner_id: &str, + ) -> Result<(), StoreError> { + let _guard = self + .append_lock + .lock() + .map_err(|_| std::io::Error::other("session store append lock was poisoned"))?; + release_uploaded_file_pin(&self.paste_artifact_dir(session_id), artifact_id, owner_id) + } + fn delete_uploaded_file( &self, session_id: SessionId, @@ -865,6 +892,94 @@ mod tests { assert!(store.read_uploaded_file(owner, &reference).is_err()); } + #[test] + fn pending_upload_pin_survives_cleanup_until_release_or_history_binding() { + let tmp = tempfile::TempDir::new().unwrap(); + let store = FsStore::new(tmp.path()).unwrap(); + let session_id = new_session_id(); + let limits = UploadedFileLimits { + max_file_bytes: 64, + max_session_bytes: 128, + }; + let pending = store + .write_uploaded_file(session_id, "pending.txt", "text/plain", b"pending", limits) + .unwrap(); + store + .pin_uploaded_file(session_id, &pending, "submission-1") + .unwrap(); + assert!(matches!( + store.pin_uploaded_file(session_id, &pending, "submission-other"), + Err(StoreError::ArtifactAlreadyCommitted) + )); + drop(store); + let store = FsStore::new(tmp.path()).unwrap(); + assert_eq!( + store.delete_uncommitted_uploaded_files(session_id).unwrap(), + 0 + ); + assert_eq!( + store + .read_uploaded_file_by_id(session_id, &pending.artifact_id) + .unwrap() + .1, + b"pending" + ); + + let fork_session_id = new_session_id(); + assert_eq!( + store + .copy_committed_uploaded_files(session_id, fork_session_id) + .unwrap(), + 1 + ); + assert_eq!( + store + .read_uploaded_file_by_id(fork_session_id, &pending.artifact_id) + .unwrap() + .1, + b"pending" + ); + + let committed = store + .bind_uploaded_file(session_id, &pending, "entry-1") + .unwrap(); + store + .create_segment( + session_id, + new_segment_id(), + &[LogEntry::InputSegmentsCheckpoint { + ts: 1, + user_segments: vec![vec![protocol::Segment::UploadedFile { + file: committed.clone(), + }]], + }], + ) + .unwrap(); + assert_eq!( + store.delete_uncommitted_uploaded_files(session_id).unwrap(), + 0 + ); + + let releasable = store + .write_uploaded_file(session_id, "cancelled.txt", "text/plain", b"cancel", limits) + .unwrap(); + store + .pin_uploaded_file(session_id, &releasable, "submission-2") + .unwrap(); + store + .release_uploaded_file_pin(session_id, &releasable.artifact_id, "submission-2") + .unwrap(); + assert_eq!( + store.delete_uncommitted_uploaded_files(session_id).unwrap(), + 1 + ); + assert!( + store + .read_uploaded_file_by_id(session_id, &releasable.artifact_id) + .is_err() + ); + } + #[test] fn uploaded_file_validation_and_shared_quota_fail_closed() { let tmp = tempfile::TempDir::new().unwrap(); diff --git a/crates/session-store/src/store.rs b/crates/session-store/src/store.rs index a22fd6c7..4c98fa33 100644 --- a/crates/session-store/src/store.rs +++ b/crates/session-store/src/store.rs @@ -226,6 +226,26 @@ pub trait Store: Send + Sync { Err(StoreError::PasteArtifactUnsupported) } + /// Retain an uploaded file while a durable pending operation owns it. + fn pin_uploaded_file( + &self, + _session_id: SessionId, + _reference: &UploadedFileRef, + _owner_id: &str, + ) -> Result<(), StoreError> { + Err(StoreError::PasteArtifactUnsupported) + } + + /// Release a pending-operation pin without changing committed ownership. + fn release_uploaded_file_pin( + &self, + _session_id: SessionId, + _artifact_id: &str, + _owner_id: &str, + ) -> Result<(), StoreError> { + Err(StoreError::PasteArtifactUnsupported) + } + /// Delete an uncommitted uploaded file owned by `session_id`. fn delete_uploaded_file( &self, diff --git a/crates/session-store/src/uploaded_file.rs b/crates/session-store/src/uploaded_file.rs index 4ce37588..a7c544ea 100644 --- a/crates/session-store/src/uploaded_file.rs +++ b/crates/session-store/src/uploaded_file.rs @@ -24,6 +24,12 @@ pub const DEFAULT_MAX_FILES_PER_SUBMISSION: usize = 8; pub const DEFAULT_MAX_SESSION_UPLOADED_FILES: u64 = 256; const MAX_FILE_NAME_CHARS: usize = 255; const MAX_MEDIA_TYPE_BYTES: usize = 127; +fn validate_pending_owner_id(owner_id: &str) -> Result<()> { + if owner_id.is_empty() || owner_id.len() > 256 { + return Err(StoreError::ArtifactIntegrityMismatch); + } + Ok(()) +} #[derive(Debug, Clone, Copy, PartialEq, Eq)] pub struct UploadedFileLimits { @@ -59,6 +65,8 @@ struct StoredUploadedFile { #[serde(default, skip_serializing_if = "Option::is_none")] source_entry_id: Option, #[serde(default, skip_serializing_if = "Option::is_none")] + pending_owner_id: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] upload_context: Option, content_base64: String, } @@ -291,6 +299,7 @@ pub(crate) fn write_uploaded_file( byte_len, sha256: sha256.clone(), source_entry_id: None, + pending_owner_id: None, upload_context: context.cloned(), content_base64: BASE64.encode(content), }; @@ -376,6 +385,72 @@ pub(crate) fn clear_uploaded_file_binding( Ok(()) } +pub(crate) fn pin_uploaded_file( + dir: &Path, + reference: &UploadedFileRef, + owner_id: &str, +) -> Result<()> { + validate_pending_owner_id(owner_id)?; + if reference.source_entry_id.is_some() { + return Err(StoreError::ArtifactAlreadyCommitted); + } + let aggregate_lock = fs::OpenOptions::new() + .create(true) + .read(true) + .write(true) + .open(dir.join(".aggregate.lock"))?; + FileExt::lock_exclusive(&aggregate_lock)?; + let path = record_path(dir, &reference.artifact_id)?; + let mut stored: StoredUploadedFile = serde_json::from_slice(&fs::read(&path)?)?; + if stored.file_name != reference.file_name + || stored.media_type != reference.media_type + || stored.created_at_ms != reference.created_at_ms + || stored.byte_len != reference.byte_len + || stored.sha256 != reference.sha256 + { + return Err(StoreError::ArtifactIntegrityMismatch); + } + if stored.source_entry_id.is_some() { + return Err(StoreError::ArtifactAlreadyCommitted); + } + if let Some(existing_owner) = stored.pending_owner_id.as_deref() { + return if existing_owner == owner_id { + Ok(()) + } else { + Err(StoreError::ArtifactAlreadyCommitted) + }; + } + stored.pending_owner_id = Some(owner_id.to_owned()); + let temp = dir.join(format!(".{}.file.pin.tmp", reference.artifact_id)); + fs::write(&temp, serde_json::to_vec(&stored)?)?; + fs::rename(temp, path)?; + Ok(()) +} + +pub(crate) fn release_uploaded_file_pin( + dir: &Path, + artifact_id: &str, + owner_id: &str, +) -> Result<()> { + validate_pending_owner_id(owner_id)?; + let aggregate_lock = fs::OpenOptions::new() + .create(true) + .read(true) + .write(true) + .open(dir.join(".aggregate.lock"))?; + FileExt::lock_exclusive(&aggregate_lock)?; + let path = record_path(dir, artifact_id)?; + let mut stored: StoredUploadedFile = serde_json::from_slice(&fs::read(&path)?)?; + if stored.source_entry_id.is_some() || stored.pending_owner_id.as_deref() != Some(owner_id) { + return Err(StoreError::ArtifactIntegrityMismatch); + } + stored.pending_owner_id = None; + let temp = dir.join(format!(".{artifact_id}.file.unpin.tmp")); + fs::write(&temp, serde_json::to_vec(&stored)?)?; + fs::rename(temp, path)?; + Ok(()) +} + pub(crate) fn bind_uploaded_file( dir: &Path, reference: &UploadedFileRef, @@ -405,6 +480,7 @@ pub(crate) fn bind_uploaded_file( return Err(StoreError::ArtifactAlreadyCommitted); } stored.source_entry_id = Some(source_entry_id.to_owned()); + stored.pending_owner_id = None; let temp = dir.join(format!(".{}.file.bind.tmp", reference.artifact_id)); fs::write(&temp, serde_json::to_vec(&stored)?)?; fs::rename(&temp, path)?; @@ -455,7 +531,7 @@ pub(crate) fn copy_committed_uploaded_files(source_dir: &Path, target_dir: &Path } let bytes = fs::read(&path)?; let stored: StoredUploadedFile = serde_json::from_slice(&bytes)?; - if stored.source_entry_id.is_none() { + if stored.source_entry_id.is_none() && stored.pending_owner_id.is_none() { continue; } let target = target_dir.join(name); @@ -499,7 +575,7 @@ pub(crate) fn delete_uncommitted_uploaded_files(dir: &Path) -> Result { continue; } let stored: StoredUploadedFile = serde_json::from_slice(&fs::read(&path)?)?; - if stored.source_entry_id.is_none() { + if stored.source_entry_id.is_none() && stored.pending_owner_id.is_none() { fs::remove_file(path)?; removed = removed .checked_add(1) @@ -523,7 +599,7 @@ pub(crate) fn delete_uploaded_file(dir: &Path, artifact_id: &str) -> Result return Ok(false), Err(error) => return Err(error.into()), }; - if stored.source_entry_id.is_some() { + if stored.source_entry_id.is_some() || stored.pending_owner_id.is_some() { return Err(StoreError::ArtifactAlreadyCommitted); } match fs::remove_file(path) { diff --git a/crates/worker-runtime/src/worker_backend.rs b/crates/worker-runtime/src/worker_backend.rs index 07800e82..408d8155 100644 --- a/crates/worker-runtime/src/worker_backend.rs +++ b/crates/worker-runtime/src/worker_backend.rs @@ -1522,6 +1522,7 @@ fn method_starts_turn(method: &Method) -> bool { Method::Submit { .. } | Method::SubmitTracked { .. } | Method::Notify { auto_run: true, .. } + | Method::NotifyTracked { auto_run: true, .. } | Method::Resume | Method::Compact ) @@ -1549,6 +1550,7 @@ fn accepted_run_state_for_method(method: &Method) -> WorkerExecutionRunState { Method::Submit { .. } | Method::SubmitTracked { .. } | Method::Notify { auto_run: true, .. } + | Method::NotifyTracked { auto_run: true, .. } | Method::Resume | Method::Compact => WorkerExecutionRunState::Busy, Method::Shutdown => WorkerExecutionRunState::Stopped, @@ -1946,6 +1948,9 @@ where input: input.segments.unwrap_or_else(|| { vec![Segment::text(input.content.trim().to_string())] }), + source: protocol::AuthenticatedInputSource::Backend { + operation_id: submission_id.clone(), + }, }, Some(submission_id), ) @@ -1966,6 +1971,7 @@ where Method::Submit { .. } | Method::SubmitTracked { .. } | Method::Notify { .. } + | Method::NotifyTracked { .. } | Method::Compact => WorkerExecutionRunState::Busy, _ => WorkerExecutionRunState::Idle, }; diff --git a/crates/worker/src/controller.rs b/crates/worker/src/controller.rs index 03a98e5e..14fd26f6 100644 --- a/crates/worker/src/controller.rs +++ b/crates/worker/src/controller.rs @@ -229,6 +229,43 @@ enum PendingRun { Resume, } +fn stage_pending_notification( + pending_submissions: &crate::worker::PendingSubmissionHandle, + 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( + pending_submissions: &crate::worker::PendingSubmissionHandle, + 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( pending_submissions: &crate::worker::PendingSubmissionHandle, notify_buffer: &NotifyBuffer, @@ -241,7 +278,12 @@ fn prepare_pending_run( 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( ); let mut pending: Option = None; let pending_submissions = worker.pending_submission_handle(); + stage_oldest_passive_notification(&pending_submissions, ¬ify_buffer); loop { // Top-of-iteration: if an event handler staged a run, fire it @@ -1347,6 +1390,8 @@ async fn controller_loop( } => 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( 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( .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(¬ification_request_id); + stage_oldest_passive_notification(&pending_submissions, ¬ify_buffer); } if !shutdown && may_drain_pending && new_status == WorkerStatus::Idle { @@ -1470,13 +1519,47 @@ async fn controller_loop( 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( message, auto_run, } => { - if auto_run { - match pending_submissions.accept_notification(notification_request_id, message) - { - Ok(true) => { - match prepare_pending_run(&pending_submissions, ¬ify_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, ¬ify_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, + ¬ify_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, ¬ify_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, + ¬ify_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 { .. }) => { diff --git a/crates/worker/src/ipc/interceptor.rs b/crates/worker/src/ipc/interceptor.rs index 5853a2f7..86f26230 100644 --- a/crates/worker/src/ipc/interceptor.rs +++ b/crates/worker/src/ipc/interceptor.rs @@ -178,14 +178,21 @@ impl WorkerInterceptor { /// matches worker-history order. fn commit_system_items_with_extensions( &self, - items: &[(SystemItem, Vec)], + items: &[( + SystemItem, + Vec, + Option, + )], ) -> 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::>(), ) } @@ -341,8 +348,11 @@ impl Interceptor for WorkerInterceptor { projection_digest: projection.catalog_digest.clone(), logical_name: "internal.notify_wrapper".to_string(), }; - let mut system_items: Vec<(SystemItem, Vec)> = - Vec::with_capacity(drained.len()); + let mut system_items: Vec<( + SystemItem, + Vec, + Option, + )> = Vec::with_capacity(drained.len()); let mut items: Vec = Vec::with_capacity(drained.len()); for entry in &drained { let system_item = match build_system_item_with_provenance( @@ -360,7 +370,7 @@ impl Interceptor 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); diff --git a/crates/worker/src/ipc/notify_buffer.rs b/crates/worker/src/ipc/notify_buffer.rs index cbaedc5d..2212dfe1 100644 --- a/crates/worker/src/ipc/notify_buffer.rs +++ b/crates/worker/src/ipc/notify_buffer.rs @@ -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, + history_provenance: Option, }, WorkerEvent { event: WorkerEvent, @@ -58,6 +59,15 @@ impl PendingNotify { PendingNotify::WorkerEvent { .. } => Vec::new(), } } + + pub(crate) fn history_provenance(&self) -> Option { + 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(); diff --git a/crates/worker/src/worker.rs b/crates/worker/src/worker.rs index 3275aab0..d24658d0 100644 --- a/crates/worker/src/worker.rs +++ b/crates/worker/src/worker.rs @@ -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, @@ -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( + ¬ification.source_namespace, + ¬ification.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( + ¬ification.source_namespace, + ¬ification.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::() + ) +} + 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 { + 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, activate_now: bool, + ) -> Result { + 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, + source_namespace: String, + provenance: WorkerHistoryProvenance, + activate_now: bool, ) -> Result { 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(¤t) { + 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 { + 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 { 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 { + 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 { + 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::>(); 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, + history_provenance: Option, ) -> Result, 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 Worker { .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 Worker { 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 Worker { &mut self, input: Vec, mut input_extensions: Vec, + input_provenance: WorkerHistoryProvenance, on_input_committed: F, ) -> Result where @@ -3560,8 +3804,12 @@ impl Worker { 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 Worker { input: &[Segment], flow_projection: Option<&PreparedFlowProjection>, entry_ids: &[SessionHistoryEntryId], + provenance: &WorkerHistoryProvenance, ) -> Vec> { if let Some(flow) = flow_projection { return input @@ -3887,10 +4136,7 @@ impl Worker { 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 Worker { .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 { diff --git a/crates/worker/tests/controller_test.rs b/crates/worker/tests/controller_test.rs index 3f6438fe..319f773d 100644 --- a/crates/worker/tests/controller_test.rs +++ b/crates/worker/tests/controller_test.rs @@ -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()); diff --git a/crates/workspace-server/src/hosts.rs b/crates/workspace-server/src/hosts.rs index f86a1d00..f7750e74 100644 --- a/crates/workspace-server/src/hosts.rs +++ b/crates/workspace-server/src/hosts.rs @@ -4987,7 +4987,7 @@ mod tests { "missing test context", ); }; - let submission_id = input.submission_id.clone(); + let submission_request_id = input.submission_request_id.clone(); let content = input.content; std::thread::spawn(move || { std::thread::sleep(std::time::Duration::from_millis(10)); @@ -5004,11 +5004,13 @@ mod tests { status: protocol::WorkerStatus::Idle, }); }); - if let Some(submission_id) = submission_id { - worker_runtime::execution::WorkerExecutionResult::accepted_input_committed( + if let Some(submission_request_id) = submission_request_id { + worker_runtime::execution::WorkerExecutionResult::accepted_submission( worker_runtime::execution::WorkerExecutionOperation::Input, WorkerExecutionRunState::Busy, - submission_id, + submission_request_id, + uuid::Uuid::now_v7().to_string(), + protocol::SubmissionDisposition::Started, ) } else { worker_runtime::execution::WorkerExecutionResult::accepted( diff --git a/crates/workspace-server/src/runtime_subscription_tests.rs b/crates/workspace-server/src/runtime_subscription_tests.rs index a42c0c02..6c14c56a 100644 --- a/crates/workspace-server/src/runtime_subscription_tests.rs +++ b/crates/workspace-server/src/runtime_subscription_tests.rs @@ -32,11 +32,13 @@ impl WorkerExecutionBackend for TestExecutionBackend { _handle: &WorkerExecutionHandle, input: worker_runtime::interaction::WorkerInput, ) -> WorkerExecutionResult { - if let Some(submission_id) = input.submission_id { - WorkerExecutionResult::accepted_input_committed( + if let Some(submission_request_id) = input.submission_request_id { + WorkerExecutionResult::accepted_submission( WorkerExecutionOperation::Input, WorkerExecutionRunState::Busy, - submission_id, + submission_request_id, + uuid::Uuid::now_v7().to_string(), + protocol::SubmissionDisposition::Started, ) } else { WorkerExecutionResult::accepted( diff --git a/crates/workspace-server/src/server.rs b/crates/workspace-server/src/server.rs index 3599f26d..d53b239d 100644 --- a/crates/workspace-server/src/server.rs +++ b/crates/workspace-server/src/server.rs @@ -8537,13 +8537,15 @@ async fn scoped_list_runtimes( async fn scoped_workspace_protocol_ws( State(api): State, + Extension(actor): Extension, AxumPath(workspace_id): AxumPath, ws: axum::extract::ws::WebSocketUpgrade, ) -> std::result::Result { validate_workspace_scope(&api, &workspace_id).map_err(|error| error.into_response())?; + let input_source = authenticated_browser_input_source(&actor); Ok(ws .on_upgrade(move |socket| { - crate::workspace_subscription::serve_workspace_subscription(api, socket) + crate::workspace_subscription::serve_workspace_subscription(api, socket, input_source) }) .into_response()) } @@ -11458,6 +11460,7 @@ async fn scoped_cancel_runtime_worker( async fn scoped_worker_protocol_ws( ws: WebSocketUpgrade, State(api): State, + Extension(actor): Extension, AxumPath(path): AxumPath, ) -> Response { if let Err(err) = validate_workspace_scope(&api, &path.workspace_id) { @@ -11465,6 +11468,7 @@ async fn scoped_worker_protocol_ws( } worker_protocol_ws( State(api), + Extension(actor), AxumPath((path.worker.runtime_id, path.worker.worker_id)), ws, ) @@ -13877,8 +13881,45 @@ async fn cancel_runtime_worker( Ok(Json(result)) } +fn authenticated_browser_input_source(actor: &RequestActor) -> protocol::AuthenticatedInputSource { + protocol::AuthenticatedInputSource::Account { + account_id: actor.account_id.clone(), + } +} + +pub(crate) fn authorize_browser_worker_method( + method: protocol::Method, + source: &protocol::AuthenticatedInputSource, +) -> std::result::Result { + match method { + protocol::Method::Submit { + submission_request_id, + input, + } => Ok(protocol::Method::SubmitTracked { + submission_request_id, + input, + source: source.clone(), + }), + protocol::Method::Notify { + notification_request_id, + message, + auto_run, + } => Ok(protocol::Method::NotifyTracked { + notification_request_id, + message, + auto_run, + source: source.clone(), + }), + protocol::Method::SubmitTracked { .. } | protocol::Method::NotifyTracked { .. } => { + Err("authenticated Worker input source is server-owned") + } + other => Ok(other), + } +} + async fn worker_protocol_ws( State(api): State, + Extension(actor): Extension, AxumPath((runtime_id, worker_id)): AxumPath<(String, String)>, ws: WebSocketUpgrade, ) -> impl IntoResponse { @@ -13902,7 +13943,8 @@ async fn worker_protocol_ws( .into_response(); } }; - ws.on_upgrade(move |socket| worker_protocol_ws_session(source, socket)) + let input_source = authenticated_browser_input_source(&actor); + ws.on_upgrade(move |socket| worker_protocol_ws_session(source, socket, input_source)) } pub(crate) struct WorkspaceWorkerProtocolConnection { @@ -14029,13 +14071,17 @@ async fn connect_embedded_worker_protocol( Ok(WorkspaceWorkerProtocolConnection { methods, events }) } -async fn worker_protocol_ws_session(source: RuntimeObservationSource, socket: WebSocket) { +async fn worker_protocol_ws_session( + source: RuntimeObservationSource, + socket: WebSocket, + input_source: protocol::AuthenticatedInputSource, +) { match source { RuntimeObservationSource::RemoteWs(config) => { - remote_worker_protocol_ws_session(config, socket).await; + remote_worker_protocol_ws_session(config, socket, input_source).await; } RuntimeObservationSource::Embedded(source) => { - embedded_worker_protocol_ws_session(source, socket).await; + embedded_worker_protocol_ws_session(source, socket, input_source).await; } } } @@ -14043,6 +14089,7 @@ async fn worker_protocol_ws_session(source: RuntimeObservationSource, socket: We async fn remote_worker_protocol_ws_session( config: RuntimeObservationSourceConfig, socket: WebSocket, + input_source: protocol::AuthenticatedInputSource, ) { let mut request = match config.endpoint.clone().into_client_request() { Ok(request) => request, @@ -14091,14 +14138,33 @@ async fn remote_worker_protocol_ws_session( inbound = client_stream.next() => { match inbound { Some(Ok(WsMessage::Text(text))) => { - if upstream_sink.send(TungsteniteMessage::Text(text.to_string().into())).await.is_err() { + let method = match protocol::stream::decode_method(text.as_ref()) { + Ok(method) => match authorize_browser_worker_method(method, &input_source) { + Ok(method) => method, + Err(message) => { + if let Ok(event) = protocol::stream::encode_event(&protocol_error_event(message)) { + let _ = client_sink.send(WsMessage::Text(event.into())).await; + } + break; + } + }, + Err(error) => { + if let Ok(event) = protocol::stream::encode_event(&protocol_error_event(error.to_string())) { + let _ = client_sink.send(WsMessage::Text(event.into())).await; + } + break; + } + }; + let Ok(method) = protocol::stream::encode_method(&method) else { break }; + if upstream_sink.send(TungsteniteMessage::Text(method.into())).await.is_err() { break; } } - Some(Ok(WsMessage::Binary(binary))) => { - if upstream_sink.send(TungsteniteMessage::Binary(binary.to_vec().into())).await.is_err() { - break; + Some(Ok(WsMessage::Binary(_))) => { + if let Ok(event) = protocol::stream::encode_event(&protocol_error_event("binary Worker methods are not accepted")) { + let _ = client_sink.send(WsMessage::Text(event.into())).await; } + break; } Some(Ok(WsMessage::Close(_))) | None => { let _ = upstream_sink.send(TungsteniteMessage::Close(None)).await; @@ -14154,6 +14220,7 @@ async fn remote_worker_protocol_ws_session( async fn embedded_worker_protocol_ws_session( source: crate::observation::EmbeddedRuntimeObservationSource, mut socket: WebSocket, + input_source: protocol::AuthenticatedInputSource, ) { let mut upstream = match RuntimeObservationClient::connect(&RuntimeObservationSource::Embedded( source.clone(), @@ -14173,24 +14240,32 @@ async fn embedded_worker_protocol_ws_session( inbound = socket.next() => { match inbound { Some(Ok(WsMessage::Text(text))) => match decode_method(&text) { - Ok(method) => match source.runtime.send_protocol_method(&source.worker_ref, method) { - Ok(events) => { - for event in events { + Ok(method) => match authorize_browser_worker_method(method, &input_source) { + Ok(method) => match source.runtime.send_protocol_method(&source.worker_ref, method) { + Ok(events) => { + for event in events { + if !send_protocol_event(&mut socket, &event).await { + return; + } + } + } + Err(error) => { + let event = protocol_error_event(error.to_string()); if !send_protocol_event(&mut socket, &event).await { return; } } - } - Err(error) => { - let event = protocol_error_event(error.to_string()); - if !send_protocol_event(&mut socket, &event).await { - return; - } + }, + Err(message) => { + let event = protocol_error_event(message); + let _ = send_protocol_event(&mut socket, &event).await; + return; } }, Err(error) => { - let event = - protocol_error_event(format!("malformed protocol method frame: {error}")); + let event = protocol_error_event(format!( + "malformed protocol method frame: {error}" + )); if !send_protocol_event(&mut socket, &event).await { return; } @@ -16580,6 +16655,29 @@ mod tests { &tail[..end] } + #[test] + fn browser_worker_methods_receive_server_owned_account_source() { + let source = protocol::AuthenticatedInputSource::Account { + account_id: "account-1".into(), + }; + let method = authorize_browser_worker_method( + protocol::Method::Submit { + submission_request_id: "request-1".into(), + input: vec![protocol::Segment::text("hello")], + }, + &source, + ) + .unwrap(); + assert!(matches!( + method, + protocol::Method::SubmitTracked { + source: protocol::AuthenticatedInputSource::Account { ref account_id }, + .. + } if account_id == "account-1" + )); + assert!(authorize_browser_worker_method(method, &source).is_err()); + } + #[test] fn merge_request_http_paths_observe_refs_through_runtime_provider_authority() { let source = include_str!("server.rs"); @@ -18825,7 +18923,7 @@ mod tests { .get(handle.worker_ref()) .cloned() .expect("execution context"); - let submission_id = input.submission_id.clone(); + let submission_request_id = input.submission_request_id.clone(); let content = input.content.clone(); std::thread::spawn(move || { std::thread::sleep(std::time::Duration::from_millis(25)); @@ -18833,11 +18931,13 @@ mod tests { text: format!("server companion echoed: {content}"), }); }); - if let Some(submission_id) = submission_id { - worker_runtime::execution::WorkerExecutionResult::accepted_input_committed( + if let Some(submission_request_id) = submission_request_id { + worker_runtime::execution::WorkerExecutionResult::accepted_submission( worker_runtime::execution::WorkerExecutionOperation::Input, worker_runtime::execution::WorkerExecutionRunState::Idle, - submission_id, + submission_request_id, + uuid::Uuid::now_v7().to_string(), + protocol::SubmissionDisposition::Started, ) } else { worker_runtime::execution::WorkerExecutionResult::accepted( @@ -27223,6 +27323,16 @@ mod tests { (runtime, worker_ref, endpoint) } + fn test_browser_request_actor() -> RequestActor { + RequestActor { + user_id: "test-user".into(), + account_id: format!("account-{TEST_WORKSPACE_ID}"), + handle: "test".into(), + display_name: "Test".into(), + auth_method: ActorAuthMethod::BrowserSession, + } + } + async fn spawn_workspace_proxy( source: RuntimeObservationSourceConfig, ) -> (String, tempfile::TempDir) { @@ -27241,11 +27351,8 @@ mod tests { .unwrap(); let app_listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); let app_addr = app_listener.local_addr().unwrap(); - tokio::spawn(async move { - axum::serve(app_listener, build_inner_router(api)) - .await - .unwrap() - }); + let app = build_inner_router(api).layer(Extension(test_browser_request_actor())); + tokio::spawn(async move { axum::serve(app_listener, app).await.unwrap() }); ( format!("ws://{app_addr}/api/runtimes/{runtime_id}/workers/{worker_id}/protocol/ws"), dir, @@ -27257,7 +27364,8 @@ mod tests { let dir = tempfile::tempdir().unwrap(); let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); let address = listener.local_addr().unwrap(); - let app = build_inner_router(test_api(dir.path()).await); + let app = build_inner_router(test_api(dir.path()).await) + .layer(Extension(test_browser_request_actor())); let server = tokio::spawn(async move { let _ = axum::serve(listener, app).await; }); @@ -27305,7 +27413,7 @@ mod tests { let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); let address = listener.local_addr().unwrap(); - let app = build_inner_router(api); + let app = build_inner_router(api).layer(Extension(test_browser_request_actor())); let server = tokio::spawn(async move { let _ = axum::serve(listener, app).await; }); diff --git a/crates/workspace-server/src/workspace_subscription.rs b/crates/workspace-server/src/workspace_subscription.rs index fe26f7f9..d72017f1 100644 --- a/crates/workspace-server/src/workspace_subscription.rs +++ b/crates/workspace-server/src/workspace_subscription.rs @@ -11,7 +11,9 @@ use tokio::sync::mpsc; use worker_runtime::identity::RuntimeWorkerRef; use crate::runtime_subscription::{BrokerSubscriptionEvent, RuntimeSubscriptionBroker}; -use crate::server::{WorkspaceApi, connect_workspace_worker_protocol}; +use crate::server::{ + WorkspaceApi, authorize_browser_worker_method, connect_workspace_worker_protocol, +}; use crate::store::WorkspaceResourceKind; const OUTBOUND_CAPACITY: usize = 256; @@ -21,7 +23,11 @@ struct ActiveSubscription { methods: Option>, } -pub(crate) async fn serve_workspace_subscription(api: WorkspaceApi, socket: WebSocket) { +pub(crate) async fn serve_workspace_subscription( + api: WorkspaceApi, + socket: WebSocket, + input_source: protocol::AuthenticatedInputSource, +) { let broker = api.runtime_subscription_broker().clone(); let (mut socket_sender, mut socket_receiver) = socket.split(); let (control_outbound, mut control_receiver) = mpsc::channel::(OUTBOUND_CAPACITY); @@ -153,7 +159,12 @@ pub(crate) async fn serve_workspace_subscription(api: WorkspaceApi, socket: WebS else { break; }; - if methods.send(message.method).await.is_err() { + let Ok(method) = + authorize_browser_worker_method(message.method, &input_source) + else { + break; + }; + if methods.send(method).await.is_err() { break; } }