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