Skip to main content

eggress_protocol_reverse/
client.rs

1use 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/// Configuration for a reverse proxy control client.
11#[derive(Clone)]
12pub struct ReverseClientConfig {
13    /// Address of the reverse server to connect to.
14    pub server_addr: SocketAddr,
15    /// Optional username for authentication.
16    pub auth_username: Option<String>,
17    /// Optional password for authentication.
18    pub auth_password: Option<String>,
19    /// Reconnect backoff initial delay in milliseconds.
20    pub reconnect_initial_ms: u64,
21    /// Reconnect backoff max delay in milliseconds.
22    pub reconnect_max_ms: u64,
23    /// Default target host (used when no target is specified in proxy request).
24    pub default_target_host: Option<String>,
25    /// Default target port.
26    pub default_target_port: Option<u16>,
27    /// Idle read timeout on the control channel in milliseconds. 0 = no
28    /// timeout.
29    pub read_timeout_ms: u64,
30    /// Grace period for drain on shutdown, in milliseconds.
31    pub drain_grace_ms: u64,
32    /// Timeout for target connect attempts in milliseconds. 0 = no timeout.
33    pub target_connect_timeout_ms: u64,
34    /// Optional TLS for the control channel. When present, the TCP control
35    /// stream is wrapped with Rustls before reverse framing/authentication.
36    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    /// Validate TLS material before any connection attempt.
62    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/// Result of resolving where to send a relayed stream.
89#[derive(Debug, Clone, PartialEq, Eq)]
90pub enum TargetResolution {
91    /// Connect to the given host:port.
92    Connect { host: String, port: u16 },
93    /// Reject the stream and close the control channel.
94    Reject { reason: String },
95}
96
97/// Trait for resolving the target of a relayed reverse stream.
98///
99/// The default implementation (used when no resolver is attached) returns
100/// the configured `default_target_host`/`default_target_port`, or rejects the
101/// stream if no default is configured. Production deployments inject a
102/// resolver that consults the route engine.
103pub trait TargetResolver: Send + Sync {
104    fn resolve(&self) -> TargetResolution;
105}
106
107/// Default resolver: uses the configured default target, or rejects.
108pub 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
133/// A reverse proxy control client.
134///
135/// Connects to a reverse server, authenticates, and services incoming proxy
136/// requests by connecting to local targets and relaying data.
137///
138/// In pproxy's backward model, each control connection carries exactly one
139/// proxy session. When the session ends, the client reconnects.
140pub 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    /// Attach metrics to this client instance.
162    pub fn set_metrics(&mut self, metrics: Arc<ReverseMetrics>) {
163        self.metrics = Some(metrics);
164    }
165
166    /// Replace the target resolver (defaults to `DefaultTargetResolver`).
167    pub fn set_resolver(&mut self, resolver: Arc<dyn TargetResolver>) {
168        self.resolver = Some(resolver);
169    }
170
171    /// Get a cancel token for external shutdown.
172    pub fn cancel_token(&self) -> CancellationToken {
173        self.cancel.clone()
174    }
175
176    /// Run the reverse client with automatic reconnection.
177    pub async fn run(&self) -> Result<(), ProtocolError> {
178        // Validate TLS material once before any dial so impossible combos
179        // fail fast instead of retrying forever.
180        self.config.validate()?;
181
182        // Build the immutable TLS client config once and reuse it across
183        // reconnect attempts rather than rebuilding roots per try.
184        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                    // Normal session end (external client disconnected)
212                    // Reset backoff and reconnect immediately
213                    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        // Drain phase: wait briefly for any pending cleanup
239        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    /// Run a single session with the server. TLS (when configured) wraps the
249    /// TCP stream before reverse framing so credentials are never sent in
250    /// plaintext. The shared client config is reused across reconnects.
251    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        // TLS handshake before any reverse bytes when configured.
277        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        // Authenticate
299        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            // No auth: just read the handshake response
316            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        // Resolve target via the route engine (or default resolver).
331        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    /// Shut down the reverse client.
413    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}