1use std::fmt;
2use std::sync::Arc;
3use std::time::Duration;
4
5use serde::{Deserialize, Serialize};
6use time::OffsetDateTime;
7
8use crate::events::{ThreadId, TurnId};
9use crate::extension::TaskExecutorId;
10use crate::processes::ProcessRegistrySink;
11use crate::remote_runner::{RemoteRunnerSession, RunnerDestination};
12use crate::{ToolSchemaPolicy, normalize_tool_schema};
13
14pub type TaskId = String;
15
16#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
17pub struct TaskSpec {
18 pub kind: String,
19 pub description: String,
20 pub input_schema: serde_json::Value,
21 #[serde(default, skip_serializing_if = "Option::is_none")]
22 pub default_timeout_seconds: Option<u64>,
23 #[serde(default)]
24 pub metadata: serde_json::Value,
25}
26
27impl TaskSpec {
28 pub fn normalized_for_model(&self, policy: ToolSchemaPolicy) -> Self {
29 let mut spec = self.clone();
30 spec.input_schema = normalize_tool_schema(&spec.kind, &spec.input_schema, policy).schema;
31 spec
32 }
33}
34
35#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
36#[serde(rename_all = "snake_case")]
37pub enum TaskState {
38 Queued,
39 Running,
40 Completed,
41 Failed,
42 Cancelled,
43}
44
45#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
46pub struct TaskHandle {
47 pub task_id: TaskId,
48 pub executor_id: TaskExecutorId,
49 pub spec: TaskSpec,
50 pub state: TaskState,
51 #[serde(with = "time::serde::rfc3339")]
52 pub created_at: OffsetDateTime,
53 #[serde(default, with = "time::serde::rfc3339::option")]
54 pub started_at: Option<OffsetDateTime>,
55 #[serde(default, with = "time::serde::rfc3339::option")]
56 pub finished_at: Option<OffsetDateTime>,
57}
58
59#[derive(Clone)]
60pub struct TaskExecutionContext {
61 pub task_id: TaskId,
62 pub thread_id: Option<ThreadId>,
63 pub turn_id: Option<TurnId>,
64 pub workspace_root: Option<String>,
65 pub runner_destination: Option<RunnerDestination>,
66 pub runner_session: Option<Arc<dyn RemoteRunnerSession>>,
67 pub deadline: Option<OffsetDateTime>,
68 pub process_grace_timeout: Duration,
70 pub process_kill_timeout: Duration,
72 pub metadata: serde_json::Value,
73 pub process_registry: Option<Arc<dyn ProcessRegistrySink>>,
74 pub output: TaskOutputSink,
75}
76
77#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
78pub struct TaskExecutionResult {
79 #[serde(default, skip_serializing_if = "Option::is_none")]
80 pub exit_code: Option<i32>,
81 #[serde(default)]
82 pub payload: serde_json::Value,
83}
84
85impl TaskExecutionResult {
86 pub fn success(payload: serde_json::Value) -> Self {
87 Self {
88 exit_code: None,
89 payload,
90 }
91 }
92}
93
94impl fmt::Debug for TaskExecutionContext {
95 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
96 f.debug_struct("TaskExecutionContext")
97 .field("task_id", &self.task_id)
98 .field("thread_id", &self.thread_id)
99 .field("turn_id", &self.turn_id)
100 .field("workspace_root", &self.workspace_root)
101 .field("runner_destination", &self.runner_destination)
102 .field(
103 "runner_session",
104 &self.runner_session.as_ref().map(|session| session.state()),
105 )
106 .field("deadline", &self.deadline)
107 .field("process_grace_timeout", &self.process_grace_timeout)
108 .field("process_kill_timeout", &self.process_kill_timeout)
109 .field("metadata", &self.metadata)
110 .field(
111 "process_registry",
112 &self.process_registry.as_ref().map(|_| "<process-registry>"),
113 )
114 .finish_non_exhaustive()
115 }
116}
117
118#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
119#[serde(rename_all = "snake_case")]
120pub enum TaskOutputStream {
121 Stdout,
122 Stderr,
123 Log,
124}
125
126#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
127pub struct TaskStarted {
128 pub task_id: TaskId,
129 pub executor_id: TaskExecutorId,
130 pub task_kind: String,
131 #[serde(default)]
132 pub queue_depth: usize,
133 #[serde(default, skip_serializing_if = "Option::is_none")]
134 pub thread_id: Option<ThreadId>,
135 #[serde(default, skip_serializing_if = "Option::is_none")]
136 pub turn_id: Option<TurnId>,
137 #[serde(with = "time::serde::rfc3339")]
138 pub timestamp: OffsetDateTime,
139}
140
141#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
142pub struct TaskOutput {
143 pub task_id: TaskId,
144 pub stream: TaskOutputStream,
145 pub chunk: String,
146 #[serde(default)]
147 pub dropped_bytes: u64,
148 #[serde(default, skip_serializing_if = "Option::is_none")]
149 pub thread_id: Option<ThreadId>,
150 #[serde(default, skip_serializing_if = "Option::is_none")]
151 pub turn_id: Option<TurnId>,
152 #[serde(with = "time::serde::rfc3339")]
153 pub timestamp: OffsetDateTime,
154}
155
156#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
157pub struct TaskCompleted {
158 pub task_id: TaskId,
159 #[serde(default, skip_serializing_if = "Option::is_none")]
160 pub exit_code: Option<i32>,
161 #[serde(default)]
162 pub payload: serde_json::Value,
163 #[serde(default, skip_serializing_if = "Option::is_none")]
164 pub thread_id: Option<ThreadId>,
165 #[serde(default, skip_serializing_if = "Option::is_none")]
166 pub turn_id: Option<TurnId>,
167 #[serde(with = "time::serde::rfc3339")]
168 pub timestamp: OffsetDateTime,
169}
170
171#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
172pub struct TaskFailed {
173 pub task_id: TaskId,
174 pub error: String,
175 #[serde(default, skip_serializing_if = "Option::is_none")]
176 pub error_kind: Option<String>,
177 #[serde(default, skip_serializing_if = "Option::is_none")]
178 pub partial_result: Option<String>,
179 #[serde(default, skip_serializing_if = "Option::is_none")]
180 pub thread_id: Option<ThreadId>,
181 #[serde(default, skip_serializing_if = "Option::is_none")]
182 pub turn_id: Option<TurnId>,
183 #[serde(with = "time::serde::rfc3339")]
184 pub timestamp: OffsetDateTime,
185}
186
187#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
188pub struct TaskCancelled {
189 pub task_id: TaskId,
190 #[serde(default, skip_serializing_if = "Option::is_none")]
191 pub reason: Option<String>,
192 #[serde(default, skip_serializing_if = "Option::is_none")]
193 pub thread_id: Option<ThreadId>,
194 #[serde(default, skip_serializing_if = "Option::is_none")]
195 pub turn_id: Option<TurnId>,
196 #[serde(with = "time::serde::rfc3339")]
197 pub timestamp: OffsetDateTime,
198}
199
200#[derive(Clone)]
201pub struct TaskOutputSink {
202 writer: Arc<dyn TaskOutputWriter>,
203}
204
205impl Default for TaskOutputSink {
206 fn default() -> Self {
207 Self {
208 writer: Arc::new(NoopTaskOutputWriter),
209 }
210 }
211}
212
213impl fmt::Debug for TaskOutputSink {
214 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
215 f.debug_struct("TaskOutputSink").finish_non_exhaustive()
216 }
217}
218
219impl TaskOutputSink {
220 pub fn new(writer: Arc<dyn TaskOutputWriter>) -> Self {
221 Self { writer }
222 }
223
224 pub async fn write(
225 &self,
226 stream: TaskOutputStream,
227 chunk: impl Into<String>,
228 ) -> anyhow::Result<()> {
229 self.writer.write(stream, chunk.into()).await
230 }
231}
232
233#[async_trait::async_trait]
234pub trait TaskOutputWriter: Send + Sync + 'static {
235 async fn write(&self, stream: TaskOutputStream, chunk: String) -> anyhow::Result<()>;
236}
237
238struct NoopTaskOutputWriter;
239
240#[async_trait::async_trait]
241impl TaskOutputWriter for NoopTaskOutputWriter {
242 async fn write(&self, _stream: TaskOutputStream, _chunk: String) -> anyhow::Result<()> {
243 Ok(())
244 }
245}
246
247#[async_trait::async_trait]
248pub trait TaskExecutor: Send + Sync + 'static {
249 fn id(&self) -> TaskExecutorId;
250
251 fn spec(&self) -> TaskSpec;
252
253 async fn execute(
254 &self,
255 ctx: TaskExecutionContext,
256 input: serde_json::Value,
257 ) -> anyhow::Result<TaskExecutionResult>;
258}
259
260#[cfg(test)]
261mod tests {
262 use std::sync::Arc;
263
264 use super::*;
265
266 struct NoopTaskExecutor;
267
268 #[async_trait::async_trait]
269 impl TaskExecutor for NoopTaskExecutor {
270 fn id(&self) -> TaskExecutorId {
271 "noop-task".to_string()
272 }
273
274 fn spec(&self) -> TaskSpec {
275 TaskSpec {
276 kind: "noop".to_string(),
277 description: "No-op task".to_string(),
278 input_schema: serde_json::json!({ "type": "object" }),
279 default_timeout_seconds: Some(30),
280 metadata: serde_json::json!({ "category": "test" }),
281 }
282 }
283
284 async fn execute(
285 &self,
286 ctx: TaskExecutionContext,
287 input: serde_json::Value,
288 ) -> anyhow::Result<TaskExecutionResult> {
289 Ok(TaskExecutionResult::success(serde_json::json!({
290 "task_id": ctx.task_id,
291 "input": input,
292 })))
293 }
294 }
295
296 #[test]
297 fn task_handle_round_trips_json() {
298 let handle = TaskHandle {
299 task_id: "task-1".to_string(),
300 executor_id: "process".to_string(),
301 spec: TaskSpec {
302 kind: "process".to_string(),
303 description: "Run a process".to_string(),
304 input_schema: serde_json::json!({ "type": "object" }),
305 default_timeout_seconds: Some(60),
306 metadata: serde_json::json!({}),
307 },
308 state: TaskState::Queued,
309 created_at: OffsetDateTime::UNIX_EPOCH,
310 started_at: None,
311 finished_at: None,
312 };
313
314 let encoded = serde_json::to_string(&handle).expect("serialize task handle");
315 let decoded: TaskHandle = serde_json::from_str(&encoded).expect("deserialize task handle");
316
317 assert_eq!(decoded, handle);
318 }
319
320 #[test]
321 fn task_events_round_trip_json() {
322 let started = TaskStarted {
323 task_id: "task-1".to_string(),
324 executor_id: "process".to_string(),
325 task_kind: "process".to_string(),
326 queue_depth: 0,
327 thread_id: Some("thread-a".to_string()),
328 turn_id: Some("turn-a".to_string()),
329 timestamp: OffsetDateTime::UNIX_EPOCH,
330 };
331 let output = TaskOutput {
332 task_id: "task-1".to_string(),
333 stream: TaskOutputStream::Stdout,
334 chunk: "hello\n".to_string(),
335 dropped_bytes: 0,
336 thread_id: Some("thread-a".to_string()),
337 turn_id: Some("turn-a".to_string()),
338 timestamp: OffsetDateTime::UNIX_EPOCH,
339 };
340
341 assert_eq!(
342 serde_json::from_value::<TaskStarted>(serde_json::to_value(&started).unwrap()).unwrap(),
343 started
344 );
345 assert_eq!(
346 serde_json::from_value::<TaskOutput>(serde_json::to_value(&output).unwrap()).unwrap(),
347 output
348 );
349 }
350
351 #[tokio::test]
352 async fn task_executor_trait_is_object_safe() {
353 let executor: Arc<dyn TaskExecutor> = Arc::new(NoopTaskExecutor);
354 let result = executor
355 .execute(
356 TaskExecutionContext {
357 task_id: "task-1".to_string(),
358 thread_id: None,
359 turn_id: None,
360 workspace_root: None,
361 runner_destination: None,
362 runner_session: None,
363 deadline: None,
364 process_grace_timeout: Duration::from_millis(250),
365 process_kill_timeout: Duration::from_secs(1),
366 metadata: serde_json::json!({}),
367 process_registry: None,
368 output: TaskOutputSink::default(),
369 },
370 serde_json::json!({ "ok": true }),
371 )
372 .await
373 .unwrap();
374
375 assert_eq!(executor.id(), "noop-task");
376 assert_eq!(executor.spec().kind, "noop");
377 assert_eq!(result.payload["task_id"], "task-1");
378 }
379}