Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 2 additions & 0 deletions Cargo.lock

Some generated files are not rendered by default. Learn more about how customized files appear on GitHub.

2 changes: 1 addition & 1 deletion components/spider-core/src/types/io.rs
Original file line number Diff line number Diff line change
Expand Up @@ -432,7 +432,7 @@ pub struct ExecutionContext {
pub task_instance_id: TaskInstanceId,
pub tdl_context: TdlContext,
pub timeout_policy: TimeoutPolicy,
pub serialized_inputs: Vec<u8>,
pub serialized_task_io: Vec<u8>,
}

#[cfg(test)]
Expand Down
3 changes: 3 additions & 0 deletions components/spider-execution-manager/Cargo.toml
Original file line number Diff line number Diff line change
Expand Up @@ -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"
132 changes: 126 additions & 6 deletions components/spider-execution-manager/src/process_pool.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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;
Expand Down Expand Up @@ -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.
Expand Down Expand Up @@ -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<Request, InternalError> {
let ExecuteRequest {
Expand All @@ -368,17 +375,130 @@ fn build_request(request: ExecuteRequest) -> Result<Request, InternalError> {
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<u8>) -> 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<TaskOutput> = 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<TaskOutput> = 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(())
}
}
2 changes: 1 addition & 1 deletion components/spider-proto-rust/src/generated/storage.rs
Original file line number Diff line number Diff line change
Expand Up @@ -86,7 +86,7 @@ pub struct ExecutionContext {
#[prost(message, optional, tag = "3")]
pub timeout_policy: ::core::option::Option<TimeoutPolicy>,
#[prost(bytes = "vec", tag = "4")]
pub serialized_inputs: ::prost::alloc::vec::Vec<u8>,
pub serialized_task_io: ::prost::alloc::vec::Vec<u8>,
}
#[derive(Clone, PartialEq, Eq, Hash, ::prost::Message)]
pub struct TdlContext {
Expand Down
8 changes: 4 additions & 4 deletions components/spider-proto-rust/src/io.rs
Original file line number Diff line number Diff line change
Expand Up @@ -27,7 +27,7 @@ impl TryFrom<storage::ExecutionContext> 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,
})
}
}
Expand All @@ -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 =
Expand All @@ -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]
Expand All @@ -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!(
Expand Down
2 changes: 1 addition & 1 deletion components/spider-proto/storage/storage.proto
Original file line number Diff line number Diff line change
Expand Up @@ -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 {
Expand Down
3 changes: 3 additions & 0 deletions components/spider-storage/src/cache/error.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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),
}
Expand Down
39 changes: 15 additions & 24 deletions components/spider-storage/src/cache/job.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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;
Expand Down Expand Up @@ -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<Vec<TaskOutput>, 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)
}
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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<ReadyQueueSenderType, DbConnectorType, TaskInstancePoolConnectorType>,
execution_manager_id: ExecutionManagerId,
Expand Down Expand Up @@ -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,
})
}

Expand Down Expand Up @@ -804,7 +795,7 @@ impl<
task_instance_id,
tdl_context,
timeout_policy,
serialized_inputs: Vec::new(),
serialized_task_io: Vec::new(),
})
}
}
Expand Down
Loading
Loading