use toride_ssh_core::{Result, privilege::PrivilegedOp, run_privileged};
use super::ast::{ConfigAst, ConfigNode, DirectiveData, Separator, parse};
const SSHD_CONFIG_PATH: &str = "/etc/ssh/sshd_config";
const SSHD_EDIT_LOCK_ENV: &str = "TORIDE_SSHD_EDIT_LOCK";
const SSHD_EDIT_LOCK_DEFAULT: &str = "/tmp/toride-sshd-config.lock";
fn edit_lock_path() -> std::path::PathBuf {
match std::env::var_os(SSHD_EDIT_LOCK_ENV) {
Some(p) if !p.is_empty() => std::path::PathBuf::from(p),
_ => std::path::PathBuf::from(SSHD_EDIT_LOCK_DEFAULT),
}
}
fn ensure_lock_file(path: &std::path::Path) {
let created = std::fs::OpenOptions::new()
.write(true)
.create_new(true)
.open(path);
if let Ok(f) = created {
#[cfg(unix)]
{
use std::os::unix::fs::PermissionsExt;
let _ = f.set_permissions(std::fs::Permissions::from_mode(0o666));
}
#[cfg(not(unix))]
{
let _ = f; }
}
#[cfg(unix)]
{
use std::os::unix::fs::PermissionsExt;
let _ = std::fs::set_permissions(path, std::fs::Permissions::from_mode(0o666));
}
}
pub async fn load() -> Result<ConfigAst> {
let path = std::path::Path::new(SSHD_CONFIG_PATH);
if !path.exists() {
return Ok(ConfigAst { nodes: Vec::new() });
}
let content = tokio::fs::read_to_string(path).await?;
Ok(parse(&content))
}
pub async fn save(ast: &ConfigAst, running_as_root: bool) -> Result<()> {
let content = ast.to_string_lossless();
run_privileged(PrivilegedOp::WriteSshdConfig { content }, running_as_root).await
}
#[cfg(test)]
fn with_edit_lock<T>(path: &std::path::Path, f: impl FnOnce() -> Result<T>) -> Result<T> {
ensure_lock_file(path);
toride_fs::with_lock(path, || {
f().map_err(|e| {
toride_fs::Error::Io(std::io::Error::other(format!(
"sshd_config edit critical section failed: {e}"
)))
})
})
.map_err(|e| {
toride_ssh_core::Error::Io(std::io::Error::other(format!(
"sshd_config edit lock failed: {e}"
)))
})
}
pub async fn edit<F>(running_as_root: bool, f: F) -> Result<()>
where
F: FnOnce(&mut ConfigAst) -> Result<()>,
{
let lock_path = edit_lock_path();
ensure_lock_file(&lock_path);
let file = std::fs::OpenOptions::new()
.read(true)
.write(true)
.create(true)
.truncate(false)
.open(&lock_path)
.map_err(|e| {
toride_ssh_core::Error::Io(std::io::Error::other(format!(
"sshd_config edit lock failed: cannot open lock file {}: {e}",
lock_path.display()
)))
})?;
let mut lock = fd_lock::RwLock::new(file);
let _guard = lock.write().map_err(|e| {
toride_ssh_core::Error::Io(std::io::Error::other(format!(
"sshd_config edit lock failed: cannot acquire lock on {}: {e}",
lock_path.display()
)))
})?;
tracing::debug!(path = %lock_path.display(), "sshd_config edit lock acquired");
let mut ast = load().await?;
f(&mut ast)?;
save(&ast, running_as_root).await
}
#[must_use]
pub fn get_allow_users(ast: &ConfigAst) -> Vec<String> {
collect_global_users(ast, "AllowUsers")
}
#[must_use]
pub fn get_deny_users(ast: &ConfigAst) -> Vec<String> {
collect_global_users(ast, "DenyUsers")
}
#[must_use]
pub fn get_allow_groups(ast: &ConfigAst) -> Vec<String> {
collect_global_users(ast, "AllowGroups")
}
#[must_use]
pub fn get_deny_groups(ast: &ConfigAst) -> Vec<String> {
collect_global_users(ast, "DenyGroups")
}
fn collect_global_users(ast: &ConfigAst, key: &str) -> Vec<String> {
let indices: Vec<usize> = ast
.nodes
.iter()
.enumerate()
.filter(
|(_, n)| matches!(n, ConfigNode::Directive(d) if d.keyword.eq_ignore_ascii_case(key)),
)
.map(|(i, _)| i)
.collect();
let leaked: Vec<usize> = directive_follows_match_or_host_set(&ast.nodes, &indices);
let leaked_set: std::collections::HashSet<usize> = leaked.into_iter().collect();
let mut out = Vec::new();
for &i in &indices {
if leaked_set.contains(&i) {
continue;
}
if let ConfigNode::Directive(d) = &ast.nodes[i] {
out.extend(d.value.split_whitespace().map(str::to_owned));
}
}
out
}
fn directive_follows_match_or_host(nodes: &[ConfigNode], directive_indices: &[usize]) -> bool {
!directive_follows_match_or_host_set(nodes, directive_indices).is_empty()
}
fn directive_follows_match_or_host_set(
nodes: &[ConfigNode],
directive_indices: &[usize],
) -> Vec<usize> {
let mut leaked = Vec::new();
for &idx in directive_indices {
let mut j = idx;
while j > 0 {
j -= 1;
match &nodes[j] {
ConfigNode::BlankLine | ConfigNode::Comment { .. } => {}
ConfigNode::MatchBlock(_) | ConfigNode::HostBlock(_) => {
leaked.push(idx);
break;
}
ConfigNode::Directive(_) => break,
}
}
}
leaked.sort_unstable();
leaked.dedup();
leaked
}
#[must_use]
pub fn has_pattern_tokens(value: &str) -> bool {
value
.split_whitespace()
.any(|tok| tok.contains('*') || tok.contains('?') || tok.contains('@'))
}
#[must_use]
pub fn directive_has_patterns(ast: &ConfigAst, key: &str) -> bool {
let indices: Vec<usize> = ast
.nodes
.iter()
.enumerate()
.filter(
|(_, n)| matches!(n, ConfigNode::Directive(d) if d.keyword.eq_ignore_ascii_case(key)),
)
.map(|(i, _)| i)
.collect();
let leaked = directive_follows_match_or_host_set(&ast.nodes, &indices);
let leaked_set: std::collections::HashSet<usize> = leaked.into_iter().collect();
indices.iter().any(|&i| {
if leaked_set.contains(&i) {
return false;
}
matches!(&ast.nodes[i], ConfigNode::Directive(d) if has_pattern_tokens(&d.value))
})
}
pub fn add_user_to_allow(ast: &mut ConfigAst, user: &str) -> Result<()> {
upsert_user_in_directive(ast, "AllowUsers", user, Action::Add)
}
pub fn remove_user_from_allow(ast: &mut ConfigAst, user: &str) -> Result<()> {
upsert_user_in_directive(ast, "AllowUsers", user, Action::Remove)
}
pub fn add_user_to_deny(ast: &mut ConfigAst, user: &str) -> Result<()> {
upsert_user_in_directive(ast, "DenyUsers", user, Action::Add)
}
pub fn remove_user_from_deny(ast: &mut ConfigAst, user: &str) -> Result<()> {
upsert_user_in_directive(ast, "DenyUsers", user, Action::Remove)
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
enum Action {
Add,
Remove,
}
fn upsert_user_in_directive(
ast: &mut ConfigAst,
key: &str,
user: &str,
action: Action,
) -> Result<()> {
let indices: Vec<usize> = ast
.nodes
.iter()
.enumerate()
.filter(
|(_, n)| matches!(n, ConfigNode::Directive(d) if d.keyword.eq_ignore_ascii_case(key)),
)
.map(|(i, _)| i)
.collect();
if directive_follows_match_or_host(&ast.nodes, &indices) {
return Err(toride_ssh_core::Error::SshdConfigInvalid(format!(
"refusing to edit {key}: directive follows a Match/Host block and its \
scope is ambiguous (unindented Match body)"
)));
}
if indices
.iter()
.any(|&i| matches!(&ast.nodes[i], ConfigNode::Directive(d) if has_pattern_tokens(&d.value)))
{
return Err(toride_ssh_core::Error::SshdConfigInvalid(format!(
"{key} uses pattern tokens (* ? @); refusing to edit by exact username"
)));
}
if indices.len() > 1 {
tracing::warn!(
"{key} appears {} times in sshd_config; merging into the first occurrence",
indices.len()
);
}
match indices.first() {
Some(&first) => {
let mut merged: Vec<String> = Vec::new();
for &i in &indices {
if let ConfigNode::Directive(d) = &ast.nodes[i] {
for tok in d.value.split_whitespace() {
if !merged.iter().any(|m| m == tok) {
merged.push(tok.to_owned());
}
}
}
}
apply_action_to_users(&mut merged, user, action);
if merged.is_empty() {
for &i in indices.iter().rev() {
ast.nodes.remove(i);
}
} else {
let joined = merged.join(" ");
let single = indices.len() == 1;
if let ConfigNode::Directive(d) = &mut ast.nodes[first] {
let value_changed = d.value != joined;
if value_changed || !single {
if d.comment.is_some() && !single {
tracing::warn!(
"{key}: discarding inline comment(s) while \
merging {} occurrences",
indices.len()
);
}
d.comment = None;
}
d.value = joined;
}
for &i in indices.iter().skip(1).rev() {
ast.nodes.remove(i);
}
}
}
None => {
if action == Action::Add {
append_global_directive(ast, key, user);
}
}
}
Ok(())
}
fn apply_action_to_users(users: &mut Vec<String>, user: &str, action: Action) {
match action {
Action::Add => {
if !users.iter().any(|u| u == user) {
users.push(user.to_owned());
}
}
Action::Remove => {
users.retain(|u| u != user);
}
}
}
fn append_global_directive(ast: &mut ConfigAst, key: &str, value: &str) {
let directive = ConfigNode::Directive(Box::new(DirectiveData {
keyword: key.to_owned(),
separator: Separator::Space,
value: value.to_owned(),
comment: None,
indent: String::new(),
}));
let insert_pos = ast
.nodes
.iter()
.position(|n| matches!(n, ConfigNode::HostBlock(_) | ConfigNode::MatchBlock(_)))
.unwrap_or(ast.nodes.len());
ast.nodes.insert(insert_pos, directive);
}
#[cfg(test)]
mod tests {
use super::*;
fn ast(input: &str) -> ConfigAst {
parse(input)
}
#[test]
fn get_allow_users_reads_global_directive() {
let a = ast("Port 22\nAllowUsers alice bob\n");
assert_eq!(get_allow_users(&a), vec!["alice", "bob"]);
}
#[test]
fn get_allow_users_ignores_match_scoped() {
let a = ast("AllowUsers alice\nMatch User carol\n AllowUsers bob\n");
assert_eq!(get_allow_users(&a), vec!["alice"]);
}
#[test]
fn get_deny_users_is_empty_when_absent() {
let a = ast("Port 22\n");
assert_eq!(get_deny_users(&a), Vec::<String>::new());
}
#[test]
fn add_user_to_allow_creates_directive() {
let mut a = ast("Port 22\n");
add_user_to_allow(&mut a, "alice").unwrap();
assert_eq!(get_allow_users(&a), vec!["alice"]);
}
#[test]
fn add_user_to_allow_is_idempotent() {
let mut a = ast("AllowUsers alice\n");
add_user_to_allow(&mut a, "alice").unwrap();
assert_eq!(get_allow_users(&a), vec!["alice"]);
}
#[test]
fn add_user_to_allow_appends_to_existing() {
let mut a = ast("AllowUsers alice\n");
add_user_to_allow(&mut a, "bob").unwrap();
assert_eq!(get_allow_users(&a), vec!["alice", "bob"]);
}
#[test]
fn remove_user_from_allow_deletes_directive_when_empty() {
let mut a = ast("AllowUsers alice\n");
remove_user_from_allow(&mut a, "alice").unwrap();
assert_eq!(get_allow_users(&a), Vec::<String>::new());
assert!(!a.to_string_lossless().contains("AllowUsers"));
}
#[test]
fn remove_user_from_allow_keeps_others() {
let mut a = ast("AllowUsers alice bob\n");
remove_user_from_allow(&mut a, "alice").unwrap();
assert_eq!(get_allow_users(&a), vec!["bob"]);
}
#[test]
fn remove_missing_user_is_noop() {
let mut a = ast("AllowUsers alice\n");
remove_user_from_allow(&mut a, "zzz").unwrap();
assert_eq!(get_allow_users(&a), vec!["alice"]);
}
#[test]
fn add_deny_then_reset_removes_from_both() {
let mut a = ast("AllowUsers alice\n");
add_user_to_deny(&mut a, "alice").unwrap();
remove_user_from_allow(&mut a, "alice").unwrap();
remove_user_from_deny(&mut a, "alice").unwrap();
assert_eq!(get_allow_users(&a), Vec::<String>::new());
assert_eq!(get_deny_users(&a), Vec::<String>::new());
}
#[test]
fn round_trip_preserves_comments_and_match_blocks() {
let input = "# top comment\nPort 22\n\nMatch User alice\n PermitRootLogin no\n";
let mut a = ast(input);
add_user_to_allow(&mut a, "bob").unwrap();
let out = a.to_string_lossless();
assert!(out.contains("# top comment"), "comment preserved");
assert!(out.contains("Match User alice"), "match block preserved");
assert!(out.contains("PermitRootLogin no"), "match body preserved");
assert!(out.contains("AllowUsers bob"));
}
#[test]
fn new_directive_inserted_before_match_block() {
let mut a = ast("Port 22\nMatch User alice\n X11Forwarding no\n");
add_user_to_allow(&mut a, "bob").unwrap();
let out = a.to_string_lossless();
let allow_pos = out.find("AllowUsers").unwrap();
let match_pos = out.find("Match User").unwrap();
assert!(
allow_pos < match_pos,
"AllowUsers must precede the Match block"
);
}
#[test]
fn has_pattern_tokens_detects_wildcards() {
assert!(has_pattern_tokens("alice * bob"));
assert!(has_pattern_tokens("ali?ce"));
assert!(has_pattern_tokens("alice@host"));
assert!(!has_pattern_tokens("alice bob carol"));
assert!(!has_pattern_tokens(""));
}
#[test]
fn directive_has_patterns_scans_global_only() {
let a = ast("AllowUsers *\n");
assert!(directive_has_patterns(&a, "AllowUsers"));
let a = ast("AllowUsers alice bob\n");
assert!(!directive_has_patterns(&a, "AllowUsers"));
let a = ast("AllowUsers alice\nMatch User carol\n AllowUsers *\n");
assert!(
!directive_has_patterns(&a, "AllowUsers"),
"Match-scoped patterns must not count"
);
}
#[test]
fn add_user_to_allow_refuses_pattern_directive() {
let mut a = ast("AllowUsers *\n");
let before = a.to_string_lossless();
let err = add_user_to_allow(&mut a, "bob").unwrap_err();
assert!(matches!(err, toride_ssh_core::Error::SshdConfigInvalid(_)));
assert_eq!(a.to_string_lossless(), before);
}
#[test]
fn remove_user_from_allow_refuses_pattern_directive() {
let mut a = ast("AllowUsers alice *@host\n");
let before = a.to_string_lossless();
let err = remove_user_from_allow(&mut a, "alice").unwrap_err();
assert!(matches!(err, toride_ssh_core::Error::SshdConfigInvalid(_)));
assert_eq!(a.to_string_lossless(), before);
}
#[test]
fn add_user_to_deny_refuses_question_pattern() {
let mut a = ast("DenyUsers ?uest\n");
let before = a.to_string_lossless();
let err = add_user_to_deny(&mut a, "bob").unwrap_err();
assert!(matches!(err, toride_ssh_core::Error::SshdConfigInvalid(_)));
assert_eq!(a.to_string_lossless(), before);
}
#[test]
fn add_merges_multiple_allow_users_lines() {
let mut a = ast("AllowUsers alice\nPort 22\nAllowUsers bob\n");
add_user_to_allow(&mut a, "carol").unwrap();
assert_eq!(get_allow_users(&a), vec!["alice", "bob", "carol"]);
assert_eq!(
a.to_string_lossless().matches("AllowUsers").count(),
1,
"exactly one AllowUsers line after merge"
);
}
#[test]
fn add_merge_dedupes_and_preserves_order() {
let mut a = ast("AllowUsers alice bob\nAllowUsers bob carol\n");
add_user_to_allow(&mut a, "alice").unwrap();
assert_eq!(get_allow_users(&a), vec!["alice", "bob", "carol"]);
assert_eq!(a.to_string_lossless().matches("AllowUsers").count(), 1);
}
#[test]
fn remove_across_multiple_occurrences() {
let mut a = ast("AllowUsers alice\nAllowUsers bob\n");
remove_user_from_allow(&mut a, "bob").unwrap();
assert_eq!(get_allow_users(&a), vec!["alice"]);
assert_eq!(a.to_string_lossless().matches("AllowUsers").count(), 1);
}
#[test]
fn remove_empties_union_deletes_all_occurrences() {
let mut a = ast("AllowUsers alice\nAllowUsers alice\n");
remove_user_from_allow(&mut a, "alice").unwrap();
assert_eq!(get_allow_users(&a), Vec::<String>::new());
assert!(
!a.to_string_lossless().contains("AllowUsers"),
"directive deleted entirely when union is empty"
);
}
#[test]
fn merge_drops_stale_inline_comment_from_first_occurrence() {
let mut a = ast("AllowUsers alice # production admins\nAllowUsers bob # contractors\n");
add_user_to_allow(&mut a, "carol").unwrap();
let out = a.to_string_lossless();
assert_eq!(get_allow_users(&a), vec!["alice", "bob", "carol"]);
assert_eq!(out.matches("AllowUsers").count(), 1, "merged to one line");
assert!(
!out.contains("#production admins"),
"stale first-occurrence comment must not survive the merge"
);
assert!(
!out.contains("#contractors"),
"dropped occurrence's comment must not leak into the merged line"
);
}
#[test]
fn single_occurrence_value_change_drops_stale_inline_comment() {
let mut a = ast("AllowUsers alice # admins\n");
add_user_to_allow(&mut a, "bob").unwrap();
let out = a.to_string_lossless();
assert_eq!(get_allow_users(&a), vec!["alice", "bob"]);
assert!(
!out.contains("#admins"),
"stale comment must be dropped when the value changes"
);
}
#[test]
fn single_occurrence_unchanged_keeps_inline_comment() {
let mut a = ast("AllowUsers alice bob # both admins\n");
add_user_to_allow(&mut a, "alice").unwrap();
let out = a.to_string_lossless();
assert_eq!(get_allow_users(&a), vec!["alice", "bob"]);
assert!(
out.contains("#both admins"),
"inline comment must survive an unchanged-value edit"
);
}
#[test]
fn get_allow_groups_reads_global() {
let a = ast("AllowGroups wheel staff\nMatch User carol\n AllowGroups extra\n");
assert_eq!(get_allow_groups(&a), vec!["wheel", "staff"]);
}
#[test]
fn get_deny_groups_reads_global_and_skips_match() {
let a = ast("DenyGroups banned\nMatch User carol\n DenyGroups scoped\n");
assert_eq!(get_deny_groups(&a), vec!["banned"]);
}
#[test]
fn get_deny_users_concatenates_multiple_lines() {
let a = ast("DenyUsers alice\nDenyUsers bob\n");
assert_eq!(get_deny_users(&a), vec!["alice", "bob"]);
}
#[test]
fn get_groups_empty_when_absent() {
let a = ast("Port 22\n");
assert_eq!(get_allow_groups(&a), Vec::<String>::new());
assert_eq!(get_deny_groups(&a), Vec::<String>::new());
}
#[test]
fn upsert_refuses_directive_unindented_after_match_block() {
let input = "Port 22\nMatch User sftpuser\nAllowUsers bob\n";
let mut a = ast(input);
let before = a.to_string_lossless();
let err = add_user_to_allow(&mut a, "carol").unwrap_err();
assert!(
matches!(err, toride_ssh_core::Error::SshdConfigInvalid(ref msg)
if msg.contains("refusing to edit")
&& msg.contains("Match/Host")
&& msg.contains("ambiguous")),
"expected SshdConfigInvalid scope-ambiguity error, got {err:?}"
);
assert_eq!(a.to_string_lossless(), before);
}
#[test]
fn remove_refuses_directive_unindented_after_match_block() {
let input = "Match User sftpuser\nAllowUsers bob carol\n";
let mut a = ast(input);
let before = a.to_string_lossless();
let err = remove_user_from_allow(&mut a, "bob").unwrap_err();
assert!(matches!(err, toride_ssh_core::Error::SshdConfigInvalid(_)));
assert_eq!(a.to_string_lossless(), before);
}
#[test]
fn upsert_refuses_directive_unindented_after_host_block() {
let input = "Host restricted\nAllowUsers bob\n";
let mut a = ast(input);
let before = a.to_string_lossless();
let err = add_user_to_allow(&mut a, "carol").unwrap_err();
assert!(matches!(err, toride_ssh_core::Error::SshdConfigInvalid(_)));
assert_eq!(a.to_string_lossless(), before);
}
#[test]
fn upsert_refuses_when_only_some_occurrences_follow_a_block() {
let input = "AllowUsers alice\nMatch User sftpuser\nAllowUsers bob\n";
let mut a = ast(input);
let before = a.to_string_lossless();
let err = add_user_to_allow(&mut a, "carol").unwrap_err();
assert!(matches!(err, toride_ssh_core::Error::SshdConfigInvalid(_)));
assert_eq!(a.to_string_lossless(), before);
}
#[test]
fn upsert_edits_normally_when_directive_precedes_match_block() {
let mut a = ast("AllowUsers alice\nMatch User sftpuser\n PermitRootLogin no\n");
add_user_to_allow(&mut a, "bob").unwrap();
assert_eq!(get_allow_users(&a), vec!["alice", "bob"]);
}
#[test]
fn upsert_edits_when_comments_and_blanks_intervene_before_block() {
let input = "Match User sftpuser\n# note\n\nAllowUsers bob\n";
let mut a = ast(input);
let before = a.to_string_lossless();
let err = add_user_to_allow(&mut a, "carol").unwrap_err();
assert!(matches!(err, toride_ssh_core::Error::SshdConfigInvalid(_)));
assert_eq!(a.to_string_lossless(), before);
}
#[test]
fn upsert_edits_when_another_directive_intervenes() {
let mut a = ast("Match User sftpuser\n PermitRootLogin no\nPort 22\nAllowUsers bob\n");
add_user_to_allow(&mut a, "carol").unwrap();
assert_eq!(get_allow_users(&a), vec!["bob", "carol"]);
}
#[test]
fn get_allow_users_excludes_leaked_directive_after_match() {
let a = ast("AllowUsers alice\nMatch User sftpuser\nAllowUsers bob\n");
assert_eq!(get_allow_users(&a), vec!["alice"]);
}
#[test]
fn directive_has_patterns_excludes_leaked_directive_after_match() {
let a = ast("AllowUsers alice\nMatch User sftpuser\nAllowUsers *\n");
assert!(
!directive_has_patterns(&a, "AllowUsers"),
"leaked Match-scoped pattern must not count as global"
);
}
#[test]
fn with_edit_lock_serializes_two_concurrent_holders() {
use std::sync::{Arc, Barrier, mpsc};
use std::thread;
let dir = tempfile::TempDir::new().expect("temp dir");
let lock_path = dir.path().join("sshd-config.lock");
let barrier = Arc::new(Barrier::new(2));
let (tx, rx) = mpsc::channel::<String>();
let make_thread = |label: &'static str,
lock_path: std::path::PathBuf,
barrier: Arc<Barrier>,
tx: mpsc::Sender<String>| {
thread::spawn(move || {
barrier.wait();
with_edit_lock(&lock_path, || -> Result<()> {
tx.send(format!("{label}-enter")).unwrap();
std::thread::sleep(std::time::Duration::from_millis(150));
tx.send(format!("{label}-exit")).unwrap();
Ok(())
})
.expect("with_edit_lock should succeed");
})
};
let t1 = make_thread("A", lock_path.clone(), barrier.clone(), tx.clone());
let t2 = make_thread("B", lock_path, barrier, tx.clone());
drop(tx);
let events: Vec<String> = rx.into_iter().collect();
t1.join().expect("thread A panicked");
t2.join().expect("thread B panicked");
let first_enter = events.first().expect("at least one event");
let winner = if first_enter.starts_with('A') {
"A"
} else {
"B"
};
let loser = if winner == "A" { "B" } else { "A" };
let winner_exit = events
.iter()
.position(|e| e == &format!("{winner}-exit"))
.expect("winner exit");
let loser_enter = events
.iter()
.position(|e| e == &format!("{loser}-enter"))
.expect("loser enter");
assert!(
winner_exit < loser_enter,
"edits must serialize: winner must EXIT ({winner}-exit at #{winner_exit}) \
before loser ENTERS ({loser}-enter at #{loser_enter}). Events: {events:?}"
);
}
#[test]
fn with_edit_lock_releases_on_closure_error() {
let dir = tempfile::TempDir::new().expect("temp dir");
let lock_path = dir.path().join("sshd-config-err.lock");
let err = with_edit_lock(&lock_path, || -> Result<()> {
Err(toride_ssh_core::Error::SshdConfigInvalid(
"simulated bad config".into(),
))
});
assert!(err.is_err(), "first call must propagate the error");
let result = with_edit_lock(&lock_path, || Ok(42));
assert_eq!(
result.expect("second call must succeed after error-path release"),
42
);
}
fn drive_edit_on_runtime(rt: &tokio::runtime::Runtime) {
let dir = tempfile::TempDir::new().expect("temp dir");
let lock = dir.path().join("sshd-edit-async.lock");
unsafe {
std::env::set_var(SSHD_EDIT_LOCK_ENV, &lock);
}
let ran = std::sync::Arc::new(std::sync::atomic::AtomicBool::new(false));
let ran_clone = ran.clone();
let result = rt.block_on(async {
edit(false, |ast| {
add_user_to_allow(ast, "toride-edit-async-test-user")?;
ran_clone.store(true, std::sync::atomic::Ordering::SeqCst);
Ok(())
})
.await
});
unsafe {
std::env::remove_var(SSHD_EDIT_LOCK_ENV);
}
assert!(
ran.load(std::sync::atomic::Ordering::SeqCst),
"edit()'s mutation closure must have executed"
);
drop(result);
}
#[test]
fn edit_does_not_panic_on_current_thread_runtime() {
let rt = tokio::runtime::Builder::new_current_thread()
.enable_all()
.build()
.expect("current-thread runtime");
drive_edit_on_runtime(&rt);
}
#[test]
fn edit_does_not_panic_on_multi_thread_runtime() {
let rt = tokio::runtime::Builder::new_multi_thread()
.worker_threads(1)
.enable_all()
.build()
.expect("multi-thread runtime");
drive_edit_on_runtime(&rt);
}
}