use std::{collections::HashMap, sync::Arc};
use praxis_core::{
config::{SimpleStrategy, SubsetFallbackPolicy},
health::ClusterHealthState,
};
use super::{
endpoint::WeightedEndpoint,
strategy::{Strategy, build_simple_strategy},
};
pub(crate) struct Subset {
subset_strategy: Option<Box<Strategy>>,
fallback_strategy: Box<Strategy>,
fallback_policy: SubsetFallbackPolicy,
subset_addresses: Vec<Arc<str>>,
}
impl Subset {
pub(crate) fn new(
endpoints: Vec<WeightedEndpoint>,
selector: &HashMap<String, String>,
inner_strategy: &SimpleStrategy,
fallback_policy: SubsetFallbackPolicy,
) -> Self {
let matched: Vec<WeightedEndpoint> = endpoints
.iter()
.filter(|ep| {
selector
.iter()
.all(|(k, v)| ep.metadata.get(k).is_some_and(|mv| mv == v))
})
.cloned()
.collect();
let subset_addresses: Vec<Arc<str>> = matched.iter().map(|ep| Arc::clone(&ep.address)).collect();
let subset_strategy = if matched.is_empty() {
None
} else {
Some(Box::new(build_simple_strategy(inner_strategy, matched)))
};
let fallback_strategy = Box::new(build_simple_strategy(inner_strategy, endpoints));
Self {
subset_strategy,
fallback_strategy,
fallback_policy,
subset_addresses,
}
}
pub(crate) fn select(
&self,
hash_key: Option<&str>,
health: Option<&ClusterHealthState>,
exclude: &[Arc<str>],
) -> Option<Arc<str>> {
if let Some(strategy) = &self.subset_strategy
&& !self.all_subset_unhealthy(health)
{
let result = strategy.select(hash_key, health, exclude);
if result.is_some() {
return result;
}
}
match self.fallback_policy {
SubsetFallbackPolicy::AnyEndpoint => self.fallback_strategy.select(hash_key, health, exclude),
SubsetFallbackPolicy::NoEndpoint => None,
}
}
fn all_subset_unhealthy(&self, health: Option<&ClusterHealthState>) -> bool {
let Some(state) = health else {
return false;
};
!self.subset_addresses.is_empty() && self.subset_addresses.iter().all(|addr| !state.is_address_healthy(addr))
}
pub(crate) fn release(&self, addr: &str) {
if let Some(strategy) = &self.subset_strategy {
strategy.release(addr);
}
self.fallback_strategy.release(addr);
}
}
#[cfg(test)]
#[expect(clippy::allow_attributes, reason = "blanket test suppressions")]
#[allow(
clippy::unwrap_used,
clippy::expect_used,
clippy::indexing_slicing,
clippy::panic,
reason = "tests"
)]
mod tests {
use std::collections::HashSet;
use praxis_core::health::{ClusterHealthEntry, EndpointHealth};
use super::*;
#[test]
fn selects_from_matching_subset() {
let endpoints = vec![
ep("10.0.0.1:80", &[("version", "stable")]),
ep("10.0.0.2:80", &[("version", "canary")]),
ep("10.0.0.3:80", &[("version", "canary")]),
];
let selector = HashMap::from([("version".to_owned(), "canary".to_owned())]);
let subset = Subset::new(
endpoints,
&selector,
&SimpleStrategy::RoundRobin,
SubsetFallbackPolicy::AnyEndpoint,
);
let mut seen = HashSet::new();
for _ in 0..10 {
let addr = subset.select(None, None, &[]).unwrap();
seen.insert(addr);
}
assert!(!seen.contains("10.0.0.1:80"), "stable endpoint should not be selected");
assert!(seen.contains("10.0.0.2:80"), "canary endpoint 2 should be selected");
assert!(seen.contains("10.0.0.3:80"), "canary endpoint 3 should be selected");
}
#[test]
fn fallback_any_endpoint_when_no_match() {
let endpoints = vec![
ep("10.0.0.1:80", &[("version", "stable")]),
ep("10.0.0.2:80", &[("version", "stable")]),
];
let selector = HashMap::from([("version".to_owned(), "canary".to_owned())]);
let subset = Subset::new(
endpoints,
&selector,
&SimpleStrategy::RoundRobin,
SubsetFallbackPolicy::AnyEndpoint,
);
let addr = subset.select(None, None, &[]);
assert!(addr.is_some(), "AnyEndpoint fallback should return an endpoint");
}
#[test]
fn fallback_no_endpoint_when_no_match() {
let endpoints = vec![
ep("10.0.0.1:80", &[("version", "stable")]),
ep("10.0.0.2:80", &[("version", "stable")]),
];
let selector = HashMap::from([("version".to_owned(), "canary".to_owned())]);
let subset = Subset::new(
endpoints,
&selector,
&SimpleStrategy::RoundRobin,
SubsetFallbackPolicy::NoEndpoint,
);
let addr = subset.select(None, None, &[]);
assert!(addr.is_none(), "NoEndpoint fallback should return None");
}
#[test]
fn multi_key_selector() {
let endpoints = vec![
ep("10.0.0.1:80", &[("version", "canary"), ("gpu", "a100")]),
ep("10.0.0.2:80", &[("version", "canary"), ("gpu", "h100")]),
ep("10.0.0.3:80", &[("version", "stable"), ("gpu", "a100")]),
];
let selector = HashMap::from([
("version".to_owned(), "canary".to_owned()),
("gpu".to_owned(), "a100".to_owned()),
]);
let subset = Subset::new(
endpoints,
&selector,
&SimpleStrategy::RoundRobin,
SubsetFallbackPolicy::AnyEndpoint,
);
for _ in 0..10 {
let addr = subset.select(None, None, &[]).unwrap();
assert_eq!(&*addr, "10.0.0.1:80", "only endpoint matching both keys");
}
}
#[test]
fn empty_selector_matches_all() {
let endpoints = vec![
ep("10.0.0.1:80", &[("version", "stable")]),
ep("10.0.0.2:80", &[("version", "canary")]),
];
let selector = HashMap::new();
let subset = Subset::new(
endpoints,
&selector,
&SimpleStrategy::RoundRobin,
SubsetFallbackPolicy::AnyEndpoint,
);
let mut seen = HashSet::new();
for _ in 0..10 {
seen.insert(subset.select(None, None, &[]).unwrap());
}
assert_eq!(seen.len(), 2, "empty selector should match all endpoints");
}
#[test]
fn fallback_when_all_subset_unhealthy() {
let endpoints = vec![
ep("10.0.0.1:80", &[("version", "canary")]),
ep("10.0.0.2:80", &[("version", "canary")]),
ep("10.0.0.3:80", &[("version", "stable")]),
ep("10.0.0.4:80", &[("version", "stable")]),
];
let selector = HashMap::from([("version".to_owned(), "canary".to_owned())]);
let subset = Subset::new(
endpoints,
&selector,
&SimpleStrategy::RoundRobin,
SubsetFallbackPolicy::AnyEndpoint,
);
let state = health_state(4);
state.endpoints()[0].mark_unhealthy();
state.endpoints()[1].mark_unhealthy();
let mut seen = HashSet::new();
for _ in 0..20 {
seen.insert(subset.select(None, Some(&state), &[]).unwrap());
}
assert!(!seen.contains("10.0.0.1:80"), "unhealthy canary should not be selected");
assert!(!seen.contains("10.0.0.2:80"), "unhealthy canary should not be selected");
assert!(
seen.contains("10.0.0.3:80") || seen.contains("10.0.0.4:80"),
"fallback should route to healthy stable endpoints"
);
}
fn ep(addr: &str, meta: &[(&str, &str)]) -> WeightedEndpoint {
WeightedEndpoint {
address: Arc::from(addr),
weight: 1,
metadata: meta.iter().map(|(k, v)| ((*k).to_owned(), (*v).to_owned())).collect(),
priority: 0,
zone: None,
}
}
fn health_state(n: usize) -> Arc<ClusterHealthEntry> {
let healths: Vec<_> = std::iter::repeat_with(EndpointHealth::new).take(n).collect();
let addrs: Vec<_> = (0..n)
.map(|i| Arc::from(format!("10.0.0.{}:80", i + 1).as_str()))
.collect();
Arc::new(ClusterHealthEntry::new(healths, addrs, None, None))
}
}