use std::collections::HashMap;
use std::path::PathBuf;
use std::sync::RwLock;
use serde::{Deserialize, Serialize};
use crate::config::config_dir;
use crate::errors::{Result, SshError};
#[derive(Default, Debug, Serialize, Deserialize)]
struct Stored {
#[serde(default)]
host: HashMap<String, Entry>,
}
#[derive(Debug, Serialize, Deserialize)]
struct Entry {
fingerprint: String,
}
pub enum KnownHostMatch {
Ok,
Mismatch {
expected: String,
},
Unknown,
Unavailable(String),
}
pub struct KnownHostsStore {
path: PathBuf,
inner: RwLock<Stored>,
flush_lock: tokio::sync::Mutex<()>,
}
impl KnownHostsStore {
pub fn open_or_create() -> Result<std::sync::Arc<Self>> {
let path = config_dir().join("known_hosts.toml");
let stored: Stored = if path.exists() {
let raw = std::fs::read_to_string(&path)?;
toml::from_str(&raw).map_err(|e| {
SshError::Config(format!(
"{}: parse failed ({e}). Move the file aside if you intend to reset.",
path.display()
))
})?
} else {
Stored::default()
};
Ok(std::sync::Arc::new(Self {
path,
inner: RwLock::new(stored),
flush_lock: tokio::sync::Mutex::new(()),
}))
}
fn endpoint_key(addr: &str, port: u16) -> String {
format!("{addr}:{port}")
}
pub fn check(&self, host: &str, addr: &str, port: u16, fingerprint: &str) -> KnownHostMatch {
let guard = match self.inner.read() {
Ok(g) => g,
Err(_) => {
return KnownHostMatch::Unavailable(
"known_hosts lock poisoned; refusing to re-pin a possibly changed key".into(),
);
}
};
let endpoint = Self::endpoint_key(addr, port);
if let Some(e) = guard.host.get(&endpoint) {
return if e.fingerprint == fingerprint {
KnownHostMatch::Ok
} else {
KnownHostMatch::Mismatch {
expected: e.fingerprint.clone(),
}
};
}
if let Some(e) = guard.host.get(host) {
return if e.fingerprint == fingerprint {
KnownHostMatch::Ok
} else {
KnownHostMatch::Mismatch {
expected: e.fingerprint.clone(),
}
};
}
KnownHostMatch::Unknown
}
pub async fn add(&self, host: &str, addr: &str, port: u16, fingerprint: &str) -> Result<()> {
{
let mut guard = self
.inner
.write()
.map_err(|_| SshError::Other("known_hosts lock poisoned".into()))?;
let endpoint = Self::endpoint_key(addr, port);
guard.host.insert(
endpoint,
Entry {
fingerprint: fingerprint.to_string(),
},
);
guard.host.remove(host);
}
self.flush().await
}
async fn flush(&self) -> Result<()> {
let _flush_guard = self.flush_lock.lock().await;
let serialized = {
let guard = self
.inner
.read()
.map_err(|_| SshError::Other("known_hosts lock poisoned".into()))?;
toml::to_string_pretty(&*guard)
.map_err(|e| SshError::Config(format!("serialize known_hosts: {e}")))?
};
let path = self.path.clone();
tokio::task::spawn_blocking(move || -> Result<()> {
if let Some(parent) = path.parent() {
std::fs::create_dir_all(parent)?;
}
let tmp = path.with_extension("toml.tmp");
std::fs::write(&tmp, serialized)?;
std::fs::rename(&tmp, &path)?;
#[cfg(unix)]
{
use std::os::unix::fs::PermissionsExt;
let _ = std::fs::set_permissions(&path, std::fs::Permissions::from_mode(0o600));
}
Ok(())
})
.await
.map_err(|e| SshError::Other(format!("known_hosts flush task: {e}")))?
}
}
#[cfg(test)]
mod tests {
use super::*;
fn store_with(entries: &[(&str, &str)]) -> std::sync::Arc<KnownHostsStore> {
let mut stored = Stored::default();
for (k, fp) in entries {
stored.host.insert(
(*k).to_string(),
Entry {
fingerprint: (*fp).to_string(),
},
);
}
std::sync::Arc::new(KnownHostsStore {
path: PathBuf::from("known_hosts.toml"),
inner: RwLock::new(stored),
flush_lock: tokio::sync::Mutex::new(()),
})
}
#[test]
fn endpoint_entry_beats_alias_entry() {
let store = store_with(&[("10.0.0.1:22", "SHA256:aaa"), ("box1", "SHA256:zzz")]);
assert!(matches!(
store.check("box1", "10.0.0.1", 22, "SHA256:aaa"),
KnownHostMatch::Ok
));
assert!(matches!(
store.check("box1", "10.0.0.1", 22, "SHA256:bbb"),
KnownHostMatch::Mismatch { .. }
));
}
#[test]
fn poisoned_lock_refuses_instead_of_reporting_unknown() {
let store = store_with(&[("10.0.0.1:22", "SHA256:aaa")]);
let poisoner = std::sync::Arc::clone(&store);
let prev = std::panic::take_hook();
std::panic::set_hook(Box::new(|_| {}));
let _ = std::thread::spawn(move || {
let _guard = poisoner.inner.write();
panic!("poison the lock");
})
.join();
std::panic::set_hook(prev);
assert!(store.inner.is_poisoned(), "lock should be poisoned");
assert!(matches!(
store.check("box1", "10.0.0.1", 22, "SHA256:aaa"),
KnownHostMatch::Unavailable(_)
));
}
}