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 query: false,
2096 body: false,
2097 },
2098 passthrough_hosts: Vec::new(),
2099 violation_action: None,
2100 require_tls_identity: require_tls,
2101 }],
2102 ..Default::default()
2103 }
2104 }
2105
2106 fn make_host_bound_secret(placeholder: &str, value: &str, host: &str) -> SecretsConfig {
2107 SecretsConfig {
2108 secrets: vec![SecretEntry {
2109 env_var: "API_KEY".into(),
2110 value: zeroize::Zeroizing::new(value.into()),
2111 source: None,
2112 placeholder: placeholder.into(),
2113 allowed_hosts: vec![HostPattern::Exact(host.into())],
2114 substitution: SecretSubstitution::default(),
2115 passthrough_hosts: Vec::new(),
2116 violation_action: None,
2117 require_tls_identity: true,
2118 }],
2119 ..Default::default()
2120 }
2121 }
2122
2123 #[test]
2124 fn sanitize_connect_headers_blocks_placeholder_metadata_header_by_default() {
2125 let secrets = make_host_bound_secret("$MSB_KEY", "real-secret-value", "example.com");
2126 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";
2127
2128 assert_eq!(
2129 sanitize_connect_headers(headers, &secrets),
2130 Err(SecretViolationAction::BlockAndLog)
2131 );
2132 }
2133
2134 #[test]
2135 fn sanitize_connect_headers_respects_block_and_terminate() {
2136 let mut secrets = make_host_bound_secret("$MSB_KEY", "real-secret-value", "example.com");
2137 secrets.violation_action = SecretViolationAction::BlockAndTerminate;
2138 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";
2139
2140 assert_eq!(
2141 sanitize_connect_headers(headers, &secrets),
2142 Err(SecretViolationAction::BlockAndTerminate)
2143 );
2144 }
2145
2146 #[test]
2147 fn sanitize_connect_headers_respects_explicit_passthrough() {
2148 let mut secrets = make_host_bound_secret("$MSB_KEY", "real-secret-value", "example.com");
2149 secrets.secrets[0].passthrough_hosts = vec![HostPattern::Any];
2150 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";
2151
2152 let sanitized = sanitize_connect_headers(headers, &secrets).unwrap();
2153
2154 assert_eq!(sanitized.as_ref(), headers);
2155 assert!(
2156 !String::from_utf8_lossy(sanitized.as_ref()).contains("real-secret-value"),
2157 "passthrough must never substitute real secrets into CONNECT metadata"
2158 );
2159 }
2160
2161 #[test]
2162 fn sanitize_connect_headers_keeps_safe_metadata_headers() {
2163 let secrets = make_host_bound_secret("$MSB_KEY", "real-secret-value", "example.com");
2164 let headers =
2165 b"CONNECT example.com:443 HTTP/1.1\r\nHost: example.com:443\r\nUser-Agent: curl\r\n\r\n";
2166
2167 let sanitized = sanitize_connect_headers(headers, &secrets).unwrap();
2168
2169 assert_eq!(sanitized.as_ref(), headers);
2170 }
2171
2172 #[test]
2173 fn sanitize_connect_headers_blocks_placeholder_in_request_line() {
2174 let secrets = make_host_bound_secret("$MSB_KEY", "real-secret-value", "example.com");
2175 let headers = b"CONNECT $MSB_KEY:443 HTTP/1.1\r\nHost: example.com:443\r\n\r\n";
2176
2177 assert_eq!(
2178 sanitize_connect_headers(headers, &secrets),
2179 Err(SecretViolationAction::BlockAndLog)
2180 );
2181 }
2182
2183 async fn spawn_sink() -> (SocketAddr, JoinHandle<Vec<u8>>) {
2184 let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
2185 let addr = listener.local_addr().unwrap();
2186 let handle = tokio::spawn(async move {
2187 let (mut stream, _) = listener.accept().await.unwrap();
2188 let mut received = Vec::new();
2189 let mut buf = vec![0u8; 4096];
2190 loop {
2191 match stream.read(&mut buf).await {
2192 Ok(0) | Err(_) => break,
2193 Ok(n) => received.extend_from_slice(&buf[..n]),
2194 }
2195 }
2196 received
2197 });
2198 (addr, handle)
2199 }
2200
2201 async fn assert_server_first_banner_is_immediate(
2202 policy: NetworkPolicy,
2203 tls_state: Option<Arc<TlsState>>,
2204 ) {
2205 let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
2206 let addr = listener.local_addr().unwrap();
2207 let server = tokio::spawn(async move {
2208 let (mut stream, _) = listener.accept().await.unwrap();
2209 stream.write_all(b"READY\n").await.unwrap();
2210 let mut received = Vec::new();
2211 stream.read_to_end(&mut received).await.unwrap();
2212 });
2213
2214 let (from_tx, from_rx) = mpsc::channel::<Bytes>(8);
2215 let (to_tx, mut to_rx) = mpsc::channel::<Bytes>(8);
2216 spawn_tcp_proxy(
2217 &tokio::runtime::Handle::current(),
2218 addr,
2219 addr,
2220 from_rx,
2221 to_tx,
2222 Arc::new(SharedState::new(4)),
2223 Arc::new(policy),
2224 Arc::new(SecretsConfig::default()),
2225 tls_state,
2226 false,
2227 Arc::new(ProxyConnectState::new()),
2228 None,
2229 );
2230
2231 let banner = tokio::time::timeout(Duration::from_secs(1), to_rx.recv())
2232 .await
2233 .expect("server-first banner was delayed by a pre-connect peek")
2234 .expect("proxy closed before relaying the server-first banner");
2235 assert_eq!(banner, b"READY\n"[..]);
2236
2237 drop(from_tx);
2238 tokio::time::timeout(Duration::from_secs(7), server)
2239 .await
2240 .expect("proxy did not close the upstream connection")
2241 .unwrap();
2242 }
2243
2244 #[tokio::test]
2245 async fn server_first_connection_skips_unrelated_domain_policy_peek() {
2246 let policy = NetworkPolicy {
2247 default_egress: Action::Allow,
2248 default_ingress: Action::Allow,
2249 rules: vec![allow_https("unused.example")],
2250 };
2251
2252 assert_server_first_banner_is_immediate(policy, None).await;
2253 }
2254
2255 #[tokio::test]
2256 async fn server_first_connection_skips_eager_connect_peek() {
2257 let _ = rustls::crypto::ring::default_provider().install_default();
2258 let tls_state = Arc::new(
2259 TlsState::new(
2260 microsandbox_types::TlsConfig::default(),
2261 crate::secrets::handle::SecretsHandle::new(SecretsConfig::default()),
2262 )
2263 .unwrap(),
2264 );
2265
2266 assert_server_first_banner_is_immediate(NetworkPolicy::default(), Some(tls_state)).await;
2267 }
2268
2269 async fn relay_through_proxy(
2270 request: Vec<u8>,
2271 secrets: SecretsConfig,
2272 handle: JoinHandle<Vec<u8>>,
2273 server_addr: SocketAddr,
2274 ) -> Vec<u8> {
2275 relay_through_proxy_with_policy(
2276 request,
2277 Arc::new(SharedState::new(4)),
2278 Arc::new(NetworkPolicy::default()),
2279 secrets,
2280 handle,
2281 server_addr,
2282 )
2283 .await
2284 }
2285
2286 async fn relay_through_proxy_with_policy(
2287 request: Vec<u8>,
2288 shared: Arc<SharedState>,
2289 policy: Arc<NetworkPolicy>,
2290 secrets: SecretsConfig,
2291 handle: JoinHandle<Vec<u8>>,
2292 server_addr: SocketAddr,
2293 ) -> Vec<u8> {
2294 relay_chunks_through_proxy_with_policy(
2295 vec![request],
2296 shared,
2297 policy,
2298 secrets,
2299 handle,
2300 server_addr,
2301 )
2302 .await
2303 }
2304
2305 async fn relay_chunks_through_proxy_with_policy(
2306 chunks: Vec<Vec<u8>>,
2307 shared: Arc<SharedState>,
2308 policy: Arc<NetworkPolicy>,
2309 secrets: SecretsConfig,
2310 handle: JoinHandle<Vec<u8>>,
2311 server_addr: SocketAddr,
2312 ) -> Vec<u8> {
2313 let (from_tx, from_rx) = mpsc::channel::<Bytes>(8);
2314 let (to_tx, _to_rx) = mpsc::channel::<Bytes>(8);
2315 let secrets = Arc::new(secrets);
2316 let proxy_connect = Arc::new(ProxyConnectState::new());
2317
2318 for chunk in chunks {
2319 from_tx.send(Bytes::from(chunk)).await.unwrap();
2320 }
2321 drop(from_tx);
2322
2323 TcpProxy::new(
2324 server_addr,
2325 UpstreamTcpTarget::direct(server_addr),
2326 from_rx,
2327 to_tx,
2328 shared,
2329 policy,
2330 secrets,
2331 None,
2332 false,
2333 proxy_connect,
2334 None,
2335 )
2336 .try_run()
2337 .await
2338 .unwrap();
2339
2340 handle.await.unwrap()
2341 }
2342
2343 #[tokio::test]
2344 async fn plain_http_domain_policy_allows_matching_host() {
2345 let (addr, sink) = spawn_sink().await;
2346 let shared = Arc::new(shared_with("allowed.example", "127.0.0.1"));
2347 let policy = Arc::new(NetworkPolicy {
2348 default_egress: Action::Deny,
2349 default_ingress: Action::Allow,
2350 rules: vec![allow_tcp("allowed.example", addr.port())],
2351 });
2352
2353 let wire = relay_through_proxy_with_policy(
2354 b"GET / HTTP/1.1\r\nHost: allowed.example\r\n\r\n".to_vec(),
2355 shared,
2356 policy,
2357 SecretsConfig::default(),
2358 sink,
2359 addr,
2360 )
2361 .await;
2362
2363 assert_eq!(wire, b"GET / HTTP/1.1\r\nHost: allowed.example\r\n\r\n");
2364 }
2365
2366 #[tokio::test]
2367 async fn plain_http_domain_policy_blocks_host_switch() {
2368 let (addr, sink) = spawn_sink().await;
2369 let shared = Arc::new(shared_with("allowed.example", "127.0.0.1"));
2370 let policy = Arc::new(NetworkPolicy {
2371 default_egress: Action::Deny,
2372 default_ingress: Action::Allow,
2373 rules: vec![allow_tcp("allowed.example", addr.port())],
2374 });
2375
2376 let wire = relay_through_proxy_with_policy(
2377 b"GET / HTTP/1.1\r\nHost: denied.example\r\n\r\n".to_vec(),
2378 shared,
2379 policy,
2380 SecretsConfig::default(),
2381 sink,
2382 addr,
2383 )
2384 .await;
2385
2386 assert!(
2387 wire.is_empty(),
2388 "switched HTTP authority must not reach upstream, got: {wire:?}"
2389 );
2390 }
2391
2392 #[tokio::test]
2393 async fn plain_http_domain_policy_blocks_keep_alive_host_switch() {
2394 let (addr, sink) = spawn_sink().await;
2395 let shared = Arc::new(shared_with("allowed.example", "127.0.0.1"));
2396 let policy = Arc::new(NetworkPolicy {
2397 default_egress: Action::Deny,
2398 default_ingress: Action::Allow,
2399 rules: vec![allow_tcp("allowed.example", addr.port())],
2400 });
2401
2402 let wire = relay_chunks_through_proxy_with_policy(
2403 vec![
2404 b"GET /one HTTP/1.1\r\nHost: allowed.example\r\n\r\n".to_vec(),
2405 b"GET /two HTTP/1.1\r\nHost: denied.example\r\n\r\n".to_vec(),
2406 ],
2407 shared,
2408 policy,
2409 SecretsConfig::default(),
2410 sink,
2411 addr,
2412 )
2413 .await;
2414
2415 assert_eq!(wire, b"GET /one HTTP/1.1\r\nHost: allowed.example\r\n\r\n");
2416 }
2417
2418 #[test]
2419 fn strict_hostname_allow_blocks_sni_authority_before_tcp_dial() {
2420 let dst = SocketAddr::new("127.0.0.1".parse().unwrap(), 443);
2421 let shared = shared_with("allowed.example", "127.0.0.1");
2422 let policy = NetworkPolicy {
2423 default_egress: Action::Deny,
2424 default_ingress: Action::Allow,
2425 rules: vec![allow_tcp("allowed.example", dst.port())],
2426 };
2427
2428 assert!(strict_hostname_allow_is_opaque(
2429 true,
2430 &policy,
2431 dst,
2432 &shared,
2433 Some("allowed.example"),
2434 &synthetic_client_hello("allowed.example"),
2435 ));
2436 }
2437
2438 #[test]
2439 fn strict_hostname_allow_blocks_tls_without_sni_before_tcp_dial() {
2440 let dst = SocketAddr::new("127.0.0.1".parse().unwrap(), 443);
2441 let shared = shared_with("allowed.example", "127.0.0.1");
2442 let policy = NetworkPolicy {
2443 default_egress: Action::Deny,
2444 default_ingress: Action::Allow,
2445 rules: vec![allow_tcp("allowed.example", dst.port())],
2446 };
2447
2448 assert!(strict_hostname_allow_is_opaque(
2449 true,
2450 &policy,
2451 dst,
2452 &shared,
2453 None,
2454 &[0x16, 0x03, 0x01],
2455 ));
2456 }
2457
2458 #[test]
2459 fn strict_hostname_allow_leaves_plain_http_for_authority_validation() {
2460 let dst = SocketAddr::new("127.0.0.1".parse().unwrap(), 80);
2461 let shared = shared_with("allowed.example", "127.0.0.1");
2462 let policy = NetworkPolicy {
2463 default_egress: Action::Deny,
2464 default_ingress: Action::Allow,
2465 rules: vec![allow_tcp("allowed.example", dst.port())],
2466 };
2467
2468 assert!(!strict_hostname_allow_is_opaque(
2469 true,
2470 &policy,
2471 dst,
2472 &shared,
2473 None,
2474 b"GET / HTTP/1.1\r\n",
2475 ));
2476 }
2477
2478 #[tokio::test]
2479 async fn strict_mode_blocks_hostname_allowed_opaque_tls() {
2480 let dst = SocketAddr::new("127.0.0.1".parse().unwrap(), 443);
2481 let shared = Arc::new(shared_with("allowed.example", "127.0.0.1"));
2482 let policy = Arc::new(NetworkPolicy {
2483 default_egress: Action::Deny,
2484 default_ingress: Action::Allow,
2485 rules: vec![allow_tcp("allowed.example", dst.port())],
2486 });
2487 let proxy_connect = Arc::new(ProxyConnectState::new());
2488 let (from_tx, from_rx) = mpsc::channel::<Bytes>(8);
2489 let (to_tx, _to_rx) = mpsc::channel::<Bytes>(8);
2490
2491 from_tx
2492 .send(Bytes::from(synthetic_client_hello("allowed.example")))
2493 .await
2494 .unwrap();
2495 drop(from_tx);
2496
2497 TcpProxy::new(
2498 dst,
2499 UpstreamTcpTarget::direct(dst),
2500 from_rx,
2501 to_tx,
2502 shared,
2503 policy,
2504 Arc::new(SecretsConfig::default()),
2505 None,
2506 true,
2507 proxy_connect.clone(),
2508 None,
2509 )
2510 .try_run()
2511 .await
2512 .unwrap();
2513
2514 assert_eq!(proxy_connect.status(), ProxyConnectStatus::PolicyDenied);
2515 }
2516
2517 #[tokio::test]
2518 async fn strict_mode_leaves_default_allowed_opaque_tls_to_policy() {
2519 let dst = SocketAddr::new("127.0.0.1".parse().unwrap(), 9);
2520 let shared = Arc::new(SharedState::new(4));
2521 let policy = Arc::new(NetworkPolicy {
2522 default_egress: Action::Allow,
2523 default_ingress: Action::Allow,
2524 rules: vec![Rule::deny_egress(Destination::Domain(
2525 "blocked.example".parse().unwrap(),
2526 ))],
2527 });
2528 let proxy_connect = Arc::new(ProxyConnectState::new());
2529 let (from_tx, from_rx) = mpsc::channel::<Bytes>(8);
2530 let (to_tx, _to_rx) = mpsc::channel::<Bytes>(8);
2531
2532 from_tx
2533 .send(Bytes::from(synthetic_client_hello("allowed.example")))
2534 .await
2535 .unwrap();
2536 drop(from_tx);
2537
2538 let result = TcpProxy::new(
2539 dst,
2540 UpstreamTcpTarget::direct(dst),
2541 from_rx,
2542 to_tx,
2543 shared,
2544 policy,
2545 Arc::new(SecretsConfig::default()),
2546 None,
2547 true,
2548 proxy_connect.clone(),
2549 None,
2550 )
2551 .try_run()
2552 .await;
2553
2554 assert!(result.is_err(), "dummy upstream should refuse the dial");
2555 assert_eq!(
2556 proxy_connect.status(),
2557 ProxyConnectStatus::UpstreamConnectFailed
2558 );
2559 }
2560
2561 #[tokio::test]
2562 async fn server_first_http_like_binary_first_flight_is_forwarded() {
2563 use tokio::net::TcpListener;
2564
2565 let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
2566 let addr = listener.local_addr().unwrap();
2567 let server = tokio::spawn(async move {
2568 let (mut stream, _) = listener.accept().await.unwrap();
2569 stream
2570 .write_all(b"binary server-first greeting")
2571 .await
2572 .unwrap();
2573 stream.flush().await.unwrap();
2574
2575 let mut received = Vec::new();
2576 stream.read_to_end(&mut received).await.unwrap();
2577 received
2578 });
2579
2580 let (from_tx, from_rx) = mpsc::channel::<Bytes>(8);
2581 let (to_tx, mut to_rx) = mpsc::channel::<Bytes>(8);
2582 spawn_tcp_proxy(
2583 &tokio::runtime::Handle::current(),
2584 addr,
2585 addr,
2586 from_rx,
2587 to_tx,
2588 Arc::new(SharedState::new(4)),
2589 Arc::new(NetworkPolicy::default()),
2590 Arc::new(make_plain_http_secret(
2591 "$MSB_UNUSED",
2592 "unused-secret-value",
2593 false,
2594 )),
2595 None,
2596 false,
2597 Arc::new(ProxyConnectState::new()),
2598 None,
2599 );
2600
2601 let greeting = to_rx.recv().await.unwrap();
2602 assert_eq!(greeting, b"binary server-first greeting"[..]);
2603
2604 let first_flight = Bytes::from_static(b"BINARY3 v1\x00\x01opaque request");
2607 from_tx.send(first_flight.clone()).await.unwrap();
2608 drop(from_tx);
2609
2610 let wire = tokio::time::timeout(Duration::from_secs(2), server)
2611 .await
2612 .expect("proxy did not finish forwarding the client first flight")
2613 .unwrap();
2614 assert_eq!(wire, first_flight);
2615 }
2616
2617 #[tokio::test]
2618 async fn plain_http_substitutes_placeholder_when_host_arrives_in_second_segment() {
2619 let (addr, sink) = spawn_sink().await;
2622 let secrets = make_plain_http_secret("$MSB_KEY", "real-secret-value", false);
2623
2624 let (from_tx, from_rx) = mpsc::channel::<Bytes>(8);
2625 let (to_tx, _to_rx) = mpsc::channel::<Bytes>(8);
2626 let proxy_connect = Arc::new(ProxyConnectState::new());
2627
2628 from_tx
2629 .send(Bytes::from_static(b"GET /api HTTP/1.1\r\n"))
2630 .await
2631 .unwrap();
2632 from_tx
2633 .send(Bytes::from_static(
2634 b"Host: example.com\r\nAuthorization: Bearer $MSB_KEY\r\n\r\n",
2635 ))
2636 .await
2637 .unwrap();
2638 drop(from_tx);
2639
2640 TcpProxy::new(
2641 addr,
2642 UpstreamTcpTarget::direct(addr),
2643 from_rx,
2644 to_tx,
2645 Arc::new(SharedState::new(4)),
2646 Arc::new(NetworkPolicy::default()),
2647 Arc::new(secrets),
2648 None,
2649 false,
2650 proxy_connect,
2651 None,
2652 )
2653 .try_run()
2654 .await
2655 .unwrap();
2656
2657 let wire = String::from_utf8(sink.await.unwrap()).unwrap();
2658 assert!(wire.contains("real-secret-value"), "got: {wire:?}");
2659 assert!(!wire.contains("$MSB_KEY"), "got: {wire:?}");
2660 }
2661
2662 #[tokio::test]
2663 async fn plain_http_passthrough_handles_a_host_in_split_headers() {
2664 let (addr, sink) = spawn_sink().await;
2668
2669 let shared = SharedState::new(4);
2670 shared.cache_resolved_hostname(
2671 "example.com",
2672 ResolvedHostnameFamily::Ipv4,
2673 ["127.0.0.1".parse::<IpAddr>().unwrap()],
2674 StdDuration::from_secs(60),
2675 );
2676
2677 let secrets = SecretsConfig {
2678 secrets: vec![SecretEntry {
2679 env_var: "API_KEY".into(),
2680 value: zeroize::Zeroizing::new("real-secret-value".into()),
2681 source: None,
2682 placeholder: "$MSB_KEY".into(),
2683 allowed_hosts: vec![HostPattern::Exact("example.com".into())],
2684 substitution: SecretSubstitution {
2685 headers: true,
2686 query: false,
2687 body: false,
2688 },
2689 passthrough_hosts: vec![HostPattern::Exact("example.com".into())],
2690 violation_action: None,
2691 require_tls_identity: true,
2692 }],
2693 ..Default::default()
2694 };
2695
2696 let (from_tx, from_rx) = mpsc::channel::<Bytes>(8);
2697 let (to_tx, _to_rx) = mpsc::channel::<Bytes>(8);
2698 let proxy_connect = Arc::new(ProxyConnectState::new());
2699
2700 from_tx
2701 .send(Bytes::from_static(b"GET /api HTTP/1.1\r\n"))
2702 .await
2703 .unwrap();
2704 from_tx
2705 .send(Bytes::from_static(
2706 b"Host: example.com\r\nAuthorization: Bearer $MSB_KEY\r\n\r\n",
2707 ))
2708 .await
2709 .unwrap();
2710 drop(from_tx);
2711
2712 TcpProxy::new(
2713 addr,
2714 UpstreamTcpTarget::direct(addr),
2715 from_rx,
2716 to_tx,
2717 Arc::new(shared),
2718 Arc::new(NetworkPolicy::default()),
2719 Arc::new(secrets),
2720 None,
2721 false,
2722 proxy_connect,
2723 None,
2724 )
2725 .try_run()
2726 .await
2727 .unwrap();
2728
2729 let wire = String::from_utf8(sink.await.unwrap()).unwrap();
2730 assert!(
2731 wire.contains("Host: example.com"),
2732 "request must reach the allowed host, got: {wire:?}"
2733 );
2734 assert!(
2735 wire.contains("$MSB_KEY"),
2736 "placeholder must be forwarded unchanged for a require_tls_identity secret, got: {wire:?}"
2737 );
2738 assert!(
2739 !wire.contains("real-secret-value"),
2740 "secret must never be substituted over plain HTTP, got: {wire:?}"
2741 );
2742 }
2743
2744 #[tokio::test]
2745 async fn plain_http_substitutes_placeholder_in_first_flight() {
2746 let (addr, sink) = spawn_sink().await;
2747
2748 let request =
2749 b"GET /api HTTP/1.1\r\nHost: example.com\r\nAuthorization: Bearer $MSB_KEY\r\n\r\n"
2750 .to_vec();
2751 let secrets = make_plain_http_secret("$MSB_KEY", "real-secret-value", false);
2752
2753 let wire =
2754 String::from_utf8(relay_through_proxy(request, secrets, sink, addr).await).unwrap();
2755 assert!(
2756 wire.contains("real-secret-value"),
2757 "real value must reach server, got: {wire:?}"
2758 );
2759 assert!(
2760 !wire.contains("$MSB_KEY"),
2761 "placeholder must not reach server, got: {wire:?}"
2762 );
2763 }
2764
2765 #[tokio::test]
2766 async fn plain_http_no_substitution_when_require_tls_identity_true() {
2767 let (addr, sink) = spawn_sink().await;
2768
2769 let request =
2770 b"GET /api HTTP/1.1\r\nHost: example.com\r\nAuthorization: Bearer $MSB_KEY\r\n\r\n"
2771 .to_vec();
2772 let mut secrets = make_plain_http_secret("$MSB_KEY", "real-secret-value", true);
2773 secrets.secrets[0].passthrough_hosts = vec![HostPattern::Any];
2774
2775 let wire =
2776 String::from_utf8_lossy(&relay_through_proxy(request, secrets, sink, addr).await)
2777 .into_owned();
2778 assert!(
2779 wire.contains("$MSB_KEY"),
2780 "placeholder must be forwarded unchanged when require_tls_identity=true, got: {wire:?}"
2781 );
2782 assert!(
2783 !wire.contains("real-secret-value"),
2784 "real value must not leak when require_tls_identity=true, got: {wire:?}"
2785 );
2786 }
2787
2788 #[tokio::test]
2789 async fn plain_http_large_body_forwarded_verbatim_in_relay_loop() {
2790 let (addr, sink) = spawn_sink().await;
2794 let secrets = make_plain_http_secret("$MSB_KEY", "real-value", false);
2795
2796 let body = "x".repeat(32_000);
2797 let header = format!(
2798 "POST /upload HTTP/1.1\r\nHost: example.com\r\nAuthorization: Bearer $MSB_KEY\r\nContent-Length: {}\r\n\r\n",
2799 body.len()
2800 );
2801
2802 let (from_tx, from_rx) = mpsc::channel::<Bytes>(8);
2803 let (to_tx, _to_rx) = mpsc::channel::<Bytes>(8);
2804 let proxy_connect = Arc::new(ProxyConnectState::new());
2805
2806 from_tx
2807 .send(Bytes::from(header.into_bytes()))
2808 .await
2809 .unwrap();
2810 from_tx
2811 .send(Bytes::from(body.clone().into_bytes()))
2812 .await
2813 .unwrap();
2814 drop(from_tx);
2815
2816 TcpProxy::new(
2817 addr,
2818 UpstreamTcpTarget::direct(addr),
2819 from_rx,
2820 to_tx,
2821 Arc::new(SharedState::new(4)),
2822 Arc::new(NetworkPolicy::default()),
2823 Arc::new(secrets),
2824 None,
2825 false,
2826 proxy_connect,
2827 None,
2828 )
2829 .try_run()
2830 .await
2831 .unwrap();
2832
2833 let wire = String::from_utf8_lossy(&sink.await.unwrap()).into_owned();
2834 assert!(wire.contains(&body), "got {} bytes", wire.len());
2835 assert!(!wire.contains("$MSB_KEY"), "got: {wire:?}");
2836 }
2837}