Skip to main content

eggress_protocol_reverse/
client.rs

1use 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/// Configuration for a reverse proxy control client.
11#[derive(Debug, 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}
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/// Result of resolving where to send a relayed stream.
54#[derive(Debug, Clone, PartialEq, Eq)]
55pub enum TargetResolution {
56    /// Connect to the given host:port.
57    Connect { host: String, port: u16 },
58    /// Reject the stream and close the control channel.
59    Reject { reason: String },
60}
61
62/// Trait for resolving the target of a relayed reverse stream.
63///
64/// The default implementation (used when no resolver is attached) returns
65/// the configured `default_target_host`/`default_target_port`, or rejects the
66/// stream if no default is configured. Production deployments inject a
67/// resolver that consults the route engine.
68pub trait TargetResolver: Send + Sync {
69    fn resolve(&self) -> TargetResolution;
70}
71
72/// Default resolver: uses the configured default target, or rejects.
73pub 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
98/// A reverse proxy control client.
99///
100/// Connects to a reverse server, authenticates, and services incoming proxy
101/// requests by connecting to local targets and relaying data.
102///
103/// In pproxy's backward model, each control connection carries exactly one
104/// proxy session. When the session ends, the client reconnects.
105pub 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    /// Attach metrics to this client instance.
127    pub fn set_metrics(&mut self, metrics: Arc<ReverseMetrics>) {
128        self.metrics = Some(metrics);
129    }
130
131    /// Replace the target resolver (defaults to `DefaultTargetResolver`).
132    pub fn set_resolver(&mut self, resolver: Arc<dyn TargetResolver>) {
133        self.resolver = Some(resolver);
134    }
135
136    /// Get a cancel token for external shutdown.
137    pub fn cancel_token(&self) -> CancellationToken {
138        self.cancel.clone()
139    }
140
141    /// Run the reverse client with automatic reconnection.
142    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                    // Normal session end (external client disconnected)
162                    // Reset backoff and reconnect immediately
163                    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        // Drain phase: wait briefly for any pending cleanup
189        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    /// Run a single session with the server.
199    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        // Authenticate
215        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            // No auth: just read the handshake response
234            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        // Resolve target via the route engine (or default resolver).
251        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    /// Shut down the reverse client.
332    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}