Skip to main content

toride_ssh_agent/
client.rs

1//! SSH agent client wrapping `ssh-agent-lib` and `ssh-add` CLI.
2//!
3//! When the `agent` feature is enabled, the native SSH agent protocol is used
4//! for listing identities. Key add/remove operations use the `ssh-add` CLI to
5//! avoid type mismatches between `ssh-agent-lib`'s `ssh-key 0.6` and our
6//! `ssh-key 0.7`.
7
8use std::path::Path;
9
10use toride_ssh_core::{CliRunner, Error, Fingerprint, KeySource, KeyType, Result, SshKey};
11
12/// Connect to the SSH agent via `SSH_AUTH_SOCK`.
13///
14/// Returns a boxed [`Session`](ssh_agent_lib::agent::Session) trait object
15/// that can be used to interact with the agent. Returns
16/// [`Error::AgentNotAvailable`] when `SSH_AUTH_SOCK` is unset or the socket
17/// does not exist.
18#[cfg(feature = "native")]
19pub async fn connect() -> Result<Box<dyn ssh_agent_lib::agent::Session>> {
20    let socket_path = std::env::var("SSH_AUTH_SOCK").map_err(|_| Error::AgentNotAvailable)?;
21
22    let path = std::path::PathBuf::from(&socket_path);
23    if !path.exists() {
24        return Err(Error::AgentNotAvailable);
25    }
26
27    // Verify SSH_AUTH_SOCK points to an actual socket, not a regular file or FIFO.
28    #[cfg(unix)]
29    {
30        use std::os::unix::fs::FileTypeExt;
31        if let Ok(meta) = std::fs::metadata(&path)
32            && !meta.file_type().is_socket()
33        {
34            return Err(Error::AgentOperationFailed(format!(
35                "SSH_AUTH_SOCK ({socket_path}) is not a Unix socket"
36            )));
37        }
38    }
39
40    // Connect directly via UnixStream to avoid pulling in service_binding.
41    let stream = tokio::task::spawn_blocking(move || {
42        std::os::unix::net::UnixStream::connect(&path)
43            .map_err(|e| Error::AgentOperationFailed(e.to_string()))
44    })
45    .await
46    .map_err(|e| Error::AgentOperationFailed(e.to_string()))??;
47
48    let tokio_stream = tokio::net::UnixStream::from_std(stream)
49        .map_err(|e| Error::AgentOperationFailed(e.to_string()))?;
50
51    Ok(Box::new(ssh_agent_lib::client::Client::new(tokio_stream)))
52}
53
54/// List all identities currently loaded in the SSH agent.
55///
56/// Uses the native agent protocol when the `agent` feature is enabled,
57/// falling back to parsing `ssh-add -l` output otherwise.
58///
59/// # Errors
60///
61/// Returns [`Error::AgentNotAvailable`] if the agent is not running, or
62/// [`Error::AgentOperationFailed`] if the agent protocol or CLI command fails.
63pub async fn list_identities(runner: &dyn CliRunner) -> Result<Vec<SshKey>> {
64    #[cfg(feature = "native")]
65    {
66        match list_identities_native().await {
67            Ok(keys) => return Ok(keys),
68            // No agent at all — propagate immediately rather than trying CLI.
69            Err(Error::AgentNotAvailable) => return Err(Error::AgentNotAvailable),
70            // Other errors (e.g. protocol mismatch) — fall through to CLI.
71            Err(_) => {}
72        }
73    }
74
75    list_identities_via_cli(runner).await
76}
77
78/// Add a private key to the SSH agent via `ssh-add`.
79pub async fn add_key(key_path: &Path, runner: &dyn CliRunner) -> Result<()> {
80    let path_str = key_path
81        .to_str()
82        .ok_or_else(|| Error::AgentOperationFailed("key path is not valid UTF-8".into()))?
83        .to_owned();
84
85    runner
86        .run("ssh-add", vec![path_str])
87        .await
88        .map_err(|e| Error::AgentOperationFailed(e.to_string()))?;
89    Ok(())
90}
91
92/// Test whether a key is usable by the SSH agent (`ssh-add -T`).
93///
94/// Returns `Ok(true)` if the key is usable (exit code 0), `Ok(false)` if not
95/// (non-zero exit code), and `Err` if the command itself could not be run.
96///
97/// This is useful for checking whether a key that requires a passphrase has
98/// already been decrypted and loaded, or whether a hardware token key is
99/// accessible.
100pub async fn test_key_usability(key_path: &Path, runner: &dyn CliRunner) -> Result<bool> {
101    let path_str = key_path
102        .to_str()
103        .ok_or_else(|| Error::AgentOperationFailed("key path is not valid UTF-8".into()))?
104        .to_owned();
105
106    match runner.run("ssh-add", vec!["-T".to_owned(), path_str]).await {
107        Ok(_) => Ok(true),
108        Err(Error::CommandFailed(_)) => Ok(false),
109        Err(e) => Err(Error::AgentOperationFailed(e.to_string())),
110    }
111}
112
113/// Add a key to the SSH agent restricted to specific destinations (`ssh-add -h`).
114///
115/// The `hosts` slice specifies the allowed destinations. Only connections to
116/// these hosts will be authorized to use the key. The hosts are joined with
117/// `>` to form the `ssh-add -h` argument.
118///
119/// # Errors
120///
121/// Returns an error if `hosts` is empty, if the key cannot be added, or if
122/// the command fails.
123pub async fn destination_constrained_add(
124    key_path: &Path,
125    hosts: &[&str],
126    runner: &dyn CliRunner,
127) -> Result<()> {
128    if hosts.is_empty() {
129        return Err(Error::AgentOperationFailed(
130            "destination-constrained add requires at least one host".into(),
131        ));
132    }
133
134    let path_str = key_path
135        .to_str()
136        .ok_or_else(|| Error::AgentOperationFailed("key path is not valid UTF-8".into()))?
137        .to_owned();
138
139    let constraint = hosts.join(">");
140    runner
141        .run("ssh-add", vec!["-h".to_owned(), constraint, path_str])
142        .await
143        .map_err(|e| Error::AgentOperationFailed(e.to_string()))?;
144    Ok(())
145}
146
147/// Remove a key from the SSH agent via `ssh-add -d`.
148pub async fn remove_key(key_path: &Path, runner: &dyn CliRunner) -> Result<()> {
149    let pub_path = key_path.with_extension("pub");
150    let path = if pub_path.exists() {
151        pub_path
152    } else {
153        key_path.to_path_buf()
154    };
155
156    let path_str = path
157        .to_str()
158        .ok_or_else(|| Error::AgentOperationFailed("key path is not valid UTF-8".into()))?
159        .to_owned();
160
161    runner
162        .run("ssh-add", vec!["-d".to_owned(), path_str])
163        .await
164        .map_err(|e| Error::AgentOperationFailed(e.to_string()))?;
165    Ok(())
166}
167
168// ---------------------------------------------------------------------------
169// Native implementation (agent feature)
170// ---------------------------------------------------------------------------
171
172#[cfg(feature = "native")]
173async fn list_identities_native() -> Result<Vec<SshKey>> {
174    let mut client = connect().await?;
175    let identities = client
176        .request_identities()
177        .await
178        .map_err(|e| Error::AgentOperationFailed(e.to_string()))?;
179
180    let mut keys = Vec::with_capacity(identities.len());
181    for identity in identities {
182        let key_data = identity.credential.key_data();
183        let alg_str = key_data.algorithm().to_string();
184        let Some(key_type) = parse_key_type_from_algorithm(&alg_str) else {
185            tracing::warn!("skipping agent key with unknown algorithm: {alg_str}");
186            continue;
187        };
188
189        let fingerprint = match encode_key_data(key_data) {
190            Ok(bytes) => Some(compute_sha256_fingerprint(&bytes, key_type)),
191            Err(e) => {
192                tracing::warn!("failed to encode agent key data: {e}");
193                None
194            }
195        };
196
197        keys.push(SshKey {
198            path: std::path::PathBuf::from(if identity.comment.is_empty() {
199                format!("agent:{key_type:?}")
200            } else {
201                format!("agent:{}", identity.comment)
202            }),
203            key_type,
204            fingerprint,
205            comment: if identity.comment.is_empty() {
206                None
207            } else {
208                Some(identity.comment)
209            },
210            encrypted: false,
211            source: KeySource::Agent,
212            permissions: None,
213            has_public_pair: false,
214            has_certificate: false,
215            last_modified: None,
216            used_by_hosts: Vec::new(),
217            key_format: None,
218        });
219    }
220
221    Ok(keys)
222}
223
224/// Encode `ssh_key 0.6` `KeyData` to bytes using the `Encode` trait.
225#[cfg(feature = "native")]
226fn encode_key_data(key_data: &ssh_agent_lib::ssh_key::public::KeyData) -> Result<Vec<u8>> {
227    use ssh_agent_lib::ssh_encoding::Encode;
228
229    let len = key_data
230        .encoded_len()
231        .map_err(|e| Error::AgentOperationFailed(format!("encoded_len failed: {e}")))?;
232    let mut buf = Vec::with_capacity(len);
233    key_data
234        .encode(&mut buf)
235        .map_err(|e| Error::AgentOperationFailed(format!("encode failed: {e}")))?;
236    Ok(buf)
237}
238
239/// Compute a SHA-256 fingerprint from raw public key bytes.
240#[allow(dead_code, reason = "public API helper, not exercised in-workspace")]
241fn compute_sha256_fingerprint(bytes: &[u8], key_type: KeyType) -> Fingerprint {
242    use base64::Engine;
243    use base64::engine::general_purpose::STANDARD_NO_PAD;
244    use ssh_key::sha2::{Digest, Sha256};
245
246    let hash = Sha256::digest(bytes);
247    Fingerprint {
248        hash: STANDARD_NO_PAD.encode(hash),
249        key_type,
250    }
251}
252
253/// Map an algorithm name string (e.g. `"ssh-ed25519"`, `"ssh-rsa"`) to [`KeyType`].
254///
255/// Returns `None` for unknown algorithms so callers can decide how to handle them
256/// rather than silently misidentifying the key type.
257#[allow(dead_code, reason = "public API helper, not exercised in-workspace")]
258pub(crate) fn parse_key_type_from_algorithm(alg: &str) -> Option<KeyType> {
259    match alg {
260        "ssh-ed25519" => Some(KeyType::Ed25519),
261        "ssh-rsa" => Some(KeyType::Rsa { bits: 0 }),
262        "ecdsa-sha2-nistp256" => Some(KeyType::EcdsaP256),
263        "ecdsa-sha2-nistp384" => Some(KeyType::EcdsaP384),
264        "ecdsa-sha2-nistp521" => Some(KeyType::EcdsaP521),
265        "ssh-dss" => Some(KeyType::Dsa),
266        "sk-ssh-ed25519@openssh.com" => Some(KeyType::SkEd25519),
267        "sk-ecdsa-sha2-nistp256@openssh.com" => Some(KeyType::SkEcdsaP256),
268        _ => {
269            tracing::warn!("unknown SSH key algorithm \"{alg}\"");
270            None
271        }
272    }
273}
274
275// ---------------------------------------------------------------------------
276// CLI fallback
277// ---------------------------------------------------------------------------
278
279/// Parse `ssh-add -l` output into a list of [`SshKey`] values.
280async fn list_identities_via_cli(runner: &dyn CliRunner) -> Result<Vec<SshKey>> {
281    let output = runner.run("ssh-add", vec!["-l".to_owned()]).await?;
282    Ok(output.lines().filter_map(parse_ssh_add_line).collect())
283}
284
285/// Parse a single `ssh-add -l` output line.
286///
287/// Format: `<bits> SHA256:<hash> <comment> (<type>)`
288pub(crate) fn parse_ssh_add_line(line: &str) -> Option<SshKey> {
289    let line = line.trim();
290    if line.is_empty() || line.contains("The agent has no identities") {
291        return None;
292    }
293
294    let (_bits, rest) = line.split_once(' ')?;
295    let rest = rest.trim();
296
297    // Extract parenthesised key type from the end.
298    let (rest, key_type_opt) = if let Some(start) = rest.rfind('(') {
299        if let Some(end) = rest[start..].find(')') {
300            let kt = &rest[start + 1..start + end];
301            (rest[..start].trim_end(), Some(kt))
302        } else {
303            (rest, None)
304        }
305    } else {
306        (rest, None)
307    };
308
309    let key_type = key_type_opt.and_then(parse_key_type_from_display)?;
310
311    // Split fingerprint from comment, keeping both as borrowed slices.
312    let (fingerprint_part, comment_part) = if let Some(space) = rest.find(' ') {
313        let (fp, c) = rest.split_at(space);
314        let c = c.trim();
315        (fp, if c.is_empty() { None } else { Some(c) })
316    } else {
317        (rest, None)
318    };
319
320    let hash = fingerprint_part
321        .strip_prefix("SHA256:")
322        .unwrap_or(fingerprint_part)
323        .to_string();
324
325    Some(SshKey {
326        path: std::path::PathBuf::from(format!("agent:{}", comment_part.unwrap_or("unknown"))),
327        key_type,
328        fingerprint: Some(Fingerprint { hash, key_type }),
329        comment: comment_part.map(str::to_owned),
330        encrypted: false,
331        source: KeySource::Agent,
332        permissions: None,
333        has_public_pair: false,
334        has_certificate: false,
335        last_modified: None,
336        used_by_hosts: Vec::new(),
337        key_format: None,
338    })
339}
340
341/// Map a display key type like "ED25519" or "RSA" to [`KeyType`].
342///
343/// These strings come from `ssh-add -l` output, e.g. `(ED25519)`, `(RSA)`,
344/// `(ECDSA)`, `(ED25519-SK)`, `(ECDSA-SK)`.
345fn parse_key_type_from_display(s: &str) -> Option<KeyType> {
346    if s.eq_ignore_ascii_case("ED25519") {
347        Some(KeyType::Ed25519)
348    } else if s.eq_ignore_ascii_case("ED25519-SK") {
349        Some(KeyType::SkEd25519)
350    } else if s.eq_ignore_ascii_case("RSA") {
351        Some(KeyType::Rsa { bits: 0 })
352    } else if s.eq_ignore_ascii_case("ECDSA") {
353        Some(KeyType::EcdsaP256)
354    } else if s.eq_ignore_ascii_case("ECDSA-SK") {
355        Some(KeyType::SkEcdsaP256)
356    } else if s.eq_ignore_ascii_case("DSA") {
357        Some(KeyType::Dsa)
358    } else {
359        None
360    }
361}
362
363#[cfg(test)]
364#[path = "client.test.rs"]
365mod tests;