1use 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
53pub 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
191pub 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}