Skip to main content

skiff_cli/openapi/
execute.rs

1//! Execute OpenAPI operations via HTTP.
2
3use std::collections::HashMap;
4use std::fs::File;
5use std::path::Path;
6
7use reqwest::blocking::{multipart, Client};
8use reqwest::header::{HeaderMap, HeaderName, HeaderValue};
9use serde_json::Value;
10
11use crate::cli::dynamic::{read_stdin_json, ParsedToolArgs};
12use crate::error::{Error, Result};
13use crate::model::{CommandDef, ParamLocation};
14use crate::output::{output_result, OutputOptions};
15
16pub struct OpenApiRequest {
17    pub path: String,
18    pub query: HashMap<String, Value>,
19    pub headers: HashMap<String, String>,
20    pub body: Option<Value>,
21    pub files: Vec<(String, String)>, // field name -> file path
22}
23
24pub fn collect_openapi_params(cmd: &CommandDef, parsed: &ParsedToolArgs) -> Result<OpenApiRequest> {
25    let mut path = cmd.path.clone().unwrap_or_default();
26    let mut query = HashMap::new();
27    let mut headers = HashMap::new();
28    let mut body: Option<Value> = None;
29    let mut files = Vec::new();
30
31    for p in &cmd.params {
32        if p.location == ParamLocation::Path {
33            if let Some(val) = parsed.values.get(&p.original_name) {
34                let s = value_to_string(val);
35                path = path.replace(&format!("{{{}}}", p.original_name), &s);
36            }
37        }
38    }
39
40    let method = cmd.method.as_deref().unwrap_or("get").to_lowercase();
41
42    if method == "get" {
43        for p in &cmd.params {
44            let Some(val) = parsed.values.get(&p.original_name) else {
45                continue;
46            };
47            match p.location {
48                ParamLocation::Query => {
49                    query.insert(p.original_name.clone(), val.clone());
50                }
51                ParamLocation::Header => {
52                    headers.insert(p.original_name.clone(), value_to_string(val));
53                }
54                _ => {}
55            }
56        }
57    } else if parsed.stdin {
58        body = Some(read_stdin_json("OpenAPI request body")?);
59        for p in &cmd.params {
60            if p.location == ParamLocation::Query {
61                if let Some(val) = parsed.values.get(&p.original_name) {
62                    query.insert(p.original_name.clone(), val.clone());
63                }
64            }
65        }
66    } else {
67        let mut body_obj = serde_json::Map::new();
68        for p in &cmd.params {
69            match p.location {
70                ParamLocation::Header => {
71                    if let Some(val) = parsed.values.get(&p.original_name) {
72                        headers.insert(p.original_name.clone(), value_to_string(val));
73                    }
74                }
75                ParamLocation::Path => {}
76                ParamLocation::File => {
77                    if let Some(val) = parsed.values.get(&p.original_name) {
78                        let fp = value_to_string(val);
79                        if !Path::new(&fp).is_file() {
80                            return Err(Error::runtime(format!("file not found: {fp}")));
81                        }
82                        files.push((p.original_name.clone(), fp));
83                    }
84                }
85                ParamLocation::Query => {
86                    if let Some(val) = parsed.values.get(&p.original_name) {
87                        query.insert(p.original_name.clone(), val.clone());
88                    }
89                }
90                ParamLocation::Body | ParamLocation::ToolInput | ParamLocation::GraphqlArg => {
91                    if let Some(val) = parsed.values.get(&p.original_name) {
92                        body_obj.insert(p.original_name.clone(), val.clone());
93                    }
94                }
95            }
96        }
97        if !body_obj.is_empty() {
98            body = Some(Value::Object(body_obj));
99        }
100    }
101
102    Ok(OpenApiRequest {
103        path,
104        query,
105        headers,
106        body,
107        files,
108    })
109}
110
111fn value_to_string(v: &Value) -> String {
112    match v {
113        Value::String(s) => s.clone(),
114        Value::Bool(b) => b.to_string(),
115        Value::Number(n) => n.to_string(),
116        other => other.to_string(),
117    }
118}
119
120pub fn execute_openapi(
121    parsed: &ParsedToolArgs,
122    base_url: &str,
123    auth_headers: &[(String, String)],
124    opts: &OutputOptions,
125) -> Result<()> {
126    let cmd = &parsed.command;
127    let req_parts = collect_openapi_params(cmd, parsed)?;
128    let url = format!("{}{}", base_url.trim_end_matches('/'), req_parts.path);
129    let method = cmd.method.as_deref().unwrap_or("get").to_uppercase();
130
131    let client = Client::builder()
132        .timeout(std::time::Duration::from_secs(60))
133        .build()
134        .map_err(|e| Error::runtime(e.to_string()))?;
135
136    let mut header_map = HeaderMap::new();
137    let is_multipart =
138        !req_parts.files.is_empty() || cmd.content_type.as_deref() == Some("multipart/form-data");
139    if !is_multipart {
140        header_map.insert(
141            reqwest::header::CONTENT_TYPE,
142            HeaderValue::from_static("application/json"),
143        );
144    }
145    for (k, v) in auth_headers {
146        header_map.insert(
147            HeaderName::from_bytes(k.as_bytes()).map_err(|e| Error::runtime(e.to_string()))?,
148            HeaderValue::from_str(v).map_err(|e| Error::runtime(e.to_string()))?,
149        );
150    }
151    for (k, v) in &req_parts.headers {
152        header_map.insert(
153            HeaderName::from_bytes(k.as_bytes()).map_err(|e| Error::runtime(e.to_string()))?,
154            HeaderValue::from_str(v).map_err(|e| Error::runtime(e.to_string()))?,
155        );
156    }
157
158    let mut request = client.request(
159        method
160            .parse()
161            .map_err(|_| Error::runtime(format!("invalid method: {method}")))?,
162        &url,
163    );
164    request = request.headers(header_map);
165
166    // Query params: stringify values
167    let mut pairs: Vec<(String, String)> = Vec::new();
168    for (k, v) in &req_parts.query {
169        pairs.push((k.clone(), value_to_query(v)));
170    }
171    if !pairs.is_empty() {
172        request = request.query(&pairs);
173    }
174
175    let response = if !req_parts.files.is_empty() {
176        let mut form = multipart::Form::new();
177        if let Some(Value::Object(map)) = &req_parts.body {
178            for (k, v) in map {
179                form = form.text(k.clone(), value_to_string(v));
180            }
181        }
182        for (field, path) in &req_parts.files {
183            let file = File::open(path).map_err(Error::from)?;
184            let filename = Path::new(path)
185                .file_name()
186                .and_then(|s| s.to_str())
187                .unwrap_or("upload")
188                .to_string();
189            let part = multipart::Part::reader(file)
190                .file_name(filename)
191                .mime_str("application/octet-stream")
192                .map_err(|e| Error::runtime(e.to_string()))?;
193            form = form.part(field.clone(), part);
194        }
195        request.multipart(form).send()
196    } else if cmd.content_type.as_deref() == Some("multipart/form-data") {
197        // form fields without files
198        let mut form = multipart::Form::new();
199        if let Some(Value::Object(map)) = &req_parts.body {
200            for (k, v) in map {
201                form = form.text(k.clone(), value_to_string(v));
202            }
203        }
204        request.multipart(form).send()
205    } else if let Some(body) = &req_parts.body {
206        request.json(body).send()
207    } else {
208        request.send()
209    };
210
211    let resp = response.map_err(|e| Error::runtime(e.to_string()))?;
212    let status = resp.status();
213    let bytes = resp.bytes().map_err(|e| Error::runtime(e.to_string()))?;
214
215    if !status.is_success() {
216        let text = String::from_utf8_lossy(&bytes);
217        return Err(Error::runtime(format!("Error {}: {text}", status.as_u16())));
218    }
219
220    if opts.json_output {
221        let data = serde_json::from_slice::<Value>(&bytes)
222            .unwrap_or_else(|_| Value::String(String::from_utf8_lossy(&bytes).into_owned()));
223        output_result(data, opts)?;
224        return Ok(());
225    }
226
227    if opts.raw {
228        use std::io::Write;
229        std::io::stdout().write_all(&bytes).map_err(Error::from)?;
230        return Ok(());
231    }
232
233    match serde_json::from_slice::<Value>(&bytes) {
234        Ok(data) => output_result(data, opts)?,
235        Err(_) => {
236            println!("{}", String::from_utf8_lossy(&bytes));
237        }
238    }
239    Ok(())
240}
241
242fn value_to_query(v: &Value) -> String {
243    match v {
244        Value::String(s) => s.clone(),
245        Value::Bool(b) => b.to_string(),
246        Value::Number(n) => n.to_string(),
247        Value::Array(_) | Value::Object(_) => {
248            serde_json::to_string(v).unwrap_or_else(|_| v.to_string())
249        }
250        Value::Null => String::new(),
251    }
252}
253
254/// Resolve base URL from --base-url, servers[], or spec URL origin.
255pub fn resolve_base_url(explicit: Option<&str>, spec: &Value, spec_source: &str) -> Result<String> {
256    if let Some(u) = explicit {
257        return Ok(u.to_string());
258    }
259    let mut base = spec
260        .get("servers")
261        .and_then(|s| s.as_array())
262        .and_then(|arr| arr.first())
263        .and_then(|s| s.get("url"))
264        .and_then(|u| u.as_str())
265        .unwrap_or("")
266        .to_string();
267
268    if base.is_empty() || !base.starts_with("http") {
269        if spec_source.starts_with("http://") || spec_source.starts_with("https://") {
270            let url =
271                reqwest::Url::parse(spec_source).map_err(|e| Error::runtime(e.to_string()))?;
272            let origin = format!(
273                "{}://{}",
274                url.scheme(),
275                url.host_str().unwrap_or("localhost")
276            );
277            let origin = if let Some(port) = url.port() {
278                format!("{origin}:{port}")
279            } else {
280                origin
281            };
282            if !base.is_empty() && !base.starts_with("http") {
283                base = format!("{origin}{base}");
284            } else {
285                base = origin;
286            }
287        } else if base.is_empty() {
288            return Err(Error::runtime("cannot determine base URL. Use --base-url."));
289        }
290    }
291    Ok(base)
292}