Skip to main content

toride_ssh_config/
lib.rs

1//! SSH config file parsing, editing, and host resolution.
2
3pub mod ast;
4pub mod cache;
5mod directives;
6mod editor;
7mod managed;
8mod parse;
9pub mod resolve;
10pub mod sshd;
11
12use std::collections::HashMap;
13use std::path::{Path, PathBuf};
14
15use toride_ssh_core::Result;
16use toride_ssh_core::SshPaths;
17use toride_ssh_core::{Diagnostic, Error, Severity};
18
19pub use resolve::ResolvedHost;
20
21/// SSH config file operations.
22pub struct ConfigService<'a> {
23    paths: &'a SshPaths,
24}
25
26impl<'a> ConfigService<'a> {
27    #[must_use]
28    pub fn new(paths: &'a SshPaths) -> Self {
29        Self { paths }
30    }
31
32    /// Load and parse the SSH config into a lossless AST (cached); returns an
33    /// empty AST when the file is missing.
34    ///
35    /// # Errors
36    /// [`Error::Io`] when the file is unreadable.
37    pub async fn load(&self) -> Result<ast::ConfigAst> {
38        let path = self.paths.config_path();
39        if !path.exists() {
40            return Ok(ast::ConfigAst { nodes: Vec::new() });
41        }
42        let path = path.to_path_buf();
43        let ast = tokio::task::spawn_blocking(move || cache::load_cached_ast(&path))
44            .await
45            .map_err(|e| Error::TaskFailed(e.to_string()))??;
46        Ok((*ast).clone())
47    }
48
49    /// Save the AST atomically with `0o600` permissions.
50    ///
51    /// # Errors
52    /// [`Error::ConfigWriteFailed`] (write/rename) or [`Error::Io`] (chmod).
53    pub async fn save(&self, ast: &ast::ConfigAst) -> Result<()> {
54        let path = self.paths.config_path();
55        let content = ast.to_string_lossless();
56
57        if path.exists() {
58            let backup_path = path.with_extension("config.bak");
59            if let Err(e) = std::fs::copy(path, &backup_path) {
60                tracing::warn!("failed to back up config to {}: {e}", backup_path.display());
61            }
62        }
63
64        let parent = path.parent().unwrap_or_else(|| std::path::Path::new("."));
65        let tmp_path = parent.join(format!(
66            ".config.tmp.{}.{}",
67            std::process::id(),
68            std::time::SystemTime::now()
69                .duration_since(std::time::UNIX_EPOCH)
70                .unwrap_or_default()
71                .as_nanos()
72        ));
73        tokio::fs::write(&tmp_path, &content).await?;
74
75        #[cfg(unix)]
76        {
77            use std::os::unix::fs::PermissionsExt;
78            let perms = std::fs::Permissions::from_mode(0o600);
79            tokio::fs::set_permissions(&tmp_path, perms).await?;
80        }
81
82        tokio::fs::rename(&tmp_path, path).await.map_err(|e| {
83            let _ = std::fs::remove_file(&tmp_path);
84            toride_ssh_core::Error::ConfigWriteFailed(format!("failed to rename config: {e}"))
85        })?;
86
87        Ok(())
88    }
89
90    /// Resolve a host alias to a [`ResolvedHost`], expanding Includes and
91    /// tokens.
92    ///
93    /// # Errors
94    /// [`Error::ConfigIncludeCycle`] or [`Error::Io`].
95    pub async fn resolve_host(&self, host: &str) -> Result<ResolvedHost> {
96        resolve::resolve(self.paths.ssh_dir(), host, None).await
97    }
98
99    /// Parse the config into an [`ssh2_config_rs::SshConfig`] (supports
100    /// `.query(host)`).
101    ///
102    /// # Errors
103    /// [`Error::ConfigParseFailed`] or [`Error::Io`].
104    pub async fn parse_typed(&self) -> Result<ssh2_config_rs::SshConfig> {
105        parse::parse_config(self.paths.config_path()).await
106    }
107
108    /// Get a directive value for a host from the AST; first match wins.
109    #[must_use]
110    pub fn get_host_directive(ast: &ast::ConfigAst, host: &str, key: &str) -> Option<String> {
111        directives::get_directive(ast, host, key)
112    }
113
114    /// Get all directives for a host.
115    #[must_use]
116    pub fn get_all_host_directives(ast: &ast::ConfigAst, host: &str) -> Vec<(String, String)> {
117        directives::get_all_directives(ast, host)
118    }
119
120    /// Add a new Host block.
121    ///
122    /// # Errors
123    /// [`Error::DuplicateHost`] if one with the given name already exists.
124    pub fn add_host(
125        ast: &mut ast::ConfigAst,
126        name: &str,
127        directives: Vec<(String, String)>,
128    ) -> Result<()> {
129        editor::add_host(ast, name, directives)
130    }
131
132    /// Remove a Host block by name.
133    ///
134    /// # Errors
135    /// [`Error::HostNotFound`] if absent.
136    pub fn remove_host(ast: &mut ast::ConfigAst, name: &str) -> Result<()> {
137        editor::remove_host(ast, name)
138    }
139
140    /// Rename a Host block.
141    ///
142    /// # Errors
143    /// [`Error::HostNotFound`] if `old_name` is absent, or
144    /// [`Error::DuplicateHost`] if `new_name` exists.
145    pub fn rename_host(ast: &mut ast::ConfigAst, old_name: &str, new_name: &str) -> Result<()> {
146        editor::rename_host(ast, old_name, new_name)
147    }
148
149    /// Add a managed block (or replace an existing one).
150    pub fn upsert_managed_block(
151        ast: &mut ast::ConfigAst,
152        name: &str,
153        directives: Vec<(String, String)>,
154    ) {
155        managed::upsert_managed_block(ast, name, directives);
156    }
157
158    /// Remove a managed block by name.
159    ///
160    /// # Errors
161    /// [`Error::ManagedBlockNotFound`] if absent.
162    pub fn remove_managed_block(ast: &mut ast::ConfigAst, name: &str) -> Result<()> {
163        managed::remove_managed_block(ast, name)
164    }
165
166    /// List all managed block names.
167    #[must_use]
168    pub fn list_managed_blocks(ast: &ast::ConfigAst) -> Vec<String> {
169        managed::list_managed_blocks(ast)
170    }
171
172    /// Create `~/.ssh` (`0o700`) and the config file (`0o600`) if missing.
173    ///
174    /// # Errors
175    /// [`Error::Io`] on failure.
176    pub async fn ensure_config_file(&self) -> Result<()> {
177        let path = self.paths.config_path();
178        if !path.exists() {
179            tokio::fs::create_dir_all(self.paths.ssh_dir()).await?;
180            tokio::fs::write(&path, "").await?;
181
182            #[cfg(unix)]
183            {
184                use std::os::unix::fs::PermissionsExt;
185                tokio::fs::set_permissions(path, std::fs::Permissions::from_mode(0o600)).await?;
186                tokio::fs::set_permissions(
187                    self.paths.ssh_dir(),
188                    std::fs::Permissions::from_mode(0o700),
189                )
190                .await?;
191            }
192        }
193        Ok(())
194    }
195
196    /// Get the path to the config file.
197    #[must_use]
198    pub fn config_path(&self) -> &Path {
199        self.paths.config_path()
200    }
201
202    /// Load, mutate via `f`, then save.
203    ///
204    /// # Errors
205    /// From any step: ensure, load, `f`, or save.
206    pub async fn edit<F>(&self, f: F) -> Result<()>
207    where
208        F: FnOnce(&mut ast::ConfigAst) -> Result<()>,
209    {
210        self.ensure_config_file().await?;
211        let mut ast = self.load().await?;
212        f(&mut ast)?;
213        self.save(&ast).await
214    }
215
216    /// Diagnose the config: proxy conflicts, duplicate aliases, `Host *`
217    /// placement, and missing or `.pub` `IdentityFile`s.
218    ///
219    /// # Errors
220    /// [`Error::Io`].
221    pub async fn diagnose(&self) -> Result<Vec<Diagnostic>> {
222        let ast = self.load().await?;
223        let ssh_dir = self.paths.ssh_dir();
224        let mut diagnostics = Vec::new();
225
226        let mut seen_patterns: HashMap<String, String> = HashMap::new();
227
228        let mut star_index: Option<usize> = None;
229        let mut last_specific_index: Option<usize> = None;
230
231        for (i, node) in ast.nodes.iter().enumerate() {
232            let ast::ConfigNode::HostBlock(b) = node else {
233                continue;
234            };
235
236            check_proxy_conflict(&b.header, &b.nodes, &mut diagnostics);
237            check_duplicate_aliases(&b.header, &b.patterns, &mut seen_patterns, &mut diagnostics);
238
239            if b.patterns.iter().any(|p| p == "*") {
240                if star_index.is_none() {
241                    star_index = Some(i);
242                }
243            } else if !b.patterns.is_empty() {
244                last_specific_index = Some(i);
245            }
246
247            check_identity_files(&b.header, &b.nodes, ssh_dir, &mut diagnostics);
248        }
249
250        check_host_star_placement(star_index, last_specific_index, &mut diagnostics);
251
252        Ok(diagnostics)
253    }
254}
255
256fn check_proxy_conflict(
257    header: &str,
258    nodes: &[ast::ConfigNode],
259    diagnostics: &mut Vec<Diagnostic>,
260) {
261    let has_proxy_command = nodes.iter().any(|n| {
262        matches!(
263            n,
264            ast::ConfigNode::Directive(d)
265                if d.keyword.eq_ignore_ascii_case("ProxyCommand")
266        )
267    });
268    let has_proxy_jump = nodes.iter().any(|n| {
269        matches!(
270            n,
271            ast::ConfigNode::Directive(d)
272                if d.keyword.eq_ignore_ascii_case("ProxyJump")
273        )
274    });
275    if has_proxy_command && has_proxy_jump {
276        diagnostics.push(Diagnostic {
277            id: "config_proxy_conflict",
278            severity: Severity::Warning,
279            message: format!("Host block '{header}' has both ProxyCommand and ProxyJump set"),
280            hint: Some(
281                "ProxyJump takes precedence over ProxyCommand; \
282                 remove one to avoid confusion"
283                    .into(),
284            ),
285            module: "config",
286        });
287    }
288}
289
290fn check_duplicate_aliases(
291    header: &str,
292    patterns: &[String],
293    seen_patterns: &mut HashMap<String, String>,
294    diagnostics: &mut Vec<Diagnostic>,
295) {
296    for pat in patterns {
297        if pat == "*" {
298            continue;
299        }
300        if let Some(first_header) = seen_patterns.get(pat) {
301            diagnostics.push(Diagnostic {
302                id: "config_duplicate_alias",
303                severity: Severity::Warning,
304                message: format!(
305                    "Host alias '{pat}' appears in both '{first_header}' and '{header}'",
306                ),
307                hint: Some(format!("Merge or remove the duplicate entry for '{pat}'")),
308                module: "config",
309            });
310        } else {
311            seen_patterns.insert(pat.clone(), header.to_owned());
312        }
313    }
314}
315
316fn check_identity_files(
317    header: &str,
318    nodes: &[ast::ConfigNode],
319    ssh_dir: &Path,
320    diagnostics: &mut Vec<Diagnostic>,
321) {
322    for child in nodes {
323        if let ast::ConfigNode::Directive(d) = child
324            && d.keyword.eq_ignore_ascii_case("IdentityFile")
325        {
326            if d.value.to_lowercase().ends_with(".pub") {
327                diagnostics.push(Diagnostic {
328                    id: "config_identity_pub",
329                    severity: Severity::Warning,
330                    message: format!(
331                        "IdentityFile '{}' in '{header}' points to a public key \
332                         (.pub file)",
333                        d.value,
334                    ),
335                    hint: Some(
336                        "IdentityFile should reference the private key, \
337                         not the .pub file"
338                            .into(),
339                    ),
340                    module: "config",
341                });
342            }
343
344            let expanded = expand_identity_path(&d.value, ssh_dir);
345            if !expanded.exists() {
346                diagnostics.push(Diagnostic {
347                    id: "config_identity_missing",
348                    severity: Severity::Warning,
349                    message: format!(
350                        "IdentityFile '{}' in '{header}' does not exist \
351                         (resolved: {})",
352                        d.value,
353                        expanded.display()
354                    ),
355                    hint: Some(format!(
356                        "Generate the missing key or update the \
357                         IdentityFile entry in '{header}'",
358                    )),
359                    module: "config",
360                });
361            }
362        }
363    }
364}
365
366fn check_host_star_placement(
367    star_index: Option<usize>,
368    last_specific_index: Option<usize>,
369    diagnostics: &mut Vec<Diagnostic>,
370) {
371    if let (Some(star), Some(last)) = (star_index, last_specific_index)
372        && star < last
373    {
374        diagnostics.push(Diagnostic {
375            id: "config_host_star_placement",
376            severity: Severity::Warning,
377            message: "'Host *' appears before specific Host blocks; \
378                 later blocks cannot override its defaults"
379                .into(),
380            hint: Some(
381                "Move 'Host *' to the end of the config file so \
382                 specific blocks take precedence"
383                    .into(),
384            ),
385            module: "config",
386        });
387    }
388}
389
390/// Expand `~` and resolve a relative `IdentityFile` value against `ssh_dir`.
391#[must_use]
392pub fn expand_identity_path(raw: &str, ssh_dir: &Path) -> PathBuf {
393    toride_ssh_core::paths::expand_path(raw, ssh_dir)
394}
395
396/// Check if a hostname matches any of the given SSH config patterns.
397pub fn host_matches(host: &str, patterns: &[impl AsRef<str>]) -> bool {
398    directives::host_matches_patterns(host, patterns)
399}
400
401/// Check if a path is inside the `~/.ssh` directory.
402#[must_use]
403pub fn is_in_ssh_dir(path: &Path, ssh_dir: &Path) -> bool {
404    path.starts_with(ssh_dir)
405}
406
407#[cfg(test)]
408#[path = "mod.test.rs"]
409mod tests;