use super::ast::{ConfigAst, ConfigNode, DirectiveData};
use toride_ssh_core::Result;
pub fn get_directive(ast: &ConfigAst, host: &str, key: &str) -> Option<String> {
let key_lower = key.to_lowercase();
for node in &ast.nodes {
if let ConfigNode::HostBlock(b) = node
&& host_matches_patterns(host, &b.patterns)
&& let Some(val) = find_directive_in_nodes(&b.nodes, &key_lower)
{
return Some(val.to_owned());
}
}
None
}
#[allow(dead_code, reason = "public API, not exercised in-workspace")]
pub fn get_accumulative_directive(ast: &ConfigAst, host: &str, key: &str) -> Vec<String> {
let key_lower = key.to_lowercase();
let mut values = Vec::new();
for node in &ast.nodes {
if let ConfigNode::HostBlock(b) = node
&& host_matches_patterns(host, &b.patterns)
{
collect_directives_in_nodes(&b.nodes, &key_lower, &mut values);
}
}
values
}
#[allow(dead_code, reason = "public API, not exercised in-workspace")]
pub fn get_directive_by_name(ast: &ConfigAst, name: &str, key: &str) -> Option<String> {
let key_lower = key.to_lowercase();
for node in &ast.nodes {
if let ConfigNode::HostBlock(b) = node
&& b.patterns.iter().any(|p| p == name || p == "*")
&& let Some(val) = find_directive_in_nodes(&b.nodes, &key_lower)
{
return Some(val.to_owned());
}
}
None
}
pub fn get_all_directives(ast: &ConfigAst, host: &str) -> Vec<(String, String)> {
let mut result = Vec::new();
let mut seen = std::collections::HashSet::new();
for node in &ast.nodes {
if let ConfigNode::HostBlock(b) = node
&& host_matches_patterns(host, &b.patterns)
{
collect_all_directives(&b.nodes, &mut result, &mut seen);
}
}
result
}
#[allow(dead_code, reason = "public API, not exercised in-workspace")]
pub fn set_directive(ast: &mut ConfigAst, host: &str, key: &str, value: &str) -> Result<()> {
let key_lower = key.to_lowercase();
for node in &mut ast.nodes {
if let ConfigNode::HostBlock(b) = node
&& host_matches_patterns(host, &b.patterns)
{
for child in &mut b.nodes {
if let ConfigNode::Directive(d) = child
&& d.keyword.eq_ignore_ascii_case(&key_lower)
{
value.clone_into(&mut d.value);
return Ok(());
}
}
b.nodes.push(ConfigNode::Directive(Box::new(DirectiveData {
keyword: key.to_owned(),
separator: super::ast::Separator::Space,
value: value.to_owned(),
comment: None,
indent: String::new(),
})));
return Ok(());
}
}
Err(toride_ssh_core::Error::HostNotFound(host.to_owned()))
}
fn find_directive_in_nodes<'a>(nodes: &'a [ConfigNode], key_lower: &str) -> Option<&'a str> {
for node in nodes {
if let ConfigNode::Directive(d) = node
&& d.keyword.eq_ignore_ascii_case(key_lower)
{
return Some(&d.value);
}
}
None
}
#[allow(
dead_code,
reason = "reachable only via public API not exercised in-workspace"
)]
fn collect_directives_in_nodes(nodes: &[ConfigNode], key_lower: &str, out: &mut Vec<String>) {
for node in nodes {
if let ConfigNode::Directive(d) = node
&& d.keyword.eq_ignore_ascii_case(key_lower)
{
out.push(d.value.clone());
}
}
}
fn collect_all_directives(
nodes: &[ConfigNode],
out: &mut Vec<(String, String)>,
seen: &mut std::collections::HashSet<String>,
) {
for node in nodes {
if let ConfigNode::Directive(d) = node {
if is_accumulative(&d.keyword) {
out.push((d.keyword.clone(), d.value.clone()));
} else {
let key_lower = d.keyword.to_ascii_lowercase();
if seen.insert(key_lower) {
out.push((d.keyword.clone(), d.value.clone()));
}
}
}
}
}
#[allow(dead_code, reason = "public API, not exercised in-workspace")]
pub fn get_preferred_authentications(ast: &ConfigAst, host: &str) -> Option<String> {
get_directive(ast, host, "PreferredAuthentications")
}
pub(crate) fn is_accumulative(keyword: &str) -> bool {
keyword.eq_ignore_ascii_case("identityfile")
|| keyword.eq_ignore_ascii_case("certificatefile")
|| keyword.eq_ignore_ascii_case("sendenv")
|| keyword.eq_ignore_ascii_case("setenv")
|| keyword.eq_ignore_ascii_case("dynamicforward")
|| keyword.eq_ignore_ascii_case("localforward")
|| keyword.eq_ignore_ascii_case("remoteforward")
|| keyword.eq_ignore_ascii_case("permitlocalcommand")
}
pub(crate) fn host_matches_patterns(host: &str, patterns: &[impl AsRef<str>]) -> bool {
let host_lower = host.to_ascii_lowercase();
let mut positive_match = false;
for pattern in patterns {
let pat_lower = pattern.as_ref().to_ascii_lowercase();
if let Some(negated) = pat_lower.strip_prefix('!') {
if glob_matches(&host_lower, negated) {
return false;
}
} else if glob_matches(&host_lower, &pat_lower) {
positive_match = true;
}
}
positive_match
}
pub(crate) fn glob_matches(text: &str, pattern: &str) -> bool {
if pattern.len() == 1 {
return match pattern {
"*" => true,
"?" => text.len() == 1,
c => text == c,
};
}
if !pattern.contains('*') && !pattern.contains('?') {
return text == pattern;
}
glob_match_recursive(text.as_bytes(), pattern.as_bytes())
}
fn glob_match_recursive(text: &[u8], pattern: &[u8]) -> bool {
let mut ti = 0;
let mut pi = 0;
let mut star_pi = None;
let mut star_text_idx = 0;
while ti < text.len() {
if pi < pattern.len() {
let pc = pattern[pi];
let tc = text[ti];
if pc == b'?' {
pi += 1;
ti += 1;
continue;
}
if pc == b'*' {
star_pi = Some(pi + 1);
star_text_idx = ti;
pi += 1;
continue;
}
if pc == tc {
pi += 1;
ti += 1;
continue;
}
}
if let Some(spi) = star_pi {
pi = spi;
star_text_idx += 1;
ti = star_text_idx;
continue;
}
return false;
}
while pi < pattern.len() && pattern[pi] == b'*' {
pi += 1;
}
pi == pattern.len()
}
#[cfg(test)]
#[path = "directives.test.rs"]
mod tests;