eggress_protocol_reverse/
client.rs1use crate::metrics::ReverseMetrics;
2use crate::{client_auth_handshake, 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(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 pub tls: Option<crate::tls::ReverseClientTlsConfig>,
37}
38
39impl std::fmt::Debug for ReverseClientConfig {
40 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
41 f.debug_struct("ReverseClientConfig")
42 .field("server_addr", &self.server_addr)
43 .field("auth_username", &self.auth_username)
44 .field(
45 "auth_password",
46 &self.auth_password.as_deref().map(|_| "****"),
47 )
48 .field("reconnect_initial_ms", &self.reconnect_initial_ms)
49 .field("reconnect_max_ms", &self.reconnect_max_ms)
50 .field("default_target_host", &self.default_target_host)
51 .field("default_target_port", &self.default_target_port)
52 .field("read_timeout_ms", &self.read_timeout_ms)
53 .field("drain_grace_ms", &self.drain_grace_ms)
54 .field("target_connect_timeout_ms", &self.target_connect_timeout_ms)
55 .field("tls", &self.tls)
56 .finish()
57 }
58}
59
60impl ReverseClientConfig {
61 pub fn validate(&self) -> Result<(), ProtocolError> {
63 if let Some(ref tls) = self.tls {
64 tls.validate()?;
65 }
66 Ok(())
67 }
68}
69
70impl Default for ReverseClientConfig {
71 fn default() -> Self {
72 Self {
73 server_addr: "127.0.0.1:0".parse().unwrap(),
74 auth_username: None,
75 auth_password: None,
76 reconnect_initial_ms: 1_000,
77 reconnect_max_ms: 30_000,
78 default_target_host: None,
79 default_target_port: None,
80 read_timeout_ms: 60_000,
81 drain_grace_ms: 5_000,
82 target_connect_timeout_ms: 10_000,
83 tls: None,
84 }
85 }
86}
87
88#[derive(Debug, Clone, PartialEq, Eq)]
90pub enum TargetResolution {
91 Connect { host: String, port: u16 },
93 Reject { reason: String },
95}
96
97pub trait TargetResolver: Send + Sync {
104 fn resolve(&self) -> TargetResolution;
105}
106
107pub struct DefaultTargetResolver {
109 pub host: Option<String>,
110 pub port: Option<u16>,
111}
112
113impl DefaultTargetResolver {
114 pub fn new(host: Option<String>, port: Option<u16>) -> Self {
115 Self { host, port }
116 }
117}
118
119impl TargetResolver for DefaultTargetResolver {
120 fn resolve(&self) -> TargetResolution {
121 match (&self.host, self.port) {
122 (Some(h), Some(p)) => TargetResolution::Connect {
123 host: h.clone(),
124 port: p,
125 },
126 _ => TargetResolution::Reject {
127 reason: "no default target configured".to_string(),
128 },
129 }
130 }
131}
132
133pub struct ReverseClient {
141 config: ReverseClientConfig,
142 cancel: CancellationToken,
143 metrics: Option<Arc<ReverseMetrics>>,
144 resolver: Option<Arc<dyn TargetResolver>>,
145}
146
147impl ReverseClient {
148 pub fn new(config: ReverseClientConfig) -> Self {
149 let resolver: Arc<dyn TargetResolver> = Arc::new(DefaultTargetResolver::new(
150 config.default_target_host.clone(),
151 config.default_target_port,
152 ));
153 Self {
154 config,
155 cancel: CancellationToken::new(),
156 metrics: None,
157 resolver: Some(resolver),
158 }
159 }
160
161 pub fn set_metrics(&mut self, metrics: Arc<ReverseMetrics>) {
163 self.metrics = Some(metrics);
164 }
165
166 pub fn set_resolver(&mut self, resolver: Arc<dyn TargetResolver>) {
168 self.resolver = Some(resolver);
169 }
170
171 pub fn cancel_token(&self) -> CancellationToken {
173 self.cancel.clone()
174 }
175
176 pub async fn run(&self) -> Result<(), ProtocolError> {
178 self.config.validate()?;
181
182 let tls_client_config: Option<(Arc<rustls::ClientConfig>, String)> =
185 match self.config.tls.as_ref() {
186 Some(tls) => {
187 let cfg = tls.build_client_config()?;
188 Some((cfg, tls.server_name.clone()))
189 }
190 None => None,
191 };
192
193 let mut backoff_ms = self.config.reconnect_initial_ms;
194
195 loop {
196 if self.cancel.is_cancelled() {
197 break;
198 }
199 let session_start = Instant::now();
200 match self.run_session(tls_client_config.as_ref()).await {
201 Ok(()) => {
202 if let Some(ref m) = self.metrics {
203 m.record_state_duration(
204 ControlState::Ready,
205 session_start.elapsed().as_millis() as u64,
206 );
207 }
208 if self.cancel.is_cancelled() {
209 break;
210 }
211 backoff_ms = self.config.reconnect_initial_ms;
214 debug!("session ended, reconnecting immediately");
215 }
216 Err(e) => {
217 if self.cancel.is_cancelled() {
218 break;
219 }
220 if let Some(ref m) = self.metrics {
221 m.record_reconnect();
222 m.record_state_duration(
223 ControlState::Connecting,
224 session_start.elapsed().as_millis() as u64,
225 );
226 }
227 warn!(error = %e, backoff_ms, "session failed, reconnecting");
228 let sleep = tokio::time::sleep(Duration::from_millis(backoff_ms));
229 tokio::select! {
230 _ = sleep => {}
231 _ = self.cancel.cancelled() => break,
232 }
233 backoff_ms = (backoff_ms * 2).min(self.config.reconnect_max_ms);
234 }
235 }
236 }
237
238 let drain_start = Instant::now();
240 tokio::time::sleep(Duration::from_millis(50)).await;
241 if let Some(ref m) = self.metrics {
242 m.record_drain(drain_start.elapsed().as_millis() as u64);
243 }
244 info!("reverse client shut down");
245 Ok(())
246 }
247
248 async fn run_session(
252 &self,
253 tls: Option<&(Arc<rustls::ClientConfig>, String)>,
254 ) -> Result<(), ProtocolError> {
255 let connecting_start = Instant::now();
256 let tcp = tokio::select! {
257 result = TcpStream::connect(&self.config.server_addr) => {
258 result?
259 }
260 _ = self.cancel.cancelled() => {
261 return Err(ProtocolError::ConnectionClosed);
262 }
263 };
264 if let Some(ref m) = self.metrics {
265 m.record_state_duration(
266 ControlState::Connecting,
267 connecting_start.elapsed().as_millis() as u64,
268 );
269 }
270 info!(
271 server = %self.config.server_addr,
272 state = ?ControlState::Connecting,
273 "connected to reverse server"
274 );
275
276 let mut boxed: eggress_core::BoxStream = if let Some((cfg, server_name)) = tls {
278 let tcp_boxed: eggress_core::BoxStream = Box::new(tcp);
279 let handshake = eggress_transport_tls::tls_connect(tcp_boxed, cfg.clone(), server_name);
280 tokio::select! {
281 result = handshake => {
282 result.map_err(|e| {
283 let msg = format!("reverse control TLS handshake failed: {e}");
284 if let Some(m) = self.metrics.as_ref() {
285 m.record_error(&msg);
286 }
287 ProtocolError::Tls(msg)
288 })?
289 }
290 _ = self.cancel.cancelled() => {
291 return Err(ProtocolError::ConnectionClosed);
292 }
293 }
294 } else {
295 Box::new(tcp)
296 };
297
298 let authenticating_start = Instant::now();
300 if let (Some(ref username), Some(ref password)) =
301 (&self.config.auth_username, &self.config.auth_password)
302 {
303 client_auth_handshake(&mut boxed, username, password).await?;
304 if let Some(ref m) = self.metrics {
305 m.record_state_duration(
306 ControlState::Authenticating,
307 authenticating_start.elapsed().as_millis() as u64,
308 );
309 }
310 info!(
311 state = ?ControlState::Authenticating,
312 "authentication successful"
313 );
314 } else {
315 crate::read_handshake(&mut boxed).await?;
317 if let Some(ref m) = self.metrics {
318 m.record_state_duration(
319 ControlState::Authenticating,
320 authenticating_start.elapsed().as_millis() as u64,
321 );
322 }
323 };
324
325 if let Some(ref m) = self.metrics {
326 m.record_stream_opened();
327 }
328 let ready_start = Instant::now();
329
330 let resolution =
332 self.resolver
333 .as_ref()
334 .map(|r| r.resolve())
335 .unwrap_or(TargetResolution::Reject {
336 reason: "no resolver configured".to_string(),
337 });
338
339 let session_result: Result<(), ProtocolError> = match resolution {
340 TargetResolution::Connect { host, port } => {
341 let target_addr = format!("{}:{}", host, port);
342 let connect_timeout = if self.config.target_connect_timeout_ms > 0 {
343 Duration::from_millis(self.config.target_connect_timeout_ms)
344 } else {
345 Duration::from_secs(30)
346 };
347 let connect_result =
348 tokio::time::timeout(connect_timeout, TcpStream::connect(&target_addr)).await;
349 match connect_result {
350 Ok(Ok(target_stream)) => {
351 info!(
352 target = %target_addr,
353 state = ?ControlState::Ready,
354 "connected to target, relaying"
355 );
356 let target_boxed: eggress_core::BoxStream = Box::new(target_stream);
357 crate::relay_bidirectional_boxed(
358 boxed,
359 target_boxed,
360 (self.config.read_timeout_ms > 0)
361 .then(|| Duration::from_millis(self.config.read_timeout_ms)),
362 )
363 .await
364 }
365 Ok(Err(e)) => {
366 warn!(
367 target = %target_addr,
368 error = %e,
369 "failed to connect to target"
370 );
371 if let Some(ref m) = self.metrics {
372 m.record_error(&format!("target connect failed: {e}"));
373 }
374 Err(ProtocolError::Io(e))
375 }
376 Err(_elapsed) => {
377 let e = std::io::Error::new(
378 std::io::ErrorKind::TimedOut,
379 format!("target connect timed out: {target_addr}"),
380 );
381 warn!(
382 target = %target_addr,
383 "target connect timed out"
384 );
385 if let Some(ref m) = self.metrics {
386 m.record_error(&format!("target connect failed: {e}"));
387 }
388 Err(ProtocolError::Io(e))
389 }
390 }
391 }
392 TargetResolution::Reject { reason } => {
393 warn!(reason = %reason, "route resolution rejected, dropping control channel");
394 if let Some(ref m) = self.metrics {
395 m.record_error(&format!("route rejected: {reason}"));
396 }
397 Err(ProtocolError::AuthFailed)
398 }
399 };
400
401 if let Some(ref m) = self.metrics {
402 m.record_state_duration(
403 ControlState::Ready,
404 ready_start.elapsed().as_millis() as u64,
405 );
406 m.record_stream_closed(0);
407 }
408
409 session_result
410 }
411
412 pub fn shutdown(&self) {
414 self.cancel.cancel();
415 }
416}
417
418#[cfg(test)]
419mod tests {
420 use super::*;
421
422 struct FixedResolver(TargetResolution);
423 impl TargetResolver for FixedResolver {
424 fn resolve(&self) -> TargetResolution {
425 self.0.clone()
426 }
427 }
428
429 #[test]
430 fn default_resolver_returns_configured_target() {
431 let r = DefaultTargetResolver::new(Some("127.0.0.1".to_string()), Some(8080));
432 assert_eq!(
433 r.resolve(),
434 TargetResolution::Connect {
435 host: "127.0.0.1".to_string(),
436 port: 8080,
437 }
438 );
439 }
440
441 #[test]
442 fn default_resolver_rejects_when_unset() {
443 let r = DefaultTargetResolver::new(None, None);
444 match r.resolve() {
445 TargetResolution::Reject { .. } => {}
446 other => panic!("expected Reject, got {other:?}"),
447 }
448 }
449
450 #[test]
451 fn default_resolver_rejects_partial() {
452 let r = DefaultTargetResolver::new(Some("127.0.0.1".to_string()), None);
453 match r.resolve() {
454 TargetResolution::Reject { .. } => {}
455 other => panic!("expected Reject, got {other:?}"),
456 }
457 }
458
459 #[test]
460 fn custom_resolver_can_reject() {
461 let r: Arc<dyn TargetResolver> = Arc::new(FixedResolver(TargetResolution::Reject {
462 reason: "policy".to_string(),
463 }));
464 match r.resolve() {
465 TargetResolution::Reject { reason } => assert_eq!(reason, "policy"),
466 other => panic!("expected Reject, got {other:?}"),
467 }
468 }
469}