volo_grpc/layer/loadbalance/
mod.rs1use 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}