1#![forbid(unsafe_code)]
4use crate::errors::{SshCliError, SshCliResult};
17use std::collections::BTreeMap;
18use std::fs::File;
19use std::path::{Path, PathBuf};
20
21#[must_use]
28fn fingerprints_eq(a: &str, b: &str) -> bool {
29 let a = a.as_bytes();
30 let b = b.as_bytes();
31 if a.len() != b.len() {
32 return false;
33 }
34 let mut diff = 0u8;
35 for (x, y) in a.iter().zip(b.iter()) {
36 diff |= x ^ y;
37 }
38 std::hint::black_box(diff) == 0
40}
41
42#[derive(Debug, Default, Clone)]
44pub struct KnownHosts {
45 entries: BTreeMap<String, String>,
46 path: PathBuf,
47}
48
49struct HostsFileLock {
51 file: File,
52}
53
54impl HostsFileLock {
55 fn acquire(kh_path: &Path) -> SshCliResult<Self> {
57 let lock_path = {
58 let mut os = kh_path.as_os_str().to_owned();
59 os.push(".lock");
60 PathBuf::from(os)
61 };
62 if let Some(parent_dir) = lock_path.parent() {
63 std::fs::create_dir_all(parent_dir)?;
64 }
65 let file = std::fs::OpenOptions::new()
66 .create(true)
67 .truncate(false)
68 .read(true)
69 .write(true)
70 .open(&lock_path)
71 .map_err(SshCliError::Io)?;
72 let _ = crate::fs_perm::set_secret_file_mode(&lock_path);
74 fs2::FileExt::lock_exclusive(&file).map_err(SshCliError::Io)?;
75 Ok(Self { file })
76 }
77}
78
79impl Drop for HostsFileLock {
80 fn drop(&mut self) {
81 let _ = fs2::FileExt::unlock(&self.file);
82 }
83}
84
85impl KnownHosts {
86 #[must_use]
88 pub fn key(host: &str, port: u16) -> String {
89 format!("{host}:{port}")
90 }
91
92 pub fn load(path: PathBuf) -> SshCliResult<Self> {
94 let mut entries = BTreeMap::new();
95 if path.exists() {
96 let text = crate::paths::read_text_capped(
97 &path,
98 crate::paths::MAX_KNOWN_HOSTS_BYTES,
99 )?;
100 for line in text.lines() {
101 let line = line.trim();
102 if line.is_empty() || line.starts_with('#') {
103 continue;
104 }
105 let mut parts = line.split_whitespace();
106 if let (Some(k), Some(fp)) = (parts.next(), parts.next()) {
107 entries.insert(k.to_string(), fp.to_string());
108 }
109 }
110 }
111 Ok(Self { entries, path })
112 }
113
114 #[must_use]
116 pub fn path_beside_config(config_toml: &Path) -> PathBuf {
117 config_toml
118 .parent()
119 .map(|p| p.join(crate::constants::KNOWN_HOSTS_FILE_NAME))
120 .unwrap_or_else(|| PathBuf::from(crate::constants::KNOWN_HOSTS_FILE_NAME))
121 }
122
123 #[must_use]
125 pub fn get(&self, host: &str, port: u16) -> Option<&str> {
126 self.entries
127 .get(&Self::key(host, port))
128 .map(String::as_str)
129 }
130
131 pub fn store(&mut self, host: &str, port: u16, fingerprint: &str) -> SshCliResult<()> {
136 let _lock = HostsFileLock::acquire(&self.path)?;
137 self.reload_from_disk_unlocked()?;
138 self.entries
139 .insert(Self::key(host, port), fingerprint.to_string());
140 self.persist_unlocked()
141 }
142
143 fn reload_from_disk_unlocked(&mut self) -> SshCliResult<()> {
144 let fresh = Self::load(self.path.clone())?;
145 self.entries = fresh.entries;
146 Ok(())
147 }
148
149 fn persist_unlocked(&self) -> SshCliResult<()> {
150 if let Some(parent_dir) = self.path.parent() {
151 std::fs::create_dir_all(parent_dir)?;
152 }
153 let mut body = String::new();
154 body.push_str("# ssh-cli known_hosts (TOFU)\n");
155 for (k, v) in &self.entries {
156 body.push_str(k);
157 body.push(' ');
158 body.push_str(v);
159 body.push('\n');
160 }
161
162 let parent_dir = self
163 .path
164 .parent()
165 .map(Path::to_path_buf)
166 .unwrap_or_else(|| PathBuf::from("."));
167 let mut tmp = tempfile::NamedTempFile::new_in(&parent_dir).map_err(SshCliError::Io)?;
168 use std::io::Write;
169 tmp.write_all(body.as_bytes())?;
170 tmp.as_file().sync_data()?;
171 tmp.persist(&self.path)
172 .map_err(|e| SshCliError::Io(e.error))?;
173
174 crate::fs_perm::set_secret_file_mode(&self.path)?;
175 Ok(())
176 }
177}
178
179pub fn verify_tofu(
191 kh: &mut KnownHosts,
192 host: &str,
193 port: u16,
194 fingerprint: &str,
195 replace: bool,
196) -> SshCliResult<bool> {
197 let _lock = HostsFileLock::acquire(&kh.path)?;
198 kh.reload_from_disk_unlocked()?;
199 match kh.get(host, port).map(str::to_string) {
200 None => {
201 kh.entries
202 .insert(KnownHosts::key(host, port), fingerprint.to_string());
203 kh.persist_unlocked()?;
204 Ok(true)
205 }
206 Some(existing) if fingerprints_eq(&existing, fingerprint) => Ok(true),
207 Some(existing) if replace => {
208 tracing::warn!(
209 host,
210 port,
211 old = %existing,
212 novo = %fingerprint,
213 "replacing host key (--replace-host-key)"
214 );
215 kh.entries
216 .insert(KnownHosts::key(host, port), fingerprint.to_string());
217 kh.persist_unlocked()?;
218 Ok(true)
219 }
220 Some(existing) => Err(SshCliError::HostKeyChanged {
221 host: host.to_string(),
222 port,
223 expected: existing,
224 obtained: fingerprint.to_string(),
225 }),
226 }
227}
228
229#[cfg(test)]
230mod tests {
231 use super::*;
232 use std::sync::{Arc, Barrier};
233 use std::thread;
234 use tempfile::TempDir;
235
236 #[test]
237 fn fingerprints_eq_matches_and_rejects() {
238 assert!(fingerprints_eq("abc", "abc"));
239 assert!(!fingerprints_eq("abc", "abd"));
240 assert!(!fingerprints_eq("abc", "ab"));
241 assert!(!fingerprints_eq("ab", "abc"));
242 }
243
244 #[test]
245 fn tofu_stores_and_accepts_same() {
246 let tmp = TempDir::new().unwrap();
247 let path = tmp.path().join("known_hosts");
248 let mut kh = KnownHosts::load(path).unwrap();
249 assert!(verify_tofu(&mut kh, "h", 22, "fp1", false).unwrap());
250 assert!(verify_tofu(&mut kh, "h", 22, "fp1", false).unwrap());
251 }
252
253 #[test]
254 fn tofu_rejects_change() {
255 let tmp = TempDir::new().unwrap();
256 let path = tmp.path().join("known_hosts");
257 let mut kh = KnownHosts::load(path).unwrap();
258 verify_tofu(&mut kh, "h", 22, "fp1", false).unwrap();
259 let err = verify_tofu(&mut kh, "h", 22, "fp2", false).unwrap_err();
260 assert!(matches!(err, SshCliError::HostKeyChanged { .. }));
261 }
262
263 #[test]
264 fn tofu_replaces_with_flag() {
265 let tmp = TempDir::new().unwrap();
266 let path = tmp.path().join("known_hosts");
267 let mut kh = KnownHosts::load(path).unwrap();
268 verify_tofu(&mut kh, "h", 22, "fp1", false).unwrap();
269 assert!(verify_tofu(&mut kh, "h", 22, "fp2", true).unwrap());
270 assert_eq!(kh.get("h", 22), Some("fp2"));
271 }
272
273 #[test]
275 fn concurrent_store_merges_both_hosts() {
276 let tmp = TempDir::new().unwrap();
277 let path = Arc::new(tmp.path().join("known_hosts"));
278 let barrier = Arc::new(Barrier::new(2));
279 let p1 = Arc::clone(&path);
280 let p2 = Arc::clone(&path);
281 let b1 = Arc::clone(&barrier);
282 let b2 = Arc::clone(&barrier);
283
284 let t1 = thread::spawn(move || {
285 let mut kh = KnownHosts::load((*p1).clone()).unwrap();
286 b1.wait();
287 verify_tofu(&mut kh, "alpha.example", 22, "fp-alpha", false).unwrap();
288 });
289 let t2 = thread::spawn(move || {
290 let mut kh = KnownHosts::load((*p2).clone()).unwrap();
291 b2.wait();
292 verify_tofu(&mut kh, "beta.example", 22, "fp-beta", false).unwrap();
293 });
294 t1.join().unwrap();
295 t2.join().unwrap();
296
297 let final_kh = KnownHosts::load((*path).clone()).unwrap();
298 assert_eq!(final_kh.get("alpha.example", 22), Some("fp-alpha"));
299 assert_eq!(final_kh.get("beta.example", 22), Some("fp-beta"));
300 assert_eq!(final_kh.entries.len(), 2);
301 }
302}