Skip to main content

skiff_cli/bake/
mod.rs

1//! Bake mode — named connection configs in `$SKIFF_CONFIG_DIR/baked.json`.
2//!
3//! `@name` expands via [`BakedTool::to_argv`]. Prefer `env:`/`file:` secrets;
4//! [`BakedTool::masked_for_display`] is used by `bake show`.
5
6use std::collections::BTreeMap;
7use std::fs;
8use std::path::{Path, PathBuf};
9use std::process::Command as StdCommand;
10
11use regex::Regex;
12use serde::{Deserialize, Serialize};
13use serde_json::Value;
14
15use crate::error::{Error, Result};
16use crate::model::BakeConfig;
17use crate::paths::{baked_file, DEFAULT_CACHE_TTL};
18
19static BAKE_NAME_RE: std::sync::LazyLock<Regex> =
20    std::sync::LazyLock::new(|| Regex::new(r"^[a-z][a-z0-9-]*$").expect("bake name regex"));
21
22/// Validate bake tool names: `[a-z][a-z0-9-]*`.
23pub fn is_valid_bake_name(name: &str) -> bool {
24    BAKE_NAME_RE.is_match(name)
25}
26
27#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
28pub struct BakedTool {
29    pub source_type: String,
30    pub source: String,
31    #[serde(default, skip_serializing_if = "Option::is_none")]
32    pub base_url: Option<String>,
33    #[serde(default)]
34    pub auth_headers: Vec<(String, String)>,
35    #[serde(default)]
36    pub env_vars: BTreeMap<String, String>,
37    #[serde(default = "default_cache_ttl")]
38    pub cache_ttl: u64,
39    #[serde(default = "default_transport")]
40    pub transport: String,
41    #[serde(default)]
42    pub oauth: bool,
43    #[serde(default, skip_serializing_if = "Option::is_none")]
44    pub oauth_client_id: Option<String>,
45    #[serde(default, skip_serializing_if = "Option::is_none")]
46    pub oauth_client_secret: Option<String>,
47    #[serde(
48        default = "default_oauth_client_name",
49        skip_serializing_if = "is_default_oauth_client_name"
50    )]
51    pub oauth_client_name: String,
52    #[serde(default, skip_serializing_if = "Option::is_none")]
53    pub oauth_scope: Option<String>,
54    #[serde(default, skip_serializing_if = "Option::is_none")]
55    pub oauth_redirect_uri: Option<String>,
56    #[serde(
57        default = "default_oauth_flow",
58        skip_serializing_if = "is_default_oauth_flow"
59    )]
60    pub oauth_flow: String,
61    /// Prefer routing through this named session daemon when present
62    #[serde(default, skip_serializing_if = "Option::is_none")]
63    pub session: Option<String>,
64    #[serde(default)]
65    pub include: Vec<String>,
66    #[serde(default)]
67    pub exclude: Vec<String>,
68    #[serde(default)]
69    pub methods: Vec<String>,
70    #[serde(default, skip_serializing_if = "String::is_empty")]
71    pub description: String,
72    /// Forward-compat for fields we don't model yet.
73    #[serde(flatten)]
74    pub extra: BTreeMap<String, Value>,
75}
76
77fn default_cache_ttl() -> u64 {
78    DEFAULT_CACHE_TTL
79}
80
81fn default_transport() -> String {
82    "auto".into()
83}
84
85fn default_oauth_client_name() -> String {
86    "skiff".into()
87}
88
89fn is_default_oauth_client_name(s: &str) -> bool {
90    s == "skiff"
91}
92
93fn default_oauth_flow() -> String {
94    "auto".into()
95}
96
97fn is_default_oauth_flow(s: &str) -> bool {
98    s == "auto"
99}
100
101impl Default for BakedTool {
102    fn default() -> Self {
103        Self {
104            source_type: String::new(),
105            source: String::new(),
106            base_url: None,
107            auth_headers: Vec::new(),
108            env_vars: BTreeMap::new(),
109            cache_ttl: DEFAULT_CACHE_TTL,
110            transport: default_transport(),
111            oauth: false,
112            oauth_client_id: None,
113            oauth_client_secret: None,
114            oauth_client_name: default_oauth_client_name(),
115            oauth_scope: None,
116            oauth_redirect_uri: None,
117            oauth_flow: default_oauth_flow(),
118            session: None,
119            include: Vec::new(),
120            exclude: Vec::new(),
121            methods: Vec::new(),
122            description: String::new(),
123            extra: BTreeMap::new(),
124        }
125    }
126}
127
128impl BakedTool {
129    pub fn bake_config(&self) -> BakeConfig {
130        BakeConfig {
131            include: self.include.clone(),
132            exclude: self.exclude.clone(),
133            methods: self.methods.clone(),
134        }
135    }
136
137    /// Reconstruct CLI argv from a baked config (Python `_baked_to_argv`).
138    pub fn to_argv(&self) -> Vec<String> {
139        let mut argv = Vec::new();
140        match self.source_type.as_str() {
141            "spec" => {
142                argv.push("--spec".into());
143                argv.push(self.source.clone());
144            }
145            "mcp" => {
146                argv.push("--mcp".into());
147                argv.push(self.source.clone());
148            }
149            "mcp_stdio" => {
150                argv.push("--mcp-stdio".into());
151                argv.push(self.source.clone());
152            }
153            "graphql" => {
154                argv.push("--graphql".into());
155                argv.push(self.source.clone());
156            }
157            other => {
158                // Unknown source types still emit a best-effort flag.
159                argv.push(format!("--{other}"));
160                argv.push(self.source.clone());
161            }
162        }
163        if let Some(base) = &self.base_url {
164            argv.push("--base-url".into());
165            argv.push(base.clone());
166        }
167        for (name, value) in &self.auth_headers {
168            argv.push("--auth-header".into());
169            argv.push(format!("{name}:{value}"));
170        }
171        for (k, v) in &self.env_vars {
172            argv.push("--env".into());
173            argv.push(format!("{k}={v}"));
174        }
175        argv.push("--cache-ttl".into());
176        argv.push(self.cache_ttl.to_string());
177        if self.transport != "auto" {
178            argv.push("--transport".into());
179            argv.push(self.transport.clone());
180        }
181        if self.oauth {
182            argv.push("--oauth".into());
183        }
184        if let Some(id) = &self.oauth_client_id {
185            argv.push("--oauth-client-id".into());
186            argv.push(id.clone());
187        }
188        if let Some(sec) = &self.oauth_client_secret {
189            argv.push("--oauth-client-secret".into());
190            argv.push(sec.clone());
191        }
192        if self.oauth_client_name != "skiff" {
193            argv.push("--oauth-client-name".into());
194            argv.push(self.oauth_client_name.clone());
195        }
196        if let Some(scope) = &self.oauth_scope {
197            argv.push("--oauth-scope".into());
198            argv.push(scope.clone());
199        }
200        if let Some(uri) = &self.oauth_redirect_uri {
201            argv.push("--oauth-redirect-uri".into());
202            argv.push(uri.clone());
203        }
204        if self.oauth_flow != "auto" {
205            argv.push("--oauth-flow".into());
206            argv.push(self.oauth_flow.clone());
207        }
208        if let Some(sess) = &self.session {
209            argv.push("--session".into());
210            argv.push(sess.clone());
211        }
212        argv
213    }
214
215    /// Copy suitable for `bake show` (literal secrets masked).
216    pub fn masked_for_display(&self) -> Value {
217        let mut display = serde_json::to_value(self).unwrap_or(Value::Null);
218        if let Some(headers) = display
219            .get_mut("auth_headers")
220            .and_then(|v| v.as_array_mut())
221        {
222            for entry in headers {
223                if let Some(arr) = entry.as_array_mut() {
224                    if arr.len() >= 2 {
225                        if let Some(val) = arr[1].as_str() {
226                            arr[1] = Value::String(mask_secret(val));
227                        }
228                    }
229                }
230            }
231        }
232        if let Some(Value::String(sec)) = display.get_mut("oauth_client_secret") {
233            *sec = mask_secret(sec);
234        }
235        display
236    }
237}
238
239fn mask_secret(val: &str) -> String {
240    if val.starts_with("env:") || val.starts_with("file:") {
241        val.to_string()
242    } else if val.len() > 4 {
243        format!("{}****", &val[..4])
244    } else {
245        "****".into()
246    }
247}
248
249pub type BakedStore = BTreeMap<String, BakedTool>;
250
251pub fn load_baked_all() -> Result<BakedStore> {
252    let path = baked_file();
253    if !path.exists() {
254        return Ok(BakedStore::new());
255    }
256    let text = match fs::read_to_string(&path) {
257        Ok(t) => t,
258        Err(_) => return Ok(BakedStore::new()),
259    };
260    match serde_json::from_str(&text) {
261        Ok(store) => Ok(store),
262        Err(_) => Ok(BakedStore::new()),
263    }
264}
265
266pub fn save_baked_all(store: &BakedStore) -> Result<()> {
267    // Pretty-print to match Python `indent=2` baked.json.
268    let path = baked_file();
269    let text = serde_json::to_string_pretty(store)? + "\n";
270    crate::fsutil::atomic_write_0600(&path, text.as_bytes())?;
271    Ok(())
272}
273
274/// Load a single baked config by name.
275pub fn get_baked(name: &str) -> Result<Option<BakedTool>> {
276    Ok(load_baked_all()?.get(name).cloned())
277}
278
279pub fn require_baked(name: &str) -> Result<BakedTool> {
280    get_baked(name)?.ok_or_else(|| Error::runtime(format!("no baked tool named '{name}'")))
281}
282
283pub fn split_csv_list(s: &str) -> Vec<String> {
284    s.split(',')
285        .map(str::trim)
286        .filter(|x| !x.is_empty())
287        .map(str::to_string)
288        .collect()
289}
290
291pub fn split_methods(s: &str) -> Vec<String> {
292    s.split(',')
293        .map(str::trim)
294        .filter(|x| !x.is_empty())
295        .map(|x| x.to_uppercase())
296        .collect()
297}
298
299/// Parse `Name:Value` auth headers without resolving secrets (stored as-is).
300pub fn parse_auth_header_raw(items: &[String]) -> Result<Vec<(String, String)>> {
301    let mut out = Vec::new();
302    for item in items {
303        let Some((k, v)) = item.split_once(':') else {
304            return Err(Error::usage(format!(
305                "invalid auth header format: {item:?}"
306            )));
307        };
308        out.push((k.trim().to_string(), v.trim().to_string()));
309    }
310    Ok(out)
311}
312
313pub fn parse_env_raw(items: &[String]) -> Result<BTreeMap<String, String>> {
314    let mut out = BTreeMap::new();
315    for item in items {
316        let Some((k, v)) = item.split_once('=') else {
317            return Err(Error::usage(format!("invalid env format: {item:?}")));
318        };
319        out.insert(k.trim().to_string(), v.to_string());
320    }
321    Ok(out)
322}
323
324pub fn create_baked(name: &str, tool: BakedTool, force: bool) -> Result<()> {
325    if !is_valid_bake_name(name) {
326        return Err(Error::usage(format!(
327            "invalid name '{name}' — must match [a-z][a-z0-9-]*"
328        )));
329    }
330    let mut store = load_baked_all()?;
331    if store.contains_key(name) && !force {
332        return Err(Error::runtime(format!(
333            "'{name}' already exists. Use --force to overwrite."
334        )));
335    }
336    store.insert(name.to_string(), tool);
337    save_baked_all(&store)
338}
339
340pub fn remove_baked(name: &str) -> Result<()> {
341    let mut store = load_baked_all()?;
342    if store.remove(name).is_none() {
343        return Err(Error::runtime(format!("no baked tool named '{name}'")));
344    }
345    save_baked_all(&store)?;
346    // Clean up default install wrapper if present.
347    if let Some(home) = home_dir() {
348        let wrapper = home.join(".local").join("bin").join(name);
349        if wrapper.exists() {
350            fs::remove_file(&wrapper)?;
351            println!("Removed installed wrapper: {}", wrapper.display());
352        }
353    }
354    Ok(())
355}
356
357pub fn update_baked(name: &str, mutator: impl FnOnce(&mut BakedTool)) -> Result<()> {
358    let mut store = load_baked_all()?;
359    let cfg = store
360        .get_mut(name)
361        .ok_or_else(|| Error::runtime(format!("no baked tool named '{name}'")))?;
362    mutator(cfg);
363    save_baked_all(&store)
364}
365
366fn home_dir() -> Option<PathBuf> {
367    std::env::var_os("HOME")
368        .or_else(|| std::env::var_os("USERPROFILE"))
369        .map(PathBuf::from)
370}
371
372fn resolve_skiff_bin() -> String {
373    if let Ok(exe) = std::env::current_exe() {
374        return exe.to_string_lossy().into_owned();
375    }
376    which_skiff().unwrap_or_else(|| "skiff".into())
377}
378
379fn which_skiff() -> Option<String> {
380    let output = StdCommand::new("which").arg("skiff").output().ok()?;
381    if !output.status.success() {
382        return None;
383    }
384    let path = String::from_utf8_lossy(&output.stdout).trim().to_string();
385    if path.is_empty() {
386        None
387    } else {
388        Some(path)
389    }
390}
391
392fn shell_quote(s: &str) -> String {
393    // POSIX single-quote escaping.
394    if s.is_empty() {
395        return "''".into();
396    }
397    format!("'{}'", s.replace('\'', "'\"'\"'"))
398}
399
400/// Install a shell wrapper that runs `skiff @name "$@"`.
401pub fn install_wrapper(name: &str, dir: Option<&Path>) -> Result<PathBuf> {
402    let _ = require_baked(name)?;
403    let bin_dir = match dir {
404        Some(d) => d.to_path_buf(),
405        None => home_dir()
406            .map(|h| h.join(".local").join("bin"))
407            .ok_or_else(|| Error::runtime("cannot determine home directory"))?,
408    };
409    fs::create_dir_all(&bin_dir)?;
410    let wrapper = bin_dir.join(name);
411    let bin = resolve_skiff_bin();
412    let content = format!("#!/bin/sh\nexec {} @{} \"$@\"\n", shell_quote(&bin), name);
413    fs::write(&wrapper, content)?;
414    #[cfg(unix)]
415    {
416        use std::os::unix::fs::PermissionsExt;
417        fs::set_permissions(&wrapper, fs::Permissions::from_mode(0o755))?;
418    }
419    Ok(wrapper)
420}
421
422pub fn default_install_dir() -> Option<PathBuf> {
423    home_dir().map(|h| h.join(".local").join("bin"))
424}
425
426#[cfg(test)]
427mod tests {
428    use super::*;
429    use crate::paths::{set_config_dir_override, TEST_PATHS_LOCK};
430    use tempfile::tempdir;
431
432    #[test]
433    fn name_validation() {
434        for name in ["petstore", "my-api", "a1", "x-y-z"] {
435            assert!(is_valid_bake_name(name), "{name} should be valid");
436        }
437        for name in ["1abc", "Abc", "a_b", "-foo", ""] {
438            assert!(!is_valid_bake_name(name), "{name} should be invalid");
439        }
440    }
441
442    #[test]
443    fn baked_to_argv_spec() {
444        let cfg = BakedTool {
445            source_type: "spec".into(),
446            source: "https://example.com/spec.json".into(),
447            base_url: Some("https://api.example.com".into()),
448            auth_headers: vec![("Authorization".into(), "env:TOKEN".into())],
449            cache_ttl: 7200,
450            ..Default::default()
451        };
452        let argv = cfg.to_argv();
453        assert!(argv.contains(&"--spec".into()));
454        assert!(argv.contains(&"https://example.com/spec.json".into()));
455        assert!(argv.contains(&"--base-url".into()));
456        assert!(argv.contains(&"--auth-header".into()));
457        assert!(argv.contains(&"Authorization:env:TOKEN".into()));
458        assert!(argv.contains(&"--cache-ttl".into()));
459        assert!(argv.contains(&"7200".into()));
460    }
461
462    #[test]
463    fn baked_to_argv_mcp_stdio() {
464        let mut env = BTreeMap::new();
465        env.insert("GH_TOKEN".into(), "abc".into());
466        let cfg = BakedTool {
467            source_type: "mcp_stdio".into(),
468            source: "npx @mcp/github".into(),
469            env_vars: env,
470            cache_ttl: 3600,
471            ..Default::default()
472        };
473        let argv = cfg.to_argv();
474        assert!(argv.contains(&"--mcp-stdio".into()));
475        assert!(argv.contains(&"npx @mcp/github".into()));
476        assert!(argv.contains(&"--env".into()));
477        assert!(argv.contains(&"GH_TOKEN=abc".into()));
478    }
479
480    #[test]
481    fn baked_to_argv_oauth() {
482        let cfg = BakedTool {
483            source_type: "mcp".into(),
484            source: "https://mcp.example.com".into(),
485            transport: "sse".into(),
486            oauth: true,
487            oauth_client_id: Some("env:CID".into()),
488            oauth_client_secret: Some("env:CSEC".into()),
489            oauth_scope: Some("read write".into()),
490            ..Default::default()
491        };
492        let argv = cfg.to_argv();
493        assert!(argv.contains(&"--oauth".into()));
494        assert!(argv.contains(&"--oauth-client-id".into()));
495        assert!(argv.contains(&"--oauth-client-secret".into()));
496        assert!(argv.contains(&"--oauth-scope".into()));
497        assert!(argv.contains(&"--transport".into()));
498        assert!(argv.contains(&"sse".into()));
499        assert!(!argv.iter().any(|a| a == "--oauth-redirect-uri"));
500    }
501
502    #[test]
503    fn baked_to_argv_oauth_redirect() {
504        let uri = "http://localhost:18080/oauth/callback";
505        let cfg = BakedTool {
506            source_type: "mcp".into(),
507            source: "https://mcp.example.com".into(),
508            oauth: true,
509            oauth_redirect_uri: Some(uri.into()),
510            ..Default::default()
511        };
512        let argv = cfg.to_argv();
513        let idx = argv
514            .iter()
515            .position(|a| a == "--oauth-redirect-uri")
516            .unwrap();
517        assert_eq!(argv[idx + 1], uri);
518    }
519
520    #[test]
521    fn baked_to_argv_session() {
522        let cfg = BakedTool {
523            source_type: "mcp_stdio".into(),
524            source: "python3 server.py".into(),
525            session: Some("warm".into()),
526            ..Default::default()
527        };
528        let argv = cfg.to_argv();
529        assert!(argv.windows(2).any(|w| w == ["--session", "warm"]));
530    }
531
532    #[test]
533    fn round_trip_store() {
534        let _g = TEST_PATHS_LOCK.lock().unwrap_or_else(|e| e.into_inner());
535        let dir = tempdir().unwrap();
536        set_config_dir_override(Some(dir.path().to_path_buf()));
537        let tool = BakedTool {
538            source_type: "spec".into(),
539            source: "https://example.com/spec.json".into(),
540            ..Default::default()
541        };
542        create_baked("test", tool.clone(), false).unwrap();
543        let loaded = require_baked("test").unwrap();
544        assert_eq!(loaded.source_type, "spec");
545        assert_eq!(loaded.source, tool.source);
546        set_config_dir_override(None);
547    }
548
549    #[test]
550    fn load_missing() {
551        let _g = TEST_PATHS_LOCK.lock().unwrap_or_else(|e| e.into_inner());
552        let dir = tempdir().unwrap();
553        set_config_dir_override(Some(dir.path().join("nope")));
554        assert!(load_baked_all().unwrap().is_empty());
555        set_config_dir_override(None);
556    }
557}