use std::collections::HashSet;
use crate::{
config::{BranchChainConfig, ChainRef, FilterChainConfig, FilterEntry, ResultMatch},
errors::ProxyError,
};
pub const MAX_BRANCH_DEPTH: usize = 10;
pub const MAX_ITERATIONS_CEILING: u32 = 100;
const MAX_BRANCHES_PER_FILTER: usize = 16;
const MAX_TOTAL_BRANCHES: usize = 256;
const MAX_RESULT_VALUE_LEN: usize = 256;
pub(crate) fn validate_branch_chains(chains: &[FilterChainConfig]) -> Result<(), ProxyError> {
let chain_names: HashSet<&str> = chains.iter().map(|chain| chain.name.as_str()).collect();
let initial_count = chain_names.len();
let mut all_names: HashSet<String> = chain_names.iter().map(|name| (*name).to_owned()).collect();
for chain in chains {
validate_filter_names_unique(&chain.filters, &chain.name)?;
collect_branch_names(&chain.filters, &mut all_names, &chain_names, 0)?;
}
let branch_count = all_names.len().saturating_sub(initial_count);
if branch_count > MAX_TOTAL_BRANCHES {
return Err(ProxyError::Config(format!(
"total branch count ({branch_count}) exceeds maximum ({MAX_TOTAL_BRANCHES})"
)));
}
Ok(())
}
pub fn validate_chain_entries_branch_chains(
chain_name: &str,
entries: &[FilterEntry],
known_chains: &HashSet<&str>,
prior_branch_count: usize,
) -> Result<usize, ProxyError> {
validate_filter_names_unique(entries, chain_name)?;
let initial_count = known_chains.len();
let mut all_names: HashSet<String> = known_chains.iter().map(|name| (*name).to_owned()).collect();
collect_branch_names(entries, &mut all_names, known_chains, 0)?;
let branch_count = all_names.len().saturating_sub(initial_count);
let total = prior_branch_count.saturating_add(branch_count);
if total > MAX_TOTAL_BRANCHES {
return Err(ProxyError::Config(format!(
"total branch count ({total}) exceeds maximum ({MAX_TOTAL_BRANCHES})"
)));
}
Ok(total)
}
pub fn count_build_branches(entries: &[FilterEntry], chains: &[&[FilterEntry]], known_chains: &HashSet<&str>) -> usize {
let initial_count = known_chains.len();
let mut all_names: HashSet<String> = known_chains.iter().map(|name| (*name).to_owned()).collect();
count_branch_names(entries, &mut all_names);
for &chain in chains {
count_branch_names(chain, &mut all_names);
}
all_names.len().saturating_sub(initial_count)
}
fn count_branch_names(entries: &[FilterEntry], all_names: &mut HashSet<String>) {
for entry in entries {
let Some(branches) = &entry.branch_chains else {
continue;
};
for branch in branches {
all_names.insert(branch.name.clone());
for chain_ref in &branch.chains {
if let ChainRef::Inline { name, filters } = chain_ref {
all_names.insert(name.clone());
count_branch_names(filters, all_names);
}
}
}
}
}
fn validate_filter_names_unique(filters: &[FilterEntry], chain_name: &str) -> Result<(), ProxyError> {
let mut seen = HashSet::new();
for entry in filters {
if let Some(name) = &entry.name {
validate_filter_name_chars(name)?;
if !seen.insert(name.as_str()) {
return Err(ProxyError::Config(format!(
"duplicate filter name '{name}' in chain '{chain_name}'"
)));
}
}
}
Ok(())
}
fn validate_filter_name_chars(name: &str) -> Result<(), ProxyError> {
if name.is_empty() {
return Err(ProxyError::Config("filter name must not be empty".into()));
}
if !name
.bytes()
.all(|byte| byte.is_ascii_alphanumeric() || byte == b'_' || byte == b'-')
{
return Err(ProxyError::Config(format!(
"filter name '{name}' must be ASCII alphanumeric, '_', or '-'"
)));
}
Ok(())
}
fn collect_branch_names(
filters: &[FilterEntry],
all_names: &mut HashSet<String>,
chain_names: &HashSet<&str>,
depth: usize,
) -> Result<(), ProxyError> {
debug_assert!(
depth <= MAX_BRANCH_DEPTH,
"collect_branch_names entered at depth {depth} > {MAX_BRANCH_DEPTH}"
);
for entry in filters {
entry.warn_config_typos();
let Some(branches) = &entry.branch_chains else {
continue;
};
if branches.len() > MAX_BRANCHES_PER_FILTER {
return Err(ProxyError::Config(format!(
"filter has {} branch chains (max {MAX_BRANCHES_PER_FILTER})",
branches.len()
)));
}
for branch in branches {
validate_branch(branch, all_names, chain_names, depth)?;
}
}
Ok(())
}
fn validate_branch(
branch: &BranchChainConfig,
all_names: &mut HashSet<String>,
chain_names: &HashSet<&str>,
depth: usize,
) -> Result<(), ProxyError> {
let bname = &branch.name;
super::validate_name_chars(bname, "branch")?;
if !all_names.insert(branch.name.clone()) {
return Err(ProxyError::Config(format!("duplicate branch name '{bname}'")));
}
let rejoin = &branch.rejoin;
if branch.rejoin.contains(':') {
return Err(ProxyError::Config(format!(
"branch '{bname}': cross-chain rejoin '{rejoin}' is not supported; \
listeners flatten all referenced chains into one pipeline, \
so use the filter's name directly (e.g. 'routing' not 'main:routing')"
)));
}
if branch.chains.is_empty() {
return Err(ProxyError::Config(format!(
"branch '{bname}' must have at least one chain"
)));
}
validate_max_iterations(branch)?;
validate_on_result_filter_name(branch)?;
validate_on_result_key_value(branch)?;
for chain_ref in &branch.chains {
validate_chain_ref(chain_ref, all_names, chain_names, depth)?;
}
Ok(())
}
fn validate_on_result_filter_name(branch: &BranchChainConfig) -> Result<(), ProxyError> {
let Some(cond) = &branch.on_result else {
return Ok(());
};
let bname = &branch.name;
let filter = &cond.filter;
if cond.filter.is_empty() {
return Err(ProxyError::Config(format!(
"branch '{bname}': on_result.filter must not be empty"
)));
}
if !cond
.filter
.bytes()
.all(|byte| byte.is_ascii_alphanumeric() || byte == b'_' || byte == b'-')
{
return Err(ProxyError::Config(format!(
"branch '{bname}': on_result.filter '{filter}' must be ASCII alphanumeric, '_', or '-'"
)));
}
Ok(())
}
fn validate_on_result_key_value(branch: &BranchChainConfig) -> Result<(), ProxyError> {
let Some(cond) = &branch.on_result else {
return Ok(());
};
let bname = &branch.name;
if cond.key.is_empty() {
return Err(ProxyError::Config(format!(
"branch '{bname}': on_result.key must not be empty"
)));
}
validate_on_result_field(&cond.key, "key", bname)?;
let (field, values) = match &cond.value {
ResultMatch::Exact(value) => ("result", std::slice::from_ref(value)),
ResultMatch::AnyOf { any_of } => ("result.any_of", any_of.as_slice()),
ResultMatch::Contains { contains } => ("result.contains", std::slice::from_ref(contains)),
ResultMatch::Not { not } => ("result.not", std::slice::from_ref(not)),
};
if values.is_empty() {
return Err(ProxyError::Config(format!(
"branch '{bname}': on_result.{field} must list at least one value"
)));
}
values
.iter()
.try_for_each(|value| validate_on_result_value(value, field, bname))
}
fn validate_on_result_value(value: &str, field: &str, branch_name: &str) -> Result<(), ProxyError> {
if value.is_empty() {
return Err(ProxyError::Config(format!(
"branch '{branch_name}': on_result.{field} must not be empty"
)));
}
if value.len() > MAX_RESULT_VALUE_LEN {
return Err(ProxyError::Config(format!(
"branch '{branch_name}': on_result.{field} must be at most {MAX_RESULT_VALUE_LEN} bytes"
)));
}
validate_on_result_field(value, field, branch_name)
}
fn validate_on_result_field(val: &str, field: &str, branch_name: &str) -> Result<(), ProxyError> {
if !val
.bytes()
.all(|byte| byte.is_ascii_alphanumeric() || byte == b'_' || byte == b'-')
{
return Err(ProxyError::Config(format!(
"branch '{branch_name}': on_result.{field} '{val}' must be ASCII alphanumeric, '_', or '-'"
)));
}
Ok(())
}
fn validate_max_iterations(branch: &BranchChainConfig) -> Result<(), ProxyError> {
let bname = &branch.name;
if let Some(max) = branch.max_iterations
&& !(1..=MAX_ITERATIONS_CEILING).contains(&max)
{
return Err(ProxyError::Config(format!(
"branch '{bname}': max_iterations must be 1-{MAX_ITERATIONS_CEILING}, got {max}"
)));
}
Ok(())
}
fn validate_chain_ref(
chain_ref: &ChainRef,
all_names: &mut HashSet<String>,
chain_names: &HashSet<&str>,
depth: usize,
) -> Result<(), ProxyError> {
match chain_ref {
ChainRef::Named(name) => {
if !chain_names.contains(name.as_str()) {
return Err(ProxyError::Config(format!("branch references unknown chain '{name}'")));
}
},
ChainRef::Inline { name, filters } => {
super::validate_name_chars(name, "inline chain")?;
if depth.saturating_add(1) > MAX_BRANCH_DEPTH {
return Err(ProxyError::Config(format!(
"branch nesting depth exceeds maximum ({MAX_BRANCH_DEPTH})"
)));
}
if filters.len() > super::filter_chain::MAX_FILTERS_PER_CHAIN {
return Err(ProxyError::Config(format!(
"inline chain '{name}' has too many filters ({}, max {})",
filters.len(),
super::filter_chain::MAX_FILTERS_PER_CHAIN
)));
}
if !all_names.insert(name.clone()) {
return Err(ProxyError::Config(format!("duplicate inline chain name '{name}'")));
}
collect_branch_names(filters, all_names, chain_names, depth.saturating_add(1))?;
},
}
Ok(())
}
#[cfg(test)]
#[expect(clippy::allow_attributes, reason = "blanket test suppressions")]
#[allow(
clippy::unwrap_used,
clippy::expect_used,
clippy::indexing_slicing,
clippy::needless_raw_strings,
clippy::needless_raw_string_hashes,
clippy::too_many_lines,
clippy::items_after_statements,
clippy::panic,
reason = "tests use unwrap/expect/indexing/raw strings/panic for brevity"
)]
mod tests {
use std::fmt::Write as _;
use crate::config::Config;
#[test]
fn reject_branch_name_with_special_chars() {
let yaml = r#"
listeners:
- name: web
address: "0.0.0.0:8080"
filter_chains: [main]
filter_chains:
- name: main
filters:
- filter: headers
branch_chains:
- name: "bad.branch"
chains:
- name: inline
filters:
- filter: headers
- filter: static_response
status: 200
"#;
let err = Config::from_yaml(yaml).unwrap_err();
assert!(
err.to_string().contains("alphanumeric"),
"branch name with dots should be rejected: {err}"
);
}
#[test]
fn reject_inline_chain_name_with_special_chars() {
let yaml = r#"
listeners:
- name: web
address: "0.0.0.0:8080"
filter_chains: [main]
filter_chains:
- name: main
filters:
- filter: headers
branch_chains:
- name: branch
chains:
- name: "bad.inline"
filters:
- filter: headers
- filter: static_response
status: 200
"#;
let err = Config::from_yaml(yaml).unwrap_err();
assert!(
err.to_string().contains("alphanumeric"),
"inline chain name with dots should be rejected: {err}"
);
}
#[test]
fn reject_empty_branch_and_inline_chain_names() {
for (branch, inline, expected) in [
("\"\"", "inline", "branch name must not be empty"),
("branch", "\"\"", "inline chain name must not be empty"),
] {
let yaml = format!(
r#"
listeners:
- name: web
address: "0.0.0.0:8080"
filter_chains: [main]
filter_chains:
- name: main
filters:
- filter: headers
branch_chains:
- name: {branch}
chains:
- name: {inline}
filters:
- filter: headers
- filter: static_response
status: 200
"#
);
let err = Config::from_yaml(&yaml).unwrap_err();
assert!(err.to_string().contains(expected), "expected '{expected}': {err}");
}
}
#[test]
fn reject_empty_on_result_key() {
let yaml = r#"
listeners:
- name: web
address: "0.0.0.0:8080"
filter_chains: [main]
filter_chains:
- name: main
filters:
- filter: headers
branch_chains:
- name: branch
on_result:
filter: cache
key: ""
result: hit
chains:
- name: inline
filters:
- filter: headers
- filter: static_response
status: 200
"#;
let err = Config::from_yaml(yaml).unwrap_err();
assert!(
err.to_string().contains("on_result.key must not be empty"),
"empty on_result key should be rejected: {err}"
);
}
#[test]
fn reject_empty_on_result_value() {
let yaml = r#"
listeners:
- name: web
address: "0.0.0.0:8080"
filter_chains: [main]
filter_chains:
- name: main
filters:
- filter: headers
branch_chains:
- name: branch
on_result:
filter: cache
result: ""
chains:
- name: inline
filters:
- filter: headers
- filter: static_response
status: 200
"#;
let err = Config::from_yaml(yaml).unwrap_err();
assert!(
err.to_string().contains("on_result.result must not be empty"),
"empty on_result value should be rejected: {err}"
);
}
#[test]
fn reject_on_result_key_with_special_chars() {
let yaml = r#"
listeners:
- name: web
address: "0.0.0.0:8080"
filter_chains: [main]
filter_chains:
- name: main
filters:
- filter: headers
branch_chains:
- name: branch
on_result:
filter: cache
key: "bad.key"
result: hit
chains:
- name: inline
filters:
- filter: headers
- filter: static_response
status: 200
"#;
let err = Config::from_yaml(yaml).unwrap_err();
assert!(
err.to_string().contains("on_result.key"),
"on_result key with dots should be rejected: {err}"
);
}
#[test]
fn valid_branch_config() {
let yaml = r#"
listeners:
- name: web
address: "0.0.0.0:8080"
filter_chains: [main]
filter_chains:
- name: utility
filters:
- filter: headers
- name: main
filters:
- filter: headers
name: pre_route
branch_chains:
- name: my_branch
chains:
- utility
- filter: static_response
status: 200
"#;
let config = Config::from_yaml(yaml).unwrap();
assert_eq!(config.filter_chains.len(), 2, "should have 2 chains");
}
#[test]
fn reject_duplicate_branch_name() {
let yaml = r#"
listeners:
- name: web
address: "0.0.0.0:8080"
filter_chains: [main]
filter_chains:
- name: main
filters:
- filter: headers
branch_chains:
- name: dup
chains:
- name: inline1
filters:
- filter: headers
- name: dup
chains:
- name: inline2
filters:
- filter: headers
- filter: static_response
status: 200
"#;
let err = Config::from_yaml(yaml).unwrap_err();
assert!(
err.to_string().contains("duplicate branch name"),
"should reject duplicate branch name: {err}"
);
}
#[test]
fn reject_duplicate_filter_name_in_chain() {
let yaml = r#"
listeners:
- name: web
address: "0.0.0.0:8080"
filter_chains: [main]
filter_chains:
- name: main
filters:
- filter: headers
name: same
- filter: cors
name: same
- filter: static_response
status: 200
"#;
let err = Config::from_yaml(yaml).unwrap_err();
assert!(
err.to_string().contains("duplicate filter name"),
"should reject duplicate filter name: {err}"
);
}
#[test]
fn reject_unknown_chain_ref() {
let yaml = r#"
listeners:
- name: web
address: "0.0.0.0:8080"
filter_chains: [main]
filter_chains:
- name: main
filters:
- filter: headers
branch_chains:
- name: my_branch
chains:
- nonexistent_chain
- filter: static_response
status: 200
"#;
let err = Config::from_yaml(yaml).unwrap_err();
assert!(
err.to_string().contains("unknown chain"),
"should reject unknown chain ref: {err}"
);
}
#[test]
fn reject_empty_branch_chains() {
let yaml = r#"
listeners:
- name: web
address: "0.0.0.0:8080"
filter_chains: [main]
filter_chains:
- name: main
filters:
- filter: headers
branch_chains:
- name: empty_branch
chains: []
- filter: static_response
status: 200
"#;
let err = Config::from_yaml(yaml).unwrap_err();
assert!(
err.to_string().contains("at least one chain"),
"should reject empty branch chains: {err}"
);
}
#[test]
fn reject_max_iterations_zero() {
let yaml = r#"
listeners:
- name: web
address: "0.0.0.0:8080"
filter_chains: [main]
filter_chains:
- name: main
filters:
- filter: headers
branch_chains:
- name: retry
max_iterations: 0
chains:
- name: inline
filters:
- filter: headers
- filter: static_response
status: 200
"#;
let err = Config::from_yaml(yaml).unwrap_err();
assert!(
err.to_string().contains("max_iterations must be 1-100"),
"should reject max_iterations=0: {err}"
);
}
#[test]
fn reject_max_iterations_too_high() {
let yaml = r#"
listeners:
- name: web
address: "0.0.0.0:8080"
filter_chains: [main]
filter_chains:
- name: main
filters:
- filter: headers
branch_chains:
- name: retry
max_iterations: 101
chains:
- name: inline
filters:
- filter: headers
- filter: static_response
status: 200
"#;
let err = Config::from_yaml(yaml).unwrap_err();
assert!(
err.to_string().contains("max_iterations must be 1-100"),
"should reject max_iterations=101: {err}"
);
}
#[test]
fn accept_max_iterations_valid_range() {
let yaml = r#"
listeners:
- name: web
address: "0.0.0.0:8080"
filter_chains: [main]
filter_chains:
- name: main
filters:
- filter: headers
branch_chains:
- name: retry
max_iterations: 3
chains:
- name: inline
filters:
- filter: headers
- filter: static_response
status: 200
"#;
Config::from_yaml(yaml).unwrap();
}
#[test]
fn reject_invalid_filter_name_chars() {
let yaml = r#"
listeners:
- name: web
address: "0.0.0.0:8080"
filter_chains: [main]
filter_chains:
- name: main
filters:
- filter: headers
name: "bad.name"
- filter: static_response
status: 200
"#;
let err = Config::from_yaml(yaml).unwrap_err();
assert!(
err.to_string().contains("alphanumeric"),
"should reject filter name with dots: {err}"
);
}
#[test]
fn accept_valid_filter_names() {
let yaml = r#"
listeners:
- name: web
address: "0.0.0.0:8080"
filter_chains: [main]
filter_chains:
- name: main
filters:
- filter: headers
name: pre-route_1
- filter: static_response
status: 200
"#;
Config::from_yaml(yaml).unwrap();
}
#[test]
fn reject_duplicate_inline_chain_name() {
let yaml = r#"
listeners:
- name: web
address: "0.0.0.0:8080"
filter_chains: [main]
filter_chains:
- name: main
filters:
- filter: headers
branch_chains:
- name: branch_a
chains:
- name: inline
filters:
- filter: headers
- name: branch_b
chains:
- name: inline
filters:
- filter: headers
- filter: static_response
status: 200
"#;
let err = Config::from_yaml(yaml).unwrap_err();
assert!(
err.to_string().contains("duplicate inline chain name"),
"should reject duplicate inline chain name: {err}"
);
}
#[test]
fn reject_nesting_depth_exceeded() {
fn nested_yaml(depth: usize) -> String {
let mut yaml = String::from(
r#"
listeners:
- name: web
address: "0.0.0.0:8080"
filter_chains: [main]
filter_chains:
- name: main
filters:
"#,
);
fn write_level(yaml: &mut String, depth: usize, current: usize, indent: usize) {
let pad = " ".repeat(indent);
let filter_name = format!("branch_{current}");
let chain_name = format!("inline_{current}");
writeln!(yaml, "{pad}- filter: headers").unwrap();
if current < depth {
writeln!(yaml, "{pad} branch_chains:").unwrap();
writeln!(yaml, "{pad} - name: {filter_name}").unwrap();
writeln!(yaml, "{pad} chains:").unwrap();
writeln!(yaml, "{pad} - name: {chain_name}").unwrap();
writeln!(yaml, "{pad} filters:").unwrap();
write_level(yaml, depth, current + 1, indent + 12);
}
}
write_level(&mut yaml, depth, 0, 6);
yaml.push_str(
r#" - filter: static_response
status: 200
"#,
);
yaml
}
let yaml = nested_yaml(11);
let err = Config::from_yaml(&yaml).unwrap_err();
assert!(
err.to_string().contains("nesting depth"),
"should reject excessive nesting depth: {err}"
);
}
#[test]
fn accept_nesting_within_limit() {
let yaml = r#"
listeners:
- name: web
address: "0.0.0.0:8080"
filter_chains: [main]
filter_chains:
- name: main
filters:
- filter: headers
branch_chains:
- name: level_0
chains:
- name: inline_0
filters:
- filter: headers
branch_chains:
- name: level_1
chains:
- name: inline_1
filters:
- filter: headers
- filter: static_response
status: 200
"#;
Config::from_yaml(yaml).unwrap();
}
#[test]
fn reject_cross_chain_rejoin() {
let yaml = r#"
listeners:
- name: web
address: "0.0.0.0:8080"
filter_chains: [main]
filter_chains:
- name: main
filters:
- filter: headers
branch_chains:
- name: cross
rejoin: "other:routing"
chains:
- name: inline
filters:
- filter: headers
- filter: static_response
status: 200
"#;
let err = Config::from_yaml(yaml).unwrap_err();
assert!(
err.to_string().contains("cross-chain rejoin"),
"should reject cross-chain rejoin syntax: {err}"
);
}
#[test]
fn reject_empty_filter_name() {
let yaml = r#"
listeners:
- name: web
address: "0.0.0.0:8080"
filter_chains: [main]
filter_chains:
- name: main
filters:
- filter: headers
name: ""
- filter: static_response
status: 200
"#;
let err = Config::from_yaml(yaml).unwrap_err();
assert!(
err.to_string().contains("must not be empty"),
"should reject empty filter name: {err}"
);
}
#[test]
fn accept_max_iterations_at_boundaries() {
let yaml_1 = r#"
listeners:
- name: web
address: "0.0.0.0:8080"
filter_chains: [main]
filter_chains:
- name: main
filters:
- filter: headers
branch_chains:
- name: retry
max_iterations: 1
chains:
- name: inline
filters:
- filter: headers
- filter: static_response
status: 200
"#;
Config::from_yaml(yaml_1).unwrap();
let yaml_100 = r#"
listeners:
- name: web
address: "0.0.0.0:8080"
filter_chains: [main]
filter_chains:
- name: main
filters:
- filter: headers
branch_chains:
- name: retry
max_iterations: 100
chains:
- name: inline
filters:
- filter: headers
- filter: static_response
status: 200
"#;
Config::from_yaml(yaml_100).unwrap();
}
#[test]
fn reject_empty_on_result_filter() {
let yaml = r#"
listeners:
- name: web
address: "0.0.0.0:8080"
filter_chains: [main]
filter_chains:
- name: main
filters:
- filter: headers
branch_chains:
- name: bad_branch
on_result:
filter: ""
result: hit
chains:
- name: inline
filters:
- filter: headers
- filter: static_response
status: 200
"#;
let err = Config::from_yaml(yaml).unwrap_err();
assert!(
err.to_string().contains("on_result.filter must not be empty"),
"should reject empty on_result.filter: {err}"
);
}
#[test]
fn reject_on_result_filter_invalid_chars() {
let yaml = r#"
listeners:
- name: web
address: "0.0.0.0:8080"
filter_chains: [main]
filter_chains:
- name: main
filters:
- filter: headers
branch_chains:
- name: bad_branch
on_result:
filter: "my.filter"
result: hit
chains:
- name: inline
filters:
- filter: headers
- filter: static_response
status: 200
"#;
let err = Config::from_yaml(yaml).unwrap_err();
assert!(
err.to_string().contains("alphanumeric"),
"should reject on_result.filter with invalid chars: {err}"
);
}
#[test]
fn accept_nesting_at_exact_max_depth() {
fn nested_yaml(depth: usize) -> String {
let mut yaml = String::from(
r#"
listeners:
- name: web
address: "0.0.0.0:8080"
filter_chains: [main]
filter_chains:
- name: main
filters:
"#,
);
fn write_level(yaml: &mut String, depth: usize, current: usize, indent: usize) {
let pad = " ".repeat(indent);
let filter_name = format!("branch_{current}");
let chain_name = format!("inline_{current}");
writeln!(yaml, "{pad}- filter: headers").unwrap();
if current < depth {
writeln!(yaml, "{pad} branch_chains:").unwrap();
writeln!(yaml, "{pad} - name: {filter_name}").unwrap();
writeln!(yaml, "{pad} chains:").unwrap();
writeln!(yaml, "{pad} - name: {chain_name}").unwrap();
writeln!(yaml, "{pad} filters:").unwrap();
write_level(yaml, depth, current + 1, indent + 12);
}
}
write_level(&mut yaml, depth, 0, 6);
yaml.push_str(
r#" - filter: static_response
status: 200
"#,
);
yaml
}
let yaml = nested_yaml(super::super::branch_chain::MAX_BRANCH_DEPTH);
Config::from_yaml(&yaml).expect("nesting at exactly MAX_BRANCH_DEPTH should pass");
}
#[test]
fn reject_too_many_branches_per_filter() {
let mut branches = String::new();
for i in 0..17 {
write!(
branches,
" - name: branch_{i}\n chains:\n \
- name: inline_{i}\n filters:\n \
- filter: headers\n"
)
.unwrap();
}
let yaml = format!(
r#"
listeners:
- name: web
address: "0.0.0.0:8080"
filter_chains: [main]
filter_chains:
- name: main
filters:
- filter: headers
branch_chains:
{branches} - filter: static_response
status: 200
"#
);
let err = Config::from_yaml(&yaml).unwrap_err();
assert!(
err.to_string().contains("branch chains"),
"should reject >16 branches per filter: {err}"
);
}
#[test]
fn accept_max_branches_per_filter() {
let mut branches = String::new();
for i in 0..16 {
write!(
branches,
" - name: branch_{i}\n chains:\n \
- name: inline_{i}\n filters:\n \
- filter: headers\n"
)
.unwrap();
}
let yaml = format!(
r#"
listeners:
- name: web
address: "0.0.0.0:8080"
filter_chains: [main]
filter_chains:
- name: main
filters:
- filter: headers
branch_chains:
{branches} - filter: static_response
status: 200
"#
);
Config::from_yaml(&yaml).expect("exactly 16 branches per filter should pass");
}
#[test]
fn reject_inline_chain_too_many_filters() {
let mut filters = String::new();
for _ in 0..101 {
filters.push_str(" - filter: headers\n");
}
let yaml = format!(
r#"
listeners:
- name: web
address: "0.0.0.0:8080"
filter_chains: [main]
filter_chains:
- name: main
filters:
- filter: headers
branch_chains:
- name: big_branch
chains:
- name: big_inline
filters:
{filters} - filter: static_response
status: 200
"#
);
let err = Config::from_yaml(&yaml).unwrap_err();
assert!(
err.to_string().contains("too many filters"),
"inline chain with >100 filters should be rejected: {err}"
);
}
#[test]
fn accept_valid_on_result_filter() {
let yaml = r#"
listeners:
- name: web
address: "0.0.0.0:8080"
filter_chains: [main]
filter_chains:
- name: main
filters:
- filter: headers
branch_chains:
- name: good_branch
on_result:
filter: cache-check_1
result: hit
chains:
- name: inline
filters:
- filter: headers
- filter: static_response
status: 200
"#;
Config::from_yaml(yaml).unwrap();
}
#[test]
fn chain_entries_reject_total_branch_count_over_ceiling() {
use std::collections::HashSet;
use crate::config::FilterEntry;
let mut yaml = String::new();
let mut n = 0;
for _ in 0..17 {
yaml.push_str("- filter: headers\n branch_chains:\n");
for _ in 0..16 {
writeln!(yaml, " - name: br_{n}\n chains: [utility]").unwrap();
n += 1;
}
}
let entries: Vec<FilterEntry> = serde_yaml::from_str(&yaml).unwrap();
let known: HashSet<&str> = HashSet::from(["utility"]);
let err = super::validate_chain_entries_branch_chains("outbound", &entries, &known, 0).unwrap_err();
assert!(
err.to_string().contains("total branch count")
&& err.to_string().contains(&super::MAX_TOTAL_BRANCHES.to_string()),
"an outbound chain whose branch total exceeds the ceiling must be rejected: {err}"
);
}
#[test]
fn chain_entries_reject_branches_over_per_filter_cap() {
use std::collections::HashSet;
use crate::config::FilterEntry;
let mut yaml = String::from("- filter: headers\n branch_chains:\n");
for n in 0..17 {
writeln!(yaml, " - name: br_{n}\n chains: [utility]").unwrap();
}
let entries: Vec<FilterEntry> = serde_yaml::from_str(&yaml).unwrap();
let known: HashSet<&str> = HashSet::from(["utility"]);
let err = super::validate_chain_entries_branch_chains("outbound", &entries, &known, 0).unwrap_err();
assert!(
err.to_string().contains("branch chains") && err.to_string().contains("16"),
"a filter exceeding the per-filter branch cap must be rejected: {err}"
);
}
#[test]
fn chain_entries_reject_invalid_on_result_condition() {
use std::collections::HashSet;
use crate::config::FilterEntry;
let entries: Vec<FilterEntry> = serde_yaml::from_str(
"
- filter: headers
branch_chains:
- name: br
on_result:
filter: \"\"
result: hit
chains: [utility]
",
)
.unwrap();
let known: HashSet<&str> = HashSet::from(["utility"]);
let err = super::validate_chain_entries_branch_chains("outbound", &entries, &known, 0).unwrap_err();
assert!(
err.to_string().contains("on_result.filter must not be empty"),
"a branch with an empty on_result filter must be rejected: {err}"
);
}
#[test]
fn count_build_branches_counts_entries_and_named_chains_ignoring_named_refs() {
use std::collections::HashSet;
use crate::config::FilterEntry;
let entries: Vec<FilterEntry> = serde_yaml::from_str(
"
- filter: headers
branch_chains:
- name: br_named
chains: [utility]
- name: br_inline
chains:
- name: sub
filters:
- filter: headers
",
)
.unwrap();
let utility: Vec<FilterEntry> = serde_yaml::from_str(
"
- filter: headers
branch_chains:
- name: u_br
chains: [utility]
",
)
.unwrap();
let chains: [&[FilterEntry]; 1] = [utility.as_slice()];
let known: HashSet<&str> = HashSet::from(["utility"]);
assert_eq!(
super::count_build_branches(&entries, &chains, &known),
4,
"config-wide counting must count entries and named chains, but not named refs"
);
}
#[test]
fn chain_entries_branch_count_accumulates_prior_and_returns_total() {
use std::collections::HashSet;
use crate::config::FilterEntry;
let mut yaml = String::from("- filter: headers\n branch_chains:\n");
for n in 0..10 {
writeln!(yaml, " - name: br_{n}\n chains: [utility]").unwrap();
}
let entries: Vec<FilterEntry> = serde_yaml::from_str(&yaml).unwrap();
let known: HashSet<&str> = HashSet::from(["utility"]);
let total = super::validate_chain_entries_branch_chains("outbound", &entries, &known, 0).unwrap();
assert_eq!(total, 10, "the cumulative total must include this chain's 10 branches");
let err =
super::validate_chain_entries_branch_chains("outbound", &entries, &known, super::MAX_TOTAL_BRANCHES - 5)
.unwrap_err();
assert!(
err.to_string().contains("total branch count")
&& err.to_string().contains(&(super::MAX_TOTAL_BRANCHES + 5).to_string()),
"prior branch count must accumulate into the cumulative total: {err}"
);
}
#[test]
fn result_operators_load_through_config() {
for result in [
"result: 0",
"result: true",
"result: {contains: unsafe}",
"result: {not: safe}",
"result: {any_of: [2, 4, 6]}",
] {
let config = Config::from_yaml(&on_result_config(result))
.unwrap_or_else(|err| panic!("{result} should load through the full config: {err}"));
assert_eq!(config.filter_chains.len(), 1, "{result}: the chain should be kept");
}
}
#[test]
fn reject_empty_result_operands() {
let cases = [
(
"result: {any_of: []}",
"on_result.result.any_of must list at least one value",
),
(
"result: {any_of: [ok, \"\"]}",
"on_result.result.any_of must not be empty",
),
(
"result: {contains: \"\"}",
"on_result.result.contains must not be empty",
),
("result: {not: \"\"}", "on_result.result.not must not be empty"),
];
for (result, expected) in cases {
let err = Config::from_yaml(&on_result_config(result)).unwrap_err();
assert!(
err.to_string().contains(expected),
"{result} should be rejected with '{expected}': {err}"
);
}
}
#[test]
fn reject_result_operand_with_special_chars() {
let cases = [
(
"result: {any_of: [ok, \"bad.value\"]}",
"on_result.result.any_of 'bad.value'",
),
("result: {contains: \"un safe\"}", "on_result.result.contains 'un safe'"),
("result: {not: \"safe!\"}", "on_result.result.not 'safe!'"),
];
for (result, expected) in cases {
let err = Config::from_yaml(&on_result_config(result)).unwrap_err();
assert!(
err.to_string().contains(expected),
"{result} should be rejected with '{expected}': {err}"
);
}
}
#[test]
fn reject_result_value_longer_than_a_filter_result() {
let at_limit = "a".repeat(super::MAX_RESULT_VALUE_LEN);
let over_limit = "a".repeat(super::MAX_RESULT_VALUE_LEN + 1);
Config::from_yaml(&on_result_config(&format!("result: {{contains: {at_limit}}}")))
.unwrap_or_else(|err| panic!("a value at the limit should load: {err}"));
let exact = Config::from_yaml(&on_result_config(&format!("result: {over_limit}"))).unwrap_err();
let any_of =
Config::from_yaml(&on_result_config(&format!("result: {{any_of: [ok, {over_limit}]}}"))).unwrap_err();
assert!(
exact.to_string().contains("on_result.result must be at most 256 bytes"),
"an exact value over the limit should be rejected: {exact}"
);
assert!(
any_of
.to_string()
.contains("on_result.result.any_of must be at most 256 bytes"),
"an any_of entry over the limit should be rejected: {any_of}"
);
}
fn on_result_config(result: &str) -> String {
format!(
r#"
listeners:
- name: web
address: "0.0.0.0:8080"
filter_chains: [main]
filter_chains:
- name: main
filters:
- filter: headers
branch_chains:
- name: branch
on_result:
filter: headers
key: verdict
{result}
chains:
- name: inline
filters:
- filter: headers
- filter: static_response
status: 200
"#
)
}
}