Skip to main content

volo_grpc/layer/loadbalance/
mod.rs

1use std::{fmt::Debug, sync::Arc};
2
3use async_broadcast::RecvError;
4use motore::Service;
5use tracing::warn;
6use volo::{
7    Layer,
8    context::Context,
9    discovery::Discover,
10    loadbalance::{LoadBalance, MkLbLayer, error::LoadBalanceError},
11};
12
13use crate::Request;
14
15#[derive(Clone, Default, Copy)]
16pub struct LoadBalanceLayer<D, LB> {
17    discover: D,
18    load_balance: LB,
19}
20
21impl<D, LB> LoadBalanceLayer<D, LB> {
22    pub fn new(discover: D, load_balance: LB) -> Self {
23        LoadBalanceLayer {
24            discover,
25            load_balance,
26        }
27    }
28}
29
30impl<D, LB, S> Layer<S> for LoadBalanceLayer<D, LB>
31where
32    D: Discover,
33    LB: LoadBalance<D>,
34{
35    type Service = LoadBalanceService<D, LB, S>;
36
37    fn layer(self, inner: S) -> Self::Service {
38        LoadBalanceService::new(self.discover, self.load_balance, inner)
39    }
40}
41#[derive(Clone)]
42pub struct LoadBalanceService<D, LB, S> {
43    discover: D,
44    load_balance: Arc<LB>,
45    service: S,
46}
47
48impl<D, LB, S> LoadBalanceService<D, LB, S>
49where
50    D: Discover,
51    LB: LoadBalance<D>,
52{
53    pub fn new(discover: D, load_balance: LB, service: S) -> Self {
54        let lb = Arc::new(load_balance);
55
56        let service = Self {
57            discover,
58            load_balance: lb.clone(),
59            service,
60        };
61
62        if let Some(mut channel) = service.discover.watch(None) {
63            tokio::spawn(async move {
64                loop {
65                    match channel.recv().await {
66                        Ok(recv) => lb.rebalance(recv),
67                        Err(err) => match err {
68                            RecvError::Closed => break,
69                            _ => warn!("[VOLO] discovering subscription error {:?}", err),
70                        },
71                    }
72                }
73            });
74        }
75        service
76    }
77}
78
79impl<Cx, T, D, LB, S> Service<Cx, Request<T>> for LoadBalanceService<D, LB, S>
80where
81    <Cx as Context>::Config: Sync,
82    Cx: 'static + Context + Send + Sync,
83    D: Discover,
84    LB: LoadBalance<D>,
85    S: Service<Cx, Request<T>> + 'static + Send + Sync,
86    LoadBalanceError: Into<S::Error>,
87    S::Error: Debug,
88    T: Send + 'static,
89{
90    type Response = S::Response;
91
92    type Error = S::Error;
93
94    async fn call(&self, cx: &mut Cx, req: Request<T>) -> Result<Self::Response, Self::Error> {
95        let callee = cx.rpc_info().callee();
96
97        let mut picker = match &callee.address {
98            None => self
99                .load_balance
100                .get_picker(callee, &self.discover)
101                .await
102                .map_err(|err| err.into())?,
103            _ => {
104                return self.service.call(cx, req).await;
105            }
106        };
107
108        if let Some(addr) = picker.next() {
109            cx.rpc_info_mut().callee_mut().address = Some(addr.clone());
110
111            return match self.service.call(cx, req).await {
112                Ok(resp) => Ok(resp),
113                Err(err) => {
114                    warn!("[VOLO] call endpoint: {:?} error: {:?}", addr, err);
115                    Err(err)
116                }
117            };
118        } else {
119            warn!("[VOLO] zero call count, call info: {:?}", cx.rpc_info());
120        }
121        Err(LoadBalanceError::Retry).map_err(|err| err.into())?
122    }
123}
124
125impl<D, LB, S> Debug for LoadBalanceService<D, LB, S>
126where
127    D: Debug,
128    LB: Debug,
129    S: Debug,
130{
131    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
132        f.debug_struct("LBService")
133            .field("discover", &self.discover)
134            .field("load_balancer", &self.load_balance)
135            .finish()
136    }
137}
138
139pub struct LbConfig<L, DISC> {
140    load_balance: L,
141    discover: DISC,
142}
143
144impl<L, DISC> LbConfig<L, DISC> {
145    pub fn new(load_balance: L, discover: DISC) -> Self {
146        LbConfig {
147            load_balance,
148            discover,
149        }
150    }
151
152    pub fn load_balance<NL>(self, load_balance: NL) -> LbConfig<NL, DISC> {
153        LbConfig {
154            load_balance,
155            discover: self.discover,
156        }
157    }
158
159    pub fn discover<NDISC>(self, discover: NDISC) -> LbConfig<L, NDISC> {
160        LbConfig {
161            load_balance: self.load_balance,
162            discover,
163        }
164    }
165}
166
167impl<LB, DISC> MkLbLayer for LbConfig<LB, DISC> {
168    type Layer = LoadBalanceLayer<DISC, LB>;
169
170    fn make(self) -> Self::Layer {
171        LoadBalanceLayer::new(self.discover, self.load_balance)
172    }
173}