1use 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
51pub 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
148pub 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}