Skip to main content

runkernel_cli_support/
lib.rs

1use runkernel::{Pipeline, PipelineGraph, RunOptions, TaskExplanation};
2use serde::{Deserialize, Serialize};
3
4const PROTOCOL_COMMAND: &str = "__runkernel";
5const PROTOCOL_VERSION: u32 = 1;
6
7pub struct RunkernelApp {
8    pipeline: Pipeline,
9}
10
11#[derive(Clone, Debug, Serialize, Deserialize, PartialEq, Eq)]
12pub struct MetadataResponse {
13    pub protocol_version: u32,
14    pub workflow_name: String,
15    pub description: Option<String>,
16    pub runkernel_version: String,
17    pub supports: ProtocolSupport,
18}
19
20#[derive(Clone, Debug, Serialize, Deserialize, PartialEq, Eq)]
21pub struct ProtocolSupport {
22    pub list: bool,
23    pub graph: bool,
24    pub explain: bool,
25    pub run_task: bool,
26    pub run_all: bool,
27}
28
29#[derive(Clone, Debug, Serialize, Deserialize, PartialEq, Eq)]
30pub struct ListResponse {
31    pub tasks: Vec<TaskListItem>,
32}
33
34#[derive(Clone, Debug, Serialize, Deserialize, PartialEq, Eq)]
35pub struct TaskListItem {
36    pub name: String,
37    pub description: Option<String>,
38    pub dependencies: Vec<String>,
39    pub cacheable: bool,
40}
41
42#[derive(Clone, Debug, Serialize, Deserialize, PartialEq, Eq)]
43pub struct ExplainResponse {
44    pub task: TaskExplanation,
45}
46
47impl RunkernelApp {
48    pub fn new(pipeline: Pipeline) -> Self {
49        Self { pipeline }
50    }
51
52    pub async fn run_from_args(self) -> anyhow::Result<()> {
53        self.run_from(std::env::args().skip(1)).await
54    }
55
56    pub async fn run_from<I, S>(self, args: I) -> anyhow::Result<()>
57    where
58        I: IntoIterator<Item = S>,
59        S: Into<String>,
60    {
61        let args: Vec<String> = args.into_iter().map(Into::into).collect();
62        if args.first().map(String::as_str) != Some(PROTOCOL_COMMAND) {
63            return finish_result(self.pipeline.run().await?);
64        }
65
66        match args.get(1).map(String::as_str) {
67            Some("metadata") => {
68                require_json_format(&args[2..])?;
69                emit_json(&metadata(&self.pipeline))
70            }
71            Some("list") => {
72                require_json_format(&args[2..])?;
73                emit_json(&list(&self.pipeline))
74            }
75            Some("graph") => {
76                require_json_format(&args[2..])?;
77                emit_json(&self.pipeline.graph()?)
78            }
79            Some("explain") => {
80                let task = args
81                    .get(2)
82                    .ok_or_else(|| anyhow::anyhow!("Missing task for __runkernel explain"))?;
83                require_json_format(&args[3..])?;
84                emit_json(&ExplainResponse {
85                    task: self.pipeline.explain(task)?,
86                })
87            }
88            Some("run") => {
89                let (task, forwarded_args) = parse_run_command(&args[2..])?;
90                let options = RunOptions {
91                    args: forwarded_args,
92                };
93                let result = match task {
94                    Some(task) => self.pipeline.run_task_with_options(&task, options).await?,
95                    None => self.pipeline.run_with_options(options).await?,
96                };
97                finish_result(result)
98            }
99            Some(other) => anyhow::bail!("Unknown __runkernel command '{}'", other),
100            None => anyhow::bail!("Missing __runkernel command"),
101        }
102    }
103}
104
105fn parse_run_command(args: &[String]) -> anyhow::Result<(Option<String>, Vec<String>)> {
106    match args {
107        [] => Ok((None, Vec::new())),
108        [separator, rest @ ..] if separator == "--" => Ok((None, rest.to_vec())),
109        [task] => Ok((Some(task.clone()), Vec::new())),
110        [task, separator, rest @ ..] if separator == "--" => Ok((Some(task.clone()), rest.to_vec())),
111        [task, unexpected, ..] => anyhow::bail!(
112            "Unexpected argument '{}' for __runkernel run task '{}'. Forwarded arguments must follow '--'",
113            unexpected,
114            task
115        ),
116    }
117}
118
119fn metadata(pipeline: &Pipeline) -> MetadataResponse {
120    MetadataResponse {
121        protocol_version: PROTOCOL_VERSION,
122        workflow_name: pipeline.name().to_string(),
123        description: None,
124        runkernel_version: env!("CARGO_PKG_VERSION").to_string(),
125        supports: ProtocolSupport {
126            list: true,
127            graph: true,
128            explain: true,
129            run_task: true,
130            run_all: true,
131        },
132    }
133}
134
135fn list(pipeline: &Pipeline) -> ListResponse {
136    let mut tasks: Vec<_> = pipeline
137        .tasks()
138        .map(|task| TaskListItem {
139            name: task.name.clone(),
140            description: task.description.clone(),
141            dependencies: task.dependencies.clone(),
142            cacheable: task.cacheable(),
143        })
144        .collect();
145    tasks.sort_by(|a, b| a.name.cmp(&b.name));
146    ListResponse { tasks }
147}
148
149fn require_json_format(args: &[String]) -> anyhow::Result<()> {
150    if args.is_empty() {
151        return Ok(());
152    }
153    if args == ["--format".to_string(), "json".to_string()] {
154        return Ok(());
155    }
156    anyhow::bail!("Only '--format json' is supported for __runkernel protocol commands")
157}
158
159fn emit_json<T>(value: &T) -> anyhow::Result<()>
160where
161    T: Serialize,
162{
163    serde_json::to_writer(std::io::stdout(), value)?;
164    println!();
165    Ok(())
166}
167
168fn finish_result(result: runkernel::PipelineResult) -> anyhow::Result<()> {
169    if result.summary.success {
170        Ok(())
171    } else {
172        anyhow::bail!("pipeline failed: {:?}", result.summary)
173    }
174}
175
176pub fn graph_to_text(graph: &PipelineGraph) -> String {
177    let mut output = String::from("Tasks:");
178    for node in &graph.nodes {
179        output.push_str("\n  ");
180        output.push_str(&node.id);
181    }
182    if !graph.edges.is_empty() {
183        output.push_str("\n\nEdges:");
184        for edge in &graph.edges {
185            output.push_str("\n  ");
186            output.push_str(&edge.from);
187            output.push_str(" -> ");
188            output.push_str(&edge.to);
189        }
190    }
191    output
192}
193
194#[cfg(test)]
195mod tests {
196    use super::*;
197    use runkernel::Task;
198    use std::sync::{Arc, Mutex};
199
200    fn pipeline() -> Pipeline {
201        let mut pipeline = Pipeline::new("support-test");
202        pipeline.add(Task::new("lint").description("Run lint").exec("true"));
203        pipeline.add(
204            Task::new("test")
205                .description("Run tests")
206                .depends_on(&["lint"])
207                .exec("true"),
208        );
209        pipeline
210    }
211
212    #[test]
213    fn test_metadata_shape() {
214        let metadata = metadata(&pipeline());
215        assert_eq!(metadata.protocol_version, 1);
216        assert_eq!(metadata.workflow_name, "support-test");
217        assert!(metadata.supports.list);
218    }
219
220    #[test]
221    fn test_list_shape() {
222        let list = list(&pipeline());
223        assert_eq!(list.tasks.len(), 2);
224        assert_eq!(list.tasks[0].name, "lint");
225        assert_eq!(list.tasks[0].description.as_deref(), Some("Run lint"));
226    }
227
228    #[test]
229    fn test_graph_to_text() {
230        let graph = pipeline().graph().unwrap();
231        assert_eq!(
232            graph_to_text(&graph),
233            "Tasks:\n  lint\n  test\n\nEdges:\n  lint -> test"
234        );
235    }
236
237    #[test]
238    fn test_graph_to_text_includes_isolated_node() {
239        let mut pipeline = Pipeline::new("isolated");
240        pipeline.add(Task::new("build").exec("true"));
241        let graph = pipeline.graph().unwrap();
242        assert_eq!(graph_to_text(&graph), "Tasks:\n  build");
243    }
244
245    #[test]
246    fn test_parse_run_command_forwards_default_workflow_args() {
247        let parsed =
248            parse_run_command(&["--".to_string(), "--target".to_string(), "prod".to_string()])
249                .unwrap();
250        assert_eq!(
251            parsed,
252            (None, vec!["--target".to_string(), "prod".to_string()])
253        );
254    }
255
256    #[test]
257    fn test_parse_run_command_forwards_task_args() {
258        let parsed = parse_run_command(&[
259            "deploy".to_string(),
260            "--".to_string(),
261            "--dry-run".to_string(),
262        ])
263        .unwrap();
264        assert_eq!(
265            parsed,
266            (Some("deploy".to_string()), vec!["--dry-run".to_string()])
267        );
268    }
269
270    #[test]
271    fn test_parse_run_command_without_forwarded_args() {
272        assert_eq!(
273            parse_run_command(&["deploy".to_string()]).unwrap(),
274            (Some("deploy".to_string()), Vec::new())
275        );
276        assert_eq!(parse_run_command(&[]).unwrap(), (None, Vec::new()));
277    }
278
279    #[test]
280    fn test_parse_run_command_rejects_unseparated_extra_args() {
281        let err = parse_run_command(&["deploy".to_string(), "--target".to_string()]).unwrap_err();
282        assert!(err
283            .to_string()
284            .contains("Forwarded arguments must follow '--'"));
285    }
286
287    #[tokio::test]
288    async fn test_run_protocol_forwards_args_to_context() {
289        let captured = Arc::new(Mutex::new(None));
290        let seen = Arc::clone(&captured);
291        let mut pipeline = Pipeline::new("support-args");
292        pipeline.add(Task::new("deploy").exec_fn(move |ctx| {
293            let seen = Arc::clone(&seen);
294            async move {
295                *seen.lock().unwrap() = Some(ctx.args().to_vec());
296                Ok(())
297            }
298        }));
299
300        RunkernelApp::new(pipeline)
301            .run_from(["__runkernel", "run", "deploy", "--", "--target", "prod"])
302            .await
303            .unwrap();
304
305        assert_eq!(
306            *captured.lock().unwrap(),
307            Some(vec!["--target".to_string(), "prod".to_string()])
308        );
309    }
310}