use core::net::IpAddr;
use alloc::string::{String, ToString};
use alloc::vec::Vec;
use super::ConfigError;
use super::glob::{HostPattern, glob_match, host_matches};
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct AddressPattern {
pub negated: bool,
pub kind: AddressKind,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum AddressKind {
Cidr {
base: IpAddr,
prefix: u8,
},
Glob(String),
}
impl AddressPattern {
pub fn parse(token: &str) -> Self {
let (negated, body) = match token.strip_prefix('!') {
Some(rest) => (true, rest),
None => (false, token),
};
let kind = parse_address_kind(body);
AddressPattern { negated, kind }
}
}
fn parse_address_kind(body: &str) -> AddressKind {
if let Some((addr_s, prefix_s)) = body.split_once('/') {
if let (Ok(base), Ok(prefix)) = (addr_s.parse::<IpAddr>(), prefix_s.parse::<u8>()) {
let max = if base.is_ipv4() { 32 } else { 128 };
if prefix <= max {
return AddressKind::Cidr { base, prefix };
}
}
return AddressKind::Glob(body.to_string());
}
if let Ok(base) = body.parse::<IpAddr>() {
let prefix = if base.is_ipv4() { 32 } else { 128 };
return AddressKind::Cidr { base, prefix };
}
AddressKind::Glob(body.to_string())
}
fn parse_address_list(s: &str) -> Vec<AddressPattern> {
s.split(',')
.filter(|t| !t.is_empty())
.map(AddressPattern::parse)
.collect()
}
fn cidr_contains(base: IpAddr, prefix: u8, addr: IpAddr) -> bool {
match (base, addr) {
(IpAddr::V4(b), IpAddr::V4(a)) => {
let bits = u32::from(b);
let abits = u32::from(a);
if prefix == 0 {
return true;
}
if prefix > 32 {
return false;
}
let mask = u32::MAX.checked_shl(32 - prefix as u32).unwrap_or(0);
(bits & mask) == (abits & mask)
}
(IpAddr::V6(b), IpAddr::V6(a)) => {
let bits = u128::from(b);
let abits = u128::from(a);
if prefix == 0 {
return true;
}
if prefix > 128 {
return false;
}
let mask = u128::MAX.checked_shl(128 - prefix as u32).unwrap_or(0);
(bits & mask) == (abits & mask)
}
_ => false,
}
}
pub fn address_matches(patterns: &[AddressPattern], addr_str: &str) -> bool {
if patterns.is_empty() {
return false;
}
let parsed = addr_str.parse::<IpAddr>().ok();
let mut any_positive = false;
let mut positive_hit = false;
for p in patterns {
let hit = match &p.kind {
AddressKind::Cidr { base, prefix } => {
matches!(parsed, Some(a) if cidr_contains(*base, *prefix, a))
}
AddressKind::Glob(g) => glob_match(g, addr_str),
};
if p.negated {
if hit {
return false;
}
} else {
any_positive = true;
if !positive_hit && hit {
positive_hit = true;
}
}
}
any_positive && positive_hit
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum MatchCondition {
Host(Vec<HostPattern>),
OriginalHost(Vec<HostPattern>),
User(Vec<HostPattern>),
LocalUser(Vec<HostPattern>),
Group(Vec<HostPattern>),
Address(Vec<AddressPattern>),
LocalAddress(Vec<AddressPattern>),
LocalPort(Vec<u16>),
Exec(String),
All,
Canonical,
Final,
}
#[derive(Debug, Clone, Default)]
pub struct MatchContext<'a> {
pub host: &'a str,
pub original_host: Option<&'a str>,
pub user: Option<&'a str>,
pub local_user: Option<&'a str>,
pub address: Option<&'a str>,
pub local_address: Option<&'a str>,
pub local_port: Option<u16>,
pub groups: Option<&'a [String]>,
}
impl MatchContext<'_> {
pub fn original_host_or_host(&self) -> &str {
self.original_host.unwrap_or(self.host)
}
}
pub fn parse_match_line(
args: &[String],
line_no: usize,
) -> Result<Vec<MatchCondition>, ConfigError> {
parse_match_line_impl(args, line_no, MatchSide::Client)
}
pub fn parse_match_line_server(
args: &[String],
line_no: usize,
) -> Result<Vec<MatchCondition>, ConfigError> {
parse_match_line_impl(args, line_no, MatchSide::Server)
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum MatchSide {
Client,
Server,
}
fn parse_match_line_impl(
args: &[String],
line_no: usize,
side: MatchSide,
) -> Result<Vec<MatchCondition>, ConfigError> {
let server = side == MatchSide::Server;
let mut out = Vec::new();
let mut i = 0;
while i < args.len() {
let kw = args[i].to_ascii_lowercase();
match kw.as_str() {
"all" => {
out.push(MatchCondition::All);
i += 1;
}
"canonical" => {
out.push(MatchCondition::Canonical);
i += 1;
}
"final" => {
out.push(MatchCondition::Final);
i += 1;
}
"host" | "originalhost" | "localuser" if server => {
return Err(ConfigError::Unsupported {
line: line_no,
msg: alloc::format!("Match {kw} is not valid in sshd_config"),
});
}
"rdomain" | "connection" => {
return Err(ConfigError::Unsupported {
line: line_no,
msg: alloc::format!("Match {kw} is not supported"),
});
}
"host" | "originalhost" | "user" | "localuser" => {
let patterns = take_pattern_list(args, &mut i, &kw, line_no)?;
let cond = match kw.as_str() {
"host" => MatchCondition::Host(patterns),
"originalhost" => MatchCondition::OriginalHost(patterns),
"user" => MatchCondition::User(patterns),
"localuser" => MatchCondition::LocalUser(patterns),
_ => unreachable!(),
};
out.push(cond);
}
"group" if server => {
let patterns = take_pattern_list(args, &mut i, &kw, line_no)?;
out.push(MatchCondition::Group(patterns));
}
"address" | "localaddress" if server => {
let raw = take_raw_arg(args, &mut i, &kw, line_no)?;
let patterns = parse_address_list(&raw);
if kw == "address" {
out.push(MatchCondition::Address(patterns));
} else {
out.push(MatchCondition::LocalAddress(patterns));
}
}
"localport" if server => {
let raw = take_raw_arg(args, &mut i, &kw, line_no)?;
let mut ports = Vec::new();
for p in raw.split(',').filter(|t| !t.is_empty()) {
let port = p.parse::<u16>().map_err(|_| ConfigError::BadValue {
line: line_no,
keyword: "match".to_string(),
msg: alloc::format!("Match localport: bad port {p:?}"),
})?;
ports.push(port);
}
if ports.is_empty() {
return Err(ConfigError::BadValue {
line: line_no,
keyword: "match".to_string(),
msg: "Match localport requires at least one port".into(),
});
}
out.push(MatchCondition::LocalPort(ports));
}
"exec" => {
if i + 1 >= args.len() {
return Err(ConfigError::BadValue {
line: line_no,
keyword: "match".to_string(),
msg: "Match exec requires a command argument".into(),
});
}
let cmd = args[i + 1..].join(" ");
out.push(MatchCondition::Exec(cmd));
i = args.len();
}
other => {
return Err(ConfigError::BadValue {
line: line_no,
keyword: "match".to_string(),
msg: alloc::format!("unknown Match criterion: {other:?}"),
});
}
}
}
if out.is_empty() {
return Err(ConfigError::BadValue {
line: line_no,
keyword: "match".to_string(),
msg: "Match requires at least one criterion".into(),
});
}
Ok(out)
}
fn take_raw_arg(
args: &[String],
i: &mut usize,
kw: &str,
line_no: usize,
) -> Result<String, ConfigError> {
if *i + 1 >= args.len() {
return Err(ConfigError::BadValue {
line: line_no,
keyword: "match".to_string(),
msg: alloc::format!("Match {kw} requires an argument"),
});
}
let v = args[*i + 1].clone();
*i += 2;
Ok(v)
}
fn take_pattern_list(
args: &[String],
i: &mut usize,
kw: &str,
line_no: usize,
) -> Result<Vec<HostPattern>, ConfigError> {
let raw = take_raw_arg(args, i, kw, line_no)?;
Ok(parse_match_pattern_list(&raw))
}
fn parse_match_pattern_list(s: &str) -> Vec<HostPattern> {
s.split(',')
.filter(|t| !t.is_empty())
.map(HostPattern::parse)
.collect()
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum ExecPolicy {
Deny,
Allow,
}
pub fn evaluate(cond: &MatchCondition, ctx: &MatchContext<'_>, policy: ExecPolicy) -> bool {
match cond {
MatchCondition::All => true,
MatchCondition::Canonical | MatchCondition::Final => false,
MatchCondition::Host(patterns) => host_matches(patterns, ctx.host),
MatchCondition::OriginalHost(patterns) => {
host_matches(patterns, ctx.original_host_or_host())
}
MatchCondition::User(patterns) => match ctx.user {
Some(u) => host_matches(patterns, u),
None => false,
},
MatchCondition::LocalUser(patterns) => match ctx.local_user {
Some(u) => host_matches(patterns, u),
None => false,
},
MatchCondition::Group(patterns) => match ctx.groups {
Some(groups) => groups.iter().any(|g| host_matches(patterns, g)),
None => false,
},
MatchCondition::Address(patterns) => match ctx.address {
Some(a) => address_matches(patterns, a),
None => false,
},
MatchCondition::LocalAddress(patterns) => match ctx.local_address {
Some(a) => address_matches(patterns, a),
None => false,
},
MatchCondition::LocalPort(ports) => match ctx.local_port {
Some(p) => ports.contains(&p),
None => false,
},
MatchCondition::Exec(cmd) => match policy {
ExecPolicy::Deny => false,
ExecPolicy::Allow => run_exec_match(cmd),
},
}
}
pub fn all_match(conds: &[MatchCondition], ctx: &MatchContext<'_>, policy: ExecPolicy) -> bool {
if conds.is_empty() {
return false;
}
conds.iter().all(|c| evaluate(c, ctx, policy))
}
#[cfg(feature = "std")]
fn run_exec_match(cmd: &str) -> bool {
use std::process::Command;
#[cfg(unix)]
let result = Command::new("/bin/sh").arg("-c").arg(cmd).status();
#[cfg(windows)]
let result = Command::new("cmd").arg("/C").arg(cmd).status();
#[cfg(not(any(unix, windows)))]
let result: Result<std::process::ExitStatus, std::io::Error> = Err(std::io::Error::new(
std::io::ErrorKind::Unsupported,
"Match exec is not supported on this platform",
));
match result {
Ok(status) => status.success(),
Err(_) => false,
}
}
#[cfg(not(feature = "std"))]
fn run_exec_match(_cmd: &str) -> bool {
false
}
#[cfg(test)]
mod tests {
use super::*;
fn ctx<'a>(host: &'a str) -> MatchContext<'a> {
MatchContext {
host,
original_host: None,
user: None,
local_user: None,
..MatchContext::default()
}
}
#[test]
fn parse_all_alone() {
let args = vec!["all".to_string()];
let conds = parse_match_line(&args, 1).unwrap();
assert_eq!(conds, vec![MatchCondition::All]);
}
#[test]
fn parse_host_with_pattern_list() {
let args = vec!["host".to_string(), "*.example.com,!secret.*".to_string()];
let conds = parse_match_line(&args, 1).unwrap();
match &conds[0] {
MatchCondition::Host(p) => {
assert_eq!(p.len(), 2);
}
_ => panic!("wrong cond"),
}
}
#[test]
fn parse_missing_argument_errors() {
let args = vec!["host".to_string()];
let err = parse_match_line(&args, 7).unwrap_err();
match err {
ConfigError::BadValue { line, .. } => assert_eq!(line, 7),
_ => panic!("wrong err: {err:?}"),
}
}
#[test]
fn parse_unknown_criterion_errors() {
let args = vec!["address".to_string(), "1.2.3.4".to_string()];
let err = parse_match_line(&args, 3).unwrap_err();
match err {
ConfigError::BadValue { line, .. } => assert_eq!(line, 3),
_ => panic!("wrong err: {err:?}"),
}
}
#[test]
fn parse_empty_errors() {
let err = parse_match_line(&[], 1).unwrap_err();
match err {
ConfigError::BadValue { .. } => {}
_ => panic!("wrong err: {err:?}"),
}
}
#[test]
fn evaluate_all_matches() {
let c = ctx("anything");
assert!(evaluate(&MatchCondition::All, &c, ExecPolicy::Deny));
}
#[test]
fn evaluate_canonical_never_matches() {
let c = ctx("anything");
assert!(!evaluate(&MatchCondition::Canonical, &c, ExecPolicy::Deny));
assert!(!evaluate(&MatchCondition::Final, &c, ExecPolicy::Deny));
}
#[test]
fn evaluate_user_missing_in_context_is_no_match() {
let conds = parse_match_line(&["user".to_string(), "alice".to_string()], 1).unwrap();
let c = ctx("h"); assert!(!all_match(&conds, &c, ExecPolicy::Deny));
}
#[test]
fn evaluate_exec_denied_by_default() {
let conds = parse_match_line(&["exec".to_string(), "true".to_string()], 1).unwrap();
let c = ctx("h");
assert!(!all_match(&conds, &c, ExecPolicy::Deny));
}
#[test]
fn parse_match_pattern_list_skips_empties() {
let pats = parse_match_pattern_list("a,,b");
assert_eq!(pats.len(), 2);
}
#[test]
fn server_parses_address_group_localport() {
let args = vec![
"address".to_string(),
"192.0.2.0/24,!192.0.2.7".to_string(),
"group".to_string(),
"admin,wheel".to_string(),
"localport".to_string(),
"22,2222".to_string(),
];
let conds = parse_match_line_server(&args, 1).unwrap();
assert_eq!(conds.len(), 3);
assert!(matches!(conds[0], MatchCondition::Address(_)));
assert!(matches!(conds[1], MatchCondition::Group(_)));
assert!(matches!(&conds[2], MatchCondition::LocalPort(p) if p == &[22, 2222]));
}
#[test]
fn server_rejects_client_only_criteria() {
for kw in ["host", "originalhost", "localuser"] {
let args = vec![kw.to_string(), "x".to_string()];
let err = parse_match_line_server(&args, 3).unwrap_err();
assert!(
matches!(err, ConfigError::Unsupported { line: 3, .. }),
"{kw}: got {err:?}"
);
}
}
#[test]
fn server_rejects_rdomain_connection() {
for kw in ["rdomain", "connection"] {
let args = vec![kw.to_string(), "x".to_string()];
let err = parse_match_line_server(&args, 5).unwrap_err();
assert!(matches!(err, ConfigError::Unsupported { line: 5, .. }));
}
}
#[test]
fn client_still_rejects_server_criteria() {
let args = vec!["address".to_string(), "1.2.3.4".to_string()];
let err = parse_match_line(&args, 1).unwrap_err();
assert!(matches!(err, ConfigError::BadValue { line: 1, .. }));
}
#[test]
fn address_cidr_v4_match() {
let pats = parse_address_list("192.0.2.0/24");
assert!(address_matches(&pats, "192.0.2.7"));
assert!(!address_matches(&pats, "192.0.3.7"));
}
#[test]
fn address_bare_v4_is_host_route() {
let pats = parse_address_list("192.0.2.7");
assert!(address_matches(&pats, "192.0.2.7"));
assert!(!address_matches(&pats, "192.0.2.8"));
}
#[test]
fn address_negation() {
let pats = parse_address_list("192.0.2.0/24,!192.0.2.7");
assert!(address_matches(&pats, "192.0.2.1"));
assert!(!address_matches(&pats, "192.0.2.7"));
}
#[test]
fn address_v6_cidr() {
let pats = parse_address_list("2001:db8::/32");
assert!(address_matches(&pats, "2001:db8::1"));
assert!(!address_matches(&pats, "2001:dba::1"));
assert!(!address_matches(&pats, "192.0.2.1"));
}
#[test]
fn address_glob_fallback() {
let pats = parse_address_list("192.0.2.*");
assert!(address_matches(&pats, "192.0.2.55"));
assert!(!address_matches(&pats, "192.0.3.55"));
}
#[test]
fn address_zero_prefix_matches_family() {
let pats = parse_address_list("0.0.0.0/0");
assert!(address_matches(&pats, "8.8.8.8"));
assert!(!address_matches(&pats, "::1"));
}
#[test]
fn evaluate_group_any_member() {
let conds =
parse_match_line_server(&["group".to_string(), "wheel".to_string()], 1).unwrap();
let groups = vec!["users".to_string(), "wheel".to_string()];
let c = MatchContext {
host: "h",
groups: Some(&groups),
..MatchContext::default()
};
assert!(all_match(&conds, &c, ExecPolicy::Deny));
let c2 = ctx("h");
assert!(!all_match(&conds, &c2, ExecPolicy::Deny));
}
}