diff --git a/benches/dependencies.rs b/benches/dependencies.rs index 3869d82..a002cfb 100644 --- a/benches/dependencies.rs +++ b/benches/dependencies.rs @@ -9,7 +9,7 @@ use criterion::{criterion_group, criterion_main, BenchmarkId, Criterion}; use pprof::criterion::{Output, PProfProfiler}; use serde::{Deserialize, Serialize}; use taskmill::{ - Domain, DomainKey, Scheduler, SchedulerEvent, TaskContext, TaskError, TaskStore, + Domain, DomainKey, DomainTaskContext, Scheduler, SchedulerEvent, TaskError, TaskStore, TaskSubmission, TypedExecutor, TypedTask, }; use tokio::runtime::Runtime; @@ -30,7 +30,11 @@ impl TypedTask for BenchTask { struct NoopExecutor; impl TypedExecutor for NoopExecutor { - async fn execute(&self, _payload: BenchTask, _ctx: &TaskContext) -> Result<(), TaskError> { + async fn execute<'a>( + &'a self, + _payload: BenchTask, + _ctx: DomainTaskContext<'a, BenchDomain>, + ) -> Result<(), TaskError> { Ok(()) } } diff --git a/benches/groups.rs b/benches/groups.rs index 1872157..7c68716 100644 --- a/benches/groups.rs +++ b/benches/groups.rs @@ -7,7 +7,7 @@ use std::time::{Duration, Instant}; use criterion::{criterion_group, criterion_main, BenchmarkId, Criterion}; use serde::{Deserialize, Serialize}; use taskmill::{ - Domain, DomainKey, Scheduler, SchedulerEvent, TaskContext, TaskError, TaskStore, + Domain, DomainKey, DomainTaskContext, Scheduler, SchedulerEvent, TaskError, TaskStore, TaskSubmission, TypedExecutor, TypedTask, }; use tokio::runtime::Runtime; @@ -28,7 +28,11 @@ impl TypedTask for BenchTask { struct NoopExecutor; impl TypedExecutor for NoopExecutor { - async fn execute(&self, _payload: BenchTask, _ctx: &TaskContext) -> Result<(), TaskError> { + async fn execute<'a>( + &'a self, + _payload: BenchTask, + _ctx: DomainTaskContext<'a, BenchDomain>, + ) -> Result<(), TaskError> { Ok(()) } } diff --git a/benches/history.rs b/benches/history.rs index dcd8917..96eeb52 100644 --- a/benches/history.rs +++ b/benches/history.rs @@ -7,7 +7,7 @@ use std::time::Duration; use criterion::{criterion_group, criterion_main, BenchmarkId, Criterion}; use serde::{Deserialize, Serialize}; use taskmill::{ - Domain, DomainKey, Scheduler, SchedulerEvent, TaskContext, TaskError, TaskStore, + Domain, DomainKey, DomainTaskContext, Scheduler, SchedulerEvent, TaskError, TaskStore, TaskSubmission, TypedExecutor, TypedTask, }; use tokio::runtime::Runtime; @@ -28,7 +28,11 @@ impl TypedTask for BenchTask { struct NoopExecutor; impl TypedExecutor for NoopExecutor { - async fn execute(&self, _payload: BenchTask, _ctx: &TaskContext) -> Result<(), TaskError> { + async fn execute<'a>( + &'a self, + _payload: BenchTask, + _ctx: DomainTaskContext<'a, BenchDomain>, + ) -> Result<(), TaskError> { Ok(()) } } diff --git a/benches/retry.rs b/benches/retry.rs index d208069..26ea421 100644 --- a/benches/retry.rs +++ b/benches/retry.rs @@ -7,7 +7,7 @@ use std::time::{Duration, Instant}; use criterion::{black_box, criterion_group, criterion_main, BenchmarkId, Criterion}; use serde::{Deserialize, Serialize}; use taskmill::{ - BackoffStrategy, Domain, DomainKey, RetryPolicy, Scheduler, SchedulerEvent, TaskContext, + BackoffStrategy, Domain, DomainKey, DomainTaskContext, RetryPolicy, Scheduler, SchedulerEvent, TaskError, TaskStore, TaskSubmission, TypedExecutor, TypedTask, }; use tokio::runtime::Runtime; @@ -34,7 +34,11 @@ impl TypedTask for FailTask { struct FailPermanentExecutor; impl TypedExecutor for FailPermanentExecutor { - async fn execute(&self, _payload: FailTask, _ctx: &TaskContext) -> Result<(), TaskError> { + async fn execute<'a>( + &'a self, + _payload: FailTask, + _ctx: DomainTaskContext<'a, BenchDomain>, + ) -> Result<(), TaskError> { Err(TaskError::permanent("bench: permanent failure")) } } @@ -42,7 +46,11 @@ impl TypedExecutor for FailPermanentExecutor { struct FailRetryableExecutor; impl TypedExecutor for FailRetryableExecutor { - async fn execute(&self, _payload: FailTask, _ctx: &TaskContext) -> Result<(), TaskError> { + async fn execute<'a>( + &'a self, + _payload: FailTask, + _ctx: DomainTaskContext<'a, BenchDomain>, + ) -> Result<(), TaskError> { Err(TaskError::retryable("bench: transient failure")) } } diff --git a/benches/scheduler.rs b/benches/scheduler.rs index 8e6f3a4..31debef 100644 --- a/benches/scheduler.rs +++ b/benches/scheduler.rs @@ -7,8 +7,8 @@ use std::time::{Duration, Instant}; use criterion::{criterion_group, criterion_main, BenchmarkId, Criterion}; use serde::{Deserialize, Serialize}; use taskmill::{ - Domain, DomainKey, Priority, Scheduler, SchedulerEvent, TaskContext, TaskError, TaskStore, - TaskSubmission, TypedExecutor, TypedTask, + Domain, DomainKey, DomainTaskContext, Priority, Scheduler, SchedulerEvent, TaskError, + TaskStore, TaskSubmission, TypedExecutor, TypedTask, }; use tokio::runtime::Runtime; use tokio_util::sync::CancellationToken; @@ -41,7 +41,11 @@ impl TypedTask for ByteTestTask { struct NoopExecutor; impl TypedExecutor for NoopExecutor { - async fn execute(&self, _payload: BenchTask, _ctx: &TaskContext) -> Result<(), TaskError> { + async fn execute<'a>( + &'a self, + _payload: BenchTask, + _ctx: DomainTaskContext<'a, BenchDomain>, + ) -> Result<(), TaskError> { Ok(()) } } @@ -53,7 +57,11 @@ struct ByteProgressExecutor { } impl TypedExecutor for ByteProgressExecutor { - async fn execute(&self, _payload: ByteTestTask, ctx: &TaskContext) -> Result<(), TaskError> { + async fn execute<'a>( + &'a self, + _payload: ByteTestTask, + ctx: DomainTaskContext<'a, BenchDomain>, + ) -> Result<(), TaskError> { ctx.set_bytes_total(self.total); let mut remaining = self.total; while remaining > 0 { diff --git a/benches/tags.rs b/benches/tags.rs index a98016a..af7212c 100644 --- a/benches/tags.rs +++ b/benches/tags.rs @@ -7,8 +7,8 @@ use std::time::{Duration, Instant}; use criterion::{criterion_group, criterion_main, BenchmarkId, Criterion}; use serde::{Deserialize, Serialize}; use taskmill::{ - Domain, DomainKey, Scheduler, TaskContext, TaskError, TaskStore, TaskSubmission, TypedExecutor, - TypedTask, + Domain, DomainKey, DomainTaskContext, Scheduler, TaskError, TaskStore, TaskSubmission, + TypedExecutor, TypedTask, }; use tokio::runtime::Runtime; @@ -27,7 +27,11 @@ impl TypedTask for BenchTask { struct NoopExecutor; impl TypedExecutor for NoopExecutor { - async fn execute(&self, _payload: BenchTask, _ctx: &TaskContext) -> Result<(), TaskError> { + async fn execute<'a>( + &'a self, + _payload: BenchTask, + _ctx: DomainTaskContext<'a, BenchDomain>, + ) -> Result<(), TaskError> { Ok(()) } } diff --git a/docs/migrating-to-0.6.md b/docs/migrating-to-0.6.md new file mode 100644 index 0000000..7a934c3 --- /dev/null +++ b/docs/migrating-to-0.6.md @@ -0,0 +1,176 @@ +# Migrating from 0.5.x to 0.6.0 + +0.6.0 replaces the untyped `&TaskContext` in `TypedExecutor` with a +domain-parameterized `DomainTaskContext<'a, D>` wrapper. It also removes the +untyped executor API (`TaskExecutor`, `TaskContext`, `raw_executor`, +`submit_raw`) from the public surface. + +### 1. Executor signature change + +All `TypedExecutor` implementations must update `execute`, `finalize`, and +`on_cancel` to accept `DomainTaskContext<'a, T::Domain>` instead of +`&'a TaskContext`. + +**Before:** + +```rust +impl TypedExecutor for ThumbnailExec { + async fn execute<'a>( + &'a self, + thumb: Thumbnail, + ctx: &'a TaskContext, + ) -> Result<(), TaskError> { + // ... + } +} +``` + +**After:** + +```rust +impl TypedExecutor for ThumbnailExec { + async fn execute<'a>( + &'a self, + thumb: Thumbnail, + ctx: DomainTaskContext<'a, Media>, + ) -> Result<(), TaskError> { + // ... + } +} +``` + +All accessor methods (`record()`, `token()`, `check_cancelled()`, `progress()`, +`state()`, `domain_state()`, `domain()`, `try_domain()`, IO tracking, and +byte-level progress) are delegated identically — no call-site changes needed. + +### 2. Generic executors + +For generic `impl TypedExecutor` implementations, use +`T::Domain` as the type parameter: + +**Before:** + +```rust +impl TypedExecutor for NoopExecutor { + async fn execute<'a>( + &'a self, + _payload: T, + _ctx: &'a TaskContext, + ) -> Result<(), TaskError> { + Ok(()) + } +} +``` + +**After:** + +```rust +impl TypedExecutor for NoopExecutor { + async fn execute<'a>( + &'a self, + _payload: T, + _ctx: DomainTaskContext<'a, T::Domain>, + ) -> Result<(), TaskError> { + Ok(()) + } +} +``` + +### 3. Child spawning (same domain) + +`spawn_child(TaskSubmission)` is no longer exposed on `DomainTaskContext`. +Use `spawn_child_with(task)` which is type-safe — only tasks where +`T::Domain == D` are accepted. + +**Before:** + +```rust +ctx.spawn_child( + TaskSubmission::new("upload-part") + .key(&part.etag) + .payload_json(&UploadPart { etag: part.etag.clone(), size: part.size })?, +).await?; +``` + +**After:** + +```rust +ctx.spawn_child_with(UploadPart { etag: part.etag.clone(), size: part.size }) + .key(&part.etag) + .await?; +``` + +`ChildSpawnBuilder` supports `.key()`, `.priority()`, `.ttl()`, and `.group()` +overrides before `.await`. + +### 4. Batch child spawning + +`spawn_children(Vec)` (mixed types) is replaced by +`spawn_children_with(tasks)` (single type `T`). + +**Before:** + +```rust +let subs: Vec = parts.iter() + .map(|p| TaskSubmission::new("upload-part").key(&p.etag).payload_json(p).unwrap()) + .collect(); +ctx.spawn_children(subs).await?; +``` + +**After:** + +```rust +ctx.spawn_children_with(parts).await?; +``` + +Mixed-type fan-out now requires separate `spawn_child_with` calls per type. +This forgoes SQL batching for that specific pattern. + +### 5. Cross-domain children + +Use the new `child_of(&ctx)` method instead of manually extracting the parent ID. + +**Before:** + +```rust +ctx.domain::() + .submit_with(ScanStartedEvent { .. }) + .parent(ctx.record().id) + .await?; +``` + +**After:** + +```rust +ctx.domain::() + .submit_with(ScanStartedEvent { .. }) + .child_of(&ctx) + .await?; +``` + +### 6. Removed public API + +The following are no longer exported from the crate root: + +| Removed | Replacement | +|---|---| +| `TaskExecutor` trait | Use `TypedExecutor` | +| `TaskContext` struct | Use `DomainTaskContext<'a, D>` | +| `Domain::raw_executor(name, exec)` | Use `Domain::task::(exec)` | +| `DomainHandle::submit_raw(sub)` | Use `DomainHandle::submit(task)` or `submit_with(task)` | + +### 7. Import changes + +**Before:** + +```rust +use taskmill::{TaskContext, TaskExecutor, TypedExecutor, /* ... */}; +``` + +**After:** + +```rust +use taskmill::{DomainTaskContext, TypedExecutor, /* ... */}; +// Optional, if using typed child spawning directly: +use taskmill::ChildSpawnBuilder; +``` diff --git a/examples/profile_dep_chain.rs b/examples/profile_dep_chain.rs index 77ce17c..8a0669a 100644 --- a/examples/profile_dep_chain.rs +++ b/examples/profile_dep_chain.rs @@ -4,8 +4,8 @@ use std::time::{Duration, Instant}; use serde::{Deserialize, Serialize}; use taskmill::{ - Domain, DomainKey, Scheduler, TaskContext, TaskError, TaskStore, TaskSubmission, TypedExecutor, - TypedTask, + Domain, DomainKey, DomainTaskContext, Scheduler, TaskError, TaskStore, TaskSubmission, + TypedExecutor, TypedTask, }; struct BenchDomain; @@ -25,7 +25,7 @@ impl TypedExecutor for NoopExecutor { async fn execute<'a>( &'a self, _payload: BenchTask, - _ctx: &'a TaskContext, + _ctx: DomainTaskContext<'a, BenchDomain>, ) -> Result<(), TaskError> { Ok(()) } diff --git a/src/domain.rs b/src/domain.rs index bdd5754..6067f7b 100644 --- a/src/domain.rs +++ b/src/domain.rs @@ -22,7 +22,7 @@ use std::time::Duration; use crate::module::{ExecutorOptions, ModuleExecutor, ModuleHandle}; use crate::priority::Priority; -use crate::registry::{ErasedExecutor, TaskContext, TaskExecutor}; +use crate::registry::{DomainTaskContext, ErasedExecutor, TaskContext, TaskExecutor}; use crate::scheduler::progress::TaskProgress; use crate::scheduler::SchedulerEvent; use crate::store::StoreError; @@ -154,19 +154,30 @@ pub struct TaskTypeOptions { // ── TypedExecutor ───────────────────────────────────────────────── -/// An executor that receives a deserialized, typed payload. +/// An executor that receives a deserialized, typed payload and a +/// domain-parameterized context. /// /// Register with [`Domain::task::(executor)`](Domain::task). The library /// wraps this in an erased adapter so the scheduler engine remains untyped /// internally. /// +/// The [`DomainTaskContext`] carries the domain identity `D` as a type +/// parameter, enabling compile-time–safe child spawning via +/// [`spawn_child_with`](DomainTaskContext::spawn_child_with). +/// /// # Example /// /// ```ignore /// impl TypedExecutor for ThumbnailExec { -/// async fn execute(&self, thumb: Thumbnail, ctx: &TaskContext) -> Result<(), TaskError> { +/// async fn execute( +/// &self, +/// thumb: Thumbnail, +/// ctx: DomainTaskContext<'_, Media>, +/// ) -> Result<(), TaskError> { /// ctx.check_cancelled()?; -/// process(thumb).await +/// ctx.spawn_child_with(ResizeTask { path: thumb.path, size: 256 }) +/// .await?; +/// Ok(()) /// } /// } /// ``` @@ -175,7 +186,7 @@ pub trait TypedExecutor: Send + Sync + 'static { fn execute<'a>( &'a self, payload: T, - ctx: &'a TaskContext, + ctx: DomainTaskContext<'a, T::Domain>, ) -> impl Future> + Send + 'a; /// Called when all child tasks spawned by this task have settled. @@ -183,7 +194,7 @@ pub trait TypedExecutor: Send + Sync + 'static { fn finalize<'a>( &'a self, _payload: T, - _ctx: &'a TaskContext, + _ctx: DomainTaskContext<'a, T::Domain>, ) -> impl Future> + Send + 'a { async { Ok(()) } } @@ -193,7 +204,7 @@ pub trait TypedExecutor: Send + Sync + 'static { fn on_cancel<'a>( &'a self, _payload: T, - _ctx: &'a TaskContext, + _ctx: DomainTaskContext<'a, T::Domain>, ) -> impl Future> + Send + 'a { async { Ok(()) } } @@ -213,17 +224,20 @@ struct TypedExecutorAdapter { impl> TaskExecutor for TypedExecutorAdapter { async fn execute<'a>(&'a self, ctx: &'a TaskContext) -> Result<(), TaskError> { let payload: T = ctx.payload()?; - self.executor.execute(payload, ctx).await + let dctx = DomainTaskContext::::new(ctx); + self.executor.execute(payload, dctx).await } async fn finalize<'a>(&'a self, ctx: &'a TaskContext) -> Result<(), TaskError> { let payload: T = ctx.payload()?; - self.executor.finalize(payload, ctx).await + let dctx = DomainTaskContext::::new(ctx); + self.executor.finalize(payload, dctx).await } async fn on_cancel<'a>(&'a self, ctx: &'a TaskContext) -> Result<(), TaskError> { let payload: T = ctx.payload()?; - self.executor.on_cancel(payload, ctx).await + let dctx = DomainTaskContext::::new(ctx); + self.executor.on_cancel(payload, dctx).await } } @@ -342,23 +356,6 @@ impl Domain { self } - /// Escape hatch: register an untyped executor by task-type string. - /// - /// Use only when [`TypedExecutor`] is not practical (e.g. dynamic - /// plugin systems). - pub fn raw_executor( - mut self, - task_type: impl Into, - executor: impl TaskExecutor, - ) -> Self { - self.executors.push(ModuleExecutor { - task_type: task_type.into(), - executor: Arc::new(executor) as Arc, - options: ExecutorOptions::default(), - }); - self - } - /// Set the domain-wide default priority. pub fn default_priority(mut self, p: Priority) -> Self { self.default_priority = Some(p); @@ -498,11 +495,6 @@ impl DomainHandle { Ok(outcomes) } - /// Escape hatch: raw untyped submission. - pub fn submit_raw(&self, sub: TaskSubmission) -> crate::task::SubmitBuilder { - self.inner.submit(sub) - } - // ── Lifecycle ─────────────────────────────────────────────────── /// Cancel a task by ID. @@ -752,10 +744,8 @@ impl DomainSubmitBuilder { } /// Set a dependency failure policy. - pub fn on_dependency_failure(self, _p: DependencyFailurePolicy) -> Self { - // DependencyFailurePolicy is set on the TaskSubmission, which is already - // resolved inside the inner SubmitBuilder. For now, delegate to the - // underlying submission mechanism. + pub fn on_dependency_failure(mut self, p: DependencyFailurePolicy) -> Self { + self.inner = self.inner.on_dependency_failure(p); self } @@ -765,6 +755,21 @@ impl DomainSubmitBuilder { self } + /// Mark this task as a child of the task currently executing in the + /// given [`DomainTaskContext`]. + /// + /// This is the idiomatic way to create cross-domain children: + /// + /// ```ignore + /// ctx.domain::() + /// .submit_with(ScanStartedEvent { .. }) + /// .child_of(&ctx) + /// .await?; + /// ``` + pub fn child_of(self, ctx: &DomainTaskContext<'_, D2>) -> Self { + self.parent(ctx.record().id) + } + /// Submit the task, returning the outcome. pub async fn submit(self) -> Result { self.inner.submit().await @@ -1018,7 +1023,7 @@ mod tests { async fn execute<'a>( &'a self, _payload: TestTask, - _ctx: &'a TaskContext, + _ctx: DomainTaskContext<'a, TestDomain>, ) -> Result<(), TaskError> { Ok(()) } diff --git a/src/lib.rs b/src/lib.rs index 314b6ad..adcc01b 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -71,7 +71,7 @@ //! - **`Supersede`** — cancel the existing task (recording it in history as //! [`HistoryStatus::Superseded`]) and replace it with the new submission. //! For running tasks the cancellation token is fired, the -//! [`on_cancel`](TaskExecutor::on_cancel) hook runs, and children are +//! [`on_cancel`](TypedExecutor::on_cancel) hook runs, and children are //! cascade-cancelled. Returns [`SubmitOutcome::Superseded`]. //! - **`Reject`** — return [`SubmitOutcome::Rejected`] without modifying the //! existing task. @@ -124,7 +124,7 @@ //! An executor can spawn child tasks via [`TaskContext::spawn_child`]. When //! children exist, the parent enters a **waiting** state after its executor //! returns. Once all children complete, the parent's -//! [`TaskExecutor::finalize`] method is called — useful for assembly work +//! [`TypedExecutor::finalize`] method is called — useful for assembly work //! like `CompleteMultipartUpload`. If any child fails and //! [`fail_fast`](TaskSubmission::fail_fast) is `true` (the default), siblings //! are cancelled and the parent fails immediately. @@ -332,7 +332,7 @@ //! struct ThumbnailExecutor; //! //! impl TypedExecutor for ThumbnailExecutor { -//! async fn execute(&self, thumb: Thumbnail, ctx: &TaskContext) -> Result<(), TaskError> { +//! async fn execute(&self, thumb: Thumbnail, ctx: DomainTaskContext<'_, Media>) -> Result<(), TaskError> { //! ctx.progress().report(0.5, Some("resizing".into())); //! // ... do work, check ctx.token().is_cancelled() ... //! ctx.record_read_bytes(4_096); @@ -387,7 +387,7 @@ //! .await?; //! //! // Inside an ingest executor — domain state checked first, then global: -//! async fn execute(&self, task: FetchTask, ctx: &TaskContext) -> Result<(), TaskError> { +//! async fn execute(&self, task: FetchTask, ctx: DomainTaskContext<'_, Ingest>) -> Result<(), TaskError> { //! let cfg = ctx.state::().expect("IngestConfig not registered"); //! let svc = ctx.state::().expect("AppServices not registered"); //! svc.db.query("...").await?; @@ -510,7 +510,7 @@ //! ``` //! //! For untyped batch submission with batch-wide defaults, use [`BatchSubmission`] -//! and [`DomainHandle::submit_raw`]. +//! and [`Scheduler::submit_built`]. //! //! A [`SchedulerEvent::BatchSubmitted`] event is emitted for observability //! whenever at least one task in the batch was inserted. @@ -524,21 +524,17 @@ //! //! ```ignore //! impl TypedExecutor for MultipartUploadExecutor { -//! async fn execute(&self, upload: MultipartUpload, ctx: &TaskContext) -> Result<(), TaskError> { +//! async fn execute(&self, upload: MultipartUpload, ctx: DomainTaskContext<'_, Uploads>) -> Result<(), TaskError> { //! for part in &upload.parts { -//! // "upload-part" is prefixed with the owning domain name automatically. -//! ctx.spawn_child( -//! TaskSubmission::new("upload-part") -//! .key(&part.etag) -//! .priority(ctx.record().priority) -//! .payload_json(part) -//! .expected_io(IoBudget::disk(part.size as i64, 0)), -//! ).await?; +//! ctx.spawn_child_with(UploadPart { etag: part.etag.clone(), size: part.size }) +//! .key(&part.etag) +//! .priority(ctx.record().priority) +//! .await?; //! } //! Ok(()) //! } //! -//! async fn finalize(&self, upload: MultipartUpload, ctx: &TaskContext) -> Result<(), TaskError> { +//! async fn finalize(&self, upload: MultipartUpload, ctx: DomainTaskContext<'_, Uploads>) -> Result<(), TaskError> { //! // All parts uploaded — complete the multipart upload. //! complete_multipart(&upload).await?; //! Ok(()) @@ -562,7 +558,7 @@ //! //! ```ignore //! impl TypedExecutor for UploadExecutor { -//! async fn execute(&self, upload: Upload, ctx: &TaskContext) -> Result<(), TaskError> { +//! async fn execute(&self, upload: Upload, ctx: DomainTaskContext<'_, Uploads>) -> Result<(), TaskError> { //! // Cooperatively check for cancellation in long loops. //! for chunk in upload.chunks() { //! ctx.check_cancelled()?; @@ -571,7 +567,7 @@ //! Ok(()) //! } //! -//! async fn on_cancel(&self, upload: Upload, ctx: &TaskContext) -> Result<(), TaskError> { +//! async fn on_cancel(&self, upload: Upload, ctx: DomainTaskContext<'_, Uploads>) -> Result<(), TaskError> { //! // Abort the in-progress multipart upload. //! abort_multipart(&upload.upload_id).await?; //! Ok(()) @@ -638,10 +634,19 @@ //! .run_after(Duration::from_secs(30)) //! .await?; //! -//! // For recurring or complex scheduling, use `submit_raw` with TaskSubmission: -//! media.submit_raw(TaskSubmission::new("cleanup") +//! // For recurring scheduling, configure it in `TypedTask::config()`: +//! // impl TypedTask for Cleanup { +//! // fn config() -> TaskTypeConfig { +//! // TaskTypeConfig::new().recurring(RecurringSchedule { +//! // interval: Duration::from_secs(6 * 3600), +//! // initial_delay: None, +//! // max_executions: None, +//! // }) +//! // } +//! // } +//! media.submit_with(Cleanup { target: "stale-uploads".into() }) //! .key("stale-uploads") -//! .recurring(Duration::from_secs(6 * 3600))).await?; +//! .await?; //! //! // Pause/resume/cancel recurring schedules via the domain handle. //! media.pause_recurring(task_id).await?; @@ -784,11 +789,11 @@ pub use domain::{ Domain, DomainHandle, DomainKey, DomainSubmitBuilder, TaskEvent, TaskTypeConfig, TaskTypeOptions, TypedEventStream, TypedExecutor, }; +pub use registry::{ChildSpawnBuilder, DomainTaskContext}; // ── Core re-exports ────────────────────────────────────────────────── pub use backpressure::{CompositePressure, PressureSource, ThrottlePolicy}; pub use priority::Priority; -pub use registry::{TaskContext, TaskExecutor}; pub use resource::network_pressure::NetworkPressure; pub use resource::sampler::SamplerConfig; pub use resource::{ResourceReader, ResourceSampler, ResourceSnapshot}; diff --git a/src/module.rs b/src/module.rs index 30c0f20..3396291 100644 --- a/src/module.rs +++ b/src/module.rs @@ -895,7 +895,7 @@ mod tests { use crate::domain::{Domain, DomainKey, TypedExecutor}; use crate::priority::Priority; - use crate::registry::TaskContext; + use crate::registry::DomainTaskContext; use crate::task::retry::{BackoffStrategy, RetryPolicy}; use crate::task::{TaskError, TypedTask}; @@ -907,7 +907,7 @@ mod tests { async fn execute<'a>( &'a self, _payload: T, - _ctx: &'a TaskContext, + _ctx: DomainTaskContext<'a, T::Domain>, ) -> Result<(), TaskError> { Ok(()) } diff --git a/src/registry/context.rs b/src/registry/context.rs index 94b8f6e..6a3d466 100644 --- a/src/registry/context.rs +++ b/src/registry/context.rs @@ -15,7 +15,7 @@ use super::child_spawner::ChildSpawner; use super::io_tracker::IoTracker; use super::state::StateSnapshot; -/// Execution context passed to a [`TaskExecutor`](super::TaskExecutor). +/// Execution context passed to a [`TypedExecutor`](crate::TypedExecutor). /// /// Provides access to the task record, cancellation token, progress reporter, /// shared application state, and domain-scoped task submission. Use the accessor @@ -218,7 +218,7 @@ impl TaskContext { /// /// ```ignore /// impl TypedExecutor for VideoProcessor { - /// async fn execute(&self, event: VideoUploaded, ctx: &TaskContext) -> Result<(), TaskError> { + /// async fn execute(&self, event: VideoUploaded, ctx: DomainTaskContext<'_, Media>) -> Result<(), TaskError> { /// let media = ctx.domain::(); /// media.submit(Thumbnail { path: event.path.clone(), size: 256 }) /// .await diff --git a/src/registry/domain_context.rs b/src/registry/domain_context.rs new file mode 100644 index 0000000..3d96405 --- /dev/null +++ b/src/registry/domain_context.rs @@ -0,0 +1,299 @@ +//! [`DomainTaskContext`] — domain-parameterized execution context wrapper. +//! +//! This module provides a zero-cost wrapper around [`TaskContext`] that carries +//! domain identity as a type parameter. This enables compile-time–safe child +//! spawning: [`spawn_child_with`](DomainTaskContext::spawn_child_with) only +//! accepts tasks belonging to the same domain `D`. +//! +//! # Why a wrapper instead of `TaskContext`? +//! +//! Directly parameterizing `TaskContext` would break type erasure. The registry +//! stores all executors in a single `HashMap>`, +//! and `ErasedExecutor` methods take `&TaskContext`. Making `TaskContext` generic +//! would require per-domain registries and force the dispatcher to know `D` at +//! runtime. The wrapper approach preserves the untyped internals while presenting +//! a typed API to executor authors. + +use std::marker::PhantomData; + +use tokio_util::sync::CancellationToken; + +use crate::domain::{DomainHandle, DomainKey}; +use crate::scheduler::ProgressReporter; +use crate::store::StoreError; +use crate::task::{SubmitOutcome, TaskError, TaskRecord, TaskSubmission, TypedTask}; + +use super::context::TaskContext; + +/// Domain-parameterized execution context passed to [`TypedExecutor`](crate::TypedExecutor). +/// +/// A zero-cost wrapper around the internal [`TaskContext`] that carries the +/// domain identity `D` as a type parameter. This enables: +/// +/// - **Compile-time–safe child spawning** via [`spawn_child_with`](Self::spawn_child_with): +/// only tasks where `T::Domain == D` are accepted. +/// - **No escape hatch**: the untyped `spawn_child(TaskSubmission)` is not +/// accessible through this wrapper. +/// +/// Obtain a handle to a different domain via [`domain`](Self::domain) for +/// cross-domain submission. +/// +/// # Example +/// +/// ```ignore +/// impl TypedExecutor for RootScanner { +/// async fn execute( +/// &self, +/// task: ScanRootTask, +/// ctx: DomainTaskContext<'_, Scanner>, +/// ) -> Result<(), TaskError> { +/// ctx.spawn_child_with(ScanL1DirTask { bucket: task.bucket, .. }) +/// .key(&format!("{}:{}", task.bucket, task.prefix)) +/// .await?; +/// Ok(()) +/// } +/// } +/// ``` +pub struct DomainTaskContext<'a, D: DomainKey> { + pub(crate) inner: &'a TaskContext, + _domain: PhantomData, +} + +impl<'a, D: DomainKey> DomainTaskContext<'a, D> { + /// Construct a new `DomainTaskContext` wrapping the given `TaskContext`. + /// + /// This is `pub(crate)` — only [`TypedExecutorAdapter`](crate::domain::TypedExecutorAdapter) + /// should construct these. + pub(crate) fn new(ctx: &'a TaskContext) -> Self { + Self { + inner: ctx, + _domain: PhantomData, + } + } + + // ── Delegated accessors ───────────────────────────────────────── + + /// The persisted task record (id, key, priority, payload, etc.). + pub fn record(&self) -> &TaskRecord { + self.inner.record() + } + + /// Cancellation token — check `token().is_cancelled()` for preemption. + pub fn token(&self) -> &CancellationToken { + self.inner.token() + } + + /// Check whether this task has been cancelled, returning a + /// [`TaskError::cancelled()`] if so. + pub fn check_cancelled(&self) -> Result<(), TaskError> { + self.inner.check_cancelled() + } + + /// Progress reporter for this task. + pub fn progress(&self) -> &ProgressReporter { + self.inner.progress() + } + + // ── Shared state ──────────────────────────────────────────────── + + /// Retrieve shared application state registered via + /// [`SchedulerBuilder::app_state`](crate::SchedulerBuilder::app_state) or + /// [`Domain::state`](crate::Domain::state). + pub fn state(&self) -> Option<&T> { + self.inner.state::() + } + + /// Retrieve domain-scoped state. + pub fn domain_state(&self) -> Option<&S> { + self.inner.domain_state::() + } + + // ── IO tracking ───────────────────────────────────────────────── + + /// Record actual bytes read during this task's execution. + pub fn record_read_bytes(&self, bytes: i64) { + self.inner.record_read_bytes(bytes); + } + + /// Record actual bytes written during this task's execution. + pub fn record_write_bytes(&self, bytes: i64) { + self.inner.record_write_bytes(bytes); + } + + /// Record actual bytes received over the network. + pub fn record_net_rx_bytes(&self, bytes: i64) { + self.inner.record_net_rx_bytes(bytes); + } + + /// Record actual bytes transmitted over the network. + pub fn record_net_tx_bytes(&self, bytes: i64) { + self.inner.record_net_tx_bytes(bytes); + } + + // ── Byte-level progress ───────────────────────────────────────── + + /// Set the total number of bytes expected for byte-level progress. + pub fn set_bytes_total(&self, total: u64) { + self.inner.set_bytes_total(total); + } + + /// Increment completed bytes by `delta` for byte-level progress. + pub fn add_bytes(&self, delta: u64) { + self.inner.add_bytes(delta); + } + + /// Set both completed and total bytes to absolute values. + pub fn report_bytes(&self, completed: u64, total: u64) { + self.inner.report_bytes(completed, total); + } + + // ── Domain access ─────────────────────────────────────────────── + + /// Returns a typed [`DomainHandle`] for cross-domain task submission. + /// + /// # Panics + /// + /// Panics if `D2` was not registered with the scheduler. + pub fn domain(&self) -> DomainHandle { + self.inner.domain::() + } + + /// Returns a typed [`DomainHandle`] for the given domain, or `None` + /// if the domain is not registered. + pub fn try_domain(&self) -> Option> { + self.inner.try_domain::() + } + + // ── Typed child spawning ──────────────────────────────────────── + + /// Spawn a same-domain child task with compile-time type safety. + /// + /// Only accepts tasks where `T::Domain == D`. Returns a + /// [`ChildSpawnBuilder`] for optional per-call overrides (`.key()`, + /// `.priority()`, etc.), then `.await` to submit. + /// + /// # Example + /// + /// ```ignore + /// ctx.spawn_child_with(ScanL1DirTask { bucket, prefix }) + /// .key(&format!("{bucket}:{prefix}")) + /// .await?; + /// ``` + pub fn spawn_child_with>( + &self, + task: T, + ) -> ChildSpawnBuilder<'a, D, T> { + ChildSpawnBuilder { + ctx: self.inner, + task, + override_key: None, + override_priority: None, + override_ttl: None, + override_group: None, + _domain: PhantomData, + } + } + + /// Spawn multiple same-domain children in one call. + /// + /// All tasks must be the same type `T`. For mixed-type fan-out, use + /// separate [`spawn_child_with`](Self::spawn_child_with) calls per type. + /// + /// # Example + /// + /// ```ignore + /// let children: Vec = prefixes.into_iter() + /// .map(|p| ScanL1DirTask { bucket: bucket.clone(), prefix: p }) + /// .collect(); + /// ctx.spawn_children_with(children).await?; + /// ``` + pub async fn spawn_children_with>( + &self, + tasks: impl IntoIterator, + ) -> Result, StoreError> { + let submissions: Vec = tasks + .into_iter() + .map(|t| TaskSubmission::from_typed(&t)) + .collect(); + self.inner.spawn_children(submissions).await + } +} + +// ── ChildSpawnBuilder ─────────────────────────────────────────────── + +/// Builder for spawning a single typed child task with optional per-call +/// overrides. +/// +/// Created by [`DomainTaskContext::spawn_child_with`]. Chain override methods +/// then `.await` to submit. +/// +/// Implements [`IntoFuture`] so bare `.await` works: +/// +/// ```ignore +/// ctx.spawn_child_with(task).key("my-key").await?; +/// ``` +pub struct ChildSpawnBuilder<'a, D: DomainKey, T: TypedTask> { + ctx: &'a TaskContext, + task: T, + override_key: Option, + override_priority: Option, + override_ttl: Option, + override_group: Option, + _domain: PhantomData, +} + +impl<'a, D: DomainKey, T: TypedTask> ChildSpawnBuilder<'a, D, T> { + /// Override the dedup key for this child task. + pub fn key(mut self, k: impl Into) -> Self { + self.override_key = Some(k.into()); + self + } + + /// Override the priority for this child task. + pub fn priority(mut self, p: crate::priority::Priority) -> Self { + self.override_priority = Some(p); + self + } + + /// Override the time-to-live for this child task. + pub fn ttl(mut self, d: std::time::Duration) -> Self { + self.override_ttl = Some(d); + self + } + + /// Override the group key for this child task. + pub fn group(mut self, key: impl Into) -> Self { + self.override_group = Some(key.into()); + self + } + + /// Submit the child task. + pub async fn submit(self) -> Result { + let mut sub = TaskSubmission::from_typed(&self.task); + if let Some(k) = self.override_key { + sub = sub.key(k); + } + if let Some(p) = self.override_priority { + sub = sub.priority(p); + } + if let Some(d) = self.override_ttl { + sub = sub.ttl(d); + } + if let Some(g) = self.override_group { + sub = sub.group(g); + } + self.ctx.spawn_child(sub).await + } +} + +impl<'a, D: DomainKey, T: TypedTask> std::future::IntoFuture + for ChildSpawnBuilder<'a, D, T> +{ + type Output = Result; + type IntoFuture = + std::pin::Pin + Send + 'a>>; + + fn into_future(self) -> Self::IntoFuture { + Box::pin(self.submit()) + } +} diff --git a/src/registry/mod.rs b/src/registry/mod.rs index 2f94a5e..b25a996 100644 --- a/src/registry/mod.rs +++ b/src/registry/mod.rs @@ -11,6 +11,7 @@ pub(crate) mod child_spawner; mod context; +mod domain_context; pub(crate) mod io_tracker; pub(crate) mod state; @@ -22,7 +23,8 @@ use crate::task::retry::RetryPolicy; use crate::task::TaskError; pub(crate) use child_spawner::{ChildSpawner, ParentContext}; -pub use context::TaskContext; +pub(crate) use context::TaskContext; +pub use domain_context::{ChildSpawnBuilder, DomainTaskContext}; pub(crate) use io_tracker::IoTracker; pub(crate) use state::{StateMap, StateSnapshot}; @@ -53,7 +55,7 @@ pub(crate) use state::{StateMap, StateSnapshot}; /// } /// } /// ``` -pub trait TaskExecutor: Send + Sync + 'static { +pub(crate) trait TaskExecutor: Send + Sync + 'static { /// Execute a task. /// /// - `ctx`: Execution context with the task record, cancellation token, @@ -163,40 +165,6 @@ impl TaskTypeRegistry { } } - /// Register an executor for a named task type. - /// - /// Panics if the name is already registered (catch configuration errors - /// at startup, not at runtime). - pub fn register(&mut self, name: &str, executor: Arc) { - if self.types.contains_key(name) { - panic!("task type '{name}' already registered"); - } - self.types - .insert(name.to_string(), executor as Arc); - } - - /// Register an executor with a per-type default TTL. - pub fn register_with_ttl( - &mut self, - name: &str, - executor: Arc, - ttl: std::time::Duration, - ) { - self.register(name, executor); - self.type_ttls.insert(name.to_string(), ttl); - } - - /// Register an executor with a per-type retry policy. - pub fn register_with_retry_policy( - &mut self, - name: &str, - executor: Arc, - policy: RetryPolicy, - ) { - self.register(name, executor); - self.type_retry_policies.insert(name.to_string(), policy); - } - /// Look up the per-type default TTL for a task type. pub fn type_ttl(&self, name: &str) -> Option<&std::time::Duration> { self.type_ttls.get(name) @@ -295,7 +263,7 @@ mod tests { async fn execute<'a>( &'a self, _payload: NoopTask, - _ctx: &'a TaskContext, + _ctx: DomainTaskContext<'a, TestDomain>, ) -> Result<(), TaskError> { Ok(()) } diff --git a/src/scheduler/builder.rs b/src/scheduler/builder.rs index daf75b2..4a030eb 100644 --- a/src/scheduler/builder.rs +++ b/src/scheduler/builder.rs @@ -149,7 +149,7 @@ impl SchedulerBuilder { self } - /// Set the timeout for [`TaskExecutor::on_cancel`](crate::TaskExecutor::on_cancel) hooks. + /// Set the timeout for [`TypedExecutor::on_cancel`](crate::TypedExecutor::on_cancel) hooks. /// Default: 30 seconds. pub fn cancel_hook_timeout(mut self, timeout: Duration) -> Self { self.config.cancel_hook_timeout = timeout; diff --git a/src/scheduler/event.rs b/src/scheduler/event.rs index 87925c3..3c41402 100644 --- a/src/scheduler/event.rs +++ b/src/scheduler/event.rs @@ -250,7 +250,7 @@ pub struct SchedulerConfig { /// When set, a background task polls active tasks' byte counters at this /// interval and emits [`TaskProgress`] events on a dedicated channel. pub progress_interval: Option, - /// Timeout for [`TaskExecutor::on_cancel`](crate::TaskExecutor::on_cancel) + /// Timeout for [`TypedExecutor::on_cancel`](crate::TypedExecutor::on_cancel) /// hooks. If a cancel hook does not complete within this duration it is /// aborted. Default: 30 seconds. pub cancel_hook_timeout: Duration, diff --git a/src/scheduler/progress.rs b/src/scheduler/progress.rs index 3334c59..57fc282 100644 --- a/src/scheduler/progress.rs +++ b/src/scheduler/progress.rs @@ -52,7 +52,7 @@ use super::SchedulerEvent; /// # Example /// /// ```ignore -/// // Inside a TaskExecutor::execute implementation: +/// // Inside a TypedExecutor::execute implementation: /// async fn execute<'a>(&'a self, ctx: &'a TaskContext) -> Result<(), TaskError> { /// let items = vec![/* ... */]; /// for (i, item) in items.iter().enumerate() { diff --git a/src/scheduler/submit.rs b/src/scheduler/submit.rs index e1da13f..15b1f93 100644 --- a/src/scheduler/submit.rs +++ b/src/scheduler/submit.rs @@ -250,7 +250,7 @@ impl Scheduler { /// /// Records the task in history as `cancelled` (instead of silently /// deleting). For running/waiting tasks, triggers the cancellation - /// token, fires the [`on_cancel`](crate::TaskExecutor::on_cancel) hook, + /// token, fires the [`on_cancel`](crate::TypedExecutor::on_cancel) hook, /// and emits a [`SchedulerEvent::Cancelled`] event. Returns `true` if /// the task was found and cancelled. pub async fn cancel(&self, task_id: i64) -> Result { diff --git a/src/scheduler/tests.rs b/src/scheduler/tests.rs index 84f945b..1834096 100644 --- a/src/scheduler/tests.rs +++ b/src/scheduler/tests.rs @@ -6,7 +6,7 @@ use tokio_util::sync::CancellationToken; use crate::domain::{Domain, DomainKey, TypedExecutor}; use crate::priority::Priority; -use crate::registry::TaskContext; +use crate::registry::DomainTaskContext; use crate::store::TaskStore; use crate::task::{ DuplicateStrategy, HistoryStatus, SubmitOutcome, TaskError, TaskSubmission, TypedTask, @@ -96,7 +96,11 @@ impl TypedTask for BetaTask { struct InstantExecutor; impl TypedExecutor for InstantExecutor { - async fn execute<'a>(&'a self, _payload: T, ctx: &'a TaskContext) -> Result<(), TaskError> { + async fn execute<'a>( + &'a self, + _payload: T, + ctx: DomainTaskContext<'a, T::Domain>, + ) -> Result<(), TaskError> { ctx.record_read_bytes(100); ctx.record_write_bytes(50); Ok(()) @@ -106,7 +110,11 @@ impl TypedExecutor for InstantExecutor { struct SlowExecutor; impl TypedExecutor for SlowExecutor { - async fn execute<'a>(&'a self, _payload: T, ctx: &'a TaskContext) -> Result<(), TaskError> { + async fn execute<'a>( + &'a self, + _payload: T, + ctx: DomainTaskContext<'a, T::Domain>, + ) -> Result<(), TaskError> { tokio::select! { _ = ctx.token().cancelled() => { Err(TaskError::new("cancelled")) @@ -129,7 +137,7 @@ impl TypedExecutor for CancelHookExecutor { async fn execute<'a>( &'a self, _payload: TestTask, - ctx: &'a TaskContext, + ctx: DomainTaskContext<'a, TestDomain>, ) -> Result<(), TaskError> { tokio::select! { _ = ctx.token().cancelled() => { @@ -144,7 +152,7 @@ impl TypedExecutor for CancelHookExecutor { async fn on_cancel<'a>( &'a self, _payload: TestTask, - _ctx: &'a TaskContext, + _ctx: DomainTaskContext<'a, TestDomain>, ) -> Result<(), TaskError> { self.cancel_called .store(true, std::sync::atomic::Ordering::SeqCst); @@ -159,7 +167,7 @@ impl TypedExecutor for FailingExecutor { async fn execute<'a>( &'a self, _payload: TestTask, - _ctx: &'a TaskContext, + _ctx: DomainTaskContext<'a, TestDomain>, ) -> Result<(), TaskError> { Err(TaskError::retryable("boom")) } @@ -458,7 +466,7 @@ async fn app_state_accessible_from_executor() { async fn execute<'a>( &'a self, _payload: TestTask, - ctx: &'a TaskContext, + ctx: DomainTaskContext<'a, TestDomain>, ) -> Result<(), TaskError> { let state = ctx.state::().expect("state should be set"); state.flag.store(true, Ordering::SeqCst); @@ -567,14 +575,14 @@ impl TypedExecutor for SpawningExecutor { async fn execute<'a>( &'a self, _payload: ParentTask, - ctx: &'a TaskContext, + ctx: DomainTaskContext<'a, ParentDomain>, ) -> Result<(), TaskError> { for i in 0..self.num_children { let sub = TaskSubmission::new("child") .key(format!("child-{i}")) .priority(ctx.record().priority) .payload_json(&ChildTask); - ctx.spawn_child(sub).await?; + ctx.inner.spawn_child(sub).await?; } Ok(()) } @@ -590,14 +598,14 @@ impl TypedExecutor for FinalizeTrackingExecutor { async fn execute<'a>( &'a self, _payload: ParentTask, - ctx: &'a TaskContext, + ctx: DomainTaskContext<'a, ParentDomain>, ) -> Result<(), TaskError> { for i in 0..self.children { let sub = TaskSubmission::new("child") .key(format!("ft-child-{i}")) .priority(ctx.record().priority) .payload_json(&ChildTask); - ctx.spawn_child(sub).await?; + ctx.inner.spawn_child(sub).await?; } Ok(()) } @@ -605,7 +613,7 @@ impl TypedExecutor for FinalizeTrackingExecutor { async fn finalize<'a>( &'a self, _payload: ParentTask, - _ctx: &'a TaskContext, + _ctx: DomainTaskContext<'a, ParentDomain>, ) -> Result<(), TaskError> { self.finalized .store(true, std::sync::atomic::Ordering::SeqCst); @@ -844,7 +852,7 @@ impl TypedExecutor for ByteProgressExecutor { async fn execute<'a>( &'a self, _payload: ByteTestTask, - ctx: &'a TaskContext, + ctx: DomainTaskContext<'a, ByteTestDomain>, ) -> Result<(), TaskError> { ctx.set_bytes_total(1_048_576); for _ in 0..1024 { @@ -1283,7 +1291,7 @@ async fn on_cancel_hook_timeout_does_not_block() { async fn execute<'a>( &'a self, _payload: TestTask, - ctx: &'a TaskContext, + ctx: DomainTaskContext<'a, TestDomain>, ) -> Result<(), TaskError> { tokio::select! { _ = ctx.token().cancelled() => Err(TaskError::new("cancelled")), @@ -1294,7 +1302,7 @@ async fn on_cancel_hook_timeout_does_not_block() { async fn on_cancel<'a>( &'a self, _payload: TestTask, - _ctx: &'a TaskContext, + _ctx: DomainTaskContext<'a, TestDomain>, ) -> Result<(), TaskError> { // Simulate a very slow cancel hook. tokio::time::sleep(Duration::from_secs(60)).await; diff --git a/src/task/submit_builder.rs b/src/task/submit_builder.rs index 6d9e910..b5c1ec9 100644 --- a/src/task/submit_builder.rs +++ b/src/task/submit_builder.rs @@ -35,6 +35,7 @@ use chrono::{DateTime, Utc}; use crate::priority::Priority; use crate::scheduler::Scheduler; use crate::store::StoreError; +use crate::task::submission::DependencyFailurePolicy; use crate::task::{SubmitOutcome, TaskSubmission, TtlFrom}; /// Module-level defaults applied to every submission through a module handle. @@ -108,6 +109,7 @@ pub struct SubmitBuilder { override_run_after: Option>, override_ttl: Option, override_depends_on: Vec, + override_on_dependency_failure: Option, override_tags: HashMap, override_parent_id: Option, } @@ -135,6 +137,7 @@ impl SubmitBuilder { override_run_after: None, override_ttl: None, override_depends_on: Vec::new(), + override_on_dependency_failure: None, override_tags: HashMap::new(), override_parent_id: None, } @@ -196,6 +199,12 @@ impl SubmitBuilder { self } + /// Set the dependency failure policy (`Cancel`, `Fail`, or `Ignore`). + pub fn on_dependency_failure(mut self, policy: DependencyFailurePolicy) -> Self { + self.override_on_dependency_failure = Some(policy); + self + } + /// Add a metadata tag. Override tags win over submission-level and /// module-level tags for the same key. pub fn tag(mut self, key: impl Into, value: impl Into) -> Self { @@ -331,6 +340,9 @@ impl SubmitBuilder { .dependencies .append(&mut self.override_depends_on); } + if let Some(p) = self.override_on_dependency_failure.take() { + self.submission.on_dependency_failure = p; + } for (k, v) in std::mem::take(&mut self.override_tags) { self.submission.tags.insert(k, v); } diff --git a/tests/integration/common.rs b/tests/integration/common.rs index d3cdec6..37d1590 100644 --- a/tests/integration/common.rs +++ b/tests/integration/common.rs @@ -6,10 +6,12 @@ use std::sync::atomic::{AtomicBool, AtomicI32, AtomicUsize, Ordering}; use std::sync::Arc; use std::time::Duration; +use std::marker::PhantomData; + use serde::{Deserialize, Serialize}; use taskmill::{ - DomainKey, PressureSource, SchedulerEvent, TaskContext, TaskError, TaskSubmission, - TypedExecutor, TypedTask, + DomainKey, DomainTaskContext, PressureSource, SchedulerEvent, TaskError, TypedExecutor, + TypedTask, }; // ── Domain Keys ──────────────────────────────────────────────────── @@ -68,7 +70,7 @@ impl DomainKey for GammaDomain { /// ``` macro_rules! define_task { ($name:ident, $domain:ty, $task_type:expr) => { - #[derive(Debug, Clone, Serialize, Deserialize)] + #[derive(Debug, Default, Clone, Serialize, Deserialize)] pub struct $name; impl TypedTask for $name { type Domain = $domain; @@ -124,7 +126,11 @@ define_task!(DomainBTask, DomainB, "task"); pub struct NoopExecutor; impl TypedExecutor for NoopExecutor { - async fn execute<'a>(&'a self, _payload: T, _ctx: &'a TaskContext) -> Result<(), TaskError> { + async fn execute<'a>( + &'a self, + _payload: T, + _ctx: DomainTaskContext<'a, T::Domain>, + ) -> Result<(), TaskError> { Ok(()) } } @@ -133,7 +139,11 @@ impl TypedExecutor for NoopExecutor { pub struct DelayExecutor(pub Duration); impl TypedExecutor for DelayExecutor { - async fn execute<'a>(&'a self, _payload: T, ctx: &'a TaskContext) -> Result<(), TaskError> { + async fn execute<'a>( + &'a self, + _payload: T, + ctx: DomainTaskContext<'a, T::Domain>, + ) -> Result<(), TaskError> { tokio::select! { _ = ctx.token().cancelled() => Err(TaskError::new("cancelled")), _ = tokio::time::sleep(self.0) => Ok(()), @@ -147,7 +157,11 @@ pub struct CountingExecutor { } impl TypedExecutor for CountingExecutor { - async fn execute<'a>(&'a self, _payload: T, _ctx: &'a TaskContext) -> Result<(), TaskError> { + async fn execute<'a>( + &'a self, + _payload: T, + _ctx: DomainTaskContext<'a, T::Domain>, + ) -> Result<(), TaskError> { self.count.fetch_add(1, Ordering::SeqCst); Ok(()) } @@ -160,7 +174,11 @@ pub struct FailNTimesExecutor { } impl TypedExecutor for FailNTimesExecutor { - async fn execute<'a>(&'a self, _payload: T, _ctx: &'a TaskContext) -> Result<(), TaskError> { + async fn execute<'a>( + &'a self, + _payload: T, + _ctx: DomainTaskContext<'a, T::Domain>, + ) -> Result<(), TaskError> { let count = self.failures.fetch_add(1, Ordering::SeqCst); if count < self.max_failures { Err(TaskError::retryable("transient failure")) @@ -170,14 +188,18 @@ impl TypedExecutor for FailNTimesExecutor { } } -/// Records IO bytes via TaskContext. +/// Records IO bytes via DomainTaskContext. pub struct IoReportingExecutor { pub read: i64, pub write: i64, } impl TypedExecutor for IoReportingExecutor { - async fn execute<'a>(&'a self, _payload: T, ctx: &'a TaskContext) -> Result<(), TaskError> { + async fn execute<'a>( + &'a self, + _payload: T, + ctx: DomainTaskContext<'a, T::Domain>, + ) -> Result<(), TaskError> { ctx.record_read_bytes(self.read); ctx.record_write_bytes(self.write); Ok(()) @@ -192,7 +214,11 @@ pub struct ConcurrencyTrackingExecutor { } impl TypedExecutor for ConcurrencyTrackingExecutor { - async fn execute<'a>(&'a self, _payload: T, ctx: &'a TaskContext) -> Result<(), TaskError> { + async fn execute<'a>( + &'a self, + _payload: T, + ctx: DomainTaskContext<'a, T::Domain>, + ) -> Result<(), TaskError> { let prev = self.current.fetch_add(1, Ordering::SeqCst); self.max_seen.fetch_max(prev + 1, Ordering::SeqCst); tokio::select! { @@ -204,44 +230,84 @@ impl TypedExecutor for ConcurrencyTrackingExecutor { } } -/// An executor that spawns N child tasks. -pub struct ChildSpawnerExecutor { - pub child_type: &'static str, +/// An executor that spawns N child tasks of type `C` in the same domain. +pub struct ChildSpawnerExecutor { pub count: usize, - pub fail_fast: bool, + pub _child: PhantomData, +} + +impl ChildSpawnerExecutor { + pub fn new(count: usize) -> Self { + Self { + count, + _child: PhantomData, + } + } } -impl TypedExecutor for ChildSpawnerExecutor { - async fn execute<'a>(&'a self, _payload: T, ctx: &'a TaskContext) -> Result<(), TaskError> { +impl TypedExecutor for ChildSpawnerExecutor +where + T: TypedTask, + C: TypedTask + Default + Send + Sync + 'static, +{ + async fn execute<'a>( + &'a self, + _payload: T, + ctx: DomainTaskContext<'a, T::Domain>, + ) -> Result<(), TaskError> { for i in 0..self.count { - let sub = TaskSubmission::new(self.child_type) + ctx.spawn_child_with(C::default()) .key(format!("child-{i}")) .priority(ctx.record().priority) - .fail_fast(self.fail_fast); - ctx.spawn_child(sub).await?; + .await + .map_err(|e| TaskError::new(e.to_string()))?; } Ok(()) } } -/// Tracks whether finalize was called. -pub struct FinalizeTracker { +/// Tracks whether finalize was called. Generic over child type `C`. +pub struct FinalizeTracker { pub child_count: usize, pub finalized: Arc, + pub _child: PhantomData, } -impl TypedExecutor for FinalizeTracker { - async fn execute<'a>(&'a self, _payload: T, ctx: &'a TaskContext) -> Result<(), TaskError> { +impl FinalizeTracker { + pub fn new(child_count: usize, finalized: Arc) -> Self { + Self { + child_count, + finalized, + _child: PhantomData, + } + } +} + +impl TypedExecutor for FinalizeTracker +where + T: TypedTask, + C: TypedTask + Default + Send + Sync + 'static, +{ + async fn execute<'a>( + &'a self, + _payload: T, + ctx: DomainTaskContext<'a, T::Domain>, + ) -> Result<(), TaskError> { for i in 0..self.child_count { - let sub = TaskSubmission::new("child") + ctx.spawn_child_with(C::default()) .key(format!("ft-child-{i}")) - .priority(ctx.record().priority); - ctx.spawn_child(sub).await?; + .priority(ctx.record().priority) + .await + .map_err(|e| TaskError::new(e.to_string()))?; } Ok(()) } - async fn finalize<'a>(&'a self, _payload: T, _ctx: &'a TaskContext) -> Result<(), TaskError> { + async fn finalize<'a>( + &'a self, + _payload: T, + _ctx: DomainTaskContext<'a, T::Domain>, + ) -> Result<(), TaskError> { self.finalized.store(true, Ordering::SeqCst); Ok(()) } @@ -251,7 +317,11 @@ impl TypedExecutor for FinalizeTracker { pub struct AlwaysFailExecutor; impl TypedExecutor for AlwaysFailExecutor { - async fn execute<'a>(&'a self, _payload: T, _ctx: &'a TaskContext) -> Result<(), TaskError> { + async fn execute<'a>( + &'a self, + _payload: T, + _ctx: DomainTaskContext<'a, T::Domain>, + ) -> Result<(), TaskError> { Err(TaskError::new("permanent failure")) } } diff --git a/tests/integration/cross_module.rs b/tests/integration/cross_module.rs index ffc015d..f2b971d 100644 --- a/tests/integration/cross_module.rs +++ b/tests/integration/cross_module.rs @@ -7,8 +7,8 @@ use std::time::Duration; use serde::{Deserialize, Serialize}; use taskmill::{ - Domain, DomainHandle, DomainKey, Scheduler, SchedulerEvent, TaskContext, TaskError, TaskStore, - TaskSubmission, TypedExecutor, TypedTask, + Domain, DomainHandle, DomainKey, DomainTaskContext, Scheduler, SchedulerEvent, TaskError, + TaskStore, TypedExecutor, TypedTask, }; use tokio_util::sync::CancellationToken; @@ -29,10 +29,11 @@ impl TypedExecutor for CrossDomainSubmitter { async fn execute<'a>( &'a self, _payload: TriggerTask, - ctx: &'a TaskContext, + ctx: DomainTaskContext<'a, DomainA>, ) -> Result<(), TaskError> { let b: DomainHandle = ctx.domain::(); - b.submit_raw(TaskSubmission::new("task").key("cross-module-child")) + b.submit_with(DomainBTask) + .key("cross-module-child") .await .map_err(|e| TaskError::new(format!("{e}")))?; self.submitted.store(true, Ordering::SeqCst); @@ -60,7 +61,7 @@ async fn ctx_domain_submits_to_other_module_with_prefix_and_defaults() { async fn execute<'a>( &'a self, _payload: DomainBTask, - _ctx: &'a TaskContext, + _ctx: DomainTaskContext<'a, DomainB>, ) -> Result<(), TaskError> { self.0.store(true, Ordering::SeqCst); Ok(()) @@ -76,7 +77,8 @@ async fn ctx_domain_submits_to_other_module_with_prefix_and_defaults() { sched .domain::() - .submit_raw(TaskSubmission::new("trigger").key("t1")) + .submit_with(TriggerTask) + .key("t1") .await .unwrap(); @@ -108,11 +110,12 @@ impl TypedExecutor for SameDomainSubmitter { async fn execute<'a>( &'a self, _payload: MediaLeaderTask, - ctx: &'a TaskContext, + ctx: DomainTaskContext<'a, MediaDomain>, ) -> Result<(), TaskError> { let media: DomainHandle = ctx.domain::(); media - .submit_raw(TaskSubmission::new("follower").key("same-module-follower")) + .submit_with(MediaFollowerTask) + .key("same-module-follower") .await .map_err(|e| TaskError::new(format!("{e}")))?; self.submitted.store(true, Ordering::SeqCst); @@ -140,7 +143,7 @@ async fn ctx_domain_self_submit_applies_owning_module_defaults() { async fn execute<'a>( &'a self, _payload: MediaFollowerTask, - _ctx: &'a TaskContext, + _ctx: DomainTaskContext<'a, MediaDomain>, ) -> Result<(), TaskError> { self.0.store(true, Ordering::SeqCst); Ok(()) @@ -158,7 +161,8 @@ async fn ctx_domain_self_submit_applies_owning_module_defaults() { sched .domain::() - .submit_raw(TaskSubmission::new("leader").key("l1")) + .submit_with(MediaLeaderTask) + .key("l1") .await .unwrap(); @@ -195,7 +199,7 @@ async fn ctx_try_domain_returns_none_for_unknown_domain() { async fn execute<'a>( &'a self, _payload: ProbeTask, - ctx: &'a TaskContext, + ctx: DomainTaskContext<'a, TestDomain>, ) -> Result<(), TaskError> { let found = ctx.try_domain::().is_some(); *self.0.lock().unwrap() = Some(found); @@ -214,7 +218,8 @@ async fn ctx_try_domain_returns_none_for_unknown_domain() { sched .domain::() - .submit_raw(TaskSubmission::new("probe").key("p1")) + .submit_with(ProbeTask) + .key("p1") .await .unwrap(); @@ -245,9 +250,10 @@ async fn spawn_child_routes_through_owning_module() { async fn execute<'a>( &'a self, _payload: SpawnerTask, - ctx: &'a TaskContext, + ctx: DomainTaskContext<'a, TestDomain>, ) -> Result<(), TaskError> { - ctx.spawn_child(TaskSubmission::new("worker").key("spawned-child")) + ctx.spawn_child_with(WorkerTask) + .key("spawned-child") .await?; Ok(()) } @@ -258,7 +264,7 @@ async fn spawn_child_routes_through_owning_module() { async fn execute<'a>( &'a self, _payload: WorkerTask, - _ctx: &'a TaskContext, + _ctx: DomainTaskContext<'a, TestDomain>, ) -> Result<(), TaskError> { self.0.store(true, Ordering::SeqCst); Ok(()) @@ -280,7 +286,8 @@ async fn spawn_child_routes_through_owning_module() { sched .domain::() - .submit_raw(TaskSubmission::new("spawner").key("s1")) + .submit_with(SpawnerTask) + .key("s1") .await .unwrap(); @@ -310,12 +317,13 @@ impl TypedExecutor for CrossDomainParentExec { async fn execute<'a>( &'a self, _payload: MediaParentTask, - ctx: &'a TaskContext, + ctx: DomainTaskContext<'a, MediaDomain>, ) -> Result<(), TaskError> { let analytics: DomainHandle = ctx.domain::(); analytics - .submit_raw(TaskSubmission::new("work").key("cross-child")) - .parent(ctx.record().id) + .submit_with(AnalyticsWorkTask) + .key("cross-child") + .child_of(&ctx) .await .map_err(|e| TaskError::new(format!("{e}")))?; self.child_submitted.store(true, Ordering::SeqCst); @@ -345,7 +353,7 @@ async fn cross_domain_parent_child_lifecycle() { async fn execute<'a>( &'a self, _payload: AnalyticsWorkTask, - _ctx: &'a TaskContext, + _ctx: DomainTaskContext<'a, AnalyticsDomain>, ) -> Result<(), TaskError> { self.0.store(true, Ordering::SeqCst); Ok(()) @@ -364,7 +372,8 @@ async fn cross_domain_parent_child_lifecycle() { sched .domain::() - .submit_raw(TaskSubmission::new("parent").key("media-parent-1")) + .submit_with(MediaParentTask) + .key("media-parent-1") .await .unwrap(); @@ -421,11 +430,8 @@ async fn cross_domain_failure_cascade() { sched .domain::() - .submit_raw( - TaskSubmission::new("parent") - .key("media-parent-cascade") - .fail_fast(true), - ) + .submit_with(MediaParentTask) + .key("media-parent-cascade") .await .unwrap(); @@ -483,7 +489,7 @@ async fn scheduler_active_tasks_returns_tasks_from_all_modules() { async fn execute<'a>( &'a self, _payload: AlphaWorkTask, - ctx: &'a TaskContext, + ctx: DomainTaskContext<'a, AlphaDomain>, ) -> Result<(), TaskError> { self.0.wait().await; tokio::select! { @@ -500,7 +506,7 @@ async fn scheduler_active_tasks_returns_tasks_from_all_modules() { async fn execute<'a>( &'a self, _payload: BetaWorkTask, - ctx: &'a TaskContext, + ctx: DomainTaskContext<'a, BetaDomain>, ) -> Result<(), TaskError> { self.0.wait().await; tokio::select! { @@ -527,12 +533,14 @@ async fn scheduler_active_tasks_returns_tasks_from_all_modules() { sched .domain::() - .submit_raw(TaskSubmission::new("work").key("a1")) + .submit_with(AlphaWorkTask) + .key("a1") .await .unwrap(); sched .domain::() - .submit_raw(TaskSubmission::new("work").key("b1")) + .submit_with(BetaWorkTask) + .key("b1") .await .unwrap(); @@ -578,28 +586,25 @@ async fn cross_domain_cancel_by_tag_via_domain_handles() { // Tagged tasks — targets for cross-module cancel. let alpha_tagged = alpha - .submit_raw( - TaskSubmission::new("work") - .key("a-tagged") - .tag("job_id", "job-1"), - ) + .submit_with(AlphaWorkTask) + .key("a-tagged") + .tag("job_id", "job-1") .await .unwrap() .id() .unwrap(); let beta_tagged = beta - .submit_raw( - TaskSubmission::new("work") - .key("b-tagged") - .tag("job_id", "job-1"), - ) + .submit_with(BetaWorkTask) + .key("b-tagged") + .tag("job_id", "job-1") .await .unwrap() .id() .unwrap(); // Untagged task — must survive. let alpha_untagged = alpha - .submit_raw(TaskSubmission::new("work").key("a-untagged")) + .submit_with(AlphaWorkTask) + .key("a-untagged") .await .unwrap() .id() @@ -664,19 +669,18 @@ async fn parent_method_inherits_ttl_and_tags() { // Submit parent with a 60-second TTL and a custom tag. let parent_outcome = media - .submit_raw( - TaskSubmission::new("parent") - .key("ttl-parent") - .ttl(Duration::from_secs(60)) - .tag("job", "pipeline-42"), - ) + .submit_with(MediaParentTask) + .key("ttl-parent") + .ttl(Duration::from_secs(60)) + .tag("job", "pipeline-42") .await .unwrap(); let parent_id = parent_outcome.id().unwrap(); // Submit child with .parent() — no explicit TTL or tags on the child. let child_outcome = media - .submit_raw(TaskSubmission::new("child").key("ttl-child")) + .submit_with(MediaChildTask) + .key("ttl-child") .parent(parent_id) .await .unwrap(); @@ -700,11 +704,9 @@ async fn parent_method_inherits_ttl_and_tags() { // Child's own tags take precedence — a tag set directly on the child // should not be overwritten by the parent tag with the same key. let child2_outcome = media - .submit_raw( - TaskSubmission::new("child") - .key("ttl-child-2") - .tag("job", "child-override"), - ) + .submit_with(MediaChildTask) + .key("ttl-child-2") + .tag("job", "child-override") .parent(parent_id) .await .unwrap(); @@ -738,7 +740,8 @@ async fn event_header_module_field_populated_from_task_type_prefix() { sched .domain::() - .submit_raw(TaskSubmission::new("thumbnail").key("thumb-1")) + .submit_with(MediaThumbnailTask) + .key("thumb-1") .await .unwrap(); @@ -790,11 +793,13 @@ async fn module_receiver_events_match_module_field() { let sync_handle = sched.domain::(); for i in 0..2 { media - .submit_raw(TaskSubmission::new("thumbnail").key(format!("t{i}"))) + .submit_with(MediaThumbnailTask) + .key(format!("t{i}")) .await .unwrap(); sync_handle - .submit_raw(TaskSubmission::new("push").key(format!("p{i}"))) + .submit_with(PushTask) + .key(format!("p{i}")) .await .unwrap(); } diff --git a/tests/integration/dependencies.rs b/tests/integration/dependencies.rs index 01c863a..aa4e6e1 100644 --- a/tests/integration/dependencies.rs +++ b/tests/integration/dependencies.rs @@ -4,11 +4,24 @@ use std::sync::atomic::{AtomicUsize, Ordering}; use std::sync::Arc; use std::time::Duration; -use taskmill::{Domain, Scheduler, TaskStore, TaskSubmission}; +use taskmill::{Domain, Scheduler, TaskStore}; use tokio_util::sync::CancellationToken; use super::common::*; +// ═══════════════════════════════════════════════════════════════════ +// Helper: build a minimal scheduler for store-level dependency tests. +// ═══════════════════════════════════════════════════════════════════ + +async fn dep_scheduler() -> Scheduler { + Scheduler::builder() + .store(TaskStore::open_memory().await.unwrap()) + .domain(Domain::::new().task::(NoopExecutor)) + .build() + .await + .unwrap() +} + // ═══════════════════════════════════════════════════════════════════ // M. Task Dependencies // ═══════════════════════════════════════════════════════════════════ @@ -17,13 +30,26 @@ use super::common::*; async fn dep_basic_blocked_then_unblocked() { // Submit A, submit B depending on A → B is blocked. // Complete A → B becomes pending. - let store = TaskStore::open_memory().await.unwrap(); + let sched = dep_scheduler().await; + let handle = sched.domain::(); + let store = sched.store(); - let sub_a = TaskSubmission::new("test").key("dep-a"); - let id_a = store.submit(&sub_a).await.unwrap().id().unwrap(); + let id_a = handle + .submit_with(TestTask) + .key("dep-a") + .await + .unwrap() + .id() + .unwrap(); - let sub_b = TaskSubmission::new("test").key("dep-b").depends_on(id_a); - let id_b = store.submit(&sub_b).await.unwrap().id().unwrap(); + let id_b = handle + .submit_with(TestTask) + .key("dep-b") + .depends_on(id_a) + .await + .unwrap() + .id() + .unwrap(); let b = store.task_by_id(id_b).await.unwrap().unwrap(); assert_eq!(b.status, taskmill::TaskStatus::Blocked); @@ -48,13 +74,26 @@ async fn dep_basic_blocked_then_unblocked() { #[tokio::test] async fn dep_fail_cancels_dependent() { // Submit A, submit B depending on A. Fail A → B moves to history as DependencyFailed. - let store = TaskStore::open_memory().await.unwrap(); + let sched = dep_scheduler().await; + let handle = sched.domain::(); + let store = sched.store(); - let sub_a = TaskSubmission::new("test").key("fail-a"); - let id_a = store.submit(&sub_a).await.unwrap().id().unwrap(); + let id_a = handle + .submit_with(TestTask) + .key("fail-a") + .await + .unwrap() + .id() + .unwrap(); - let sub_b = TaskSubmission::new("test").key("fail-b").depends_on(id_a); - let id_b = store.submit(&sub_b).await.unwrap().id().unwrap(); + let id_b = handle + .submit_with(TestTask) + .key("fail-b") + .depends_on(id_a) + .await + .unwrap() + .id() + .unwrap(); // Fail A permanently. let a = store.pop_next().await.unwrap().unwrap(); @@ -84,18 +123,34 @@ async fn dep_fail_cancels_dependent() { #[tokio::test] async fn dep_fan_in() { // C depends on both A and B. Complete A → C still blocked. Complete B → C pending. - let store = TaskStore::open_memory().await.unwrap(); + let sched = dep_scheduler().await; + let handle = sched.domain::(); + let store = sched.store(); - let sub_a = TaskSubmission::new("test").key("fi-a"); - let id_a = store.submit(&sub_a).await.unwrap().id().unwrap(); + let id_a = handle + .submit_with(TestTask) + .key("fi-a") + .await + .unwrap() + .id() + .unwrap(); - let sub_b = TaskSubmission::new("test").key("fi-b"); - let id_b = store.submit(&sub_b).await.unwrap().id().unwrap(); + let id_b = handle + .submit_with(TestTask) + .key("fi-b") + .await + .unwrap() + .id() + .unwrap(); - let sub_c = TaskSubmission::new("test") + let id_c = handle + .submit_with(TestTask) .key("fi-c") - .depends_on_all([id_a, id_b]); - let id_c = store.submit(&sub_c).await.unwrap().id().unwrap(); + .depends_on_all([id_a, id_b]) + .await + .unwrap() + .id() + .unwrap(); let c = store.task_by_id(id_c).await.unwrap().unwrap(); assert_eq!(c.status, taskmill::TaskStatus::Blocked); @@ -128,16 +183,35 @@ async fn dep_fan_in() { #[tokio::test] async fn dep_fan_out() { // B and C both depend on A. Complete A → both become pending. - let store = TaskStore::open_memory().await.unwrap(); + let sched = dep_scheduler().await; + let handle = sched.domain::(); + let store = sched.store(); - let sub_a = TaskSubmission::new("test").key("fo-a"); - let id_a = store.submit(&sub_a).await.unwrap().id().unwrap(); + let id_a = handle + .submit_with(TestTask) + .key("fo-a") + .await + .unwrap() + .id() + .unwrap(); - let sub_b = TaskSubmission::new("test").key("fo-b").depends_on(id_a); - let id_b = store.submit(&sub_b).await.unwrap().id().unwrap(); + let id_b = handle + .submit_with(TestTask) + .key("fo-b") + .depends_on(id_a) + .await + .unwrap() + .id() + .unwrap(); - let sub_c = TaskSubmission::new("test").key("fo-c").depends_on(id_a); - let id_c = store.submit(&sub_c).await.unwrap().id().unwrap(); + let id_c = handle + .submit_with(TestTask) + .key("fo-c") + .depends_on(id_a) + .await + .unwrap() + .id() + .unwrap(); // Complete A. let a = store.pop_next().await.unwrap().unwrap(); @@ -155,26 +229,48 @@ async fn dep_fan_out() { #[tokio::test] async fn dep_cycle_detection_direct() { // A depends on B, B depends on A → CyclicDependency error. - let store = TaskStore::open_memory().await.unwrap(); + let sched = dep_scheduler().await; + let handle = sched.domain::(); - let sub_a = TaskSubmission::new("test").key("cyc-a"); - let id_a = store.submit(&sub_a).await.unwrap().id().unwrap(); + let id_a = handle + .submit_with(TestTask) + .key("cyc-a") + .await + .unwrap() + .id() + .unwrap(); - let sub_b = TaskSubmission::new("test").key("cyc-b").depends_on(id_a); - let id_b = store.submit(&sub_b).await.unwrap().id().unwrap(); + let id_b = handle + .submit_with(TestTask) + .key("cyc-b") + .depends_on(id_a) + .await + .unwrap() + .id() + .unwrap(); // Try to make A depend on B (cycle). // We need to submit a new task that depends on B and somehow forms a cycle. // Actually, since A is already inserted, we can't make it depend on B. // The cycle detection works at submission time. Let's test A→B→C→A. - let sub_c = TaskSubmission::new("test").key("cyc-c").depends_on(id_b); - let _id_c = store.submit(&sub_c).await.unwrap().id().unwrap(); + let _id_c = handle + .submit_with(TestTask) + .key("cyc-c") + .depends_on(id_b) + .await + .unwrap() + .id() + .unwrap(); // Now try to submit D that depends on C and A, where A already has B depending on it. // That's not a cycle. Let's test an actual self-dependency. - let sub_self = TaskSubmission::new("test").key("cyc-self").depends_on(id_a); // This shouldn't cause issues because cyc-self doesn't have anyone depending on it. - let _ = store.submit(&sub_self).await.unwrap(); + let _ = handle + .submit_with(TestTask) + .key("cyc-self") + .depends_on(id_a) + .await + .unwrap(); // The true cycle test: submit a task that would create A→B→...→A. // This is tricky because we can only declare deps at submission time. @@ -199,10 +295,17 @@ async fn dep_cycle_detection_direct() { #[tokio::test] async fn dep_already_completed() { // Depend on already-completed task → task starts as pending immediately. - let store = TaskStore::open_memory().await.unwrap(); + let sched = dep_scheduler().await; + let handle = sched.domain::(); + let store = sched.store(); - let sub_a = TaskSubmission::new("test").key("done-a"); - let id_a = store.submit(&sub_a).await.unwrap().id().unwrap(); + let id_a = handle + .submit_with(TestTask) + .key("done-a") + .await + .unwrap() + .id() + .unwrap(); // Complete A. let a = store.pop_next().await.unwrap().unwrap(); @@ -212,8 +315,14 @@ async fn dep_already_completed() { .unwrap(); // Submit B depending on A (already completed). - let sub_b = TaskSubmission::new("test").key("done-b").depends_on(id_a); - let id_b = store.submit(&sub_b).await.unwrap().id().unwrap(); + let id_b = handle + .submit_with(TestTask) + .key("done-b") + .depends_on(id_a) + .await + .unwrap() + .id() + .unwrap(); let b = store.task_by_id(id_b).await.unwrap().unwrap(); assert_eq!(b.status, taskmill::TaskStatus::Pending); @@ -222,10 +331,17 @@ async fn dep_already_completed() { #[tokio::test] async fn dep_already_failed() { // Depend on already-failed task → DependencyFailed error at submission. - let store = TaskStore::open_memory().await.unwrap(); + let sched = dep_scheduler().await; + let handle = sched.domain::(); + let store = sched.store(); - let sub_a = TaskSubmission::new("test").key("af-a"); - let id_a = store.submit(&sub_a).await.unwrap().id().unwrap(); + let id_a = handle + .submit_with(TestTask) + .key("af-a") + .await + .unwrap() + .id() + .unwrap(); let a = store.pop_next().await.unwrap().unwrap(); store @@ -240,18 +356,27 @@ async fn dep_already_failed() { .await .unwrap(); - let sub_b = TaskSubmission::new("test").key("af-b").depends_on(id_a); - let err = store.submit(&sub_b).await.unwrap_err(); + let err = handle + .submit_with(TestTask) + .key("af-b") + .depends_on(id_a) + .await + .unwrap_err(); assert!(matches!(err, taskmill::StoreError::DependencyFailed(_))); } #[tokio::test] async fn dep_nonexistent() { // Depend on nonexistent task → InvalidDependency error. - let store = TaskStore::open_memory().await.unwrap(); + let sched = dep_scheduler().await; + let handle = sched.domain::(); - let sub = TaskSubmission::new("test").key("ne").depends_on(99999); - let err = store.submit(&sub).await.unwrap_err(); + let err = handle + .submit_with(TestTask) + .key("ne") + .depends_on(99999) + .await + .unwrap_err(); assert!(matches!( err, taskmill::StoreError::InvalidDependency(99999) @@ -261,13 +386,26 @@ async fn dep_nonexistent() { #[tokio::test] async fn dep_cancel_cascades() { // Cancel a task with dependents → dependents cascade-fail. - let store = TaskStore::open_memory().await.unwrap(); + let sched = dep_scheduler().await; + let handle = sched.domain::(); + let store = sched.store(); - let sub_a = TaskSubmission::new("test").key("cc-a"); - let id_a = store.submit(&sub_a).await.unwrap().id().unwrap(); + let id_a = handle + .submit_with(TestTask) + .key("cc-a") + .await + .unwrap() + .id() + .unwrap(); - let sub_b = TaskSubmission::new("test").key("cc-b").depends_on(id_a); - let id_b = store.submit(&sub_b).await.unwrap().id().unwrap(); + let id_b = handle + .submit_with(TestTask) + .key("cc-b") + .depends_on(id_a) + .await + .unwrap() + .id() + .unwrap(); store.cancel_to_history(id_a).await.unwrap(); @@ -285,16 +423,27 @@ async fn dep_cancel_cascades() { #[tokio::test] async fn dep_ignore_policy_unblocks() { // DependencyFailurePolicy::Ignore → dependent unblocked despite dep failure. - let store = TaskStore::open_memory().await.unwrap(); + let sched = dep_scheduler().await; + let handle = sched.domain::(); + let store = sched.store(); - let sub_a = TaskSubmission::new("test").key("ig-a"); - let id_a = store.submit(&sub_a).await.unwrap().id().unwrap(); + let id_a = handle + .submit_with(TestTask) + .key("ig-a") + .await + .unwrap() + .id() + .unwrap(); - let sub_b = TaskSubmission::new("test") + let id_b = handle + .submit_with(TestTask) .key("ig-b") .depends_on(id_a) - .on_dependency_failure(taskmill::DependencyFailurePolicy::Ignore); - let id_b = store.submit(&sub_b).await.unwrap().id().unwrap(); + .on_dependency_failure(taskmill::DependencyFailurePolicy::Ignore) + .await + .unwrap() + .id() + .unwrap(); let b = store.task_by_id(id_b).await.unwrap().unwrap(); assert_eq!(b.status, taskmill::TaskStatus::Blocked); @@ -324,18 +473,34 @@ async fn dep_ignore_policy_unblocks() { #[tokio::test] async fn dep_query_methods() { // Verify task_dependencies() and task_dependents() return correct edges. - let store = TaskStore::open_memory().await.unwrap(); + let sched = dep_scheduler().await; + let handle = sched.domain::(); + let store = sched.store(); - let sub_a = TaskSubmission::new("test").key("qm-a"); - let id_a = store.submit(&sub_a).await.unwrap().id().unwrap(); + let id_a = handle + .submit_with(TestTask) + .key("qm-a") + .await + .unwrap() + .id() + .unwrap(); - let sub_b = TaskSubmission::new("test").key("qm-b"); - let id_b = store.submit(&sub_b).await.unwrap().id().unwrap(); + let id_b = handle + .submit_with(TestTask) + .key("qm-b") + .await + .unwrap() + .id() + .unwrap(); - let sub_c = TaskSubmission::new("test") + let id_c = handle + .submit_with(TestTask) .key("qm-c") - .depends_on_all([id_a, id_b]); - let id_c = store.submit(&sub_c).await.unwrap().id().unwrap(); + .depends_on_all([id_a, id_b]) + .await + .unwrap() + .id() + .unwrap(); let deps = store.task_dependencies(id_c).await.unwrap(); assert_eq!(deps.len(), 2); @@ -356,21 +521,44 @@ async fn dep_query_methods() { #[tokio::test] async fn dep_diamond_chain() { // Diamond: A→B, A→C, B→D, C→D. Complete A, then B and C, then D. - let store = TaskStore::open_memory().await.unwrap(); + let sched = dep_scheduler().await; + let handle = sched.domain::(); + let store = sched.store(); - let sub_a = TaskSubmission::new("test").key("d-a"); - let id_a = store.submit(&sub_a).await.unwrap().id().unwrap(); + let id_a = handle + .submit_with(TestTask) + .key("d-a") + .await + .unwrap() + .id() + .unwrap(); - let sub_b = TaskSubmission::new("test").key("d-b").depends_on(id_a); - let id_b = store.submit(&sub_b).await.unwrap().id().unwrap(); + let id_b = handle + .submit_with(TestTask) + .key("d-b") + .depends_on(id_a) + .await + .unwrap() + .id() + .unwrap(); - let sub_c = TaskSubmission::new("test").key("d-c").depends_on(id_a); - let id_c = store.submit(&sub_c).await.unwrap().id().unwrap(); + let id_c = handle + .submit_with(TestTask) + .key("d-c") + .depends_on(id_a) + .await + .unwrap() + .id() + .unwrap(); - let sub_d = TaskSubmission::new("test") + let id_d = handle + .submit_with(TestTask) .key("d-d") - .depends_on_all([id_b, id_c]); - let id_d = store.submit(&sub_d).await.unwrap().id().unwrap(); + .depends_on_all([id_b, id_c]) + .await + .unwrap() + .id() + .unwrap(); // All B, C, D should be blocked. assert_eq!( @@ -437,18 +625,13 @@ async fn dep_blocked_count_in_snapshot() { let handle = sched.domain::(); - let outcome_a = handle - .submit_raw(TaskSubmission::new("test::test").key("snap-a")) - .await - .unwrap(); + let outcome_a = handle.submit_with(TestTask).key("snap-a").await.unwrap(); let id_a = outcome_a.id().unwrap(); handle - .submit_raw( - TaskSubmission::new("test::test") - .key("snap-b") - .depends_on(id_a), - ) + .submit_with(TestTask) + .key("snap-b") + .depends_on(id_a) .await .unwrap(); @@ -488,19 +671,24 @@ async fn dep_full_chain_with_scheduler() { }); let outcome_a = domain_handle - .submit_raw(TaskSubmission::new("step").key("chain-a")) + .submit_with(StepTask) + .key("chain-a") .await .unwrap(); let id_a = outcome_a.id().unwrap(); let outcome_b = domain_handle - .submit_raw(TaskSubmission::new("step").key("chain-b").depends_on(id_a)) + .submit_with(StepTask) + .key("chain-b") + .depends_on(id_a) .await .unwrap(); let id_b = outcome_b.id().unwrap(); let outcome_c = domain_handle - .submit_raw(TaskSubmission::new("step").key("chain-c").depends_on(id_b)) + .submit_with(StepTask) + .key("chain-c") + .depends_on(id_b) .await .unwrap(); let _id_c = outcome_c.id().unwrap(); @@ -525,13 +713,26 @@ async fn dep_full_chain_with_scheduler() { #[tokio::test] async fn dep_blocked_tasks_survive_across_store_reopen() { // Blocked tasks and their dep edges are persisted in SQLite. - let store = TaskStore::open_memory().await.unwrap(); + let sched = dep_scheduler().await; + let handle = sched.domain::(); + let store = sched.store(); - let sub_a = TaskSubmission::new("test").key("rec-a"); - let id_a = store.submit(&sub_a).await.unwrap().id().unwrap(); + let id_a = handle + .submit_with(TestTask) + .key("rec-a") + .await + .unwrap() + .id() + .unwrap(); - let sub_b = TaskSubmission::new("test").key("rec-b").depends_on(id_a); - let id_b = store.submit(&sub_b).await.unwrap().id().unwrap(); + let id_b = handle + .submit_with(TestTask) + .key("rec-b") + .depends_on(id_a) + .await + .unwrap() + .id() + .unwrap(); // B should be blocked with dep edges persisted. let b = store.task_by_id(id_b).await.unwrap().unwrap(); diff --git a/tests/integration/module_features.rs b/tests/integration/module_features.rs index d607718..4314aa6 100644 --- a/tests/integration/module_features.rs +++ b/tests/integration/module_features.rs @@ -6,7 +6,7 @@ use std::sync::Arc; use std::time::Duration; use taskmill::{ - Domain, Priority, Scheduler, SchedulerEvent, TaskContext, TaskError, TaskStore, TaskSubmission, + Domain, DomainTaskContext, Priority, Scheduler, SchedulerEvent, TaskError, TaskStore, TaskTypeConfig, TypedExecutor, TypedTask, }; use tokio_util::sync::CancellationToken; @@ -133,7 +133,8 @@ async fn module_cap_limits_concurrency_to_2() { let media = sched.domain::(); for i in 0..5 { media - .submit_raw(TaskSubmission::new("work").key(format!("t{i}"))) + .submit_with(MediaWorkTask) + .key(format!("t{i}")) .await .unwrap(); } @@ -194,11 +195,9 @@ async fn module_cap_and_group_cap_are_independent() { // Submit 6 tasks all in the "gpu" group — group cap is the binding constraint. for i in 0..6 { media - .submit_raw( - TaskSubmission::new("work") - .key(format!("t{i}")) - .group("gpu"), - ) + .submit_with(MediaWorkTask) + .key(format!("t{i}")) + .group("gpu") .await .unwrap(); } @@ -256,7 +255,8 @@ async fn ungrouped_task_respects_module_cap() { let media = sched.domain::(); for i in 0..7 { media - .submit_raw(TaskSubmission::new("work").key(format!("t{i}"))) + .submit_with(MediaWorkTask) + .key(format!("t{i}")) .await .unwrap(); } @@ -325,10 +325,12 @@ async fn global_cap_is_hard_ceiling_over_module_caps() { let sync = sched.domain::(); for i in 0..5 { media - .submit_raw(TaskSubmission::new("work").key(format!("m{i}"))) + .submit_with(MediaWorkTask) + .key(format!("m{i}")) .await .unwrap(); - sync.submit_raw(TaskSubmission::new("work").key(format!("s{i}"))) + sync.submit_with(SyncWorkTask) + .key(format!("s{i}")) .await .unwrap(); } @@ -395,7 +397,8 @@ async fn set_max_concurrency_changes_dispatch_behavior() { for i in 0..6 { media - .submit_raw(TaskSubmission::new("work").key(format!("t{i}"))) + .submit_with(MediaWorkTask) + .key(format!("t{i}")) .await .unwrap(); } @@ -443,7 +446,11 @@ async fn module_state_is_scoped_to_module() { no_b: Arc, } impl TypedExecutor for CheckerExec { - async fn execute<'a>(&'a self, _payload: T, ctx: &'a TaskContext) -> Result<(), TaskError> { + async fn execute<'a>( + &'a self, + _payload: T, + ctx: DomainTaskContext<'a, T::Domain>, + ) -> Result<(), TaskError> { self.saw_a .store(ctx.state::().is_some(), Ordering::SeqCst); if ctx.state::().is_some() { @@ -475,7 +482,8 @@ async fn module_state_is_scoped_to_module() { sched .domain::() - .submit_raw(TaskSubmission::new("task").key("t1")) + .submit_with(DomainATask) + .key("t1") .await .unwrap(); @@ -518,7 +526,11 @@ async fn global_state_accessible_from_all_modules() { struct GlobalChecker(Arc); impl TypedExecutor for GlobalChecker { - async fn execute<'a>(&'a self, _payload: T, ctx: &'a TaskContext) -> Result<(), TaskError> { + async fn execute<'a>( + &'a self, + _payload: T, + ctx: DomainTaskContext<'a, T::Domain>, + ) -> Result<(), TaskError> { self.0 .store(ctx.state::().is_some(), Ordering::SeqCst); Ok(()) @@ -537,12 +549,14 @@ async fn global_state_accessible_from_all_modules() { sched .domain::() - .submit_raw(TaskSubmission::new("task").key("ta")) + .submit_with(DomainATask) + .key("ta") .await .unwrap(); sched .domain::() - .submit_raw(TaskSubmission::new("task").key("tb")) + .submit_with(DomainBTask) + .key("tb") .await .unwrap(); @@ -583,7 +597,11 @@ async fn module_state_shadows_global_state() { struct ValueCapture(Arc>); impl TypedExecutor for ValueCapture { - async fn execute<'a>(&'a self, _payload: T, ctx: &'a TaskContext) -> Result<(), TaskError> { + async fn execute<'a>( + &'a self, + _payload: T, + ctx: DomainTaskContext<'a, T::Domain>, + ) -> Result<(), TaskError> { if let Some(cfg) = ctx.state::() { *self.0.lock().unwrap() = cfg.0.clone(); } @@ -607,12 +625,14 @@ async fn module_state_shadows_global_state() { sched .domain::() - .submit_raw(TaskSubmission::new("task").key("ta")) + .submit_with(DomainATask) + .key("ta") .await .unwrap(); sched .domain::() - .submit_raw(TaskSubmission::new("task").key("tb")) + .submit_with(DomainBTask) + .key("tb") .await .unwrap(); diff --git a/tests/integration/retry_policy.rs b/tests/integration/retry_policy.rs index 54bf142..75a17b8 100644 --- a/tests/integration/retry_policy.rs +++ b/tests/integration/retry_policy.rs @@ -4,8 +4,8 @@ use std::time::Duration; use serde::{Deserialize, Serialize}; use taskmill::{ - BackoffStrategy, Domain, RetryPolicy, Scheduler, SchedulerEvent, TaskContext, TaskError, - TaskStore, TaskSubmission, TaskTypeConfig, TypedExecutor, TypedTask, + BackoffStrategy, Domain, DomainTaskContext, RetryPolicy, Scheduler, SchedulerEvent, TaskError, + TaskStore, TaskTypeConfig, TypedExecutor, TypedTask, }; use tokio_util::sync::CancellationToken; @@ -89,7 +89,11 @@ define_task!(RetryOverrideTask, TestDomain, "retry-override"); struct AlwaysRetryableTypedExec; impl TypedExecutor for AlwaysRetryableTypedExec { - async fn execute<'a>(&'a self, _payload: T, _ctx: &'a TaskContext) -> Result<(), TaskError> { + async fn execute<'a>( + &'a self, + _payload: T, + _ctx: DomainTaskContext<'a, T::Domain>, + ) -> Result<(), TaskError> { Err(TaskError::retryable("transient")) } } @@ -101,7 +105,7 @@ impl TypedExecutor for AlwaysRetryableExecutor { async fn execute<'a>( &'a self, _payload: LegacyTask, - _ctx: &'a TaskContext, + _ctx: DomainTaskContext<'a, TestDomain>, ) -> Result<(), TaskError> { Err(TaskError::retryable("transient")) } @@ -114,7 +118,7 @@ impl TypedExecutor for RetryAfterExecutor { async fn execute<'a>( &'a self, _payload: RetryOverrideTask, - _ctx: &'a TaskContext, + _ctx: DomainTaskContext<'a, TestDomain>, ) -> Result<(), TaskError> { Err(TaskError::retryable("rate limited").retry_after(self.0)) } @@ -349,8 +353,10 @@ async fn failed_event_includes_executor_retry_after_override() { async move { s.run(t).await } }); - sched - .submit(&TaskSubmission::new("test::retry-override").key("ro1")) + let test_handle = sched.domain::(); + test_handle + .submit_with(RetryOverrideTask) + .key("ro1") .await .unwrap(); @@ -403,8 +409,10 @@ async fn null_max_retries_uses_global_default() { async move { s.run(t).await } }); - sched - .submit(&TaskSubmission::new("test::legacy").key("leg1")) + let test_handle = sched.domain::(); + test_handle + .submit_with(LegacyTask) + .key("leg1") .await .unwrap(); diff --git a/tests/integration/scheduler_core.rs b/tests/integration/scheduler_core.rs index e005062..ee5c7e9 100644 --- a/tests/integration/scheduler_core.rs +++ b/tests/integration/scheduler_core.rs @@ -535,11 +535,7 @@ async fn fail_fast_cancels_siblings_on_child_failure() { .store(TaskStore::open_memory().await.unwrap()) .domain( Domain::::new() - .task::(ChildSpawnerExecutor { - child_type: "child", - count: 3, - fail_fast: true, - }) + .task::(ChildSpawnerExecutor::::new(3)) .task::(AlwaysFailExecutor), ) .max_concurrency(4) @@ -593,10 +589,7 @@ async fn non_fail_fast_waits_for_all_children() { .store(TaskStore::open_memory().await.unwrap()) .domain( Domain::::new() - .task::(FinalizeTracker { - child_count: 2, - finalized: finalized.clone(), - }) + .task::(FinalizeTracker::::new(2, finalized.clone())) .task::(NoopExecutor), ) .max_concurrency(4) diff --git a/tests/integration/typed_events.rs b/tests/integration/typed_events.rs index e749f26..ff3cc2b 100644 --- a/tests/integration/typed_events.rs +++ b/tests/integration/typed_events.rs @@ -7,8 +7,8 @@ use std::time::Duration; use serde::{Deserialize, Serialize}; use taskmill::{ - Domain, DomainHandle, Scheduler, TaskContext, TaskError, TaskEvent, TaskStore, TypedExecutor, - TypedTask, + Domain, DomainHandle, DomainTaskContext, Scheduler, TaskError, TaskEvent, TaskStore, + TypedExecutor, TypedTask, }; use tokio_util::sync::CancellationToken; @@ -32,7 +32,7 @@ impl TypedExecutor for ThumbnailExec { async fn execute<'a>( &'a self, _payload: Thumbnail, - _ctx: &'a TaskContext, + _ctx: DomainTaskContext<'a, MediaDomain>, ) -> Result<(), TaskError> { Ok(()) } @@ -54,7 +54,7 @@ impl TypedExecutor for AlwaysFailTypedExec { async fn execute<'a>( &'a self, _payload: FailingTask, - _ctx: &'a TaskContext, + _ctx: DomainTaskContext<'a, MediaDomain>, ) -> Result<(), TaskError> { Err(TaskError::new("permanent failure")) } @@ -247,7 +247,7 @@ impl TypedExecutor for CrossDomainTypedExec { async fn execute<'a>( &'a self, thumb: Thumbnail, - ctx: &'a TaskContext, + ctx: DomainTaskContext<'a, MediaDomain>, ) -> Result<(), TaskError> { let sync: DomainHandle = ctx.domain::(); sync.submit(UploadTask { @@ -268,7 +268,7 @@ impl TypedExecutor for UploadExec { async fn execute<'a>( &'a self, _payload: UploadTask, - _ctx: &'a TaskContext, + _ctx: DomainTaskContext<'a, SyncDomain>, ) -> Result<(), TaskError> { self.ran.store(true, Ordering::SeqCst); Ok(()) @@ -336,7 +336,7 @@ impl TypedExecutor for DomainStateExec { async fn execute<'a>( &'a self, _payload: Thumbnail, - ctx: &'a TaskContext, + ctx: DomainTaskContext<'a, MediaDomain>, ) -> Result<(), TaskError> { let cfg = ctx .domain_state::()