eggress_protocol_reverse/
client.rs1use crate::metrics::ReverseMetrics;
2use crate::{client_auth_handshake, relay_bidirectional_with_timeout, ControlState, ProtocolError};
3use std::net::SocketAddr;
4use std::sync::Arc;
5use std::time::{Duration, Instant};
6use tokio::net::TcpStream;
7use tokio_util::sync::CancellationToken;
8use tracing::{debug, info, warn};
9
10#[derive(Debug, Clone)]
12pub struct ReverseClientConfig {
13 pub server_addr: SocketAddr,
15 pub auth_username: Option<String>,
17 pub auth_password: Option<String>,
19 pub reconnect_initial_ms: u64,
21 pub reconnect_max_ms: u64,
23 pub default_target_host: Option<String>,
25 pub default_target_port: Option<u16>,
27 pub read_timeout_ms: u64,
30 pub drain_grace_ms: u64,
32 pub target_connect_timeout_ms: u64,
34}
35
36impl Default for ReverseClientConfig {
37 fn default() -> Self {
38 Self {
39 server_addr: "127.0.0.1:0".parse().unwrap(),
40 auth_username: None,
41 auth_password: None,
42 reconnect_initial_ms: 1_000,
43 reconnect_max_ms: 30_000,
44 default_target_host: None,
45 default_target_port: None,
46 read_timeout_ms: 60_000,
47 drain_grace_ms: 5_000,
48 target_connect_timeout_ms: 10_000,
49 }
50 }
51}
52
53#[derive(Debug, Clone, PartialEq, Eq)]
55pub enum TargetResolution {
56 Connect { host: String, port: u16 },
58 Reject { reason: String },
60}
61
62pub trait TargetResolver: Send + Sync {
69 fn resolve(&self) -> TargetResolution;
70}
71
72pub struct DefaultTargetResolver {
74 pub host: Option<String>,
75 pub port: Option<u16>,
76}
77
78impl DefaultTargetResolver {
79 pub fn new(host: Option<String>, port: Option<u16>) -> Self {
80 Self { host, port }
81 }
82}
83
84impl TargetResolver for DefaultTargetResolver {
85 fn resolve(&self) -> TargetResolution {
86 match (&self.host, self.port) {
87 (Some(h), Some(p)) => TargetResolution::Connect {
88 host: h.clone(),
89 port: p,
90 },
91 _ => TargetResolution::Reject {
92 reason: "no default target configured".to_string(),
93 },
94 }
95 }
96}
97
98pub struct ReverseClient {
106 config: ReverseClientConfig,
107 cancel: CancellationToken,
108 metrics: Option<Arc<ReverseMetrics>>,
109 resolver: Option<Arc<dyn TargetResolver>>,
110}
111
112impl ReverseClient {
113 pub fn new(config: ReverseClientConfig) -> Self {
114 let resolver: Arc<dyn TargetResolver> = Arc::new(DefaultTargetResolver::new(
115 config.default_target_host.clone(),
116 config.default_target_port,
117 ));
118 Self {
119 config,
120 cancel: CancellationToken::new(),
121 metrics: None,
122 resolver: Some(resolver),
123 }
124 }
125
126 pub fn set_metrics(&mut self, metrics: Arc<ReverseMetrics>) {
128 self.metrics = Some(metrics);
129 }
130
131 pub fn set_resolver(&mut self, resolver: Arc<dyn TargetResolver>) {
133 self.resolver = Some(resolver);
134 }
135
136 pub fn cancel_token(&self) -> CancellationToken {
138 self.cancel.clone()
139 }
140
141 pub async fn run(&self) -> Result<(), ProtocolError> {
143 let mut backoff_ms = self.config.reconnect_initial_ms;
144
145 loop {
146 if self.cancel.is_cancelled() {
147 break;
148 }
149 let session_start = Instant::now();
150 match self.run_session().await {
151 Ok(()) => {
152 if let Some(ref m) = self.metrics {
153 m.record_state_duration(
154 ControlState::Ready,
155 session_start.elapsed().as_millis() as u64,
156 );
157 }
158 if self.cancel.is_cancelled() {
159 break;
160 }
161 backoff_ms = self.config.reconnect_initial_ms;
164 debug!("session ended, reconnecting immediately");
165 }
166 Err(e) => {
167 if self.cancel.is_cancelled() {
168 break;
169 }
170 if let Some(ref m) = self.metrics {
171 m.record_reconnect();
172 m.record_state_duration(
173 ControlState::Connecting,
174 session_start.elapsed().as_millis() as u64,
175 );
176 }
177 warn!(error = %e, backoff_ms, "session failed, reconnecting");
178 let sleep = tokio::time::sleep(Duration::from_millis(backoff_ms));
179 tokio::select! {
180 _ = sleep => {}
181 _ = self.cancel.cancelled() => break,
182 }
183 backoff_ms = (backoff_ms * 2).min(self.config.reconnect_max_ms);
184 }
185 }
186 }
187
188 let drain_start = Instant::now();
190 tokio::time::sleep(Duration::from_millis(50)).await;
191 if let Some(ref m) = self.metrics {
192 m.record_drain(drain_start.elapsed().as_millis() as u64);
193 }
194 info!("reverse client shut down");
195 Ok(())
196 }
197
198 async fn run_session(&self) -> Result<(), ProtocolError> {
200 let connecting_start = Instant::now();
201 let stream = TcpStream::connect(&self.config.server_addr).await?;
202 if let Some(ref m) = self.metrics {
203 m.record_state_duration(
204 ControlState::Connecting,
205 connecting_start.elapsed().as_millis() as u64,
206 );
207 }
208 info!(
209 server = %self.config.server_addr,
210 state = ?ControlState::Connecting,
211 "connected to reverse server"
212 );
213
214 let authenticating_start = Instant::now();
216 let stream = if let (Some(ref username), Some(ref password)) =
217 (&self.config.auth_username, &self.config.auth_password)
218 {
219 let mut s = stream;
220 client_auth_handshake(&mut s, username, password).await?;
221 if let Some(ref m) = self.metrics {
222 m.record_state_duration(
223 ControlState::Authenticating,
224 authenticating_start.elapsed().as_millis() as u64,
225 );
226 }
227 info!(
228 state = ?ControlState::Authenticating,
229 "authentication successful"
230 );
231 s
232 } else {
233 let mut s = stream;
235 crate::read_handshake(&mut s).await?;
236 if let Some(ref m) = self.metrics {
237 m.record_state_duration(
238 ControlState::Authenticating,
239 authenticating_start.elapsed().as_millis() as u64,
240 );
241 }
242 s
243 };
244
245 if let Some(ref m) = self.metrics {
246 m.record_stream_opened();
247 }
248 let ready_start = Instant::now();
249
250 let resolution =
252 self.resolver
253 .as_ref()
254 .map(|r| r.resolve())
255 .unwrap_or(TargetResolution::Reject {
256 reason: "no resolver configured".to_string(),
257 });
258
259 let session_result: Result<(), ProtocolError> = match resolution {
260 TargetResolution::Connect { host, port } => {
261 let target_addr = format!("{}:{}", host, port);
262 let connect_timeout = if self.config.target_connect_timeout_ms > 0 {
263 Duration::from_millis(self.config.target_connect_timeout_ms)
264 } else {
265 Duration::from_secs(30)
266 };
267 let connect_result =
268 tokio::time::timeout(connect_timeout, TcpStream::connect(&target_addr)).await;
269 match connect_result {
270 Ok(Ok(target_stream)) => {
271 info!(
272 target = %target_addr,
273 state = ?ControlState::Ready,
274 "connected to target, relaying"
275 );
276 relay_bidirectional_with_timeout(
277 stream,
278 target_stream,
279 (self.config.read_timeout_ms > 0)
280 .then(|| Duration::from_millis(self.config.read_timeout_ms)),
281 )
282 .await
283 }
284 Ok(Err(e)) => {
285 warn!(
286 target = %target_addr,
287 error = %e,
288 "failed to connect to target"
289 );
290 if let Some(ref m) = self.metrics {
291 m.record_error(&format!("target connect failed: {e}"));
292 }
293 Err(ProtocolError::Io(e))
294 }
295 Err(_elapsed) => {
296 let e = std::io::Error::new(
297 std::io::ErrorKind::TimedOut,
298 format!("target connect timed out: {target_addr}"),
299 );
300 warn!(
301 target = %target_addr,
302 "target connect timed out"
303 );
304 if let Some(ref m) = self.metrics {
305 m.record_error(&format!("target connect failed: {e}"));
306 }
307 Err(ProtocolError::Io(e))
308 }
309 }
310 }
311 TargetResolution::Reject { reason } => {
312 warn!(reason = %reason, "route resolution rejected, dropping control channel");
313 if let Some(ref m) = self.metrics {
314 m.record_error(&format!("route rejected: {reason}"));
315 }
316 Err(ProtocolError::AuthFailed)
317 }
318 };
319
320 if let Some(ref m) = self.metrics {
321 m.record_state_duration(
322 ControlState::Ready,
323 ready_start.elapsed().as_millis() as u64,
324 );
325 m.record_stream_closed(0);
326 }
327
328 session_result
329 }
330
331 pub fn shutdown(&self) {
333 self.cancel.cancel();
334 }
335}
336
337#[cfg(test)]
338mod tests {
339 use super::*;
340
341 struct FixedResolver(TargetResolution);
342 impl TargetResolver for FixedResolver {
343 fn resolve(&self) -> TargetResolution {
344 self.0.clone()
345 }
346 }
347
348 #[test]
349 fn default_resolver_returns_configured_target() {
350 let r = DefaultTargetResolver::new(Some("127.0.0.1".to_string()), Some(8080));
351 assert_eq!(
352 r.resolve(),
353 TargetResolution::Connect {
354 host: "127.0.0.1".to_string(),
355 port: 8080,
356 }
357 );
358 }
359
360 #[test]
361 fn default_resolver_rejects_when_unset() {
362 let r = DefaultTargetResolver::new(None, None);
363 match r.resolve() {
364 TargetResolution::Reject { .. } => {}
365 other => panic!("expected Reject, got {other:?}"),
366 }
367 }
368
369 #[test]
370 fn default_resolver_rejects_partial() {
371 let r = DefaultTargetResolver::new(Some("127.0.0.1".to_string()), None);
372 match r.resolve() {
373 TargetResolution::Reject { .. } => {}
374 other => panic!("expected Reject, got {other:?}"),
375 }
376 }
377
378 #[test]
379 fn custom_resolver_can_reject() {
380 let r: Arc<dyn TargetResolver> = Arc::new(FixedResolver(TargetResolution::Reject {
381 reason: "policy".to_string(),
382 }));
383 match r.resolve() {
384 TargetResolution::Reject { reason } => assert_eq!(reason, "policy"),
385 other => panic!("expected Reject, got {other:?}"),
386 }
387 }
388}