Skip to main content

roder_api/
tasks.rs

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    /// Bounded local-process graceful-stop budget selected by the task host.
69    pub process_grace_timeout: Duration,
70    /// Bounded local-process forced-kill/reap budget selected by the task host.
71    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}