fix: reclaim failed worker creations

This commit is contained in:
2026-09-07 04:44:19 +09:00
parent f5c5ea5a0b
commit cc27d57e4a
4 changed files with 762 additions and 141 deletions
+76 -12
View File
@@ -153,6 +153,7 @@ impl Drop for RuntimeEventSelectorSubscription {
#[derive(Clone, Debug)]
pub struct Runtime {
inner: Arc<Mutex<RuntimeState>>,
worker_operations: Arc<Mutex<BTreeMap<WorkerId, Arc<Mutex<()>>>>>,
}
impl Runtime {
@@ -166,6 +167,7 @@ impl Runtime {
let state = RuntimeState::new(options.display_name);
Self {
inner: Arc::new(Mutex::new(state)),
worker_operations: Arc::new(Mutex::new(BTreeMap::new())),
}
}
@@ -222,6 +224,7 @@ impl Runtime {
state.execution_backend = execution_backend;
let runtime = Self {
inner: Arc::new(Mutex::new(state)),
worker_operations: Arc::new(Mutex::new(BTreeMap::new())),
};
runtime.restore_persisted_worker_executions()?;
Ok(runtime)
@@ -781,6 +784,10 @@ impl Runtime {
request: CreateWorkerRequest,
scope: Option<&RuntimeWorkspaceScope>,
) -> Result<WorkerDetail, RuntimeError> {
let operation_lock = self.worker_operation_lock(request.worker_id)?;
let _operation_guard = operation_lock
.lock()
.map_err(|_| RuntimeError::StatePoisoned)?;
if let Some(existing) = self.existing_worker_for_create(&request, scope)? {
return Ok(existing);
}
@@ -893,8 +900,7 @@ impl Runtime {
initial_input.submission_request_id = Some(expected_submission_id.clone());
let dispatch_result = backend.dispatch_input(&handle, initial_input.clone());
if !dispatch_result.is_accepted() {
let _ = backend.stop_worker(&handle);
self.rollback_failed_create(&worker_ref)?;
self.cleanup_connected_failed_create(&backend, &worker_ref, &handle)?;
return Err(RuntimeError::WorkerExecutionRejected {
worker_id: worker_ref.worker_id.clone(),
operation: dispatch_result.operation,
@@ -908,8 +914,7 @@ impl Runtime {
.as_ref()
.is_some_and(|ack| ack.submission_request_id == expected_submission_id);
if !has_commit_ack {
let _ = backend.stop_worker(&handle);
self.rollback_failed_create(&worker_ref)?;
self.cleanup_connected_failed_create(&backend, &worker_ref, &handle)?;
let result = WorkerExecutionResult::rejected(
WorkerExecutionOperation::Input,
"execution backend accepted initial input without a durable session commit acknowledgement",
@@ -922,21 +927,36 @@ impl Runtime {
result,
});
}
let detail = self.commit_created_worker(
let detail = match self.commit_created_worker(
&worker_ref,
handle,
handle.clone(),
working_directory,
dispatch_result,
)?;
self.record_input_observation(&worker_ref, initial_input)?;
) {
Ok(detail) => detail,
Err(error) => {
self.cleanup_connected_failed_create(&backend, &worker_ref, &handle)?;
return Err(error);
}
};
if let Err(error) = self.record_input_observation(&worker_ref, initial_input) {
self.cleanup_connected_failed_create(&backend, &worker_ref, &handle)?;
return Err(error);
}
Ok(detail)
} else {
self.commit_created_worker(
match self.commit_created_worker(
&worker_ref,
handle,
handle.clone(),
working_directory,
WorkerExecutionResult::accepted(WorkerExecutionOperation::Spawn),
)
) {
Ok(detail) => Ok(detail),
Err(error) => {
self.cleanup_connected_failed_create(&backend, &worker_ref, &handle)?;
Err(error)
}
}
}
}
@@ -1639,14 +1659,43 @@ impl Runtime {
Ok(detail)
}
fn cleanup_connected_failed_create(
&self,
backend: &WorkerExecutionBackendRef,
worker_ref: &WorkerRef,
handle: &WorkerExecutionHandle,
) -> Result<(), RuntimeError> {
let stop_result = backend.stop_worker(handle);
if stop_result.is_accepted() {
return self.rollback_failed_create(worker_ref);
}
let mut state = self.lock()?;
let record = state.worker_mut(worker_ref)?;
record.execution_handle = Some(handle.clone());
record.worker_state = stop_result.worker_state.clone();
state.persist_runtime_snapshot()?;
state.persist_worker(&worker_ref.worker_id)?;
Ok(())
}
fn rollback_failed_create(&self, worker_ref: &WorkerRef) -> Result<(), RuntimeError> {
let mut state = self.lock()?;
if let Some(record) = state.workers.remove(&worker_ref.worker_id) {
if state.workers.contains_key(&worker_ref.worker_id) {
state.delete_worker_snapshot(&worker_ref.worker_id)?;
let record = state
.workers
.remove(&worker_ref.worker_id)
.expect("Worker existence checked before failed-create rollback");
let workspace_id = record.workspace_id.clone();
if let Some(workspace_id) = workspace_id.as_deref() {
state.forget_workspace_owner_if_unused(workspace_id);
}
#[cfg(feature = "ws-server")]
state
.observation_events
.retain(|event| event.worker_ref != *worker_ref);
state.publish_worker_removed(worker_ref.worker_id, workspace_id.as_deref())?;
state.persist_runtime_snapshot()?;
}
Ok(())
}
@@ -1800,6 +1849,10 @@ impl Runtime {
&self,
worker_ref: &WorkerRef,
) -> Result<WorkerDeleteResult, RuntimeError> {
let operation_lock = self.worker_operation_lock(worker_ref.worker_id)?;
let _operation_guard = operation_lock
.lock()
.map_err(|_| RuntimeError::StatePoisoned)?;
let mut state = self.lock()?;
state.ensure_running()?;
state.ensure_worker_ref(worker_ref)?;
@@ -2260,6 +2313,17 @@ impl Runtime {
Ok(result)
}
fn worker_operation_lock(&self, worker_id: WorkerId) -> Result<Arc<Mutex<()>>, RuntimeError> {
let mut operations = self
.worker_operations
.lock()
.map_err(|_| RuntimeError::StatePoisoned)?;
Ok(operations
.entry(worker_id)
.or_insert_with(|| Arc::new(Mutex::new(())))
.clone())
}
fn lock(&self) -> Result<MutexGuard<'_, RuntimeState>, RuntimeError> {
self.inner.lock().map_err(|_| RuntimeError::StatePoisoned)
}
+143 -16
View File
@@ -1216,6 +1216,7 @@ pub struct WorkerRuntimeExecutionBackend<F = ProfileRuntimeWorkerFactory> {
working_directory_materializer: Option<Arc<dyn WorkingDirectoryMaterializer>>,
runtime: Mutex<Option<Runtime>>,
workers: Mutex<HashMap<crate::identity::WorkerRef, RuntimeWorkerExecution>>,
spawn_restore_timeout: Duration,
}
impl WorkerRuntimeExecutionBackend<ProfileRuntimeWorkerFactory> {
@@ -1243,6 +1244,7 @@ where
working_directory_materializer: None,
runtime: Mutex::new(Some(runtime)),
workers: Mutex::new(HashMap::new()),
spawn_restore_timeout: RUNTIME_TASK_TIMEOUT,
})
}
@@ -1259,6 +1261,12 @@ where
self
}
#[cfg(test)]
fn with_spawn_restore_timeout(mut self, timeout: Duration) -> Self {
self.spawn_restore_timeout = timeout;
self
}
fn wait_for_runtime_task<T>(receiver: mpsc::Receiver<Result<T, String>>) -> Result<T, String> {
receiver
.recv_timeout(RUNTIME_TASK_TIMEOUT)
@@ -1297,6 +1305,39 @@ where
Self::wait_for_runtime_task(rx)
}
fn run_spawn_restore_on_adapter_runtime<T, Fut>(&self, task: Fut) -> Result<T, String>
where
T: Send + 'static,
Fut: Future<Output = Result<T, String>> + Send + 'static,
{
let timeout = self.spawn_restore_timeout;
let (tx, rx) = mpsc::sync_channel(1);
self.spawn_on_adapter_runtime(async move {
let mut handle = tokio::spawn(task);
let result = tokio::select! {
biased;
result = &mut handle => match result {
Ok(result) => result,
Err(err) => Err(format!("worker adapter task failed: {err}")),
},
_ = tokio::time::sleep(timeout) => {
handle.abort();
match handle.await {
Ok(result) => result,
Err(err) if err.is_cancelled() => Err(format!(
"worker adapter task did not complete within {} seconds and was cancelled",
timeout.as_secs_f64()
)),
Err(err) => Err(format!("worker adapter task failed: {err}")),
}
}
};
let _ = tx.send(result);
})?;
rx.recv()
.map_err(|err| format!("worker adapter task did not complete: {err}"))?
}
fn get_execution(
&self,
handle: &WorkerExecutionHandle,
@@ -1756,8 +1797,9 @@ where
let factory = self.factory.clone();
let bridge_context = request.context.clone();
let worker_ref = request.worker_ref.clone();
let spawn_result =
self.run_on_adapter_runtime(async move { factory.spawn_controller(request).await });
let spawn_result = self.run_spawn_restore_on_adapter_runtime(async move {
factory.spawn_controller(request).await
});
let controller = match spawn_result {
Ok(controller) => controller,
@@ -1859,8 +1901,9 @@ where
let factory = self.factory.clone();
let bridge_context = request.context.clone();
let worker_ref = request.worker_ref.clone();
let restore_result =
self.run_on_adapter_runtime(async move { factory.restore_controller(request).await });
let restore_result = self.run_spawn_restore_on_adapter_runtime(async move {
factory.restore_controller(request).await
});
let controller = match restore_result {
Ok(controller) => controller,
@@ -2100,21 +2143,24 @@ where
if let Err(message) = shutdown_wait {
return WorkerExecutionResult::errored(WorkerExecutionOperation::Stop, message);
}
if let Err(error) = artifact_cleanup.delete_uncommitted_uploaded_files() {
return WorkerExecutionResult::errored(
WorkerExecutionOperation::Stop,
format!("uploaded_file_cleanup_failed: {error}"),
);
}
let artifact_cleanup_error = artifact_cleanup
.delete_uncommitted_uploaded_files()
.err()
.map(|error| format!("uploaded_file_cleanup_failed: {error}"));
match self.workers.lock() {
Ok(mut workers) => {
workers.remove(handle.worker_ref());
result
}
Err(_) => WorkerExecutionResult::errored(
WorkerExecutionOperation::Stop,
"worker adapter registry lock is poisoned after shutdown",
),
Err(poisoned) => {
poisoned.into_inner().remove(handle.worker_ref());
}
}
if let Some(message) = artifact_cleanup_error {
let mut result = result;
result.message = Some(message);
result
} else {
result
}
}
@@ -2178,7 +2224,7 @@ mod tests {
use std::fs;
use std::pin::Pin;
use std::process::Command;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering};
use crate::Runtime as EmbeddedRuntime;
use crate::catalog::{
@@ -2489,6 +2535,87 @@ mod tests {
WorkerExecutionContext::new(worker_ref)
}
struct DelayedFactory {
completed: Arc<AtomicBool>,
delay: Duration,
}
#[async_trait]
impl RuntimeWorkerFactory for DelayedFactory {
async fn spawn_controller(
&self,
_request: WorkerExecutionSpawnRequest,
) -> Result<RuntimeWorkerController, String> {
tokio::time::sleep(self.delay).await;
self.completed.store(true, Ordering::SeqCst);
Err("delayed factory completed".to_string())
}
async fn restore_controller(
&self,
_request: WorkerExecutionRestoreRequest,
) -> Result<RuntimeWorkerController, String> {
tokio::time::sleep(self.delay).await;
self.completed.store(true, Ordering::SeqCst);
Err("delayed factory completed".to_string())
}
}
#[test]
fn create_timeout_cancels_factory_and_removes_persisted_worker() {
let root = tempfile::tempdir().unwrap();
let runtime_store_dir = root.path().join("runtime");
let completed = Arc::new(AtomicBool::new(false));
let backend = Arc::new(
WorkerRuntimeExecutionBackend::new(DelayedFactory {
completed: completed.clone(),
delay: Duration::from_millis(200),
})
.unwrap()
.with_spawn_restore_timeout(Duration::from_millis(20)),
);
let runtime = EmbeddedRuntime::with_fs_store_and_execution_backend(
crate::fs_store::FsRuntimeStoreOptions {
root: runtime_store_dir.clone(),
runtime_id: "create-timeout-runtime".to_string(),
display_name: None,
},
backend.clone(),
)
.unwrap();
runtime.store_config_bundle(test_bundle()).unwrap();
let request = create_request("create timeout");
let worker_id = request.worker_id;
let create_runtime = runtime.clone();
let create = std::thread::spawn(move || create_runtime.create_worker(request));
std::thread::sleep(Duration::from_millis(5));
let delete_error = runtime
.delete_worker(&crate::identity::WorkerRef::new(worker_id))
.unwrap_err();
let error = create.join().unwrap().unwrap_err();
assert!(matches!(
delete_error,
crate::error::RuntimeError::WorkerNotFound { worker_id: missing } if missing == worker_id
));
assert!(error.to_string().contains("was cancelled"));
std::thread::sleep(Duration::from_millis(250));
assert!(
!completed.load(Ordering::SeqCst),
"timed out factory future must not resume after create returns"
);
assert!(runtime.list_workers().unwrap().is_empty());
assert!(backend.workers.lock().unwrap().is_empty());
assert!(
!runtime_store_dir
.join("workers")
.join(worker_id.to_string())
.exists(),
"failed create must remove its persisted Worker aggregate"
);
}
struct MockFactory {
client: MockClient,
runtime_base: PathBuf,