Skip to main content

eggress_server/
advanced.rs

1//! Dedicated listener handlers for multiplexed and upgraded transports.
2
3use eggress_core::{BoxStream, ClientIdentity, TargetAddr};
4use subtle::ConstantTimeEq;
5use tokio::io::{AsyncRead, AsyncWrite, ReadBuf};
6
7use crate::accept::{
8    auth_credentials, cached_identity, record_authenticated, AcceptedSession, PendingTunnel,
9    ReplyContext, TunnelProtocol,
10};
11use crate::auth::parse_basic_auth;
12use crate::ConnectionConfig;
13
14struct H2StreamAdapter {
15    reader: eggress_protocol_http::H2StreamRead,
16    writer: eggress_protocol_http::H2StreamWrite,
17}
18
19impl AsyncRead for H2StreamAdapter {
20    fn poll_read(
21        mut self: std::pin::Pin<&mut Self>,
22        cx: &mut std::task::Context<'_>,
23        buf: &mut ReadBuf<'_>,
24    ) -> std::task::Poll<std::io::Result<()>> {
25        std::pin::Pin::new(&mut self.reader).poll_read(cx, buf)
26    }
27}
28
29impl AsyncWrite for H2StreamAdapter {
30    fn poll_write(
31        mut self: std::pin::Pin<&mut Self>,
32        cx: &mut std::task::Context<'_>,
33        buf: &[u8],
34    ) -> std::task::Poll<std::io::Result<usize>> {
35        std::pin::Pin::new(&mut self.writer).poll_write(cx, buf)
36    }
37
38    fn poll_flush(
39        mut self: std::pin::Pin<&mut Self>,
40        cx: &mut std::task::Context<'_>,
41    ) -> std::task::Poll<std::io::Result<()>> {
42        std::pin::Pin::new(&mut self.writer).poll_flush(cx)
43    }
44
45    fn poll_shutdown(
46        mut self: std::pin::Pin<&mut Self>,
47        cx: &mut std::task::Context<'_>,
48    ) -> std::task::Poll<std::io::Result<()>> {
49        std::pin::Pin::new(&mut self.writer).poll_shutdown(cx)
50    }
51}
52
53/// Serve an HTTP/2 prior-knowledge or TLS/ALPN listener. Each CONNECT stream
54/// is routed independently while the parent connection continues accepting
55/// unrelated streams. Stream tasks are spawned onto `tasks` and cancelled via
56/// `cancel` so shutdown coordination can track and drain live tunnels.
57pub async fn serve_h2_connection(
58    client: BoxStream,
59    config: ConnectionConfig,
60    tasks: &tokio_util::task::TaskTracker,
61    cancel: tokio_util::sync::CancellationToken,
62) -> Result<(), String> {
63    let peer_ip = config.context.source.map(|peer| peer.ip());
64    let mut connection = h2::server::handshake(client)
65        .await
66        .map_err(|error| format!("H2 handshake failed: {error}"))?;
67
68    while let Some(result) = connection.accept().await {
69        let (request, mut response) = match result {
70            Ok(pair) => pair,
71            Err(error) => {
72                tracing::warn!(%error, "H2 accept error");
73                continue;
74            }
75        };
76        if request.method() != http::Method::CONNECT {
77            response.send_reset(h2::Reason::PROTOCOL_ERROR);
78            continue;
79        }
80
81        let target = match h2_target(request.uri()) {
82            Ok(target) => target,
83            Err(_) => {
84                let reply = http::Response::builder()
85                    .status(400)
86                    .body(())
87                    .expect("static response builds");
88                if let Err(error) = response.send_response(reply, true) {
89                    tracing::warn!(%error, "H2 send 400 response failed");
90                }
91                continue;
92            }
93        };
94
95        let cached = cached_identity(&config.authentication, peer_ip);
96        let authenticated = if cached.is_some() {
97            true
98        } else if let Some((username, password, _)) = auth_credentials(&config.authentication) {
99            matches!(
100                request
101                    .headers()
102                    .get(http::header::PROXY_AUTHORIZATION)
103                    .and_then(|value| value.to_str().ok())
104                    .and_then(parse_basic_auth),
105                Some((user, pass))
106                    if (user.as_bytes().ct_eq(username.as_bytes())
107                        & pass.as_bytes().ct_eq(password.as_bytes()))
108                        .unwrap_u8()
109                        == 1
110            )
111        } else {
112            true
113        };
114
115        if !authenticated {
116            let reply = http::Response::builder()
117                .status(407)
118                .header(http::header::PROXY_AUTHENTICATE, "Basic realm=\"eggress\"")
119                .body(())
120                .expect("static response builds");
121            if let Err(error) = response.send_response(reply, true) {
122                tracing::warn!(%error, "H2 send 407 response failed");
123            }
124            continue;
125        }
126
127        let identity = cached.unwrap_or_else(|| {
128            let identity = request
129                .headers()
130                .get(http::header::PROXY_AUTHORIZATION)
131                .and_then(|value| value.to_str().ok())
132                .and_then(parse_basic_auth)
133                .map(|(user, _)| ClientIdentity::Username(user))
134                .unwrap_or(ClientIdentity::Anonymous);
135            record_authenticated(&config.authentication, peer_ip, &identity);
136            identity
137        });
138
139        let send_stream = match response.send_response(
140            http::Response::builder()
141                .status(200)
142                .body(())
143                .expect("static response builds"),
144            false,
145        ) {
146            Ok(stream) => stream,
147            Err(error) => {
148                tracing::warn!(%error, "H2 send 200 response failed");
149                continue;
150            }
151        };
152        let client_stream: BoxStream = Box::new(H2StreamAdapter {
153            reader: eggress_protocol_http::H2StreamRead::new(request.into_body()),
154            writer: eggress_protocol_http::H2StreamWrite::new(send_stream),
155        });
156        let stream_config = config.clone();
157        let stream_cancel = cancel.child_token();
158        let pending = PendingTunnel {
159            target,
160            client: client_stream,
161            protocol: TunnelProtocol::Http2,
162            reply_context: ReplyContext::Http2,
163            identity,
164        };
165        let target_for_cancel = pending.target.to_string();
166        tasks.spawn(async move {
167            if let Some(metrics) = &stream_config.metrics {
168                metrics.record_session_start();
169            }
170            let report = tokio::select! {
171                _ = stream_cancel.cancelled() => {
172                    crate::execute::SessionReport::cancelled(
173                        Some("h2".to_string()),
174                        Some(target_for_cancel),
175                        "h2-cancelled".to_string(),
176                    )
177                }
178                result = crate::execute::execute(AcceptedSession::Tunnel(pending), &stream_config) => {
179                    result
180                }
181            };
182            if let Some(metrics) = &stream_config.metrics {
183                metrics.record_session(&report);
184            }
185        });
186    }
187
188    Ok(())
189}
190
191/// Serve a WebSocket listener. WebSocket listeners use a fixed target because
192/// the upgrade request itself carries no proxy CONNECT authority.
193pub async fn serve_websocket_connection(
194    client: BoxStream,
195    config: ConnectionConfig,
196    fixed_target: TargetAddr,
197) -> Result<(), String> {
198    let peer_ip = config.context.source.map(|peer| peer.ip());
199    let cached = cached_identity(&config.authentication, peer_ip);
200    let credentials = if cached.is_some() {
201        None
202    } else {
203        auth_credentials(&config.authentication).map(|(user, pass, _)| (user, pass))
204    };
205    let (client, authenticated_user) =
206        eggress_protocol_websocket::accept_upgrade_with_auth(client, credentials)
207            .await
208            .map_err(|error| error.to_string())?;
209    let identity = cached.unwrap_or_else(|| {
210        let identity = authenticated_user
211            .map(ClientIdentity::Username)
212            .unwrap_or(ClientIdentity::Anonymous);
213        record_authenticated(&config.authentication, peer_ip, &identity);
214        identity
215    });
216    let pending = PendingTunnel {
217        target: fixed_target.clone(),
218        client,
219        protocol: TunnelProtocol::WebSocket,
220        reply_context: ReplyContext::WebSocket,
221        identity,
222    };
223    if let Some(metrics) = &config.metrics {
224        metrics.record_session_start();
225    }
226    let report = crate::execute::execute(AcceptedSession::Tunnel(pending), &config).await;
227    if let Some(metrics) = &config.metrics {
228        metrics.record_session(&report);
229    }
230    Ok(())
231}
232
233fn h2_target(uri: &http::Uri) -> Result<TargetAddr, String> {
234    let authority = uri
235        .authority()
236        .map(|authority| authority.as_str().to_string())
237        .or_else(|| (!uri.path().is_empty()).then(|| uri.path().to_string()))
238        .ok_or_else(|| "missing H2 CONNECT authority".to_string())?;
239    if authority.contains(':') {
240        authority.parse()
241    } else {
242        format!("{authority}:443").parse()
243    }
244}
245
246#[cfg(test)]
247mod tests {
248    use super::*;
249    use bytes::Bytes;
250    use std::sync::Arc;
251    use std::time::Duration;
252
253    use eggress_routing::{RouteActionSpec, RouteService, Router};
254    use futures_util::{SinkExt, StreamExt};
255
256    fn config(peer: std::net::SocketAddr, protocol: eggress_core::ProtocolId) -> ConnectionConfig {
257        ConnectionConfig {
258            routing: Arc::new(Router::new(vec![], RouteActionSpec::Direct))
259                as Arc<dyn RouteService>,
260            context: crate::ConnectionContext {
261                source: Some(peer),
262                listener: "advanced-test".to_string(),
263                generation: 0,
264            },
265            handshake_timeout: Duration::from_secs(5),
266            connect_timeout: Duration::from_secs(5),
267            protocols: Arc::from([protocol]),
268            authentication: crate::accept::InboundAuthentication::None,
269            metrics: None,
270            udp: None,
271            tls_client_config: None,
272            shadowsocks: None,
273            shadowsocks_metrics: Some(Arc::new(
274                eggress_protocol_shadowsocks::ShadowsocksMetrics::new(),
275            )),
276            trojan: None,
277            fixed_target: None,
278            local_bind: None,
279        }
280    }
281
282    #[tokio::test]
283    async fn h2_listener_routes_connect_stream_to_local_target() {
284        let (echo_addr, echo_task) = eggress_testkit::start_echo_server().await;
285        let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
286        let listener_addr = listener.local_addr().unwrap();
287        let tasks = tokio_util::task::TaskTracker::new();
288        let cancel = tokio_util::sync::CancellationToken::new();
289        let server_tasks = tasks.clone();
290        let server = tokio::spawn(async move {
291            let (stream, peer) = listener.accept().await.unwrap();
292            let result = serve_h2_connection(
293                Box::new(stream),
294                config(peer, eggress_core::ProtocolId::Http2),
295                &server_tasks,
296                cancel,
297            )
298            .await;
299            tasks.close();
300            result
301        });
302
303        let stream = tokio::net::TcpStream::connect(listener_addr).await.unwrap();
304        let (mut sender, connection) = h2::client::handshake(stream).await.unwrap();
305        let driver = tokio::spawn(connection);
306        let request = http::Request::builder()
307            .method(http::Method::CONNECT)
308            .uri(echo_addr.to_string())
309            .body(())
310            .unwrap();
311        let (response, mut send) = sender.send_request(request, false).unwrap();
312        let response = match tokio::time::timeout(Duration::from_secs(5), response).await {
313            Ok(response) => response.unwrap(),
314            Err(_) => {
315                let server_done = server.is_finished();
316                server.abort();
317                panic!(
318                    "H2 response timed out (client driver done: {}, server done: {server_done})",
319                    driver.is_finished(),
320                );
321            }
322        };
323        assert_eq!(response.status(), http::StatusCode::OK);
324        send.send_data(Bytes::from_static(b"h2 listener"), true)
325            .unwrap();
326
327        let mut body = response.into_body();
328        let mut received = Vec::new();
329        while let Some(chunk) = body.data().await {
330            received.extend_from_slice(&chunk.unwrap());
331        }
332        assert_eq!(received, b"h2 listener");
333        drop(sender);
334        driver.abort();
335        server.abort();
336        echo_task.abort();
337    }
338
339    #[tokio::test]
340    async fn websocket_listener_routes_binary_to_local_target() {
341        let (echo_addr, echo_task) = eggress_testkit::start_echo_server().await;
342        let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
343        let listener_addr = listener.local_addr().unwrap();
344        let target = echo_addr.to_string().parse().unwrap();
345        let server = tokio::spawn(async move {
346            let (stream, peer) = listener.accept().await.unwrap();
347            serve_websocket_connection(
348                Box::new(stream),
349                config(peer, eggress_core::ProtocolId::WebSocket),
350                target,
351            )
352            .await
353        });
354
355        let (mut client, _) = tokio_tungstenite::connect_async(format!("ws://{listener_addr}"))
356            .await
357            .unwrap();
358        client
359            .send(tokio_tungstenite::tungstenite::Message::Binary(
360                b"websocket listener".to_vec().into(),
361            ))
362            .await
363            .unwrap();
364        let response = tokio::time::timeout(Duration::from_secs(5), client.next())
365            .await
366            .unwrap()
367            .unwrap()
368            .unwrap();
369        assert_eq!(&response.into_data()[..], b"websocket listener");
370        let _ = client.close(None).await;
371        server.abort();
372        echo_task.abort();
373    }
374}