Skip to main content

warrant_core/
paths.rs

1use std::path::{Path, PathBuf};
2
3use crate::error::{Error, Result};
4
5#[derive(Debug, Clone, PartialEq, Eq)]
6pub struct ToolId(String);
7
8impl ToolId {
9    pub fn parse(input: &str) -> Result<Self> {
10        if input.is_empty() {
11            return Err(Error::InvalidToolId(input.to_string()));
12        }
13        if input == "." || input == ".." {
14            return Err(Error::InvalidToolId(input.to_string()));
15        }
16        if Path::new(input).is_absolute() {
17            return Err(Error::InvalidToolId(input.to_string()));
18        }
19        if input.contains('/') || input.contains('\\') {
20            return Err(Error::InvalidToolId(input.to_string()));
21        }
22        if !input
23            .chars()
24            .all(|c| c.is_ascii_alphanumeric() || matches!(c, '.' | '_' | '-'))
25        {
26            return Err(Error::InvalidToolId(input.to_string()));
27        }
28
29        Ok(Self(input.to_string()))
30    }
31
32    pub fn as_str(&self) -> &str {
33        &self.0
34    }
35}
36
37#[derive(Debug, Clone)]
38pub struct ToolPaths {
39    pub tool_id: ToolId,
40    pub installed_warrant_path: PathBuf,
41    pub version_state_path: PathBuf,
42    pub signing_private_key_path: PathBuf,
43    pub signing_public_key_path: PathBuf,
44    pub host_secret_path: PathBuf,
45    pub session_dir_path: PathBuf,
46}
47
48impl ToolPaths {
49    pub fn for_tool(tool: &str) -> Result<Self> {
50        let tool_id = ToolId::parse(tool)?;
51        if cfg!(target_os = "macos") {
52            let base = PathBuf::from("/Library/Application Support").join(tool_id.as_str());
53            Ok(Self {
54                tool_id,
55                installed_warrant_path: base.join("warrant.toml"),
56                version_state_path: base.join("signing").join("version"),
57                signing_private_key_path: base.join("signing").join("private.key"),
58                signing_public_key_path: base.join("signing").join("public.key"),
59                host_secret_path: base.join("host.key"),
60                session_dir_path: base.join("sessions"),
61            })
62        } else if cfg!(target_os = "windows") {
63            let base = PathBuf::from(r"C:\ProgramData").join(tool_id.as_str());
64            Ok(Self {
65                tool_id,
66                installed_warrant_path: base.join("warrant.toml"),
67                version_state_path: base.join("signing").join("version"),
68                signing_private_key_path: base.join("signing").join("private.key"),
69                signing_public_key_path: base.join("signing").join("public.key"),
70                host_secret_path: base.join("host.key"),
71                session_dir_path: base.join("sessions"),
72            })
73        } else {
74            let base = PathBuf::from("/etc").join(tool_id.as_str());
75            let session_dir_path = PathBuf::from("/run").join(tool_id.as_str());
76            Ok(Self {
77                tool_id,
78                installed_warrant_path: base.join("warrant.toml"),
79                version_state_path: base.join("signing").join("version"),
80                signing_private_key_path: base.join("signing").join("private.key"),
81                signing_public_key_path: base.join("signing").join("public.key"),
82                host_secret_path: base.join("host.key"),
83                session_dir_path,
84            })
85        }
86    }
87}
88
89#[cfg(test)]
90mod tests {
91    use super::ToolPaths;
92
93    #[test]
94    fn tool_path_rejects_traversal_inputs() {
95        for invalid in ["../../etc/passwd", "/etc/shadow", "..", "foo/bar", ""] {
96            let err = ToolPaths::for_tool(invalid).expect_err("invalid tool id must fail");
97            assert!(err.to_string().contains("invalid tool identifier"));
98        }
99    }
100
101    #[test]
102    fn tool_path_accepts_valid_inputs() {
103        ToolPaths::for_tool("valid-tool").expect("valid tool id");
104        ToolPaths::for_tool("my_tool.v2").expect("valid tool id");
105    }
106}