Skip to main content

satex_load_balancer/resolver/
static.rs

1use crate::discovery::StaticFixedDiscovery;
2use crate::health_check::tcp::TcpHealthCheck;
3use crate::resolver::make::MakeLoadBalancerResolver;
4use crate::resolver::LoadBalancerResolver;
5use crate::selector::{BoxSelector, Consistent, Random, RoundRobin};
6use crate::{Backend, Backends, LoadBalancer};
7use satex_core::background::background_task;
8use satex_core::component::{Args, Configurable};
9use satex_core::Error;
10use satex_macro::make;
11use serde::Deserialize;
12use std::collections::{BTreeSet, HashMap};
13use std::net::SocketAddr;
14use std::sync::Arc;
15use tokio::spawn;
16
17pub struct StaticLoadBalancerResolver {
18    load_balancers: HashMap<String, Arc<LoadBalancer>>,
19}
20
21impl LoadBalancerResolver for StaticLoadBalancerResolver {
22    fn find(&self, name: &str) -> Option<Arc<LoadBalancer>> {
23        self.load_balancers.get(name).cloned()
24    }
25}
26
27#[derive(Deserialize)]
28struct Upstream {
29    name: String,
30    #[serde(default)]
31    policy: Policy,
32    addrs: Vec<SocketAddr>,
33    #[serde(default, rename = "health-check")]
34    health_check: HealthCheck,
35}
36
37#[derive(Deserialize, Default)]
38enum Policy {
39    #[default]
40    RoundRobin,
41    Random,
42    Consistent,
43}
44
45#[derive(Deserialize)]
46struct HealthCheck {
47    enabled: bool,
48}
49
50impl Default for HealthCheck {
51    fn default() -> Self {
52        Self { enabled: true }
53    }
54}
55
56#[make(kind = "Static", shortcut_mode = Sequence)]
57pub struct MakeStaticLoadBalancerResolver {
58    upstreams: Vec<Upstream>,
59}
60
61impl MakeLoadBalancerResolver for MakeStaticLoadBalancerResolver {
62    type Resolver = StaticLoadBalancerResolver;
63
64    fn make(&self, args: Args) -> Result<Self::Resolver, Error> {
65        let config = Config::with_args(args)?;
66        let load_balancers = config
67            .upstreams
68            .into_iter()
69            .map(|upstream| {
70                let backends = upstream
71                    .addrs
72                    .into_iter()
73                    .map(Backend::new)
74                    .collect::<BTreeSet<_>>();
75
76                let selector = match upstream.policy {
77                    Policy::RoundRobin => BoxSelector::new(RoundRobin::new(&backends)),
78                    Policy::Random => BoxSelector::new(Random::new(&backends)),
79                    Policy::Consistent => BoxSelector::new(Consistent::new(&backends)),
80                };
81
82                let backends = Backends::new(StaticFixedDiscovery::new(backends))
83                    .with_health_check(TcpHealthCheck::default());
84                let load_balancer = Arc::new(LoadBalancer::new(backends, selector));
85
86                if upstream.health_check.enabled {
87                    let task = background_task(
88                        format!("LoadBalancer - {}", upstream.name),
89                        load_balancer.clone(),
90                    );
91                    spawn(task);
92                }
93
94                (upstream.name, load_balancer)
95            })
96            .collect::<HashMap<_, _>>();
97        Ok(StaticLoadBalancerResolver { load_balancers })
98    }
99}