diff --git a/Cargo.lock b/Cargo.lock index 43fd39429..fb506b56d 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -1045,6 +1045,7 @@ dependencies = [ name = "integration-test-tasks" version = "0.1.0" dependencies = [ + "rmp-serde", "serde", "spider-tdl", ] @@ -1983,6 +1984,7 @@ dependencies = [ name = "spider-execution-manager" version = "0.1.0" dependencies = [ + "anyhow", "async-trait", "bincode", "bytes", diff --git a/components/spider-core/src/types/io.rs b/components/spider-core/src/types/io.rs index fc4ae717c..9485c8cb1 100644 --- a/components/spider-core/src/types/io.rs +++ b/components/spider-core/src/types/io.rs @@ -432,7 +432,7 @@ pub struct ExecutionContext { pub task_instance_id: TaskInstanceId, pub tdl_context: TdlContext, pub timeout_policy: TimeoutPolicy, - pub serialized_inputs: Vec, + pub serialized_task_io: Vec, } #[cfg(test)] diff --git a/components/spider-execution-manager/Cargo.toml b/components/spider-execution-manager/Cargo.toml index 3f6aa1a26..c1f083abe 100644 --- a/components/spider-execution-manager/Cargo.toml +++ b/components/spider-execution-manager/Cargo.toml @@ -45,3 +45,6 @@ tokio = { tokio-util = { version = "0.7", features = ["codec", "rt"] } tonic = "0.14.6" tracing = { version = "0.1.41", default-features = false, features = ["std"] } + +[dev-dependencies] +anyhow = "1.0.98" diff --git a/components/spider-execution-manager/src/process_pool.rs b/components/spider-execution-manager/src/process_pool.rs index 8251dc530..90ce5d8ac 100644 --- a/components/spider-execution-manager/src/process_pool.rs +++ b/components/spider-execution-manager/src/process_pool.rs @@ -19,6 +19,7 @@ use spider_task_executor::protocol::ExecutorOutcome; use spider_task_executor::protocol::Request; use spider_task_executor::protocol::Response; use spider_tdl::TaskContext; +use spider_tdl::TdlError; use spider_utils::wire::WireError; use tokio::process::Child; use tokio::process::ChildStdin; @@ -102,6 +103,10 @@ pub enum InternalError { /// Failed to wire-format-encode the task inputs when building the executor request. #[error("failed to encode task inputs: {0}")] EncodeTaskInputs(#[from] WireError), + + /// Failed to construct the [`TaskContext`] when building the executor request. + #[error("failed to build task context: {0}")] + BuildTaskContext(#[from] TdlError), } /// The process pool of pre-forked task executor subprocesses ready for task execution. @@ -349,13 +354,15 @@ impl ExecutorHandle { /// /// # Returns /// -/// A populated [`Request::Execute`] with `raw_ctx` set to the msgpack-encoded [`TaskContext`] and -/// `raw_inputs` set to the serialized execution inputs on success. +/// A populated [`Request::Execute`] with `raw_ctx` set to the msgpack-encoded [`TaskContext`] on +/// success. For a commit task the serialized payload is routed into the [`TaskContext`]'s +/// task-graph outputs and `raw_inputs` is left empty; for any other task it is set as `raw_inputs`. /// /// # Errors /// /// Returns an error if: /// +/// * Forwards [`TaskContext::new`]'s return values on failure. /// * Forwards [`rmp_serde::to_vec`]'s return values on failure. fn build_request(request: ExecuteRequest) -> Result { let ExecuteRequest { @@ -368,17 +375,130 @@ fn build_request(request: ExecuteRequest) -> Result { task_instance_id, tdl_context, timeout_policy: _, - serialized_inputs, + serialized_task_io, } = ctx; - let raw_ctx = rmp_serde::to_vec(&TaskContext { + let (raw_inputs, serialized_task_graph_outputs) = if task_id == TaskId::Commit { + (Vec::new(), Some(serialized_task_io)) + } else { + (serialized_task_io, None) + }; + let raw_ctx = rmp_serde::to_vec(&TaskContext::new( job_id, task_id, task_instance_id, resource_group_id, - })?; + serialized_task_graph_outputs, + )?)?; Ok(Request::Execute { tdl_context, raw_ctx, - raw_inputs: serialized_inputs, + raw_inputs, }) } + +#[cfg(test)] +mod tests { + use spider_core::task::TdlContext; + use spider_core::task::TimeoutPolicy; + use spider_core::types::io::SerializedTaskOutputs; + use spider_core::types::io::TaskOutput; + + use super::*; + + /// Builds an [`ExecutionContext`] with dummy TDL and timeout metadata wrapping the given task + /// IO buffer. + /// + /// # Returns + /// + /// The constructed [`ExecutionContext`]. + fn make_execution_context(serialized_task_io: Vec) -> ExecutionContext { + ExecutionContext { + task_instance_id: 1, + tdl_context: TdlContext { + package: "test-package".to_owned(), + task_func: "test-task".to_owned(), + }, + timeout_policy: TimeoutPolicy { + soft_timeout_ms: 1_000, + hard_timeout_ms: 2_000, + }, + serialized_task_io, + } + } + + #[test] + fn build_request_routes_commit_outputs_into_context() -> anyhow::Result<()> { + let outputs: Vec = vec![vec![1, 2, 3], vec![4, 5, 6]]; + let serialized = SerializedTaskOutputs::serialize_with_size_hint(&outputs)?.to_raw(); + let request = ExecuteRequest { + job_id: JobId::random(), + task_id: TaskId::Commit, + resource_group_id: ResourceGroupId::random(), + ctx: make_execution_context(serialized), + }; + + let Request::Execute { + raw_ctx, + raw_inputs, + .. + } = build_request(request)? + else { + panic!("build_request must produce a Request::Execute"); + }; + + assert!(raw_inputs.is_empty()); + let ctx: TaskContext = rmp_serde::from_slice(&raw_ctx)?; + assert_eq!(ctx.get_task_graph_outputs()?, Some(outputs)); + Ok(()) + } + + #[test] + fn build_request_routes_empty_commit_outputs() -> anyhow::Result<()> { + let outputs: Vec = Vec::new(); + let serialized = SerializedTaskOutputs::serialize_with_size_hint(&outputs)?.to_raw(); + let request = ExecuteRequest { + job_id: JobId::random(), + task_id: TaskId::Commit, + resource_group_id: ResourceGroupId::random(), + ctx: make_execution_context(serialized), + }; + + let Request::Execute { + raw_ctx, + raw_inputs, + .. + } = build_request(request)? + else { + panic!("build_request must produce a Request::Execute"); + }; + + assert!(raw_inputs.is_empty()); + let ctx: TaskContext = rmp_serde::from_slice(&raw_ctx)?; + assert_eq!(ctx.get_task_graph_outputs()?, Some(Vec::new())); + Ok(()) + } + + #[test] + fn build_request_routes_regular_task_inputs() -> anyhow::Result<()> { + let request = ExecuteRequest { + job_id: JobId::random(), + task_id: TaskId::Index(0), + resource_group_id: ResourceGroupId::random(), + ctx: make_execution_context(vec![10, 20, 30]), + }; + + let Request::Execute { + raw_ctx, + raw_inputs, + .. + } = build_request(request)? + else { + panic!("build_request must produce a Request::Execute"); + }; + + assert_eq!(raw_inputs, vec![10, 20, 30]); + let ctx: TaskContext = rmp_serde::from_slice(&raw_ctx)?; + assert!(ctx.get_task_graph_outputs()?.is_none()); + Ok(()) + } +} diff --git a/components/spider-proto-rust/src/generated/storage.rs b/components/spider-proto-rust/src/generated/storage.rs index b4d4fcc5e..54b4a3ef4 100644 --- a/components/spider-proto-rust/src/generated/storage.rs +++ b/components/spider-proto-rust/src/generated/storage.rs @@ -86,7 +86,7 @@ pub struct ExecutionContext { #[prost(message, optional, tag = "3")] pub timeout_policy: ::core::option::Option, #[prost(bytes = "vec", tag = "4")] - pub serialized_inputs: ::prost::alloc::vec::Vec, + pub serialized_task_io: ::prost::alloc::vec::Vec, } #[derive(Clone, PartialEq, Eq, Hash, ::prost::Message)] pub struct TdlContext { diff --git a/components/spider-proto-rust/src/io.rs b/components/spider-proto-rust/src/io.rs index a7596c2f0..aa20934b0 100644 --- a/components/spider-proto-rust/src/io.rs +++ b/components/spider-proto-rust/src/io.rs @@ -27,7 +27,7 @@ impl TryFrom for ExecutionContext { soft_timeout_ms: timeout_policy.soft_timeout_ms, hard_timeout_ms: timeout_policy.hard_timeout_ms, }, - serialized_inputs: execution_context.serialized_inputs, + serialized_task_io: execution_context.serialized_task_io, }) } } @@ -48,7 +48,7 @@ mod tests { soft_timeout_ms: 100, hard_timeout_ms: 200, }), - serialized_inputs: vec![1, 2, 3], + serialized_task_io: vec![1, 2, 3], }; let execution_context = @@ -59,7 +59,7 @@ mod tests { assert_eq!(execution_context.tdl_context.task_func, "func"); assert_eq!(execution_context.timeout_policy.soft_timeout_ms, 100); assert_eq!(execution_context.timeout_policy.hard_timeout_ms, 200); - assert_eq!(execution_context.serialized_inputs, vec![1, 2, 3]); + assert_eq!(execution_context.serialized_task_io, vec![1, 2, 3]); } #[test] @@ -71,7 +71,7 @@ mod tests { soft_timeout_ms: 100, hard_timeout_ms: 200, }), - serialized_inputs: Vec::new(), + serialized_task_io: Vec::new(), }; assert!(matches!( diff --git a/components/spider-proto/storage/storage.proto b/components/spider-proto/storage/storage.proto index 11caf6e3f..314824364 100644 --- a/components/spider-proto/storage/storage.proto +++ b/components/spider-proto/storage/storage.proto @@ -107,7 +107,7 @@ message ExecutionContext { uint64 task_instance_id = 1; TdlContext tdl_context = 2; TimeoutPolicy timeout_policy = 3; - bytes serialized_inputs = 4; + bytes serialized_task_io = 4; } message TdlContext { diff --git a/components/spider-storage/src/cache/error.rs b/components/spider-storage/src/cache/error.rs index 25a6e6ac3..482364038 100644 --- a/components/spider-storage/src/cache/error.rs +++ b/components/spider-storage/src/cache/error.rs @@ -97,6 +97,9 @@ pub enum InternalError { #[error(transparent)] WireError(#[from] WireError), + #[error(transparent)] + TaskOutputs(#[from] spider_core::types::io::TaskOutputsError), + #[error(transparent)] Db(#[from] crate::db::DbError), } diff --git a/components/spider-storage/src/cache/job.rs b/components/spider-storage/src/cache/job.rs index 10d6043b6..139983a75 100644 --- a/components/spider-storage/src/cache/job.rs +++ b/components/spider-storage/src/cache/job.rs @@ -13,6 +13,7 @@ use spider_core::types::id::ResourceGroupId; use spider_core::types::id::TaskId; use spider_core::types::id::TaskInstanceId; use spider_core::types::io::ExecutionContext; +use spider_core::types::io::SerializedTaskOutputs; use spider_core::types::io::TaskOutput; use tokio::sync::RwLockReadGuard; use tokio::sync::RwLockWriteGuard; @@ -219,20 +220,11 @@ impl< /// Returns an error if: /// /// * [`InternalError::UnexpectedJobState`] if the job is not in [`JobState::Succeeded`]. - /// * [`InternalError::TaskInputNotReady`] if any output has no value. + /// * Forwards [`TaskGraph::read_output_payloads`]'s return values on failure. pub async fn get_outputs(&self) -> Result, CacheError> { let jcb = &self.inner; let job = jcb.job_execution_state.read_succeeded().await?; - let mut outputs = Vec::new(); - for output_reader in job.task_graph.get_outputs() { - let payload = output_reader - .read() - .await - .as_ref() - .ok_or(InternalError::TaskInputNotReady)? - .clone(); - outputs.push(payload); - } + let outputs = job.task_graph.read_output_payloads().await?; drop(job); Ok(outputs) } @@ -376,7 +368,7 @@ impl< /// * Forwards [`ReadyQueueSender::send_task_ready`]'s return values on failure. /// * Forwards [`ReadyQueueSender::send_commit_ready`]'s return values on failure. /// * Forwards [`SharedJobControlBlock::commit_outputs`]'s return values on failure. - /// * Forwards [`OutputReader::read_as_task_output`]'s return values on failure. + /// * Forwards [`TaskGraph::read_output_payloads`]'s return values on failure. /// * Forwards [`InternalJobOrchestration::commit_outputs`]'s return values on failure. pub async fn succeed_task_instance( &self, @@ -417,16 +409,7 @@ impl< // Release the read lock prior to acquiring a write lock for committing job outputs. drop(job); let mut job = jcb.job_execution_state.write_running().await?; - let mut job_outputs = Vec::new(); - for output_reader in job.task_graph.get_outputs() { - let payload = output_reader - .read() - .await - .as_ref() - .ok_or(InternalError::TaskInputNotReady)? - .clone(); - job_outputs.push(payload); - } + let job_outputs = job.task_graph.read_output_payloads().await?; let has_commit_task = job.task_graph.has_commit_task(); job.db_connector .commit_outputs(jcb.id, job_outputs, has_commit_task) @@ -720,6 +703,8 @@ impl< /// failure. /// * Forwards [`TaskInstancePoolConnector::register_termination_task_instance`]'s return values /// on failure. + /// * Forwards [`TaskGraph::read_output_payloads`]'s return values on failure. + /// * Forwards [`SerializedTaskOutputs::serialize_with_size_hint`]'s return values on failure. async fn create_commit_task_instance( jcb: &JobControlBlock, execution_manager_id: ExecutionManagerId, @@ -747,12 +732,18 @@ impl< .register_termination_task_instance(commit_tcb.clone(), registration) .await?; + let task_graph_outputs = job.task_graph.read_output_payloads().await?; + let serialized_task_io = + SerializedTaskOutputs::serialize_with_size_hint(&task_graph_outputs) + .map_err(InternalError::TaskOutputs)? + .to_raw(); + drop(job); Ok(ExecutionContext { task_instance_id, tdl_context, timeout_policy, - serialized_inputs: Vec::new(), + serialized_task_io, }) } @@ -804,7 +795,7 @@ impl< task_instance_id, tdl_context, timeout_policy, - serialized_inputs: Vec::new(), + serialized_task_io: Vec::new(), }) } } diff --git a/components/spider-storage/src/cache/task.rs b/components/spider-storage/src/cache/task.rs index 26efd0535..912ac74bd 100644 --- a/components/spider-storage/src/cache/task.rs +++ b/components/spider-storage/src/cache/task.rs @@ -222,6 +222,31 @@ impl TaskGraph { &self.outputs } + /// Reads the payloads of all outputs in the task graph. + /// + /// # Returns + /// + /// The payloads of all task-graph outputs on success. + /// + /// # Errors + /// + /// Returns an error if: + /// + /// * [`InternalError::TaskInputNotReady`] if any output has no value. + pub async fn read_output_payloads(&self) -> Result, InternalError> { + let mut outputs = Vec::new(); + for output_reader in &self.outputs { + let payload = output_reader + .read() + .await + .as_ref() + .ok_or(InternalError::TaskInputNotReady)? + .clone(); + outputs.push(payload); + } + Ok(outputs) + } + #[must_use] pub const fn has_commit_task(&self) -> bool { self.commit_task.is_some() @@ -267,7 +292,7 @@ impl SharedTaskControlBlock { task_instance_id: instance_id, tdl_context: tcb.base.tdl_context.clone(), timeout_policy: tcb.base.timeout_policy.clone(), - serialized_inputs: tcb.fetch_inputs().await?, + serialized_task_io: tcb.fetch_inputs().await?, }) }; result.map_err(CacheError::from) @@ -1317,7 +1342,7 @@ mod tests { .iter() .map(|v| TaskInput::ValuePayload(v.clone())) .collect(); - let actual_inputs = deserialize_task_inputs(&ctx.serialized_inputs); + let actual_inputs = deserialize_task_inputs(&ctx.serialized_task_io); assert_eq!(actual_inputs, expected, "task {task_index} inputs mismatch"); let outputs = compute_outputs(&actual_inputs); tcb.succeed_task_instance(id, outputs) @@ -1348,7 +1373,7 @@ mod tests { .iter() .map(|v| TaskInput::ValuePayload(v.clone())) .collect(); - let actual_inputs = deserialize_task_inputs(&ctx.serialized_inputs); + let actual_inputs = deserialize_task_inputs(&ctx.serialized_task_io); assert_eq!(actual_inputs, expected, "task inputs mismatch"); barrier.wait().await; let outputs = compute_outputs(&actual_inputs); diff --git a/components/spider-storage/src/grpc.rs b/components/spider-storage/src/grpc.rs index 03ff724e8..db5a29f4a 100644 --- a/components/spider-storage/src/grpc.rs +++ b/components/spider-storage/src/grpc.rs @@ -607,7 +607,7 @@ impl< soft_timeout_ms: execution_context.timeout_policy.soft_timeout_ms, hard_timeout_ms: execution_context.timeout_policy.hard_timeout_ms, }), - serialized_inputs: execution_context.serialized_inputs, + serialized_task_io: execution_context.serialized_task_io, }), })) } diff --git a/components/spider-storage/tests/scheduling_infra.rs b/components/spider-storage/tests/scheduling_infra.rs index 3e762d5da..eb9124150 100644 --- a/components/spider-storage/tests/scheduling_infra.rs +++ b/components/spider-storage/tests/scheduling_infra.rs @@ -91,6 +91,7 @@ use spider_core::types::id::ResourceGroupId; use spider_core::types::id::TaskId; use spider_core::types::id::TaskInstanceId; use spider_core::types::io::ExecutionContext; +use spider_core::types::io::SerializedTaskOutputs; use spider_core::types::io::TaskOutput; use spider_storage::cache::error::CacheError; use spider_storage::cache::error::InternalError; @@ -1061,6 +1062,13 @@ async fn process_task( /// /// * Forwards [`SharedJobControlBlock::create_task_instance`]'s return values on failure. /// * Forwards [`SharedJobControlBlock::succeed_commit_task_instance`]'s return values of failure. +/// +/// # Panics +/// +/// Panics if: +/// +/// * The commit task's execution context does not carry decodable task-graph outputs. +/// * The commit task's execution context carries an empty task-graph output list. async fn process_commit( ctx: &EmContext, ) -> Result<(), CacheError> { @@ -1068,6 +1076,13 @@ async fn process_commit( .jcb .create_task_instance(TaskId::Commit, ctx.execution_manager_id) .await?; + let task_graph_outputs = + SerializedTaskOutputs::deserialize_from_raw(&exec_ctx.serialized_task_io) + .expect("commit task execution context should carry decodable task-graph outputs"); + assert!( + !task_graph_outputs.is_empty(), + "task-graph outputs should not be empty" + ); let state = ctx .jcb .succeed_commit_task_instance(exec_ctx.task_instance_id) diff --git a/components/spider-tdl/src/error.rs b/components/spider-tdl/src/error.rs index 27664ffbe..85907e67d 100644 --- a/components/spider-tdl/src/error.rs +++ b/components/spider-tdl/src/error.rs @@ -20,6 +20,9 @@ pub enum TdlError { #[error("execution error: {0}")] ExecutionError(String), + #[error("invalid task context: {0}")] + InvalidTaskContext(String), + #[error("{0}")] Custom(String), } @@ -35,6 +38,7 @@ mod tests { TdlError::DeserializationError("deserialization_error".to_owned()), TdlError::SerializationError("serialization_error".to_owned()), TdlError::ExecutionError("execution_error".to_owned()), + TdlError::InvalidTaskContext("invalid_task_context".to_owned()), TdlError::Custom("custom".to_owned()), ]; for error in errors_to_test { diff --git a/components/spider-tdl/src/task.rs b/components/spider-tdl/src/task.rs index d235dfeeb..8cb75931d 100644 --- a/components/spider-tdl/src/task.rs +++ b/components/spider-tdl/src/task.rs @@ -252,12 +252,14 @@ mod tests { } fn make_encoded_ctx() -> Vec { - let ctx = TaskContext { - job_id: JobId::random(), - task_id: TaskId::Index(0), - task_instance_id: 1, - resource_group_id: ResourceGroupId::random(), - }; + let ctx = TaskContext::new( + JobId::random(), + TaskId::Index(0), + 1, + ResourceGroupId::random(), + None, + ) + .expect("failed to build `TaskContext`"); rmp_serde::to_vec(&ctx).expect("failed to serialize `TaskContext`") } diff --git a/components/spider-tdl/src/task_context.rs b/components/spider-tdl/src/task_context.rs index 6e0d648d7..d0b4895d9 100644 --- a/components/spider-tdl/src/task_context.rs +++ b/components/spider-tdl/src/task_context.rs @@ -8,6 +8,10 @@ use spider_core::types::id::JobId; use spider_core::types::id::ResourceGroupId; use spider_core::types::id::TaskId; use spider_core::types::id::TaskInstanceId; +use spider_core::types::io::SerializedTaskOutputs; +use spider_core::types::io::TaskOutput; + +use crate::error::TdlError; /// Runtime metadata about the current task execution. /// @@ -22,6 +26,77 @@ pub struct TaskContext { pub task_id: TaskId, pub task_instance_id: TaskInstanceId, pub resource_group_id: ResourceGroupId, + serialized_task_graph_outputs: Option>, +} + +impl TaskContext { + /// Creates a new [`TaskContext`], validating that the job's task-graph outputs are present for + /// a commit task and omitted for any other task. + /// + /// The raw serialized outputs are stored as-is and only deserialized on demand by + /// [`Self::get_task_graph_outputs`]. + /// + /// # Returns + /// + /// The constructed [`TaskContext`] on success. + /// + /// # Errors + /// + /// Returns an error if: + /// + /// * [`TdlError::InvalidTaskContext`] if: + /// * The task is a commit task but no task-graph outputs are provided, or, + /// * The task is not a commit task but task-graph outputs are provided. + pub fn new( + job_id: JobId, + task_id: TaskId, + task_instance_id: TaskInstanceId, + resource_group_id: ResourceGroupId, + serialized_task_graph_outputs: Option>, + ) -> Result { + let is_commit_task = task_id == TaskId::Commit; + if is_commit_task && serialized_task_graph_outputs.is_none() { + return Err(TdlError::InvalidTaskContext( + "task-graph outputs are required for a commit task but were not provided" + .to_owned(), + )); + } + if !is_commit_task && serialized_task_graph_outputs.is_some() { + return Err(TdlError::InvalidTaskContext(format!( + "task-graph outputs are only valid for a commit task, but the current task is \ + {task_id}" + ))); + } + Ok(Self { + job_id, + task_id, + task_instance_id, + resource_group_id, + serialized_task_graph_outputs, + }) + } + + /// Deserializes the job's task-graph outputs carried by this context. + /// + /// # Returns + /// + /// The job's task-graph outputs (if any) or `None` on success. + /// + /// # Errors + /// + /// Returns an error if: + /// + /// * Forwards [`SerializedTaskOutputs::deserialize_from_raw`]'s return values on failure. + pub fn get_task_graph_outputs(&self) -> Result>, TdlError> { + match &self.serialized_task_graph_outputs { + Some(bytes) => { + let task_graph_outputs = SerializedTaskOutputs::deserialize_from_raw(bytes) + .map_err(|e| TdlError::DeserializationError(e.to_string()))?; + Ok(Some(task_graph_outputs)) + } + None => Ok(None), + } + } } #[cfg(test)] @@ -29,20 +104,137 @@ mod tests { use spider_core::types::id::JobId; use spider_core::types::id::ResourceGroupId; use spider_core::types::id::TaskId; + use spider_core::types::io::SerializedTaskOutputs; + use spider_core::types::io::TaskOutput; use super::TaskContext; + use crate::error::TdlError; + + /// Serializes the given task outputs into their raw wire buffer for constructing a + /// [`TaskContext`]. + /// + /// # Returns + /// + /// The raw serialized task-outputs buffer. + /// + /// # Panics + /// + /// Panics if [`SerializedTaskOutputs::serialize_with_size_hint`] returns an error. + fn serialize_outputs(outputs: &[TaskOutput]) -> Vec { + SerializedTaskOutputs::serialize_with_size_hint(outputs) + .expect("task outputs must serialize") + .to_raw() + } #[test] fn round_trip_msgpack() -> anyhow::Result<()> { - let ctx = TaskContext { - job_id: JobId::random(), - task_id: TaskId::Index(0), - task_instance_id: 13, - resource_group_id: ResourceGroupId::random(), - }; + let ctx = TaskContext::new( + JobId::random(), + TaskId::Index(0), + 13, + ResourceGroupId::random(), + None, + )?; let encoded = rmp_serde::to_vec(&ctx)?; let decoded: TaskContext = rmp_serde::from_slice(&encoded)?; assert_eq!(decoded, ctx); Ok(()) } + + #[test] + fn get_task_graph_outputs_returns_commit_outputs() -> anyhow::Result<()> { + let outputs: Vec = vec![vec![1, 2, 3], vec![4, 5, 6]]; + let ctx = TaskContext::new( + JobId::random(), + TaskId::Commit, + 1, + ResourceGroupId::random(), + Some(serialize_outputs(&outputs)), + )?; + assert_eq!(ctx.get_task_graph_outputs()?, Some(outputs)); + Ok(()) + } + + #[test] + fn get_task_graph_outputs_returns_empty_commit_outputs() -> anyhow::Result<()> { + let ctx = TaskContext::new( + JobId::random(), + TaskId::Commit, + 1, + ResourceGroupId::random(), + Some(serialize_outputs(&[])), + )?; + assert_eq!(ctx.get_task_graph_outputs()?, Some(Vec::new())); + Ok(()) + } + + #[test] + fn get_task_graph_outputs_returns_none_for_non_commit_task() -> anyhow::Result<()> { + let ctx = TaskContext::new( + JobId::random(), + TaskId::Index(0), + 1, + ResourceGroupId::random(), + None, + )?; + assert!(ctx.get_task_graph_outputs()?.is_none()); + Ok(()) + } + + #[test] + fn new_rejects_outputs_for_non_commit_task() { + let outputs: Vec = vec![vec![1, 2, 3], vec![4, 5, 6]]; + let result = TaskContext::new( + JobId::random(), + TaskId::Index(0), + 1, + ResourceGroupId::random(), + Some(serialize_outputs(&outputs)), + ); + assert!(matches!(result, Err(TdlError::InvalidTaskContext(_)))); + } + + #[test] + fn new_rejects_missing_outputs_for_commit_task() { + let result = TaskContext::new( + JobId::random(), + TaskId::Commit, + 1, + ResourceGroupId::random(), + None, + ); + assert!(matches!(result, Err(TdlError::InvalidTaskContext(_)))); + } + + #[test] + fn get_task_graph_outputs_rejects_invalid_bytes() -> anyhow::Result<()> { + let ctx = TaskContext::new( + JobId::random(), + TaskId::Commit, + 1, + ResourceGroupId::random(), + Some(vec![0xff, 0x00, 0x13, 0x37]), + )?; + assert!(matches!( + ctx.get_task_graph_outputs(), + Err(TdlError::DeserializationError(_)) + )); + Ok(()) + } + + #[test] + fn get_task_graph_outputs_survives_msgpack_round_trip() -> anyhow::Result<()> { + let outputs: Vec = vec![vec![1, 2, 3], vec![4, 5, 6]]; + let ctx = TaskContext::new( + JobId::random(), + TaskId::Commit, + 1, + ResourceGroupId::random(), + Some(serialize_outputs(&outputs)), + )?; + let encoded = rmp_serde::to_vec(&ctx)?; + let decoded: TaskContext = rmp_serde::from_slice(&encoded)?; + assert_eq!(decoded.get_task_graph_outputs()?, Some(outputs)); + Ok(()) + } } diff --git a/components/spider-tdl/tests/test_task_macro.rs b/components/spider-tdl/tests/test_task_macro.rs index 3f772267a..16e49bac9 100644 --- a/components/spider-tdl/tests/test_task_macro.rs +++ b/components/spider-tdl/tests/test_task_macro.rs @@ -79,12 +79,14 @@ fn translate(_ctx: TaskContext, p: Point, dx: int32, dy: int32) -> Result<(Point /// /// A mocked encoded task context for testing. fn make_encoded_ctx() -> Vec { - let ctx = TaskContext { - job_id: JobId::random(), - task_id: TaskId::Index(0), - task_instance_id: 1, - resource_group_id: ResourceGroupId::random(), - }; + let ctx = TaskContext::new( + JobId::random(), + TaskId::Index(0), + 1, + ResourceGroupId::random(), + None, + ) + .expect("failed to build `TaskContext`"); rmp_serde::to_vec(&ctx).expect("failed to serialize `TaskContext`") } @@ -301,12 +303,13 @@ fn direct_execute_call_round_trips() -> anyhow::Result<()> { const OPERAND_B: int32 = 35; const EXPECTED_SUM: int32 = OPERAND_A + OPERAND_B; - let ctx = TaskContext { - job_id: JobId::random(), - task_id: TaskId::Index(0), - task_instance_id: 1, - resource_group_id: ResourceGroupId::random(), - }; + let ctx = TaskContext::new( + JobId::random(), + TaskId::Index(0), + 1, + ResourceGroupId::random(), + None, + )?; let mut inputs = TaskInputsSerializer::new(); append_value(&mut inputs, &OPERAND_A)?; diff --git a/tests/huntsman/em-runtime/tests/test_runtime.rs b/tests/huntsman/em-runtime/tests/test_runtime.rs index c4d5d8251..31f6d7410 100644 --- a/tests/huntsman/em-runtime/tests/test_runtime.rs +++ b/tests/huntsman/em-runtime/tests/test_runtime.rs @@ -92,7 +92,7 @@ fn execution_context(task_func: &str, inputs: Vec) -> ExecutionContex soft_timeout_ms: 1_000, hard_timeout_ms: 5_000, }, - serialized_inputs: serializer.release(), + serialized_task_io: serializer.release(), } } diff --git a/tests/huntsman/integration-test-tasks/Cargo.toml b/tests/huntsman/integration-test-tasks/Cargo.toml index 0c77122e8..ec871b647 100644 --- a/tests/huntsman/integration-test-tasks/Cargo.toml +++ b/tests/huntsman/integration-test-tasks/Cargo.toml @@ -12,5 +12,6 @@ name = "integration_test_tasks" path = "src/lib.rs" [dependencies] +rmp-serde = "1.3.1" serde = { version = "1.0.228", features = ["derive"] } spider-tdl = { path = "../../../components/spider-tdl", features = ["derive"] } diff --git a/tests/huntsman/integration-test-tasks/src/lib.rs b/tests/huntsman/integration-test-tasks/src/lib.rs index c432bcb83..23bf00533 100644 --- a/tests/huntsman/integration-test-tasks/src/lib.rs +++ b/tests/huntsman/integration-test-tasks/src/lib.rs @@ -1,6 +1,6 @@ //! Test TDL package used by the `task-executor` integration tests. //! -//! Exposes four tasks that exercise distinct executor code paths: +//! Exposes five tasks that exercise distinct executor code paths: //! //! * [`task_decl::fibonacci`] — basic compute + correctness. //! * [`task_decl::always_fail`] — in-task error reporting. @@ -9,6 +9,8 @@ //! duration then echoes its `Vec` payload back. Used by the overhead bench so the //! non-sleep portion of the executor's reported FFI time isolates the in-executor input/output //! serde cost, while the parent-side delta isolates IPC framing cost. +//! * [`task_decl::assert_outputs_sum_zero`] — commit task: reads the job's task-graph outputs from +//! its [`TaskContext`](spider_tdl::TaskContext) and asserts the `i64` outputs sum to zero. /// The constant sleep duration used by [`task_decl::sleep_and_echo`]. /// @@ -65,6 +67,32 @@ mod task_decl { sleep(Duration::from_micros(INSTRUMENT_SLEEP_US)); Ok(items) } + + /// Commit task that reads the job's task-graph outputs and asserts that the `i64` values they + /// carry sum to zero. + #[task(name = "assert_outputs_sum_zero")] + pub fn assert_outputs_sum_zero(ctx: TaskContext) -> Result<(), TdlError> { + let outputs = ctx.get_task_graph_outputs()?.ok_or_else(|| { + TdlError::ExecutionError("assert_outputs_sum_zero must run as a commit task".to_owned()) + })?; + if outputs.is_empty() { + return Err(TdlError::ExecutionError( + "assert_outputs_sum_zero: task-graph outputs empty".to_owned(), + )); + } + let mut sum: i64 = 0; + for output in &outputs { + let value: i64 = rmp_serde::from_slice(output) + .map_err(|e| TdlError::DeserializationError(e.to_string()))?; + sum += value; + } + if sum != 0 { + return Err(TdlError::ExecutionError(format!( + "task-graph outputs do not sum to zero: got {sum}" + ))); + } + Ok(()) + } } spider_tdl::register_tdl_package! { @@ -74,5 +102,6 @@ spider_tdl::register_tdl_package! { task_decl::always_fail, task_decl::always_panic, task_decl::sleep_and_echo, + task_decl::assert_outputs_sum_zero, ], } diff --git a/tests/huntsman/task-executor/tests/test_executor.rs b/tests/huntsman/task-executor/tests/test_executor.rs index acc0b3903..1825a5efe 100644 --- a/tests/huntsman/task-executor/tests/test_executor.rs +++ b/tests/huntsman/task-executor/tests/test_executor.rs @@ -8,11 +8,32 @@ use spider_task_executor::protocol::ExecutorOutcome; use spider_task_executor::protocol::Response; use spider_tdl::TdlError; use test_utils::ExecutorHandle; +use test_utils::commit_execute_request; use test_utils::decode_single_output; use test_utils::encode_no_inputs; use test_utils::encode_single_input; use test_utils::execute_request; +/// Asserts that `outcome` is a task [`TdlError::ExecutionError`] whose message contains `needle`. +fn assert_task_execution_error(outcome: ExecutorOutcome, needle: &str) { + match outcome { + ExecutorOutcome::Success { outputs } => { + panic!("expected Failure, got Success with {} bytes", outputs.len()); + } + ExecutorOutcome::Failure { error } => { + let err: ExecutorError = + rmp_serde::from_slice(&error).expect("decode ExecutorError payload"); + let ExecutorError::TaskError(TdlError::ExecutionError(message)) = &err else { + panic!("expected TaskError(ExecutionError), got {err:?}"); + }; + assert!( + message.contains(needle), + "unexpected error message: {message}", + ); + } + } +} + #[tokio::test] #[ignore = "requires `integration-test-tasks` cdylib and `spider-task-executor` binary"] async fn fibonacci_returns_correct_value() { @@ -85,3 +106,70 @@ async fn always_panic_crashes_the_process() { "expected non-zero exit after panic, got {status:?}", ); } + +#[tokio::test] +#[ignore = "requires `integration-test-tasks` cdylib and `spider-task-executor` binary"] +async fn assert_outputs_sum_zero_succeeds_when_outputs_sum_to_zero() { + let mut handle = ExecutorHandle::spawn(); + handle + .send(&commit_execute_request( + "assert_outputs_sum_zero", + &[3_i64, -1, -2], + )) + .await; + let Response::Result { outcome, .. } = handle.recv().await; + match outcome { + ExecutorOutcome::Success { .. } => {} + ExecutorOutcome::Failure { error } => { + let err: ExecutorError = + rmp_serde::from_slice(&error).expect("decode ExecutorError payload"); + panic!("expected Success for zero-sum outputs, got Failure: {err:?}"); + } + } + handle.shutdown_clean().await; +} + +#[tokio::test] +#[ignore = "requires `integration-test-tasks` cdylib and `spider-task-executor` binary"] +async fn assert_outputs_sum_zero_fails_when_outputs_do_not_sum_to_zero() { + let mut handle = ExecutorHandle::spawn(); + handle + .send(&commit_execute_request( + "assert_outputs_sum_zero", + &[1_i64, 2, 3], + )) + .await; + let Response::Result { outcome, .. } = handle.recv().await; + assert_task_execution_error(outcome, "sum to zero"); + handle.shutdown_clean().await; +} + +#[tokio::test] +#[ignore = "requires `integration-test-tasks` cdylib and `spider-task-executor` binary"] +async fn assert_outputs_sum_zero_fails_when_outputs_empty() { + let mut handle = ExecutorHandle::spawn(); + handle + .send(&commit_execute_request::( + "assert_outputs_sum_zero", + &[], + )) + .await; + let Response::Result { outcome, .. } = handle.recv().await; + assert_task_execution_error(outcome, "empty"); + handle.shutdown_clean().await; +} + +#[tokio::test] +#[ignore = "requires `integration-test-tasks` cdylib and `spider-task-executor` binary"] +async fn assert_outputs_sum_zero_fails_for_non_commit_task() { + let mut handle = ExecutorHandle::spawn(); + handle + .send(&execute_request( + "assert_outputs_sum_zero", + encode_no_inputs(), + )) + .await; + let Response::Result { outcome, .. } = handle.recv().await; + assert_task_execution_error(outcome, "commit task"); + handle.shutdown_clean().await; +} diff --git a/tests/huntsman/task-executor/tests/test_process_pool.rs b/tests/huntsman/task-executor/tests/test_process_pool.rs index 26a6c3ad7..2b9898438 100644 --- a/tests/huntsman/task-executor/tests/test_process_pool.rs +++ b/tests/huntsman/task-executor/tests/test_process_pool.rs @@ -95,7 +95,7 @@ fn make_request(task_func: &str, inputs: Vec) -> ExecuteRequest { soft_timeout_ms: 100, hard_timeout_ms: 1000, }, - serialized_inputs: serializer.release(), + serialized_task_io: serializer.release(), }, } } diff --git a/tests/huntsman/tdl-integration/tests/complex.rs b/tests/huntsman/tdl-integration/tests/complex.rs index 2c5a69aef..f35235d6b 100644 --- a/tests/huntsman/tdl-integration/tests/complex.rs +++ b/tests/huntsman/tdl-integration/tests/complex.rs @@ -32,12 +32,14 @@ fn lib_path() -> std::path::PathBuf { /// /// An encoded task context for testing. fn encode_ctx() -> Vec { - let ctx = TaskContext { - job_id: JobId::random(), - task_id: TaskId::Index(0), - task_instance_id: 1, - resource_group_id: ResourceGroupId::random(), - }; + let ctx = TaskContext::new( + JobId::random(), + TaskId::Index(0), + 1, + ResourceGroupId::random(), + None, + ) + .expect("failed to build `TaskContext`"); rmp_serde::to_vec(&ctx).expect("failed to serialize `TaskContext`") } diff --git a/tests/huntsman/test-utils/src/executor.rs b/tests/huntsman/test-utils/src/executor.rs index 0b93f8351..f113f18f3 100644 --- a/tests/huntsman/test-utils/src/executor.rs +++ b/tests/huntsman/test-utils/src/executor.rs @@ -24,8 +24,10 @@ use spider_core::task::TdlContext; use spider_core::types::id::JobId; use spider_core::types::id::ResourceGroupId; use spider_core::types::id::TaskId; +use spider_core::types::io::SerializedTaskOutputs; use spider_core::types::io::TaskInput; use spider_core::types::io::TaskInputsSerializer; +use spider_core::types::io::TaskOutput; use spider_core::types::io::TaskOutputsSerializer; use spider_task_executor::protocol::Request; use spider_task_executor::protocol::Response; @@ -192,15 +194,17 @@ pub fn tdl_package_dir() -> PathBuf { /// /// # Panics /// -/// Panics if msgpack encoding fails. +/// Panics if constructing or msgpack-encoding the [`TaskContext`] fails. #[must_use] pub fn build_ctx() -> Vec { - let ctx = TaskContext { - job_id: JobId::random(), - task_id: TaskId::Index(0), - task_instance_id: 1, - resource_group_id: ResourceGroupId::random(), - }; + let ctx = TaskContext::new( + JobId::random(), + TaskId::Index(0), + 1, + ResourceGroupId::random(), + None, + ) + .expect("build TaskContext"); rmp_serde::to_vec(&ctx).expect("serialize TaskContext") } @@ -300,3 +304,63 @@ pub fn execute_request(task_func: &str, raw_inputs: Vec) -> Request { raw_inputs, } } + +/// Builds a msgpack-encoded commit [`TaskContext`] carrying `outputs` as the job's task-graph +/// outputs. +/// +/// Each value in `outputs` is msgpack-encoded into one task-graph output payload, mirroring the +/// shape the execution manager ships to a commit task. +/// +/// # Type Parameters +/// +/// * `ValueType` - The Serde-serializable type of each task-graph output value. +/// +/// # Returns +/// +/// A msgpack-encoded commit [`TaskContext`] whose task-graph outputs decode back to `outputs`. +/// +/// # Panics +/// +/// Panics if msgpack encoding, output serialization, or [`TaskContext`] construction fails. +#[must_use] +pub fn build_commit_ctx(outputs: &[ValueType]) -> Vec { + let task_outputs: Vec = outputs + .iter() + .map(|value| rmp_serde::to_vec(value).expect("msgpack encode task-graph output")) + .collect(); + let serialized_outputs = SerializedTaskOutputs::serialize_with_size_hint(&task_outputs) + .expect("serialize task-graph outputs") + .to_raw(); + let ctx = TaskContext::new( + JobId::random(), + TaskId::Commit, + 1, + ResourceGroupId::random(), + Some(serialized_outputs), + ) + .expect("build commit TaskContext"); + rmp_serde::to_vec(&ctx).expect("serialize TaskContext") +} + +/// # Type Parameters +/// +/// * `ValueType` - The Serde-serializable type of each task-graph output value. +/// +/// # Returns +/// +/// A [`Request::Execute`] targeting `task_func` in the integration package, carrying a commit +/// [`TaskContext`] whose task-graph outputs are `outputs` and no task inputs. +#[must_use] +pub fn commit_execute_request( + task_func: &str, + outputs: &[ValueType], +) -> Request { + Request::Execute { + tdl_context: TdlContext { + package: PACKAGE_NAME.to_owned(), + task_func: task_func.to_owned(), + }, + raw_ctx: build_commit_ctx(outputs), + raw_inputs: encode_no_inputs(), + } +}