use std::fmt;
use std::net::{IpAddr, Ipv4Addr, Ipv6Addr, SocketAddr};
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Default)]
pub enum AddressFamily {
#[default]
Any,
V4,
V6,
}
impl AddressFamily {
pub fn from_flags(ipv4: bool, ipv6: bool) -> Option<Self> {
match (ipv4, ipv6) {
(true, _) => Some(Self::V4),
(_, true) => Some(Self::V6),
_ => None,
}
}
pub fn from_config_value(value: &str) -> Self {
match value.trim().to_ascii_lowercase().as_str() {
"any" => Self::Any,
"inet" => Self::V4,
"inet6" => Self::V6,
other => {
tracing::warn!(
"Unrecognized AddressFamily value '{other}' in SSH config; expected any, inet, or inet6. Falling back to 'any'"
);
Self::Any
}
}
}
pub fn resolve(ipv4_flag: bool, ipv6_flag: bool, config_value: Option<&str>) -> Self {
Self::from_flags(ipv4_flag, ipv6_flag)
.or_else(|| config_value.map(Self::from_config_value))
.unwrap_or_default()
}
pub const fn is_forced(self) -> bool {
!matches!(self, Self::Any)
}
pub const fn matches(self, addr: &SocketAddr) -> bool {
match self {
Self::Any => true,
Self::V4 => addr.is_ipv4(),
Self::V6 => addr.is_ipv6(),
}
}
pub fn filter(self, addrs: Vec<SocketAddr>) -> Vec<SocketAddr> {
if !self.is_forced() {
return addrs;
}
addrs.into_iter().filter(|a| self.matches(a)).collect()
}
pub fn first_match(self, addrs: &[SocketAddr]) -> Option<SocketAddr> {
addrs.iter().copied().find(|a| self.matches(a))
}
pub const fn loopback(self) -> IpAddr {
match self {
Self::V6 => IpAddr::V6(Ipv6Addr::LOCALHOST),
_ => IpAddr::V4(Ipv4Addr::LOCALHOST),
}
}
pub const fn unspecified(self) -> IpAddr {
match self {
Self::V6 => IpAddr::V6(Ipv6Addr::UNSPECIFIED),
_ => IpAddr::V4(Ipv4Addr::UNSPECIFIED),
}
}
}
impl fmt::Display for AddressFamily {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
let text = match self {
Self::Any => "any",
Self::V4 => "IPv4",
Self::V6 => "IPv6",
};
f.write_str(text)
}
}
#[cfg(test)]
mod tests {
use super::*;
fn v4(s: &str) -> SocketAddr {
s.parse().expect("valid IPv4 socket address")
}
fn v6(s: &str) -> SocketAddr {
s.parse().expect("valid IPv6 socket address")
}
fn mixed() -> Vec<SocketAddr> {
vec![
v6("[2001:db8::1]:22"),
v4("192.0.2.10:22"),
v6("[2001:db8::2]:22"),
v4("192.0.2.11:22"),
]
}
#[test]
fn ipv4_filter_keeps_only_ipv4_in_resolver_order() {
let filtered = AddressFamily::V4.filter(mixed());
assert_eq!(filtered, vec![v4("192.0.2.10:22"), v4("192.0.2.11:22")]);
}
#[test]
fn ipv6_filter_keeps_only_ipv6_in_resolver_order() {
let filtered = AddressFamily::V6.filter(mixed());
assert_eq!(
filtered,
vec![v6("[2001:db8::1]:22"), v6("[2001:db8::2]:22")]
);
}
#[test]
fn any_filter_returns_the_original_list_unchanged_and_in_order() {
let original = mixed();
let filtered = AddressFamily::Any.filter(original.clone());
assert_eq!(filtered, original);
}
#[test]
fn filter_can_empty_the_candidate_list() {
let only_v4 = vec![v4("192.0.2.10:22")];
assert!(AddressFamily::V6.filter(only_v4).is_empty());
let only_v6 = vec![v6("[2001:db8::1]:22")];
assert!(AddressFamily::V4.filter(only_v6).is_empty());
}
#[test]
fn first_match_picks_the_first_candidate_of_the_forced_family() {
let addrs = mixed();
assert_eq!(
AddressFamily::V4.first_match(&addrs),
Some(v4("192.0.2.10:22"))
);
assert_eq!(
AddressFamily::V6.first_match(&addrs),
Some(v6("[2001:db8::1]:22"))
);
assert_eq!(
AddressFamily::Any.first_match(&addrs),
Some(v6("[2001:db8::1]:22"))
);
assert_eq!(AddressFamily::V6.first_match(&[v4("192.0.2.10:22")]), None);
}
#[test]
fn config_values_are_case_insensitive() {
assert_eq!(AddressFamily::from_config_value("any"), AddressFamily::Any);
assert_eq!(AddressFamily::from_config_value("ANY"), AddressFamily::Any);
assert_eq!(AddressFamily::from_config_value("inet"), AddressFamily::V4);
assert_eq!(AddressFamily::from_config_value("INet"), AddressFamily::V4);
assert_eq!(AddressFamily::from_config_value("inet6"), AddressFamily::V6);
assert_eq!(AddressFamily::from_config_value("INET6"), AddressFamily::V6);
assert_eq!(
AddressFamily::from_config_value(" inet6 "),
AddressFamily::V6
);
}
#[test]
fn unrecognized_config_value_falls_back_to_any() {
assert_eq!(
AddressFamily::from_config_value("ipv6"),
AddressFamily::Any,
"an unrecognized value must not hard-fail a config OpenSSH tolerates"
);
assert_eq!(AddressFamily::from_config_value(""), AddressFamily::Any);
}
#[test]
fn command_line_flag_wins_over_config_keyword() {
assert_eq!(
AddressFamily::resolve(true, false, Some("inet6")),
AddressFamily::V4
);
assert_eq!(
AddressFamily::resolve(false, true, Some("inet")),
AddressFamily::V6
);
assert_eq!(
AddressFamily::resolve(true, false, Some("any")),
AddressFamily::V4
);
}
#[test]
fn config_keyword_applies_when_no_flag_is_given() {
assert_eq!(
AddressFamily::resolve(false, false, Some("inet")),
AddressFamily::V4
);
assert_eq!(
AddressFamily::resolve(false, false, Some("inet6")),
AddressFamily::V6
);
assert_eq!(
AddressFamily::resolve(false, false, Some("any")),
AddressFamily::Any
);
}
#[test]
fn default_is_any_when_neither_flag_nor_keyword_is_set() {
assert_eq!(
AddressFamily::resolve(false, false, None),
AddressFamily::Any
);
assert_eq!(AddressFamily::default(), AddressFamily::Any);
assert!(!AddressFamily::default().is_forced());
}
#[test]
fn default_bind_addresses_follow_the_forced_family() {
assert_eq!(
AddressFamily::V6.loopback(),
IpAddr::V6(Ipv6Addr::LOCALHOST)
);
assert_eq!(
AddressFamily::V4.loopback(),
IpAddr::V4(Ipv4Addr::LOCALHOST)
);
assert_eq!(
AddressFamily::Any.loopback(),
IpAddr::V4(Ipv4Addr::LOCALHOST),
"the no-flag default must stay IPv4 loopback"
);
assert_eq!(
AddressFamily::V6.unspecified(),
IpAddr::V6(Ipv6Addr::UNSPECIFIED)
);
assert_eq!(
AddressFamily::Any.unspecified(),
IpAddr::V4(Ipv4Addr::UNSPECIFIED)
);
}
#[test]
fn display_names_match_the_user_facing_error_text() {
assert_eq!(AddressFamily::V4.to_string(), "IPv4");
assert_eq!(AddressFamily::V6.to_string(), "IPv6");
assert_eq!(AddressFamily::Any.to_string(), "any");
}
}