1use std::borrow::Cow;
9use std::io;
10use std::net::{IpAddr, SocketAddr};
11use std::sync::Arc;
12use std::time::Duration;
13
14use bytes::Bytes;
15use tokio::io::{AsyncReadExt, AsyncWriteExt};
16use tokio::net::TcpStream;
17use tokio::sync::mpsc;
18
19use super::connection::ProxyConnectState;
20#[cfg(test)]
21use super::connection::ProxyConnectStatus;
22use super::upstream::UpstreamTcpTarget;
23use crate::engine::http_deny::{classify_http_request, http_forbidden_response};
24use crate::engine::secrets::config::SecretsConfigExt;
25use crate::engine::tls::proxy::TlsProxy;
26use crate::engine::tls::sni;
27use crate::engine::tls::state::TlsState;
28use crate::netstack::shared::SharedState;
29use crate::policy::{EgressEvaluation, HostnameSource, NetworkPolicy, Protocol};
30use crate::proxy::ResolvedOutboundProxy;
31use crate::secrets::config::{SecretViolationAction, SecretsConfig};
32use crate::secrets::handler::{
33 SecretsHandler, first_line_is_not_http_request, looks_like_http_request_prefix,
34};
35
36const SERVER_READ_BUF_SIZE: usize = 16384;
42
43const CONNECT_RESP_LIMIT: usize = 8192;
45
46pub(crate) const PEEK_BUF_SIZE: usize = 16384;
48
49pub(crate) const PEEK_BUDGET: Duration = Duration::from_secs(5);
52
53#[derive(Debug)]
58struct ConnectRequest {
59 bytes: Vec<u8>,
60 header_end: usize,
61 target: ConnectTarget,
62}
63
64#[derive(Debug, Clone, PartialEq, Eq)]
65struct ConnectTarget {
66 host: String,
67 port: u16,
68 expected_sni: Option<String>,
69}
70
71pub(crate) struct TcpProxy {
73 guest_dst: SocketAddr,
74 connect_target: UpstreamTcpTarget,
75 from_smoltcp: mpsc::Receiver<Bytes>,
76 to_smoltcp: mpsc::Sender<Bytes>,
77 shared: Arc<SharedState>,
78 network_policy: Arc<NetworkPolicy>,
79 secrets: Arc<SecretsConfig>,
80 tls_state: Option<Arc<TlsState>>,
81 strict: bool,
82 proxy_connect: Arc<ProxyConnectState>,
83 outbound_proxy: Option<Arc<ResolvedOutboundProxy>>,
84}
85
86impl ConnectRequest {
91 fn header_bytes(&self) -> &[u8] {
92 &self.bytes[..self.header_end]
93 }
94
95 fn post_header_bytes(&self) -> &[u8] {
96 &self.bytes[self.header_end..]
97 }
98}
99
100impl ConnectTarget {
101 fn is_intercepted(&self, tls_state: &TlsState) -> bool {
102 tls_state.config.intercepted_ports.contains(&self.port)
103 }
104
105 fn guest_dst(&self, fallback: SocketAddr, shared: &SharedState) -> SocketAddr {
106 if let Ok(ip) = self.host.parse::<IpAddr>() {
107 return SocketAddr::new(ip, self.port);
108 }
109
110 if self.host.eq_ignore_ascii_case(crate::HOST_ALIAS) {
111 match fallback.ip() {
112 IpAddr::V4(_) => {
113 if let Some(ip) = shared.gateway_ipv4() {
114 return SocketAddr::new(IpAddr::V4(ip), self.port);
115 }
116 }
117 IpAddr::V6(_) => {
118 if let Some(ip) = shared.gateway_ipv6() {
119 return SocketAddr::new(IpAddr::V6(ip), self.port);
120 }
121 }
122 }
123 if let Some(ip) = shared.gateway_ipv4() {
124 return SocketAddr::new(IpAddr::V4(ip), self.port);
125 }
126 if let Some(ip) = shared.gateway_ipv6() {
127 return SocketAddr::new(IpAddr::V6(ip), self.port);
128 }
129 }
130
131 SocketAddr::new(fallback.ip(), self.port)
132 }
133}
134
135impl TcpProxy {
136 #[allow(clippy::too_many_arguments)]
138 pub(crate) fn new(
139 guest_dst: SocketAddr,
140 connect_target: UpstreamTcpTarget,
141 from_smoltcp: mpsc::Receiver<Bytes>,
142 to_smoltcp: mpsc::Sender<Bytes>,
143 shared: Arc<SharedState>,
144 network_policy: Arc<NetworkPolicy>,
145 secrets: Arc<SecretsConfig>,
146 tls_state: Option<Arc<TlsState>>,
147 strict: bool,
148 proxy_connect: Arc<ProxyConnectState>,
149 outbound_proxy: Option<Arc<ResolvedOutboundProxy>>,
150 ) -> Self {
151 Self {
152 guest_dst,
153 connect_target,
154 from_smoltcp,
155 to_smoltcp,
156 shared,
157 network_policy,
158 secrets,
159 tls_state,
160 strict,
161 proxy_connect,
162 outbound_proxy,
163 }
164 }
165
166 pub(crate) async fn run(self) {
168 let guest_dst = self.guest_dst;
169 let connect_dst = self.connect_target.primary();
170
171 if let Err(error) = self.try_run().await {
172 tracing::debug!(
173 dst = %connect_dst,
174 %guest_dst,
175 %error,
176 "TCP proxy task ended",
177 );
178 }
179 }
180
181 async fn try_run(self) -> io::Result<()> {
183 let Self {
184 guest_dst,
185 connect_target,
186 mut from_smoltcp,
187 to_smoltcp,
188 shared,
189 network_policy,
190 secrets,
191 tls_state,
192 strict,
193 proxy_connect,
194 outbound_proxy,
195 } = self;
196
197 let hostname_policy_deferred = match network_policy.evaluate_egress_with_source(
201 guest_dst,
202 Protocol::Tcp,
203 &shared,
204 HostnameSource::Deferred,
205 ) {
206 EgressEvaluation::Allow => false,
207 EgressEvaluation::DeferUntilHostname => true,
208 EgressEvaluation::Deny => network_policy.has_domain_rules(),
211 };
212
213 let peek_started = tokio::time::Instant::now();
219 let (mut initial_buf, sni) = if hostname_policy_deferred {
220 peek_for_sni(&mut from_smoltcp, PEEK_BUF_SIZE, PEEK_BUDGET).await
221 } else {
222 (Vec::new(), None)
223 };
224
225 if hostname_policy_deferred {
231 let source = match sni.as_deref() {
232 Some(name) => HostnameSource::Sni(name),
233 None => HostnameSource::CacheOnly,
234 };
235 match network_policy.evaluate_egress_with_source(
236 guest_dst,
237 Protocol::Tcp,
238 &shared,
239 source,
240 ) {
241 EgressEvaluation::Allow => {
242 if strict_hostname_allow_is_opaque(
243 strict,
244 &network_policy,
245 guest_dst,
246 &shared,
247 sni.as_deref(),
248 &initial_buf,
249 ) {
250 tracing::debug!(
251 sni = sni.as_deref(),
252 dst = %guest_dst,
253 "TCP egress denied by strict hostname policy",
254 );
255 proxy_connect.mark_policy_denied();
256 shared.proxy_wake.wake();
257 return Ok(());
258 }
259 }
260 EgressEvaluation::Deny => {
261 tracing::debug!(
262 dst = %guest_dst,
263 source = source.label(),
264 "TCP egress denied by domain policy",
265 );
266 if shared.http_deny_response_enabled() {
267 initial_buf = peek_for_http_request(
268 &mut from_smoltcp,
269 initial_buf,
270 PEEK_BUF_SIZE,
271 PEEK_BUDGET.saturating_sub(peek_started.elapsed()),
272 )
273 .await;
274 }
275 return deny_http_or_close(
276 guest_dst,
277 sni.as_deref(),
278 &initial_buf,
279 to_smoltcp,
280 &shared,
281 &proxy_connect,
282 )
283 .await;
284 }
285 EgressEvaluation::DeferUntilHostname => {
286 debug_assert!(false, "DeferUntilHostname leaked into TCP proxy task");
287 if shared.http_deny_response_enabled() {
288 initial_buf = peek_for_http_request(
289 &mut from_smoltcp,
290 initial_buf,
291 PEEK_BUF_SIZE,
292 PEEK_BUDGET.saturating_sub(peek_started.elapsed()),
293 )
294 .await;
295 }
296 return deny_http_or_close(
297 guest_dst,
298 sni.as_deref(),
299 &initial_buf,
300 to_smoltcp,
301 &shared,
302 &proxy_connect,
303 )
304 .await;
305 }
306 }
307 }
308
309 if let Some(tls_state) = tls_state.clone()
313 && !initial_buf.is_empty()
314 && could_be_connect_request(&initial_buf)
315 {
316 return handle_connect_tunnel(
317 guest_dst,
318 connect_target,
319 initial_buf,
320 from_smoltcp,
321 to_smoltcp,
322 shared,
323 network_policy,
324 tls_state,
325 strict,
326 proxy_connect,
327 outbound_proxy,
328 None,
329 )
330 .await;
331 }
332
333 let stream = connect_target
338 .connect(&proxy_connect, &shared, outbound_proxy.as_deref())
339 .await?;
340 let connect_dst = stream.peer_addr().unwrap_or(connect_target.primary());
341 let (mut server_rx, mut server_tx) = stream.into_split();
342
343 let enforce_http_authority = network_policy.has_domain_rules();
349 let want_headers = enforce_http_authority
350 || secrets.has_plain_http_candidates()
351 || secrets.has_host_scoped_secrets();
352 let (initial_buf, is_tls) = if want_headers {
353 classify_first_flight(
354 initial_buf,
355 &mut from_smoltcp,
356 &mut server_rx,
357 &to_smoltcp,
358 &shared,
359 want_headers,
360 PEEK_BUF_SIZE,
361 PEEK_BUDGET,
362 )
363 .await?
364 } else {
365 (initial_buf, false)
366 };
367
368 if let Some(tls_state) = tls_state.clone()
369 && could_be_connect_request(&initial_buf)
370 {
371 let proxy_stream = server_rx
376 .reunite(server_tx)
377 .map_err(|_| io::Error::other("failed to reunite proxy stream halves"))?;
378 return handle_connect_tunnel(
379 guest_dst,
380 connect_target,
381 initial_buf,
382 from_smoltcp,
383 to_smoltcp,
384 shared,
385 network_policy,
386 tls_state,
387 strict,
388 proxy_connect,
389 outbound_proxy,
390 Some(proxy_stream),
391 )
392 .await;
393 }
394
395 let mut late_connect_state = tls_state;
396 let mut secrets_handler: Option<SecretsHandler> = if is_tls {
397 None
398 } else if enforce_http_authority {
399 let host = extract_http_host(&initial_buf).unwrap_or_default();
400 Some(SecretsHandler::new_plain_http_policy(
401 &secrets,
402 &host,
403 guest_dst,
404 network_policy.clone(),
405 shared.clone(),
406 ))
407 } else if !secrets.secrets.is_empty() {
408 Some(match extract_http_host(&initial_buf) {
409 Some(host) => {
410 SecretsHandler::new_plain_http(&secrets, &host, guest_dst.ip(), &shared)
411 }
412 None => SecretsHandler::new_plain_http_invalid_host(&secrets),
413 })
414 } else {
415 None
416 };
417
418 if !initial_buf.is_empty() {
420 let out: Cow<[u8]> = match secrets_handler.as_mut() {
421 Some(h) => match h.substitute(&initial_buf) {
422 Ok(cow) => cow,
425 Err(action) => {
426 if matches!(action, SecretViolationAction::BlockAndTerminate) {
427 shared.trigger_termination();
428 }
429 return Ok(());
430 }
431 },
432 None => Cow::Borrowed(&initial_buf),
433 };
434 if !out.is_empty() {
435 if let Err(e) = server_tx.write_all(&out).await {
436 tracing::debug!(dst = %connect_dst, error = %e, "replay of buffered first flight failed");
437 return Ok(());
438 }
439 if let Err(e) = server_tx.flush().await {
440 tracing::debug!(dst = %connect_dst, error = %e, "flush after first flight failed");
441 return Ok(());
442 }
443 }
444 }
445
446 let mut server_buf = vec![0u8; SERVER_READ_BUF_SIZE];
447
448 let mut guest_eof = false;
453 loop {
454 tokio::select! {
455 data = from_smoltcp.recv(), if !guest_eof => {
457 match data {
458 Some(bytes) => {
459 if let Some(tls_state) = late_connect_state.take()
460 && could_be_connect_request(&bytes)
461 {
462 let proxy_stream = server_rx
467 .reunite(server_tx)
468 .map_err(|_| io::Error::other("failed to reunite proxy stream halves"))?;
469 return handle_connect_tunnel(
470 guest_dst,
471 connect_target,
472 bytes.to_vec(),
473 from_smoltcp,
474 to_smoltcp,
475 shared,
476 network_policy,
477 tls_state,
478 strict,
479 proxy_connect,
480 outbound_proxy,
481 Some(proxy_stream),
482 )
483 .await;
484 }
485 let out: Cow<[u8]> = match secrets_handler.as_mut() {
488 Some(h) => match h.substitute(&bytes) {
489 Ok(cow) => cow,
490 Err(action) => {
491 if matches!(action, SecretViolationAction::BlockAndTerminate)
492 {
493 shared.trigger_termination();
494 }
495 break;
496 }
497 },
498 None => Cow::Borrowed(&bytes),
499 };
500 if !out.is_empty() {
501 if let Err(e) = server_tx.write_all(&out).await {
502 tracing::debug!(dst = %connect_dst, error = %e, "write to server failed");
503 break;
504 }
505 if let Err(e) = server_tx.flush().await {
506 tracing::debug!(dst = %connect_dst, error = %e, "flush to server failed");
507 break;
508 }
509 }
510 }
511 None => {
516 guest_eof = true;
517 if server_tx.shutdown().await.is_err() {
518 break;
519 }
520 }
521 }
522 }
523
524 result = server_rx.read(&mut server_buf) => {
526 match result {
527 Ok(0) => break, Ok(n) => {
529 late_connect_state = None;
532 let data = Bytes::copy_from_slice(&server_buf[..n]);
533 if to_smoltcp.send(data).await.is_err() {
534 break;
536 }
537 shared.proxy_wake.wake();
540 }
541 Err(e) => {
542 tracing::debug!(dst = %connect_dst, error = %e, "read from server failed");
543 break;
544 }
545 }
546 }
547 }
548 }
549
550 Ok(())
551 }
552}
553
554#[allow(clippy::too_many_arguments)]
567pub fn spawn_tcp_proxy(
568 handle: &tokio::runtime::Handle,
569 guest_dst: SocketAddr,
570 connect_dst: SocketAddr,
571 from_smoltcp: mpsc::Receiver<Bytes>,
572 to_smoltcp: mpsc::Sender<Bytes>,
573 shared: Arc<SharedState>,
574 network_policy: Arc<NetworkPolicy>,
575 secrets: Arc<SecretsConfig>,
576 tls_state: Option<Arc<TlsState>>,
577 strict: bool,
578 proxy_connect: Arc<ProxyConnectState>,
579 outbound_proxy: Option<Arc<ResolvedOutboundProxy>>,
580) {
581 let proxy = TcpProxy::new(
582 guest_dst,
583 UpstreamTcpTarget::direct(connect_dst),
584 from_smoltcp,
585 to_smoltcp,
586 shared,
587 network_policy,
588 secrets,
589 tls_state,
590 strict,
591 proxy_connect,
592 outbound_proxy,
593 );
594
595 handle.spawn(proxy.run());
596}
597
598fn strict_hostname_allow_is_opaque(
599 strict: bool,
600 network_policy: &NetworkPolicy,
601 guest_dst: SocketAddr,
602 shared: &SharedState,
603 sni: Option<&str>,
604 initial_buf: &[u8],
605) -> bool {
606 if !strict {
607 return false;
608 }
609
610 let source = if let Some(name) = sni {
611 HostnameSource::Sni(name)
612 } else if initial_buf.is_empty() || initial_buf.first() == Some(&0x16) {
613 HostnameSource::CacheOnly
614 } else {
615 return false;
616 };
617
618 network_policy.allows_egress_via_hostname(guest_dst, Protocol::Tcp, shared, source)
619}
620
621#[allow(clippy::too_many_arguments)]
627async fn handle_connect_tunnel(
628 guest_dst: SocketAddr,
629 proxy_target: UpstreamTcpTarget,
630 initial_buf: Vec<u8>,
631 mut from_smoltcp: mpsc::Receiver<Bytes>,
632 to_smoltcp: mpsc::Sender<Bytes>,
633 shared: Arc<SharedState>,
634 network_policy: Arc<NetworkPolicy>,
635 tls_state: Arc<TlsState>,
636 strict: bool,
637 proxy_connect: Arc<ProxyConnectState>,
638 outbound_proxy: Option<Arc<ResolvedOutboundProxy>>,
639 preconnected_proxy: Option<TcpStream>,
640) -> io::Result<()> {
641 let connect_req =
642 parse_connect_request(buffer_connect_request(initial_buf, &mut from_smoltcp).await?)?;
643
644 let connect_headers =
645 match sanitize_connect_headers(connect_req.header_bytes(), &tls_state.secrets.load()) {
646 Ok(headers) => headers,
647 Err(action) => {
648 if matches!(action, SecretViolationAction::BlockAndTerminate) {
649 shared.trigger_termination();
650 }
651 return Ok(());
652 }
653 };
654
655 let mut proxy_stream = match preconnected_proxy {
657 Some(stream) => stream,
658 None => {
659 proxy_target
660 .connect(&proxy_connect, &shared, outbound_proxy.as_deref())
661 .await?
662 }
663 };
664
665 if !connect_req.target.is_intercepted(&tls_state) {
666 let tunnel_dst = connect_req.target.guest_dst(guest_dst, &shared);
667 if strict
668 && let Some(expected_sni) = connect_req.target.expected_sni.as_deref()
669 && network_policy.allows_egress_via_hostname(
670 tunnel_dst,
671 Protocol::Tcp,
672 &shared,
673 HostnameSource::Sni(expected_sni),
674 )
675 {
676 tracing::debug!(
677 sni = %expected_sni,
678 dst = %tunnel_dst,
679 "CONNECT tunnel denied by strict hostname policy",
680 );
681 proxy_connect.mark_policy_denied();
682 shared.proxy_wake.wake();
683 return Ok(());
684 }
685 proxy_stream.write_all(&connect_headers).await?;
686 proxy_stream.flush().await?;
687 let (proxy_resp, header_end) = read_connect_response_headers(&mut proxy_stream).await?;
688 if to_smoltcp
689 .send(Bytes::copy_from_slice(&proxy_resp[..header_end]))
690 .await
691 .is_err()
692 {
693 return Ok(());
694 }
695 if !proxy_resp[header_end..].is_empty()
696 && to_smoltcp
697 .send(Bytes::copy_from_slice(&proxy_resp[header_end..]))
698 .await
699 .is_err()
700 {
701 return Ok(());
702 }
703 shared.proxy_wake.wake();
704 if !connect_response_is_success(&proxy_resp[..header_end]) {
705 proxy_connect.mark_connected();
706 return Ok(());
707 }
708 if !connect_req.post_header_bytes().is_empty() {
709 proxy_stream
710 .write_all(connect_req.post_header_bytes())
711 .await?;
712 }
713 proxy_stream.flush().await?;
714 proxy_connect.mark_connected();
715 return relay_connected_stream(proxy_stream, from_smoltcp, to_smoltcp, shared).await;
716 }
717
718 proxy_stream.write_all(&connect_headers).await?;
719 proxy_stream.flush().await?;
720
721 let (proxy_resp, header_end) = read_connect_response_headers(&mut proxy_stream).await?;
722 if !connect_response_is_success(&proxy_resp[..header_end]) {
723 return Err(io::Error::new(
724 io::ErrorKind::ConnectionRefused,
725 "proxy rejected CONNECT",
726 ));
727 }
728 if !proxy_resp[header_end..].is_empty() {
729 return Err(io::Error::new(
730 io::ErrorKind::InvalidData,
731 "proxy sent unexpected bytes after CONNECT response headers",
732 ));
733 }
734 proxy_connect.mark_connected();
735
736 if to_smoltcp
737 .send(Bytes::copy_from_slice(&proxy_resp[..header_end]))
738 .await
739 .is_err()
740 {
741 return Ok(());
742 }
743 shared.proxy_wake.wake();
744
745 let tls_seed = connect_req.post_header_bytes().to_vec();
746 let tls_guest_dst = connect_req.target.guest_dst(guest_dst, &shared);
747 let expected_sni = connect_req.target.expected_sni.clone();
748
749 TlsProxy::new(
750 tls_guest_dst,
751 proxy_target,
752 from_smoltcp,
753 to_smoltcp,
754 shared,
755 tls_state,
756 network_policy,
757 strict,
758 proxy_connect,
759 None,
763 )
764 .with_upstream(proxy_stream)
765 .with_expected_sni(expected_sni)
766 .with_initial_buf(tls_seed)
767 .try_run()
768 .await
769}
770
771async fn relay_connected_stream(
773 stream: TcpStream,
774 mut from_smoltcp: mpsc::Receiver<Bytes>,
775 to_smoltcp: mpsc::Sender<Bytes>,
776 shared: Arc<SharedState>,
777) -> io::Result<()> {
778 let (mut server_rx, mut server_tx) = stream.into_split();
779 let mut server_buf = vec![0u8; SERVER_READ_BUF_SIZE];
780
781 let mut guest_eof = false;
782 loop {
783 tokio::select! {
784 data = from_smoltcp.recv(), if !guest_eof => {
785 match data {
786 Some(bytes) => {
787 server_tx.write_all(&bytes).await?;
788 server_tx.flush().await?;
789 }
790 None => {
793 guest_eof = true;
794 if server_tx.shutdown().await.is_err() {
795 break;
796 }
797 }
798 }
799 }
800 result = server_rx.read(&mut server_buf) => {
801 match result {
802 Ok(0) => break,
803 Ok(n) => {
804 if to_smoltcp
805 .send(Bytes::copy_from_slice(&server_buf[..n]))
806 .await
807 .is_err()
808 {
809 break;
810 }
811 shared.proxy_wake.wake();
812 }
813 Err(e) => return Err(e),
814 }
815 }
816 }
817 }
818
819 Ok(())
820}
821
822async fn buffer_connect_request(
823 mut buf: Vec<u8>,
824 from_smoltcp: &mut mpsc::Receiver<Bytes>,
825) -> io::Result<Vec<u8>> {
826 let timeout_fut = tokio::time::sleep(PEEK_BUDGET);
827 tokio::pin!(timeout_fut);
828
829 loop {
830 if !could_be_connect_request(&buf) {
831 return Err(io::Error::new(
832 io::ErrorKind::InvalidData,
833 "malformed CONNECT request prefix",
834 ));
835 }
836 if headers_end(&buf).is_some() {
837 return Ok(buf);
838 }
839 if buf.len() >= PEEK_BUF_SIZE {
840 return Err(io::Error::new(
841 io::ErrorKind::InvalidData,
842 "CONNECT request headers too large",
843 ));
844 }
845
846 tokio::select! {
847 biased;
848 _ = &mut timeout_fut => {
849 return Err(io::Error::new(
850 io::ErrorKind::TimedOut,
851 "timed out waiting for complete CONNECT request headers",
852 ));
853 }
854 data = from_smoltcp.recv() => match data {
855 Some(bytes) => {
856 buf.extend_from_slice(&bytes);
857 }
858 None => {
859 return Err(io::Error::new(
860 io::ErrorKind::UnexpectedEof,
861 "channel closed before complete CONNECT request headers",
862 ));
863 }
864 }
865 }
866 }
867}
868
869async fn read_connect_response_headers(stream: &mut TcpStream) -> io::Result<(Vec<u8>, usize)> {
870 tokio::time::timeout(PEEK_BUDGET, async {
871 let mut proxy_resp = Vec::with_capacity(256);
872 let mut buf = [0u8; 4096];
873 loop {
874 let n = stream.read(&mut buf).await?;
875 if n == 0 {
876 return Err(io::Error::new(
877 io::ErrorKind::UnexpectedEof,
878 "proxy closed before sending CONNECT response",
879 ));
880 }
881 proxy_resp.extend_from_slice(&buf[..n]);
882 if let Some(end) = headers_end(&proxy_resp) {
883 return Ok((proxy_resp, end));
884 }
885 if proxy_resp.len() > CONNECT_RESP_LIMIT {
886 return Err(io::Error::new(
887 io::ErrorKind::InvalidData,
888 "proxy CONNECT response too large",
889 ));
890 }
891 }
892 })
893 .await
894 .map_err(|_| {
895 io::Error::new(
896 io::ErrorKind::TimedOut,
897 "timed out waiting for proxy CONNECT response",
898 )
899 })?
900}
901
902fn sanitize_connect_headers<'a>(
903 header_bytes: &'a [u8],
904 secrets: &SecretsConfig,
905) -> Result<Cow<'a, [u8]>, SecretViolationAction> {
906 if secrets.secrets.is_empty() {
907 return Ok(Cow::Borrowed(header_bytes));
908 }
909
910 let mut handler = SecretsHandler::new_plain_http_untrusted_metadata(secrets);
911 handler.substitute(header_bytes)
912}
913
914fn headers_end(buf: &[u8]) -> Option<usize> {
916 buf.windows(4).position(|w| w == b"\r\n\r\n").map(|p| p + 4)
917}
918
919fn could_be_connect_request(buf: &[u8]) -> bool {
920 const PREFIX: &[u8] = b"CONNECT ";
921 if buf.is_empty() {
922 return false;
923 }
924 let n = buf.len().min(PREFIX.len());
925 buf[..n].eq_ignore_ascii_case(&PREFIX[..n])
926}
927
928fn parse_connect_request(bytes: Vec<u8>) -> io::Result<ConnectRequest> {
929 let header_end = headers_end(&bytes).ok_or_else(|| {
930 io::Error::new(
931 io::ErrorKind::InvalidData,
932 "incomplete CONNECT request headers",
933 )
934 })?;
935 let target = {
936 let request_line = bytes[..header_end]
937 .split(|&b| b == b'\n')
938 .next()
939 .unwrap_or(&[]);
940 let request_line = std::str::from_utf8(request_line)
941 .map_err(|_| io::Error::new(io::ErrorKind::InvalidData, "CONNECT line is not UTF-8"))?
942 .trim_end_matches('\r');
943 let mut parts = request_line.split_ascii_whitespace();
944 let method = parts.next().unwrap_or_default();
945 let authority = parts.next().unwrap_or_default();
946 let version = parts.next().unwrap_or_default();
947 if !method.eq_ignore_ascii_case("CONNECT")
948 || authority.is_empty()
949 || !is_http_version(version)
950 || parts.next().is_some()
951 {
952 return Err(io::Error::new(
953 io::ErrorKind::InvalidData,
954 "malformed CONNECT request line",
955 ));
956 }
957 parse_connect_target(authority)?
958 };
959
960 Ok(ConnectRequest {
961 bytes,
962 header_end,
963 target,
964 })
965}
966
967fn parse_connect_target(authority: &str) -> io::Result<ConnectTarget> {
968 let authority = authority.trim();
969 let (host, port) = if let Some(rest) = authority.strip_prefix('[') {
970 let (host, rest) = rest.split_once(']').ok_or_else(|| {
971 io::Error::new(
972 io::ErrorKind::InvalidData,
973 "malformed CONNECT IPv6 authority",
974 )
975 })?;
976 let port = rest.strip_prefix(':').ok_or_else(|| {
977 io::Error::new(io::ErrorKind::InvalidData, "CONNECT authority missing port")
978 })?;
979 (host, port)
980 } else {
981 let (host, port) = authority.rsplit_once(':').ok_or_else(|| {
982 io::Error::new(io::ErrorKind::InvalidData, "CONNECT authority missing port")
983 })?;
984 if host.contains(':') {
985 return Err(io::Error::new(
986 io::ErrorKind::InvalidData,
987 "CONNECT IPv6 authority must be bracketed",
988 ));
989 }
990 (host, port)
991 };
992 let host = host.trim().trim_end_matches('.');
993 if host.is_empty() {
994 return Err(io::Error::new(
995 io::ErrorKind::InvalidData,
996 "CONNECT authority missing host",
997 ));
998 }
999 let port = port
1000 .parse::<u16>()
1001 .map_err(|_| io::Error::new(io::ErrorKind::InvalidData, "invalid CONNECT port"))?;
1002 let expected_sni = host
1003 .parse::<IpAddr>()
1004 .is_err()
1005 .then(|| host.to_ascii_lowercase());
1006
1007 Ok(ConnectTarget {
1008 host: host.to_ascii_lowercase(),
1009 port,
1010 expected_sni,
1011 })
1012}
1013
1014fn is_http_version(version: &str) -> bool {
1015 let Some(version) = version.strip_prefix("HTTP/") else {
1016 return false;
1017 };
1018 let Some((major, minor)) = version.split_once('.') else {
1019 return false;
1020 };
1021 !major.is_empty()
1022 && !minor.is_empty()
1023 && major.bytes().all(|b| b.is_ascii_digit())
1024 && minor.bytes().all(|b| b.is_ascii_digit())
1025}
1026
1027fn connect_response_is_success(headers: &[u8]) -> bool {
1028 let Some(status_line) = headers.split(|&b| b == b'\n').next() else {
1029 return false;
1030 };
1031 let Ok(status_line) = std::str::from_utf8(status_line) else {
1032 return false;
1033 };
1034 let mut parts = status_line.trim_end_matches('\r').split_ascii_whitespace();
1035 let version = parts.next().unwrap_or_default();
1036 let status = parts.next().unwrap_or_default();
1037 is_http_version(version)
1038 && status.len() == 3
1039 && status
1040 .parse::<u16>()
1041 .is_ok_and(|code| (200..300).contains(&code))
1042}
1043
1044pub(crate) async fn deny_http_or_close(
1050 guest_dst: SocketAddr,
1051 sni: Option<&str>,
1052 initial_buf: &[u8],
1053 to_smoltcp: mpsc::Sender<Bytes>,
1054 shared: &SharedState,
1055 proxy_connect: &ProxyConnectState,
1056) -> io::Result<()> {
1057 let answer = shared.http_deny_response_enabled() && first_flight_is_http(initial_buf);
1059 if answer {
1060 let host = denied_host_label(sni, initial_buf, guest_dst);
1061 let body = shared.http_deny_body(&host);
1062 let _ = to_smoltcp
1063 .send(Bytes::from(http_forbidden_response(&body)))
1064 .await;
1065 shared.proxy_wake.wake();
1066 }
1067 proxy_connect.mark_policy_denied();
1068 shared.proxy_wake.wake();
1069 Ok(())
1070}
1071
1072fn first_flight_is_http(buf: &[u8]) -> bool {
1073 classify_http_request(buf) == Some(true)
1074}
1075
1076fn denied_host_label(sni: Option<&str>, buf: &[u8], guest_dst: SocketAddr) -> String {
1077 if let Some(name) = sni.filter(|name| !name.is_empty()) {
1078 return name.to_string();
1079 }
1080 if let Some(host) = extract_http_host(buf) {
1081 return host;
1082 }
1083 guest_dst.ip().to_string()
1084}
1085
1086fn extract_http_host(buf: &[u8]) -> Option<String> {
1096 if buf.first() == Some(&0x16) {
1097 return None;
1098 }
1099 let mut headers = vec![httparse::EMPTY_HEADER; (buf.len() / 4).max(16)];
1105 let mut req = httparse::Request::new(&mut headers);
1106 req.parse(buf).ok()?;
1107 req.headers
1108 .iter()
1109 .find(|h| h.name.eq_ignore_ascii_case("host"))
1110 .and_then(|h| std::str::from_utf8(h.value).ok())
1111 .map(|v| {
1112 let host = v.trim();
1113 host.rsplit_once(':')
1115 .map(|(h, _)| h)
1116 .unwrap_or(host)
1117 .to_ascii_lowercase()
1118 })
1119 .filter(|h| !h.is_empty())
1120}
1121
1122#[allow(clippy::too_many_arguments)]
1139async fn classify_first_flight(
1140 mut buf: Vec<u8>,
1141 from_smoltcp: &mut mpsc::Receiver<Bytes>,
1142 server_rx: &mut tokio::net::tcp::OwnedReadHalf,
1143 to_smoltcp: &mpsc::Sender<Bytes>,
1144 shared: &SharedState,
1145 want_headers: bool,
1146 max: usize,
1147 budget: Duration,
1148) -> io::Result<(Vec<u8>, bool)> {
1149 let mut server_buf = vec![0u8; SERVER_READ_BUF_SIZE];
1150 let timeout_fut = tokio::time::sleep(budget);
1151 tokio::pin!(timeout_fut);
1152
1153 loop {
1154 if !buf.is_empty() {
1160 let is_tls = buf.first() == Some(&0x16);
1161 let not_http = !is_tls
1162 && (!looks_like_http_request_prefix(&buf) || first_line_is_not_http_request(&buf));
1163 let done = !want_headers
1164 || is_tls
1165 || not_http
1166 || buf.len() >= max
1167 || buf.windows(4).any(|w| w == b"\r\n\r\n");
1168 if done {
1169 return Ok((buf, is_tls));
1170 }
1171 }
1172
1173 tokio::select! {
1174 biased;
1175 _ = &mut timeout_fut => {
1176 let is_tls = buf.first() == Some(&0x16);
1177 return Ok((buf, is_tls));
1178 }
1179 guest = from_smoltcp.recv() => match guest {
1182 Some(bytes) => buf.extend_from_slice(&bytes),
1183 None => {
1184 let is_tls = buf.first() == Some(&0x16);
1185 return Ok((buf, is_tls));
1186 }
1187 },
1188 server = server_rx.read(&mut server_buf) => match server {
1191 Ok(0) => {
1192 let is_tls = buf.first() == Some(&0x16);
1193 return Ok((buf, is_tls));
1194 }
1195 Ok(n) => {
1196 let data = Bytes::copy_from_slice(&server_buf[..n]);
1197 if to_smoltcp.send(data).await.is_err() {
1198 let is_tls = buf.first() == Some(&0x16);
1199 return Ok((buf, is_tls));
1200 }
1201 shared.proxy_wake.wake();
1202 }
1203 Err(e) => return Err(e),
1204 },
1205 }
1206 }
1207}
1208
1209pub(crate) async fn peek_for_http_request(
1218 rx: &mut mpsc::Receiver<Bytes>,
1219 mut buf: Vec<u8>,
1220 max: usize,
1221 budget: Duration,
1222) -> Vec<u8> {
1223 buf.truncate(max);
1224 let timeout_fut = tokio::time::sleep(budget);
1225 tokio::pin!(timeout_fut);
1226
1227 while buf.len() < max {
1228 match classify_http_request(&buf) {
1229 Some(false) => break,
1230 Some(true) => {
1231 let mut request = buf.as_slice();
1232 while let Some(rest) = request.strip_prefix(b"\r\n") {
1233 request = rest;
1234 }
1235 if request.windows(4).any(|bytes| bytes == b"\r\n\r\n") {
1236 break;
1237 }
1238 }
1239 None => {}
1240 }
1241 tokio::select! {
1242 biased;
1243 _ = &mut timeout_fut => break,
1244 data = rx.recv() => match data {
1245 Some(bytes) => {
1246 let remaining = max - buf.len();
1247 buf.extend_from_slice(&bytes[..bytes.len().min(remaining)]);
1248 }
1249 None => break,
1250 }
1251 }
1252 }
1253 buf
1254}
1255
1256pub(crate) async fn peek_for_sni(
1266 rx: &mut mpsc::Receiver<Bytes>,
1267 max: usize,
1268 budget: Duration,
1269) -> (Vec<u8>, Option<String>) {
1270 let mut buf = Vec::with_capacity(PEEK_BUF_SIZE.min(8192));
1271 let timeout_fut = tokio::time::sleep(budget);
1272 tokio::pin!(timeout_fut);
1273
1274 let raw_sni = loop {
1275 tokio::select! {
1276 biased;
1277 _ = &mut timeout_fut => break None,
1278 data = rx.recv() => {
1279 match data {
1280 Some(bytes) => {
1281 buf.extend_from_slice(&bytes);
1282 if buf.first() != Some(&0x16) {
1287 break None;
1288 }
1289 if let Some(name) = sni::extract_sni(&buf) {
1290 break Some(name);
1291 }
1292 if buf.len() >= max {
1293 break None;
1294 }
1295 }
1296 None => break None,
1297 }
1298 }
1299 }
1300 };
1301
1302 let canonical = raw_sni.map(|s| s.trim_end_matches('.').to_ascii_lowercase());
1303 (buf, canonical)
1304}
1305
1306#[cfg(test)]
1311mod tests {
1312 use super::*;
1313
1314 fn synthetic_client_hello(sni: &str) -> Vec<u8> {
1318 let host_bytes = sni.as_bytes();
1321 let host_len = host_bytes.len() as u16;
1322 let server_name_list_len = 3 + host_len; let extension_data_len = 2 + server_name_list_len; let extensions_total = 4 + extension_data_len; let mut body = Vec::new();
1327 body.extend_from_slice(&[0x03, 0x03]);
1329 body.extend_from_slice(&[0u8; 32]);
1331 body.push(0);
1333 body.extend_from_slice(&[0x00, 0x02, 0x00, 0x2f]);
1335 body.extend_from_slice(&[0x01, 0x00]);
1337 body.extend_from_slice(&extensions_total.to_be_bytes());
1339 body.extend_from_slice(&[0x00, 0x00]);
1341 body.extend_from_slice(&extension_data_len.to_be_bytes());
1342 body.extend_from_slice(&server_name_list_len.to_be_bytes());
1343 body.push(0x00); body.extend_from_slice(&host_len.to_be_bytes());
1345 body.extend_from_slice(host_bytes);
1346
1347 let handshake_len = body.len() as u32;
1348 let mut hs = Vec::new();
1349 hs.push(0x01); hs.extend_from_slice(&handshake_len.to_be_bytes()[1..]); hs.extend_from_slice(&body);
1352
1353 let record_len = hs.len() as u16;
1354 let mut record = Vec::new();
1355 record.extend_from_slice(&[0x16, 0x03, 0x01]); record.extend_from_slice(&record_len.to_be_bytes());
1357 record.extend_from_slice(&hs);
1358
1359 record
1360 }
1361
1362 #[tokio::test]
1363 async fn connect_upstream_dials_target_directly_without_outbound_proxy() {
1364 use tokio::net::TcpListener;
1365
1366 let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
1367 let addr = listener.local_addr().unwrap();
1368 let accept = tokio::spawn(async move {
1369 let (mut sock, _) = listener.accept().await.unwrap();
1370 let mut buf = [0u8; 5];
1371 sock.read_exact(&mut buf).await.unwrap();
1372 assert_eq!(&buf, b"hello");
1373 });
1374
1375 let shared = SharedState::new(4);
1376 let proxy_connect = ProxyConnectState::new();
1377 let mut stream = UpstreamTcpTarget::direct(addr)
1378 .connect(&proxy_connect, &shared, None)
1379 .await
1380 .unwrap();
1381 stream.write_all(b"hello").await.unwrap();
1382
1383 accept.await.unwrap();
1384 assert!(matches!(
1385 proxy_connect.status(),
1386 ProxyConnectStatus::Connected
1387 ));
1388 }
1389
1390 #[tokio::test]
1391 async fn early_http_connect_dials_proxy_through_configured_socks5_proxy() {
1392 let _ = rustls::crypto::ring::default_provider().install_default();
1393
1394 let http_proxy_addr: SocketAddr = "93.184.216.34:3128".parse().unwrap();
1397 let socks_listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
1398 let outbound_proxy = ResolvedOutboundProxy::Socks5 {
1399 address: socks_listener.local_addr().unwrap(),
1400 credentials: None,
1401 };
1402 let socks_task = tokio::spawn(async move {
1403 let (mut client, _) = socks_listener.accept().await.unwrap();
1404
1405 let mut greeting = [0u8; 3];
1406 client.read_exact(&mut greeting).await.unwrap();
1407 assert_eq!(greeting, [0x05, 0x01, 0x00]);
1408 client.write_all(&[0x05, 0x00]).await.unwrap();
1409
1410 let mut socks_request = [0u8; 10];
1411 client.read_exact(&mut socks_request).await.unwrap();
1412 assert_eq!(socks_request[0..4], [0x05, 0x01, 0x00, 0x01]);
1413 assert_eq!(&socks_request[4..8], &[93, 184, 216, 34]);
1414 assert_eq!(
1415 u16::from_be_bytes([socks_request[8], socks_request[9]]),
1416 3128
1417 );
1418 client
1419 .write_all(&[0x05, 0x00, 0x00, 0x01, 0, 0, 0, 0, 0, 0])
1420 .await
1421 .unwrap();
1422
1423 let expected_connect =
1424 b"CONNECT example.com:80 HTTP/1.1\r\nHost: example.com:80\r\n\r\n";
1425 let mut connect_request = vec![0u8; expected_connect.len()];
1426 client.read_exact(&mut connect_request).await.unwrap();
1427 assert_eq!(&connect_request, expected_connect);
1428 client
1429 .write_all(b"HTTP/1.1 200 Connection Established\r\n\r\n")
1430 .await
1431 .unwrap();
1432 });
1433
1434 let connect_request =
1435 b"CONNECT example.com:80 HTTP/1.1\r\nHost: example.com:80\r\n\r\n".to_vec();
1436 let (from_tx, from_rx) = mpsc::channel(1);
1437 let (to_tx, mut to_rx) = mpsc::channel(1);
1438 drop(from_tx);
1439
1440 let tls_state = Arc::new(
1441 TlsState::new(
1442 microsandbox_types::TlsConfig::default(),
1443 crate::secrets::handle::SecretsHandle::new(SecretsConfig::default()),
1444 )
1445 .unwrap(),
1446 );
1447 let proxy_connect = Arc::new(ProxyConnectState::new());
1448
1449 handle_connect_tunnel(
1450 http_proxy_addr,
1451 UpstreamTcpTarget::direct(http_proxy_addr),
1452 connect_request,
1453 from_rx,
1454 to_tx,
1455 Arc::new(SharedState::new(4)),
1456 Arc::new(NetworkPolicy::default()),
1457 tls_state,
1458 false,
1459 proxy_connect.clone(),
1460 Some(Arc::new(outbound_proxy)),
1461 None,
1462 )
1463 .await
1464 .unwrap();
1465
1466 let response = to_rx.recv().await.unwrap();
1467 assert_eq!(
1468 &response[..],
1469 b"HTTP/1.1 200 Connection Established\r\n\r\n"
1470 );
1471 socks_task.await.unwrap();
1472 assert!(matches!(
1473 proxy_connect.status(),
1474 ProxyConnectStatus::Connected
1475 ));
1476 }
1477
1478 #[test]
1479 fn could_be_connect_request_matches_split_prefixes_only() {
1480 assert!(could_be_connect_request(b"C"));
1481 assert!(could_be_connect_request(b"connect "));
1482 assert!(could_be_connect_request(b"CONNECT example.com:443"));
1483 assert!(!could_be_connect_request(b"CLIENT"));
1484 assert!(!could_be_connect_request(b"GET / HTTP/1.1\r\n"));
1485 }
1486
1487 #[test]
1488 fn first_flight_http_accepts_partial_and_complete_http() {
1489 assert!(!first_flight_is_http(b"GET /index.html"));
1490 assert!(first_flight_is_http(b"GET /x HTTP/1.1\r\nHost: a\r\n"));
1491 assert!(first_flight_is_http(b"\r\nGET /x HTTP/1.0\r\n"));
1492 assert!(!first_flight_is_http(b"PRI * HTTP/2.0\r\n\r\nSM\r\n\r\n"));
1494 assert!(first_flight_is_http(b"QUERY /x HTTP/1.1\r\n"));
1496 }
1497
1498 #[test]
1499 fn first_flight_http_rejects_split_non_http_banners() {
1500 assert!(!first_flight_is_http(b""));
1501 assert!(!first_flight_is_http(&synthetic_client_hello(
1502 "example.com"
1503 )));
1504 assert!(!first_flight_is_http(b"GE"));
1505 assert!(!first_flight_is_http(b"PRI"));
1506 assert!(!first_flight_is_http(b"SSH-2.0-OpenSSH_9.9"));
1509 assert!(!first_flight_is_http(b"EHLO mail.example.com"));
1510 assert!(!first_flight_is_http(b"QUERY /x"));
1511 }
1512
1513 #[tokio::test]
1514 async fn buffer_connect_request_reads_split_headers() {
1515 let (tx, mut rx) = mpsc::channel(4);
1516 tx.send(Bytes::from_static(b"NECT example.com:443 HTTP/1.1\r\n"))
1517 .await
1518 .unwrap();
1519 tx.send(Bytes::from_static(b"Host: example.com\r\n\r\n"))
1520 .await
1521 .unwrap();
1522 drop(tx);
1523
1524 let buffered = buffer_connect_request(b"CON".to_vec(), &mut rx)
1525 .await
1526 .unwrap();
1527 let parsed = parse_connect_request(buffered).unwrap();
1528
1529 assert_eq!(parsed.target.host, "example.com");
1530 assert_eq!(parsed.target.port, 443);
1531 assert_eq!(parsed.target.expected_sni.as_deref(), Some("example.com"));
1532 assert!(parsed.post_header_bytes().is_empty());
1533 }
1534
1535 #[test]
1536 fn parse_connect_request_preserves_post_header_tls_seed() {
1537 let mut request = b"CONNECT example.com:443 HTTP/1.1\r\nHost: example.com\r\n\r\n".to_vec();
1538 request.extend_from_slice(b"\x16\x03\x01client-hello");
1539
1540 let parsed = parse_connect_request(request).unwrap();
1541
1542 assert_eq!(
1543 parsed.header_bytes(),
1544 b"CONNECT example.com:443 HTTP/1.1\r\nHost: example.com\r\n\r\n"
1545 );
1546 assert_eq!(parsed.post_header_bytes(), b"\x16\x03\x01client-hello");
1547 }
1548
1549 #[test]
1550 fn parse_connect_target_requires_authority_port() {
1551 assert!(parse_connect_target("example.com").is_err());
1552 assert!(parse_connect_target("2001:db8::1:443").is_err());
1553
1554 let target = parse_connect_target("[2001:db8::1]:8443").unwrap();
1555 assert_eq!(target.host, "2001:db8::1");
1556 assert_eq!(target.port, 8443);
1557 assert_eq!(target.expected_sni, None);
1558 }
1559
1560 #[test]
1561 fn connect_response_success_requires_exact_2xx_status_code() {
1562 assert!(connect_response_is_success(
1563 b"HTTP/1.1 200 Connection Established\r\n\r\n"
1564 ));
1565 assert!(connect_response_is_success(
1566 b"HTTP/1.1 204 Connection Established\r\n\r\n"
1567 ));
1568 assert!(!connect_response_is_success(b"HTTP/1.1 2000 Weird\r\n\r\n"));
1569 assert!(!connect_response_is_success(b"HTTP/1.1 199 Nope\r\n\r\n"));
1570 assert!(!connect_response_is_success(b"NOTHTTP 200 OK\r\n\r\n"));
1571 }
1572
1573 #[tokio::test(start_paused = true)]
1574 async fn peek_for_http_request_joins_a_fragmented_method() {
1575 let (tx, mut rx) = mpsc::channel(4);
1576 tx.send(Bytes::from_static(b"GE")).await.unwrap();
1577 tx.send(Bytes::from_static(b"T / HTTP/1.1\r\nHost: x\r\n\r\n"))
1578 .await
1579 .unwrap();
1580 drop(tx);
1581
1582 let buf = peek_for_http_request(&mut rx, Vec::new(), PEEK_BUF_SIZE, PEEK_BUDGET).await;
1583 assert_eq!(buf, b"GET / HTTP/1.1\r\nHost: x\r\n\r\n");
1584 assert!(first_flight_is_http(&buf));
1585 }
1586
1587 #[tokio::test(start_paused = true)]
1588 async fn peek_for_http_request_joins_a_fragmented_non_http_line() {
1589 let (tx, mut rx) = mpsc::channel(4);
1590 tx.send(Bytes::from_static(b"EH")).await.unwrap();
1591 tx.send(Bytes::from_static(b"LO mail.example.com\r\n"))
1592 .await
1593 .unwrap();
1594 drop(tx);
1595
1596 let buf = peek_for_http_request(&mut rx, Vec::new(), PEEK_BUF_SIZE, PEEK_BUDGET).await;
1597 assert_eq!(buf, b"EHLO mail.example.com\r\n");
1598 assert!(!first_flight_is_http(&buf));
1599 }
1600
1601 #[tokio::test(start_paused = true)]
1602 async fn denied_http_peek_preserves_seed_and_split_leading_crlf() {
1603 let (tx, mut rx) = mpsc::channel(4);
1604 tx.send(Bytes::from_static(b"\nGE")).await.unwrap();
1605 tx.send(Bytes::from_static(b"T / HTTP/1.1\r"))
1606 .await
1607 .unwrap();
1608 tx.send(Bytes::from_static(b"\nHost: blocked.example\r\n\r\n"))
1609 .await
1610 .unwrap();
1611 drop(tx);
1612 let buf = peek_for_http_request(&mut rx, b"\r".to_vec(), PEEK_BUF_SIZE, PEEK_BUDGET).await;
1613 assert!(first_flight_is_http(&buf));
1614 assert_eq!(extract_http_host(&buf).as_deref(), Some("blocked.example"));
1615 }
1616
1617 #[tokio::test(start_paused = true)]
1618 async fn denied_http_peek_bounds_incomplete_requests() {
1619 let (tx, mut rx) = mpsc::channel(1);
1620 tx.send(Bytes::from_static(b"GET /an-overlong-request"))
1621 .await
1622 .unwrap();
1623 let buf = peek_for_http_request(&mut rx, Vec::new(), 8, PEEK_BUDGET).await;
1624 assert_eq!(buf.len(), 8);
1625 assert!(!first_flight_is_http(&buf));
1626 let buf = peek_for_http_request(
1627 &mut rx,
1628 b"GE".to_vec(),
1629 PEEK_BUF_SIZE,
1630 Duration::from_millis(1),
1631 )
1632 .await;
1633 assert_eq!(buf, b"GE");
1634 assert!(!first_flight_is_http(&buf));
1635 }
1636
1637 #[tokio::test]
1638 async fn domain_denial_joins_fragmented_request_without_dialing_upstream() {
1639 for enabled in [false, true] {
1640 let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
1641 let dst = listener.local_addr().unwrap();
1642 let shared = Arc::new(shared_with("blocked.example", "127.0.0.1"));
1643 shared.set_http_config(microsandbox_types::HttpConfig {
1644 deny_response: enabled,
1645 deny_message: Some("blocked {host}".into()),
1646 });
1647 let policy = Arc::new(NetworkPolicy {
1648 default_egress: Action::Deny,
1649 default_ingress: Action::Allow,
1650 rules: vec![allow_tcp("allowed.example", dst.port())],
1651 });
1652 let status = Arc::new(ProxyConnectState::new());
1653 let (from_tx, from_rx) = mpsc::channel(4);
1654 let (to_tx, mut to_rx) = mpsc::channel(4);
1655 from_tx.send(Bytes::from_static(b"GE")).await.unwrap();
1656 from_tx
1657 .send(Bytes::from_static(b"T / HTTP/1.1\r\n"))
1658 .await
1659 .unwrap();
1660 from_tx
1661 .send(Bytes::from_static(b"Host: blocked.example\r\n\r\n"))
1662 .await
1663 .unwrap();
1664 drop(from_tx);
1665 TcpProxy::new(
1666 dst,
1667 UpstreamTcpTarget::direct(dst),
1668 from_rx,
1669 to_tx,
1670 shared,
1671 policy,
1672 Arc::new(SecretsConfig::default()),
1673 None,
1674 false,
1675 status.clone(),
1676 None,
1677 )
1678 .try_run()
1679 .await
1680 .unwrap();
1681 let response = to_rx.recv().await;
1682 if enabled {
1683 let response = response.unwrap();
1684 assert!(response.starts_with(b"HTTP/1.1 403 Forbidden\r\n"));
1685 assert!(String::from_utf8_lossy(&response).contains("blocked.example"));
1686 } else {
1687 assert!(response.is_none(), "disabled responses must close silently");
1688 }
1689 assert_eq!(status.status(), ProxyConnectStatus::PolicyDenied);
1690 assert!(
1691 tokio::time::timeout(Duration::from_millis(20), listener.accept())
1692 .await
1693 .is_err()
1694 );
1695 }
1696 }
1697
1698 #[tokio::test]
1699 async fn peek_for_sni_extracts_and_canonicalizes() {
1700 let (tx, mut rx) = mpsc::channel(4);
1701 let hello = synthetic_client_hello("Example.COM");
1702 tx.send(Bytes::from(hello.clone())).await.unwrap();
1703 drop(tx); let (buf, sni) = peek_for_sni(&mut rx, PEEK_BUF_SIZE, PEEK_BUDGET).await;
1706 assert_eq!(sni.as_deref(), Some("example.com"));
1707 assert_eq!(buf, hello);
1708 }
1709
1710 #[tokio::test]
1711 async fn peek_for_sni_returns_none_on_channel_close_without_data() {
1712 let (tx, mut rx) = mpsc::channel::<Bytes>(1);
1713 drop(tx);
1714 let (buf, sni) = peek_for_sni(&mut rx, PEEK_BUF_SIZE, PEEK_BUDGET).await;
1715 assert!(buf.is_empty());
1716 assert_eq!(sni, None);
1717 }
1718
1719 #[tokio::test]
1720 async fn peek_for_sni_returns_none_on_non_tls_data() {
1721 let (tx, mut rx) = mpsc::channel(4);
1722 tx.send(Bytes::from_static(
1724 b"GET / HTTP/1.1\r\nHost: example.com\r\n\r\n",
1725 ))
1726 .await
1727 .unwrap();
1728 drop(tx);
1729 let (buf, sni) = peek_for_sni(&mut rx, PEEK_BUF_SIZE, PEEK_BUDGET).await;
1730 assert!(
1731 !buf.is_empty(),
1732 "buffered bytes must be returned for replay"
1733 );
1734 assert_eq!(sni, None);
1735 }
1736
1737 #[tokio::test]
1738 async fn peek_for_sni_falls_back_on_timeout() {
1739 let (tx, mut rx) = mpsc::channel::<Bytes>(1);
1740 let (buf, sni) = peek_for_sni(&mut rx, PEEK_BUF_SIZE, Duration::from_millis(50)).await;
1742 drop(tx);
1743 assert!(buf.is_empty());
1744 assert_eq!(sni, None);
1745 }
1746
1747 #[tokio::test]
1748 async fn peek_for_sni_caps_at_max_bytes() {
1749 let (tx, mut rx) = mpsc::channel(4);
1750 let mut first = vec![0u8; 8192];
1754 first[0] = 0x16;
1755 tx.send(Bytes::from(first)).await.unwrap();
1756 tx.send(Bytes::from(vec![0u8; 8192])).await.unwrap();
1757 tx.send(Bytes::from(vec![0u8; 8192])).await.unwrap();
1758 drop(tx);
1759
1760 let (buf, sni) = peek_for_sni(&mut rx, PEEK_BUF_SIZE, PEEK_BUDGET).await;
1761 assert_eq!(sni, None, "no SNI in non-TLS data");
1762 assert!(
1763 buf.len() >= PEEK_BUF_SIZE,
1764 "buffer must hit the cap before bail-out: got {}",
1765 buf.len()
1766 );
1767 }
1768
1769 #[tokio::test]
1770 async fn peek_for_sni_bails_immediately_on_non_tls_first_byte() {
1771 let (tx, mut rx) = mpsc::channel(4);
1772 tx.send(Bytes::from_static(b"GET / HTTP/1.1\r\nHost: x\r\n\r\n"))
1774 .await
1775 .unwrap();
1776 drop(tx);
1777
1778 let started = std::time::Instant::now();
1781 let (buf, sni) = peek_for_sni(&mut rx, PEEK_BUF_SIZE, PEEK_BUDGET).await;
1782 let elapsed = started.elapsed();
1783 assert_eq!(sni, None);
1784 assert!(buf.starts_with(b"GET"));
1785 assert!(
1786 elapsed < Duration::from_millis(500),
1787 "non-TLS bail must be fast: took {elapsed:?}"
1788 );
1789 }
1790
1791 use std::net::IpAddr;
1796 use std::time::Duration as StdDuration;
1797
1798 use crate::netstack::shared::{ResolvedHostnameFamily, SharedState};
1799 use crate::policy::{Action, Destination, NetworkPolicy, PortRange, Rule};
1800
1801 const SHARED_FASTLY_IP: &str = "151.101.0.223";
1802
1803 fn shared_with(host: &str, ip: &str) -> SharedState {
1804 let shared = SharedState::new(4);
1805 shared.cache_resolved_hostname(
1806 host,
1807 ResolvedHostnameFamily::Ipv4,
1808 [ip.parse::<IpAddr>().unwrap()],
1809 StdDuration::from_secs(60),
1810 );
1811 shared
1812 }
1813
1814 fn allow_https(domain: &str) -> Rule {
1815 Rule {
1816 direction: crate::policy::Direction::Egress,
1817 destination: Destination::Domain(domain.parse().unwrap()),
1818 protocols: vec![Protocol::Tcp],
1819 ports: vec![PortRange::single(443)],
1820 action: Action::Allow,
1821 }
1822 }
1823
1824 fn allow_tcp(domain: &str, port: u16) -> Rule {
1825 Rule {
1826 direction: crate::policy::Direction::Egress,
1827 destination: Destination::Domain(domain.parse().unwrap()),
1828 protocols: vec![Protocol::Tcp],
1829 ports: vec![PortRange::single(port)],
1830 action: Action::Allow,
1831 }
1832 }
1833
1834 #[tokio::test]
1837 async fn integration_sni_overrides_cache_for_over_allow() {
1838 let shared = shared_with("pypi.org", SHARED_FASTLY_IP);
1839 let policy = NetworkPolicy {
1840 default_egress: Action::Deny,
1841 default_ingress: Action::Allow,
1842 rules: vec![allow_https("pypi.org")],
1843 };
1844 let dst = SocketAddr::new(SHARED_FASTLY_IP.parse().unwrap(), 443);
1845
1846 let (tx, mut rx) = mpsc::channel(4);
1847 tx.send(Bytes::from(synthetic_client_hello("evil.com")))
1848 .await
1849 .unwrap();
1850 drop(tx);
1851
1852 let (initial_buf, sni) = peek_for_sni(&mut rx, PEEK_BUF_SIZE, PEEK_BUDGET).await;
1853 assert_eq!(sni.as_deref(), Some("evil.com"));
1854 assert!(!initial_buf.is_empty());
1855
1856 let source = sni
1857 .as_deref()
1858 .map(HostnameSource::Sni)
1859 .unwrap_or(HostnameSource::CacheOnly);
1860 let eval = policy.evaluate_egress_with_source(dst, Protocol::Tcp, &shared, source);
1861 assert_eq!(
1862 eval,
1863 EgressEvaluation::Deny,
1864 "SNI=evil.com must not piggy-back on the cached pypi.org match",
1865 );
1866 }
1867
1868 #[tokio::test]
1871 async fn integration_sni_overrides_cache_for_over_block() {
1872 let shared = shared_with("ads.example.com", SHARED_FASTLY_IP);
1873 let policy = NetworkPolicy {
1874 default_egress: Action::Allow,
1875 default_ingress: Action::Allow,
1876 rules: vec![Rule::deny_egress(Destination::Domain(
1877 "ads.example.com".parse().unwrap(),
1878 ))],
1879 };
1880 let dst = SocketAddr::new(SHARED_FASTLY_IP.parse().unwrap(), 443);
1881
1882 let (tx, mut rx) = mpsc::channel(4);
1883 tx.send(Bytes::from(synthetic_client_hello("api.example.com")))
1884 .await
1885 .unwrap();
1886 drop(tx);
1887
1888 let (_initial_buf, sni) = peek_for_sni(&mut rx, PEEK_BUF_SIZE, PEEK_BUDGET).await;
1889 assert_eq!(sni.as_deref(), Some("api.example.com"));
1890
1891 let source = sni
1892 .as_deref()
1893 .map(HostnameSource::Sni)
1894 .unwrap_or(HostnameSource::CacheOnly);
1895 let eval = policy.evaluate_egress_with_source(dst, Protocol::Tcp, &shared, source);
1896 assert_eq!(
1897 eval,
1898 EgressEvaluation::Allow,
1899 "SNI=api.example.com must not be caught by the deny on ads.example.com",
1900 );
1901 }
1902
1903 #[tokio::test]
1906 async fn integration_non_tls_falls_back_to_cache() {
1907 let shared = shared_with("pypi.org", SHARED_FASTLY_IP);
1908 let policy = NetworkPolicy {
1909 default_egress: Action::Deny,
1910 default_ingress: Action::Allow,
1911 rules: vec![allow_https("pypi.org")],
1912 };
1913 let dst = SocketAddr::new(SHARED_FASTLY_IP.parse().unwrap(), 443);
1914
1915 let (tx, mut rx) = mpsc::channel(4);
1916 tx.send(Bytes::from_static(
1918 b"GET / HTTP/1.1\r\nHost: pypi.org\r\n\r\n",
1919 ))
1920 .await
1921 .unwrap();
1922 drop(tx);
1923
1924 let (initial_buf, sni) = peek_for_sni(&mut rx, PEEK_BUF_SIZE, PEEK_BUDGET).await;
1925 assert_eq!(sni, None, "non-TLS data → no SNI");
1926 assert!(
1927 !initial_buf.is_empty(),
1928 "buffered bytes must survive for replay"
1929 );
1930
1931 let source = sni
1932 .as_deref()
1933 .map(HostnameSource::Sni)
1934 .unwrap_or(HostnameSource::CacheOnly);
1935 let eval = policy.evaluate_egress_with_source(dst, Protocol::Tcp, &shared, source);
1936 assert_eq!(
1937 eval,
1938 EgressEvaluation::Allow,
1939 "cache-only fallback must still allow the cached hostname's IP",
1940 );
1941 }
1942
1943 #[tokio::test]
1946 async fn integration_sni_matches_domain_suffix_with_cache_binding() {
1947 let shared = shared_with("files.pythonhosted.org", SHARED_FASTLY_IP);
1948 let policy = NetworkPolicy {
1949 default_egress: Action::Deny,
1950 default_ingress: Action::Allow,
1951 rules: vec![Rule {
1952 direction: crate::policy::Direction::Egress,
1953 destination: Destination::DomainSuffix(".pythonhosted.org".parse().unwrap()),
1954 protocols: vec![Protocol::Tcp],
1955 ports: vec![PortRange::single(443)],
1956 action: Action::Allow,
1957 }],
1958 };
1959 let dst = SocketAddr::new(SHARED_FASTLY_IP.parse().unwrap(), 443);
1960
1961 let (tx, mut rx) = mpsc::channel(4);
1962 tx.send(Bytes::from(synthetic_client_hello(
1963 "files.pythonhosted.org",
1964 )))
1965 .await
1966 .unwrap();
1967 drop(tx);
1968
1969 let (_buf, sni) = peek_for_sni(&mut rx, PEEK_BUF_SIZE, PEEK_BUDGET).await;
1970 let source = sni
1971 .as_deref()
1972 .map(HostnameSource::Sni)
1973 .unwrap_or(HostnameSource::CacheOnly);
1974 let eval = policy.evaluate_egress_with_source(dst, Protocol::Tcp, &shared, source);
1975 assert_eq!(eval, EgressEvaluation::Allow);
1976 }
1977
1978 #[tokio::test]
1983 async fn integration_sni_denies_domain_suffix_without_cache_binding() {
1984 let shared = SharedState::new(4); let policy = NetworkPolicy {
1986 default_egress: Action::Deny,
1987 default_ingress: Action::Allow,
1988 rules: vec![Rule {
1989 direction: crate::policy::Direction::Egress,
1990 destination: Destination::DomainSuffix(".pythonhosted.org".parse().unwrap()),
1991 protocols: vec![Protocol::Tcp],
1992 ports: vec![PortRange::single(443)],
1993 action: Action::Allow,
1994 }],
1995 };
1996 let dst = SocketAddr::new(SHARED_FASTLY_IP.parse().unwrap(), 443);
1997
1998 let (tx, mut rx) = mpsc::channel(4);
1999 tx.send(Bytes::from(synthetic_client_hello(
2000 "files.pythonhosted.org",
2001 )))
2002 .await
2003 .unwrap();
2004 drop(tx);
2005
2006 let (_buf, sni) = peek_for_sni(&mut rx, PEEK_BUF_SIZE, PEEK_BUDGET).await;
2007 let source = sni
2008 .as_deref()
2009 .map(HostnameSource::Sni)
2010 .unwrap_or(HostnameSource::CacheOnly);
2011 let eval = policy.evaluate_egress_with_source(dst, Protocol::Tcp, &shared, source);
2012 assert_eq!(eval, EgressEvaluation::Deny);
2013 }
2014
2015 #[test]
2018 fn extract_http_host_basic() {
2019 let buf = b"GET / HTTP/1.1\r\nHost: example.com\r\n\r\n";
2020 assert_eq!(extract_http_host(buf), Some("example.com".into()));
2021 }
2022
2023 #[test]
2024 fn extract_http_host_strips_port() {
2025 let buf = b"POST /api HTTP/1.1\r\nHost: api.company.com:8080\r\n\r\n";
2026 assert_eq!(extract_http_host(buf), Some("api.company.com".into()));
2027 }
2028
2029 #[test]
2030 fn extract_http_host_case_insensitive_lowercased() {
2031 let buf = b"GET / HTTP/1.1\r\nhost: Example.COM\r\n\r\n";
2032 assert_eq!(extract_http_host(buf), Some("example.com".into()));
2033 }
2034
2035 #[test]
2036 fn extract_http_host_no_host_header() {
2037 let buf = b"GET / HTTP/1.1\r\nX-Other: foo\r\n\r\n";
2038 assert_eq!(extract_http_host(buf), None);
2039 }
2040
2041 #[test]
2042 fn extract_http_host_incomplete_headers() {
2043 let buf = b"GET / HTTP/1.1\r\nHost: x";
2044 assert_eq!(extract_http_host(buf), None);
2045 }
2046
2047 #[test]
2048 fn extract_http_host_tls_first_byte() {
2049 let buf = [0x16u8, 0x03, 0x01, 0x00, 0x01];
2050 assert_eq!(extract_http_host(&buf), None);
2051 }
2052
2053 #[test]
2054 fn http_403_answers_only_confirmed_http1() {
2055 let get = b"GET / HTTP/1.1\r\nHost: example.com\r\n\r\n";
2056 assert!(first_flight_is_http(get));
2057 assert!(!first_flight_is_http(b""));
2058 assert!(!first_flight_is_http(&[0x16, 0x03, 0x01]));
2059 assert!(!first_flight_is_http(b"\x00\x01binary"));
2060 }
2061
2062 #[test]
2063 fn extract_http_host_with_many_headers() {
2064 let mut req = Vec::from(&b"GET / HTTP/1.1\r\n"[..]);
2067 for i in 0..100 {
2068 req.extend_from_slice(format!("X-Pad-{i}: v\r\n").as_bytes());
2069 }
2070 req.extend_from_slice(b"Host: example.com\r\n\r\n");
2071 assert_eq!(extract_http_host(&req), Some("example.com".into()));
2072 }
2073
2074 use std::sync::Arc;
2077 use tokio::io::AsyncReadExt;
2078 use tokio::net::TcpListener;
2079 use tokio::task::JoinHandle;
2080
2081 use crate::secrets::config::{
2082 HostPattern, SecretEntry, SecretSubstitution, SecretViolationAction, SecretsConfig,
2083 };
2084
2085 fn make_plain_http_secret(placeholder: &str, value: &str, require_tls: bool) -> SecretsConfig {
2086 SecretsConfig {
2087 secrets: vec![SecretEntry {
2088 env_var: "API_KEY".into(),
2089 value: zeroize::Zeroizing::new(value.into()),
2090 source: None,
2091 placeholder: placeholder.into(),
2092 allowed_hosts: vec![HostPattern::Any],
2093 substitution: SecretSubstitution {
2094 headers: true,
2095 header_fields: Vec::new(),
2096 query: false,
2097 body: false,
2098 },
2099 passthrough_hosts: Vec::new(),
2100 violation_action: None,
2101 require_tls_identity: require_tls,
2102 }],
2103 ..Default::default()
2104 }
2105 }
2106
2107 fn make_host_bound_secret(placeholder: &str, value: &str, host: &str) -> SecretsConfig {
2108 SecretsConfig {
2109 secrets: vec![SecretEntry {
2110 env_var: "API_KEY".into(),
2111 value: zeroize::Zeroizing::new(value.into()),
2112 source: None,
2113 placeholder: placeholder.into(),
2114 allowed_hosts: vec![HostPattern::Exact(host.into())],
2115 substitution: SecretSubstitution::default(),
2116 passthrough_hosts: Vec::new(),
2117 violation_action: None,
2118 require_tls_identity: true,
2119 }],
2120 ..Default::default()
2121 }
2122 }
2123
2124 #[test]
2125 fn sanitize_connect_headers_blocks_placeholder_metadata_header_by_default() {
2126 let secrets = make_host_bound_secret("$MSB_KEY", "real-secret-value", "example.com");
2127 let headers = b"CONNECT example.com:443 HTTP/1.1\r\nHost: example.com:443\r\nProxy-Authorization: Bearer $MSB_KEY\r\nUser-Agent: curl\r\n\r\n";
2128
2129 assert_eq!(
2130 sanitize_connect_headers(headers, &secrets),
2131 Err(SecretViolationAction::BlockAndLog)
2132 );
2133 }
2134
2135 #[test]
2136 fn sanitize_connect_headers_respects_block_and_terminate() {
2137 let mut secrets = make_host_bound_secret("$MSB_KEY", "real-secret-value", "example.com");
2138 secrets.violation_action = SecretViolationAction::BlockAndTerminate;
2139 let headers = b"CONNECT example.com:443 HTTP/1.1\r\nHost: example.com:443\r\nProxy-Authorization: Bearer $MSB_KEY\r\n\r\n";
2140
2141 assert_eq!(
2142 sanitize_connect_headers(headers, &secrets),
2143 Err(SecretViolationAction::BlockAndTerminate)
2144 );
2145 }
2146
2147 #[test]
2148 fn sanitize_connect_headers_respects_explicit_passthrough() {
2149 let mut secrets = make_host_bound_secret("$MSB_KEY", "real-secret-value", "example.com");
2150 secrets.secrets[0].passthrough_hosts = vec![HostPattern::Any];
2151 let headers = b"CONNECT example.com:443 HTTP/1.1\r\nHost: example.com:443\r\nProxy-Authorization: Bearer $MSB_KEY\r\n\r\n";
2152
2153 let sanitized = sanitize_connect_headers(headers, &secrets).unwrap();
2154
2155 assert_eq!(sanitized.as_ref(), headers);
2156 assert!(
2157 !String::from_utf8_lossy(sanitized.as_ref()).contains("real-secret-value"),
2158 "passthrough must never substitute real secrets into CONNECT metadata"
2159 );
2160 }
2161
2162 #[test]
2163 fn sanitize_connect_headers_keeps_safe_metadata_headers() {
2164 let secrets = make_host_bound_secret("$MSB_KEY", "real-secret-value", "example.com");
2165 let headers =
2166 b"CONNECT example.com:443 HTTP/1.1\r\nHost: example.com:443\r\nUser-Agent: curl\r\n\r\n";
2167
2168 let sanitized = sanitize_connect_headers(headers, &secrets).unwrap();
2169
2170 assert_eq!(sanitized.as_ref(), headers);
2171 }
2172
2173 #[test]
2174 fn sanitize_connect_headers_blocks_placeholder_in_request_line() {
2175 let secrets = make_host_bound_secret("$MSB_KEY", "real-secret-value", "example.com");
2176 let headers = b"CONNECT $MSB_KEY:443 HTTP/1.1\r\nHost: example.com:443\r\n\r\n";
2177
2178 assert_eq!(
2179 sanitize_connect_headers(headers, &secrets),
2180 Err(SecretViolationAction::BlockAndLog)
2181 );
2182 }
2183
2184 async fn spawn_sink() -> (SocketAddr, JoinHandle<Vec<u8>>) {
2185 let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
2186 let addr = listener.local_addr().unwrap();
2187 let handle = tokio::spawn(async move {
2188 let (mut stream, _) = listener.accept().await.unwrap();
2189 let mut received = Vec::new();
2190 let mut buf = vec![0u8; 4096];
2191 loop {
2192 match stream.read(&mut buf).await {
2193 Ok(0) | Err(_) => break,
2194 Ok(n) => received.extend_from_slice(&buf[..n]),
2195 }
2196 }
2197 received
2198 });
2199 (addr, handle)
2200 }
2201
2202 async fn assert_server_first_banner_is_immediate(
2203 policy: NetworkPolicy,
2204 tls_state: Option<Arc<TlsState>>,
2205 ) {
2206 let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
2207 let addr = listener.local_addr().unwrap();
2208 let server = tokio::spawn(async move {
2209 let (mut stream, _) = listener.accept().await.unwrap();
2210 stream.write_all(b"READY\n").await.unwrap();
2211 let mut received = Vec::new();
2212 stream.read_to_end(&mut received).await.unwrap();
2213 });
2214
2215 let (from_tx, from_rx) = mpsc::channel::<Bytes>(8);
2216 let (to_tx, mut to_rx) = mpsc::channel::<Bytes>(8);
2217 spawn_tcp_proxy(
2218 &tokio::runtime::Handle::current(),
2219 addr,
2220 addr,
2221 from_rx,
2222 to_tx,
2223 Arc::new(SharedState::new(4)),
2224 Arc::new(policy),
2225 Arc::new(SecretsConfig::default()),
2226 tls_state,
2227 false,
2228 Arc::new(ProxyConnectState::new()),
2229 None,
2230 );
2231
2232 let banner = tokio::time::timeout(Duration::from_secs(1), to_rx.recv())
2233 .await
2234 .expect("server-first banner was delayed by a pre-connect peek")
2235 .expect("proxy closed before relaying the server-first banner");
2236 assert_eq!(banner, b"READY\n"[..]);
2237
2238 drop(from_tx);
2239 tokio::time::timeout(Duration::from_secs(7), server)
2240 .await
2241 .expect("proxy did not close the upstream connection")
2242 .unwrap();
2243 }
2244
2245 #[tokio::test]
2246 async fn server_first_connection_skips_unrelated_domain_policy_peek() {
2247 let policy = NetworkPolicy {
2248 default_egress: Action::Allow,
2249 default_ingress: Action::Allow,
2250 rules: vec![allow_https("unused.example")],
2251 };
2252
2253 assert_server_first_banner_is_immediate(policy, None).await;
2254 }
2255
2256 #[tokio::test]
2257 async fn server_first_connection_skips_eager_connect_peek() {
2258 let _ = rustls::crypto::ring::default_provider().install_default();
2259 let tls_state = Arc::new(
2260 TlsState::new(
2261 microsandbox_types::TlsConfig::default(),
2262 crate::secrets::handle::SecretsHandle::new(SecretsConfig::default()),
2263 )
2264 .unwrap(),
2265 );
2266
2267 assert_server_first_banner_is_immediate(NetworkPolicy::default(), Some(tls_state)).await;
2268 }
2269
2270 async fn relay_through_proxy(
2271 request: Vec<u8>,
2272 secrets: SecretsConfig,
2273 handle: JoinHandle<Vec<u8>>,
2274 server_addr: SocketAddr,
2275 ) -> Vec<u8> {
2276 relay_through_proxy_with_policy(
2277 request,
2278 Arc::new(SharedState::new(4)),
2279 Arc::new(NetworkPolicy::default()),
2280 secrets,
2281 handle,
2282 server_addr,
2283 )
2284 .await
2285 }
2286
2287 async fn relay_through_proxy_with_policy(
2288 request: Vec<u8>,
2289 shared: Arc<SharedState>,
2290 policy: Arc<NetworkPolicy>,
2291 secrets: SecretsConfig,
2292 handle: JoinHandle<Vec<u8>>,
2293 server_addr: SocketAddr,
2294 ) -> Vec<u8> {
2295 relay_chunks_through_proxy_with_policy(
2296 vec![request],
2297 shared,
2298 policy,
2299 secrets,
2300 handle,
2301 server_addr,
2302 )
2303 .await
2304 }
2305
2306 async fn relay_chunks_through_proxy_with_policy(
2307 chunks: Vec<Vec<u8>>,
2308 shared: Arc<SharedState>,
2309 policy: Arc<NetworkPolicy>,
2310 secrets: SecretsConfig,
2311 handle: JoinHandle<Vec<u8>>,
2312 server_addr: SocketAddr,
2313 ) -> Vec<u8> {
2314 let (from_tx, from_rx) = mpsc::channel::<Bytes>(8);
2315 let (to_tx, _to_rx) = mpsc::channel::<Bytes>(8);
2316 let secrets = Arc::new(secrets);
2317 let proxy_connect = Arc::new(ProxyConnectState::new());
2318
2319 for chunk in chunks {
2320 from_tx.send(Bytes::from(chunk)).await.unwrap();
2321 }
2322 drop(from_tx);
2323
2324 TcpProxy::new(
2325 server_addr,
2326 UpstreamTcpTarget::direct(server_addr),
2327 from_rx,
2328 to_tx,
2329 shared,
2330 policy,
2331 secrets,
2332 None,
2333 false,
2334 proxy_connect,
2335 None,
2336 )
2337 .try_run()
2338 .await
2339 .unwrap();
2340
2341 handle.await.unwrap()
2342 }
2343
2344 #[tokio::test]
2345 async fn plain_http_domain_policy_allows_matching_host() {
2346 let (addr, sink) = spawn_sink().await;
2347 let shared = Arc::new(shared_with("allowed.example", "127.0.0.1"));
2348 let policy = Arc::new(NetworkPolicy {
2349 default_egress: Action::Deny,
2350 default_ingress: Action::Allow,
2351 rules: vec![allow_tcp("allowed.example", addr.port())],
2352 });
2353
2354 let wire = relay_through_proxy_with_policy(
2355 b"GET / HTTP/1.1\r\nHost: allowed.example\r\n\r\n".to_vec(),
2356 shared,
2357 policy,
2358 SecretsConfig::default(),
2359 sink,
2360 addr,
2361 )
2362 .await;
2363
2364 assert_eq!(wire, b"GET / HTTP/1.1\r\nHost: allowed.example\r\n\r\n");
2365 }
2366
2367 #[tokio::test]
2368 async fn plain_http_domain_policy_blocks_host_switch() {
2369 let (addr, sink) = spawn_sink().await;
2370 let shared = Arc::new(shared_with("allowed.example", "127.0.0.1"));
2371 let policy = Arc::new(NetworkPolicy {
2372 default_egress: Action::Deny,
2373 default_ingress: Action::Allow,
2374 rules: vec![allow_tcp("allowed.example", addr.port())],
2375 });
2376
2377 let wire = relay_through_proxy_with_policy(
2378 b"GET / HTTP/1.1\r\nHost: denied.example\r\n\r\n".to_vec(),
2379 shared,
2380 policy,
2381 SecretsConfig::default(),
2382 sink,
2383 addr,
2384 )
2385 .await;
2386
2387 assert!(
2388 wire.is_empty(),
2389 "switched HTTP authority must not reach upstream, got: {wire:?}"
2390 );
2391 }
2392
2393 #[tokio::test]
2394 async fn plain_http_domain_policy_blocks_keep_alive_host_switch() {
2395 let (addr, sink) = spawn_sink().await;
2396 let shared = Arc::new(shared_with("allowed.example", "127.0.0.1"));
2397 let policy = Arc::new(NetworkPolicy {
2398 default_egress: Action::Deny,
2399 default_ingress: Action::Allow,
2400 rules: vec![allow_tcp("allowed.example", addr.port())],
2401 });
2402
2403 let wire = relay_chunks_through_proxy_with_policy(
2404 vec![
2405 b"GET /one HTTP/1.1\r\nHost: allowed.example\r\n\r\n".to_vec(),
2406 b"GET /two HTTP/1.1\r\nHost: denied.example\r\n\r\n".to_vec(),
2407 ],
2408 shared,
2409 policy,
2410 SecretsConfig::default(),
2411 sink,
2412 addr,
2413 )
2414 .await;
2415
2416 assert_eq!(wire, b"GET /one HTTP/1.1\r\nHost: allowed.example\r\n\r\n");
2417 }
2418
2419 async fn relay_h2c_until_proxy_closes(chunks: Vec<Vec<u8>>) -> (Vec<u8>, bool) {
2429 let (addr, mut sink) = spawn_sink().await;
2430 let shared = Arc::new(shared_with("allowed.example", "127.0.0.1"));
2431 let terminated = Arc::new(std::sync::atomic::AtomicBool::new(false));
2432 let flag = terminated.clone();
2433 shared.set_termination_hook(Arc::new(move || {
2434 flag.store(true, std::sync::atomic::Ordering::SeqCst);
2435 }));
2436 let policy = Arc::new(NetworkPolicy {
2437 default_egress: Action::Deny,
2438 default_ingress: Action::Allow,
2439 rules: vec![allow_tcp("allowed.example", addr.port())],
2440 });
2441 let (from_tx, from_rx) = mpsc::channel::<Bytes>(chunks.len());
2442 let (to_tx, _to_rx) = mpsc::channel::<Bytes>(8);
2443 for chunk in chunks {
2444 from_tx.send(Bytes::from(chunk)).await.unwrap();
2445 }
2446
2447 let proxy = TcpProxy::new(
2448 addr,
2449 UpstreamTcpTarget::direct(addr),
2450 from_rx,
2451 to_tx,
2452 shared,
2453 policy,
2454 Arc::new(SecretsConfig::default()),
2455 None,
2456 false,
2457 Arc::new(ProxyConnectState::new()),
2458 None,
2459 )
2460 .try_run();
2461 tokio::time::timeout(Duration::from_secs(5), proxy)
2462 .await
2463 .expect("proxy kept the guest connection open")
2464 .unwrap();
2465 drop(from_tx);
2466
2467 let wire = match tokio::time::timeout(Duration::from_secs(5), &mut sink).await {
2468 Ok(wire) => wire.unwrap(),
2469 Err(_) => {
2470 sink.abort();
2471 panic!(
2472 "upstream sink never accepted a connection; the proxy refused the guest before dialing"
2473 );
2474 }
2475 };
2476 (wire, terminated.load(std::sync::atomic::Ordering::SeqCst))
2477 }
2478
2479 fn h2c_first_request(authority: &[u8]) -> Vec<u8> {
2483 let mut block = vec![0x82, 0x86, 0x84, 0x41, authority.len() as u8];
2484 block.extend_from_slice(authority);
2485 let mut first = H2_PREFACE.to_vec();
2486 first.extend(h2_frame(H2_SETTINGS, 0, 0, &[]));
2487 first.extend(h2_frame(
2488 H2_HEADERS,
2489 H2_END_STREAM | H2_END_HEADERS,
2490 1,
2491 &block,
2492 ));
2493 first
2494 }
2495
2496 fn h2c_first_request_wire(authority: &[u8]) -> Vec<u8> {
2499 let mut expected_block = Vec::new();
2500 for (name, value) in [
2501 (&b":method"[..], &b"GET"[..]),
2502 (b":scheme", b"http"),
2503 (b":path", b"/"),
2504 (b":authority", authority),
2505 ] {
2506 expected_block.extend_from_slice(&[0x10, name.len() as u8]);
2507 expected_block.extend_from_slice(name);
2508 expected_block.push(value.len() as u8);
2509 expected_block.extend_from_slice(value);
2510 }
2511 let mut expected = H2_PREFACE.to_vec();
2512 expected.extend(h2_frame(H2_SETTINGS, 0, 0, &[]));
2513 expected.extend(h2_frame(
2514 H2_HEADERS,
2515 H2_END_STREAM | H2_END_HEADERS,
2516 1,
2517 &expected_block,
2518 ));
2519 expected
2520 }
2521
2522 fn h2_frame(kind: u8, flags: u8, stream_id: u32, payload: &[u8]) -> Vec<u8> {
2524 let mut frame = (payload.len() as u32).to_be_bytes()[1..].to_vec();
2525 frame.extend_from_slice(&[kind, flags]);
2526 frame.extend_from_slice(&stream_id.to_be_bytes());
2527 frame.extend_from_slice(payload);
2528 frame
2529 }
2530
2531 const H2_PREFACE: &[u8] = b"PRI * HTTP/2.0\r\n\r\nSM\r\n\r\n";
2532 const H2_HEADERS: u8 = 0x1;
2533 const H2_SETTINGS: u8 = 0x4;
2534 const H2_CONTINUATION: u8 = 0x9;
2535 const H2_END_STREAM: u8 = 0x1;
2536 const H2_END_HEADERS: u8 = 0x4;
2537
2538 #[tokio::test]
2539 async fn h2c_malformed_hpack_block_is_blocked_under_domain_policy() {
2540 let mut request = H2_PREFACE.to_vec();
2544 request.extend(h2_frame(
2545 H2_HEADERS,
2546 H2_END_STREAM | H2_END_HEADERS,
2547 1,
2548 &[0xff],
2549 ));
2550
2551 let (wire, terminated) = relay_h2c_until_proxy_closes(vec![request]).await;
2552
2553 assert!(
2555 wire.is_empty(),
2556 "malformed block reached upstream: {wire:02x?}"
2557 );
2558 assert!(!terminated);
2559 }
2560
2561 #[tokio::test]
2562 async fn h2c_late_hpack_error_closes_connection_after_valid_blocks() {
2563 let authority = b"allowed.example";
2568 let first = h2c_first_request(authority);
2569 let flags = H2_END_STREAM | H2_END_HEADERS;
2570 let second = h2_frame(H2_HEADERS, flags, 3, &[0x82, 0x86, 0x84, 0xbe, 0xff]);
2571 let third = h2_frame(H2_HEADERS, flags, 5, &[0x82, 0x86, 0x84, 0xbe]);
2572
2573 let (wire, terminated) = relay_h2c_until_proxy_closes(vec![first, second, third]).await;
2574
2575 assert_eq!(wire, h2c_first_request_wire(authority));
2578 assert!(!terminated);
2579 }
2580
2581 #[tokio::test]
2582 async fn h2c_decoder_error_after_partial_insertion_closes_connection() {
2583 let authority = b"allowed.example";
2594 let first = h2c_first_request(authority);
2595 let flags = H2_END_STREAM | H2_END_HEADERS;
2596 let failing = h2_frame(
2597 H2_HEADERS,
2598 flags,
2599 3,
2600 &[
2601 0x82, 0x86, 0x84, 0xbe, 0x40, 0x03, 0x78, 0x2d, 0x61, 0x01, 0x31, 0x80,
2602 ],
2603 );
2604 let third = h2_frame(H2_HEADERS, flags, 5, &[0x82, 0x86, 0x84, 0xbf]);
2605
2606 let (wire, terminated) = relay_h2c_until_proxy_closes(vec![first, failing, third]).await;
2607
2608 assert_eq!(wire, h2c_first_request_wire(authority));
2609 assert!(!terminated);
2610 }
2611
2612 #[tokio::test]
2613 async fn h2c_fragmented_malformed_hpack_block_is_blocked() {
2614 let settings = h2_frame(H2_SETTINGS, 0, 0, &[]);
2617 let headers = h2_frame(H2_HEADERS, H2_END_STREAM, 1, &[0x7f]);
2618 let continuation = h2_frame(H2_CONTINUATION, H2_END_HEADERS, 1, &[0xc5]);
2619 let (preface_head, preface_tail) = H2_PREFACE.split_at(18);
2620 let chunks = vec![
2621 preface_head.to_vec(),
2622 [preface_tail, &settings[..4]].concat(),
2623 [&settings[4..], &headers[..9]].concat(),
2624 headers[9..].to_vec(),
2625 continuation[..5].to_vec(),
2626 continuation[5..].to_vec(),
2627 ];
2628
2629 let (wire, terminated) = relay_h2c_until_proxy_closes(chunks).await;
2630
2631 assert_eq!(wire, [H2_PREFACE, &settings[..]].concat());
2632 assert!(!terminated);
2633 }
2634
2635 #[test]
2636 fn strict_hostname_allow_blocks_sni_authority_before_tcp_dial() {
2637 let dst = SocketAddr::new("127.0.0.1".parse().unwrap(), 443);
2638 let shared = shared_with("allowed.example", "127.0.0.1");
2639 let policy = NetworkPolicy {
2640 default_egress: Action::Deny,
2641 default_ingress: Action::Allow,
2642 rules: vec![allow_tcp("allowed.example", dst.port())],
2643 };
2644
2645 assert!(strict_hostname_allow_is_opaque(
2646 true,
2647 &policy,
2648 dst,
2649 &shared,
2650 Some("allowed.example"),
2651 &synthetic_client_hello("allowed.example"),
2652 ));
2653 }
2654
2655 #[test]
2656 fn strict_hostname_allow_blocks_tls_without_sni_before_tcp_dial() {
2657 let dst = SocketAddr::new("127.0.0.1".parse().unwrap(), 443);
2658 let shared = shared_with("allowed.example", "127.0.0.1");
2659 let policy = NetworkPolicy {
2660 default_egress: Action::Deny,
2661 default_ingress: Action::Allow,
2662 rules: vec![allow_tcp("allowed.example", dst.port())],
2663 };
2664
2665 assert!(strict_hostname_allow_is_opaque(
2666 true,
2667 &policy,
2668 dst,
2669 &shared,
2670 None,
2671 &[0x16, 0x03, 0x01],
2672 ));
2673 }
2674
2675 #[test]
2676 fn strict_hostname_allow_leaves_plain_http_for_authority_validation() {
2677 let dst = SocketAddr::new("127.0.0.1".parse().unwrap(), 80);
2678 let shared = shared_with("allowed.example", "127.0.0.1");
2679 let policy = NetworkPolicy {
2680 default_egress: Action::Deny,
2681 default_ingress: Action::Allow,
2682 rules: vec![allow_tcp("allowed.example", dst.port())],
2683 };
2684
2685 assert!(!strict_hostname_allow_is_opaque(
2686 true,
2687 &policy,
2688 dst,
2689 &shared,
2690 None,
2691 b"GET / HTTP/1.1\r\n",
2692 ));
2693 }
2694
2695 #[tokio::test]
2696 async fn strict_mode_blocks_hostname_allowed_opaque_tls() {
2697 let dst = SocketAddr::new("127.0.0.1".parse().unwrap(), 443);
2698 let shared = Arc::new(shared_with("allowed.example", "127.0.0.1"));
2699 let policy = Arc::new(NetworkPolicy {
2700 default_egress: Action::Deny,
2701 default_ingress: Action::Allow,
2702 rules: vec![allow_tcp("allowed.example", dst.port())],
2703 });
2704 let proxy_connect = Arc::new(ProxyConnectState::new());
2705 let (from_tx, from_rx) = mpsc::channel::<Bytes>(8);
2706 let (to_tx, _to_rx) = mpsc::channel::<Bytes>(8);
2707
2708 from_tx
2709 .send(Bytes::from(synthetic_client_hello("allowed.example")))
2710 .await
2711 .unwrap();
2712 drop(from_tx);
2713
2714 TcpProxy::new(
2715 dst,
2716 UpstreamTcpTarget::direct(dst),
2717 from_rx,
2718 to_tx,
2719 shared,
2720 policy,
2721 Arc::new(SecretsConfig::default()),
2722 None,
2723 true,
2724 proxy_connect.clone(),
2725 None,
2726 )
2727 .try_run()
2728 .await
2729 .unwrap();
2730
2731 assert_eq!(proxy_connect.status(), ProxyConnectStatus::PolicyDenied);
2732 }
2733
2734 #[tokio::test]
2735 async fn strict_mode_leaves_default_allowed_opaque_tls_to_policy() {
2736 let dst = SocketAddr::new("127.0.0.1".parse().unwrap(), 9);
2737 let shared = Arc::new(SharedState::new(4));
2738 let policy = Arc::new(NetworkPolicy {
2739 default_egress: Action::Allow,
2740 default_ingress: Action::Allow,
2741 rules: vec![Rule::deny_egress(Destination::Domain(
2742 "blocked.example".parse().unwrap(),
2743 ))],
2744 });
2745 let proxy_connect = Arc::new(ProxyConnectState::new());
2746 let (from_tx, from_rx) = mpsc::channel::<Bytes>(8);
2747 let (to_tx, _to_rx) = mpsc::channel::<Bytes>(8);
2748
2749 from_tx
2750 .send(Bytes::from(synthetic_client_hello("allowed.example")))
2751 .await
2752 .unwrap();
2753 drop(from_tx);
2754
2755 let result = TcpProxy::new(
2756 dst,
2757 UpstreamTcpTarget::direct(dst),
2758 from_rx,
2759 to_tx,
2760 shared,
2761 policy,
2762 Arc::new(SecretsConfig::default()),
2763 None,
2764 true,
2765 proxy_connect.clone(),
2766 None,
2767 )
2768 .try_run()
2769 .await;
2770
2771 assert!(result.is_err(), "dummy upstream should refuse the dial");
2772 assert_eq!(
2773 proxy_connect.status(),
2774 ProxyConnectStatus::UpstreamConnectFailed
2775 );
2776 }
2777
2778 #[tokio::test]
2779 async fn server_first_http_like_binary_first_flight_is_forwarded() {
2780 use tokio::net::TcpListener;
2781
2782 let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
2783 let addr = listener.local_addr().unwrap();
2784 let server = tokio::spawn(async move {
2785 let (mut stream, _) = listener.accept().await.unwrap();
2786 stream
2787 .write_all(b"binary server-first greeting")
2788 .await
2789 .unwrap();
2790 stream.flush().await.unwrap();
2791
2792 let mut received = Vec::new();
2793 stream.read_to_end(&mut received).await.unwrap();
2794 received
2795 });
2796
2797 let (from_tx, from_rx) = mpsc::channel::<Bytes>(8);
2798 let (to_tx, mut to_rx) = mpsc::channel::<Bytes>(8);
2799 spawn_tcp_proxy(
2800 &tokio::runtime::Handle::current(),
2801 addr,
2802 addr,
2803 from_rx,
2804 to_tx,
2805 Arc::new(SharedState::new(4)),
2806 Arc::new(NetworkPolicy::default()),
2807 Arc::new(make_plain_http_secret(
2808 "$MSB_UNUSED",
2809 "unused-secret-value",
2810 false,
2811 )),
2812 None,
2813 false,
2814 Arc::new(ProxyConnectState::new()),
2815 None,
2816 );
2817
2818 let greeting = to_rx.recv().await.unwrap();
2819 assert_eq!(greeting, b"binary server-first greeting"[..]);
2820
2821 let first_flight = Bytes::from_static(b"BINARY3 v1\x00\x01opaque request");
2824 from_tx.send(first_flight.clone()).await.unwrap();
2825 drop(from_tx);
2826
2827 let wire = tokio::time::timeout(Duration::from_secs(2), server)
2828 .await
2829 .expect("proxy did not finish forwarding the client first flight")
2830 .unwrap();
2831 assert_eq!(wire, first_flight);
2832 }
2833
2834 #[tokio::test]
2835 async fn plain_http_substitutes_placeholder_when_host_arrives_in_second_segment() {
2836 let (addr, sink) = spawn_sink().await;
2839 let secrets = make_plain_http_secret("$MSB_KEY", "real-secret-value", false);
2840
2841 let (from_tx, from_rx) = mpsc::channel::<Bytes>(8);
2842 let (to_tx, _to_rx) = mpsc::channel::<Bytes>(8);
2843 let proxy_connect = Arc::new(ProxyConnectState::new());
2844
2845 from_tx
2846 .send(Bytes::from_static(b"GET /api HTTP/1.1\r\n"))
2847 .await
2848 .unwrap();
2849 from_tx
2850 .send(Bytes::from_static(
2851 b"Host: example.com\r\nAuthorization: Bearer $MSB_KEY\r\n\r\n",
2852 ))
2853 .await
2854 .unwrap();
2855 drop(from_tx);
2856
2857 TcpProxy::new(
2858 addr,
2859 UpstreamTcpTarget::direct(addr),
2860 from_rx,
2861 to_tx,
2862 Arc::new(SharedState::new(4)),
2863 Arc::new(NetworkPolicy::default()),
2864 Arc::new(secrets),
2865 None,
2866 false,
2867 proxy_connect,
2868 None,
2869 )
2870 .try_run()
2871 .await
2872 .unwrap();
2873
2874 let wire = String::from_utf8(sink.await.unwrap()).unwrap();
2875 assert!(wire.contains("real-secret-value"), "got: {wire:?}");
2876 assert!(!wire.contains("$MSB_KEY"), "got: {wire:?}");
2877 }
2878
2879 #[tokio::test]
2880 async fn plain_http_passthrough_handles_a_host_in_split_headers() {
2881 let (addr, sink) = spawn_sink().await;
2885
2886 let shared = SharedState::new(4);
2887 shared.cache_resolved_hostname(
2888 "example.com",
2889 ResolvedHostnameFamily::Ipv4,
2890 ["127.0.0.1".parse::<IpAddr>().unwrap()],
2891 StdDuration::from_secs(60),
2892 );
2893
2894 let secrets = SecretsConfig {
2895 secrets: vec![SecretEntry {
2896 env_var: "API_KEY".into(),
2897 value: zeroize::Zeroizing::new("real-secret-value".into()),
2898 source: None,
2899 placeholder: "$MSB_KEY".into(),
2900 allowed_hosts: vec![HostPattern::Exact("example.com".into())],
2901 substitution: SecretSubstitution {
2902 headers: true,
2903 header_fields: Vec::new(),
2904 query: false,
2905 body: false,
2906 },
2907 passthrough_hosts: vec![HostPattern::Exact("example.com".into())],
2908 violation_action: None,
2909 require_tls_identity: true,
2910 }],
2911 ..Default::default()
2912 };
2913
2914 let (from_tx, from_rx) = mpsc::channel::<Bytes>(8);
2915 let (to_tx, _to_rx) = mpsc::channel::<Bytes>(8);
2916 let proxy_connect = Arc::new(ProxyConnectState::new());
2917
2918 from_tx
2919 .send(Bytes::from_static(b"GET /api HTTP/1.1\r\n"))
2920 .await
2921 .unwrap();
2922 from_tx
2923 .send(Bytes::from_static(
2924 b"Host: example.com\r\nAuthorization: Bearer $MSB_KEY\r\n\r\n",
2925 ))
2926 .await
2927 .unwrap();
2928 drop(from_tx);
2929
2930 TcpProxy::new(
2931 addr,
2932 UpstreamTcpTarget::direct(addr),
2933 from_rx,
2934 to_tx,
2935 Arc::new(shared),
2936 Arc::new(NetworkPolicy::default()),
2937 Arc::new(secrets),
2938 None,
2939 false,
2940 proxy_connect,
2941 None,
2942 )
2943 .try_run()
2944 .await
2945 .unwrap();
2946
2947 let wire = String::from_utf8(sink.await.unwrap()).unwrap();
2948 assert!(
2949 wire.contains("Host: example.com"),
2950 "request must reach the allowed host, got: {wire:?}"
2951 );
2952 assert!(
2953 wire.contains("$MSB_KEY"),
2954 "placeholder must be forwarded unchanged for a require_tls_identity secret, got: {wire:?}"
2955 );
2956 assert!(
2957 !wire.contains("real-secret-value"),
2958 "secret must never be substituted over plain HTTP, got: {wire:?}"
2959 );
2960 }
2961
2962 #[tokio::test]
2963 async fn plain_http_substitutes_placeholder_in_first_flight() {
2964 let (addr, sink) = spawn_sink().await;
2965
2966 let request =
2967 b"GET /api HTTP/1.1\r\nHost: example.com\r\nAuthorization: Bearer $MSB_KEY\r\n\r\n"
2968 .to_vec();
2969 let secrets = make_plain_http_secret("$MSB_KEY", "real-secret-value", false);
2970
2971 let wire =
2972 String::from_utf8(relay_through_proxy(request, secrets, sink, addr).await).unwrap();
2973 assert!(
2974 wire.contains("real-secret-value"),
2975 "real value must reach server, got: {wire:?}"
2976 );
2977 assert!(
2978 !wire.contains("$MSB_KEY"),
2979 "placeholder must not reach server, got: {wire:?}"
2980 );
2981 }
2982
2983 #[tokio::test]
2984 async fn plain_http_no_substitution_when_require_tls_identity_true() {
2985 let (addr, sink) = spawn_sink().await;
2986
2987 let request =
2988 b"GET /api HTTP/1.1\r\nHost: example.com\r\nAuthorization: Bearer $MSB_KEY\r\n\r\n"
2989 .to_vec();
2990 let mut secrets = make_plain_http_secret("$MSB_KEY", "real-secret-value", true);
2991 secrets.secrets[0].passthrough_hosts = vec![HostPattern::Any];
2992
2993 let wire =
2994 String::from_utf8_lossy(&relay_through_proxy(request, secrets, sink, addr).await)
2995 .into_owned();
2996 assert!(
2997 wire.contains("$MSB_KEY"),
2998 "placeholder must be forwarded unchanged when require_tls_identity=true, got: {wire:?}"
2999 );
3000 assert!(
3001 !wire.contains("real-secret-value"),
3002 "real value must not leak when require_tls_identity=true, got: {wire:?}"
3003 );
3004 }
3005
3006 #[tokio::test]
3007 async fn plain_http_large_body_forwarded_verbatim_in_relay_loop() {
3008 let (addr, sink) = spawn_sink().await;
3012 let secrets = make_plain_http_secret("$MSB_KEY", "real-value", false);
3013
3014 let body = "x".repeat(32_000);
3015 let header = format!(
3016 "POST /upload HTTP/1.1\r\nHost: example.com\r\nAuthorization: Bearer $MSB_KEY\r\nContent-Length: {}\r\n\r\n",
3017 body.len()
3018 );
3019
3020 let (from_tx, from_rx) = mpsc::channel::<Bytes>(8);
3021 let (to_tx, _to_rx) = mpsc::channel::<Bytes>(8);
3022 let proxy_connect = Arc::new(ProxyConnectState::new());
3023
3024 from_tx
3025 .send(Bytes::from(header.into_bytes()))
3026 .await
3027 .unwrap();
3028 from_tx
3029 .send(Bytes::from(body.clone().into_bytes()))
3030 .await
3031 .unwrap();
3032 drop(from_tx);
3033
3034 TcpProxy::new(
3035 addr,
3036 UpstreamTcpTarget::direct(addr),
3037 from_rx,
3038 to_tx,
3039 Arc::new(SharedState::new(4)),
3040 Arc::new(NetworkPolicy::default()),
3041 Arc::new(secrets),
3042 None,
3043 false,
3044 proxy_connect,
3045 None,
3046 )
3047 .try_run()
3048 .await
3049 .unwrap();
3050
3051 let wire = String::from_utf8_lossy(&sink.await.unwrap()).into_owned();
3052 assert!(wire.contains(&body), "got {} bytes", wire.len());
3053 assert!(!wire.contains("$MSB_KEY"), "got: {wire:?}");
3054 }
3055}