Skip to main content

microsandbox_network/engine/tcp/
proxy.rs

1//! Bidirectional TCP proxy: smoltcp socket ↔ channels ↔ tokio socket.
2//!
3//! Each outbound guest TCP connection gets a proxy task that opens a real
4//! TCP connection to the destination via tokio and relays data between the
5//! channel pair (connected to the smoltcp socket in the poll loop) and the
6//! real server.
7
8use 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
36//--------------------------------------------------------------------------------------------------
37// Constants
38//--------------------------------------------------------------------------------------------------
39
40/// Buffer size for reading from the real server.
41const SERVER_READ_BUF_SIZE: usize = 16384;
42
43/// Max bytes buffered while reading the proxy's CONNECT response headers.
44const CONNECT_RESP_LIMIT: usize = 8192;
45
46/// Max bytes to buffer while peeking for the ClientHello's SNI.
47pub(crate) const PEEK_BUF_SIZE: usize = 16384;
48
49/// Upper bound on time spent buffering the first flight before
50/// falling back to a cache-only egress decision.
51pub(crate) const PEEK_BUDGET: Duration = Duration::from_secs(5);
52
53//--------------------------------------------------------------------------------------------------
54// Types
55//--------------------------------------------------------------------------------------------------
56
57#[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
71/// Per-connection TCP proxy task and the state it owns.
72pub(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
86//--------------------------------------------------------------------------------------------------
87// Methods
88//--------------------------------------------------------------------------------------------------
89
90impl 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    /// Build a proxy for a newly established guest TCP connection.
137    #[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    /// Run the TCP proxy task to completion.
167    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    /// Drive the TCP proxy to completion, returning operational failures.
182    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        // Mirror the SYN-time policy walk so only flows that actually reached a
198        // hostname rule wait for client bytes. A domain rule elsewhere in the
199        // policy must not stall unrelated or server-first traffic.
200        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            // Preserve the existing fail-closed path if a DNS-cache binding
209            // expires between the SYN evaluation and proxy startup.
210            EgressEvaluation::Deny => network_policy.has_domain_rules(),
211        };
212
213        // Pre-connect peek is only for domain policy: the hostname has to be known
214        // before we dial upstream so a Deny never opens a connection. Secrets do
215        // *not* gate the connect, so they no longer force a peek here — that work is
216        // deferred to `classify_first_flight` after the socket is open, where it can
217        // run without stalling server-first protocols (see below).
218        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        // Re-evaluate egress against the *guest* dst — the address the
226        // guest dialed, not the post-rewrite host-side address. SNI
227        // refines over-allow when the cache matched a shared CDN IP;
228        // CacheOnly is the non-TLS fallback path so Domain rules still
229        // gate plain HTTP / SSH / etc.
230        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        // A policy-required peek may already have captured a CONNECT request.
310        // Otherwise the post-connect paths below classify it without delaying
311        // server-first protocols.
312        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        // Connect upstream *before* finishing the secrets-side classification. A
334        // server-first protocol (SSH, SMTP, a database) sends nothing until it has
335        // seen the server's banner; with the socket already open we can relay that
336        // banner while we wait, instead of burning the peek budget pre-connect.
337        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        // Finish classifying the first flight (TLS vs plain HTTP) and, for
344        // plain-HTTP candidates, gather a full header block — without blocking the
345        // server→guest direction. When domain rules already peeked, `initial_buf`
346        // is reused and this is cheap; with no secrets it is skipped entirely
347        // (`is_tls` only matters for deciding whether to build the handler).
348        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            // A policy-required peek can miss a client whose first bytes arrive
372            // after we dial upstream. Once classify_first_flight has captured the
373            // request, rejoin the already-open proxy socket and use the CONNECT path
374            // so intercepted tunnels still get TLS substitution and policy checks.
375            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        // Replay the buffered first flight — run through secrets handler first.
419        if !initial_buf.is_empty() {
420            let out: Cow<[u8]> = match secrets_handler.as_mut() {
421                Some(h) => match h.substitute(&initial_buf) {
422                    // Borrow the input when nothing was substituted; only a chunk
423                    // that actually carries a placeholder is reallocated.
424                    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        // Bidirectional relay using tokio::select!.
449        //
450        // guest → server: receive from channel, write to server socket.
451        // server → guest: read from server socket, send via channel + wake poll.
452        let mut guest_eof = false;
453        loop {
454            tokio::select! {
455                // Guest → server: substitute placeholders before forwarding.
456                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                                // The first guest bytes can arrive after both peek
463                                // windows have completed. Nothing has been written
464                                // to the proxy socket yet, so this is still a valid
465                                // point to switch into CONNECT tunnel handling.
466                                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                            // No handler (no secrets / TLS) is the common path: forward
486                            // the chunk borrowed, with no per-chunk allocation or copy.
487                            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                        // Channel closed — the guest half-closed (FIN) or the
512                        // connection was torn down. Propagate the half-close:
513                        // stop sending upstream but keep relaying server →
514                        // guest until the server closes.
515                        None => {
516                            guest_eof = true;
517                            if server_tx.shutdown().await.is_err() {
518                                break;
519                            }
520                        }
521                    }
522                }
523
524                // Server → guest: no substitution — server never sends placeholders.
525                result = server_rx.read(&mut server_buf) => {
526                    match result {
527                        Ok(0) => break, // Server closed connection.
528                        Ok(n) => {
529                            // A server-first byte means this is not an HTTP CONNECT
530                            // tunnel to a proxy. Keep relaying normally afterward.
531                            late_connect_state = None;
532                            let data = Bytes::copy_from_slice(&server_buf[..n]);
533                            if to_smoltcp.send(data).await.is_err() {
534                                // Channel closed — poll loop dropped the receiver.
535                                break;
536                            }
537                            // Wake the poll thread so it writes data to the
538                            // smoltcp socket.
539                            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//--------------------------------------------------------------------------------------------------
555// Functions
556//--------------------------------------------------------------------------------------------------
557
558/// Spawn a TCP proxy task for a newly established connection.
559///
560/// `guest_dst` is what the guest dialed — the address policy rules match
561/// against. `connect_dst` is the host-side address tokio actually dials.
562///
563/// `proxy_connect` is updated before the task exits so the connection
564/// tracker can decide between FIN (clean close) and RST (upstream
565/// connect failure).
566#[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/// Forward an HTTP CONNECT tunnel: dial the proxy, splice the handshake,
622/// then hand the established stream to [`TlsProxy`] for TLS MITM.
623///
624/// `guest_dst` is what the guest dialed; `proxy_target` contains the rewritten
625/// loopback address the gateway actually connects to and its optional fallback.
626#[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    // Dial the proxy and forward the CONNECT request so it opens the tunnel.
656    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        // Unused: `upstream_stream` is already `Some` below, so the
760        // outbound proxy (already applied when dialing `proxy_stream`
761        // above) is never consulted again.
762        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
771/// Relay an established TCP stream without inspecting or substituting bytes.
772async 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                    // Guest half-closed (FIN): stop sending upstream but
791                    // keep relaying server → guest until the server closes.
792                    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
914/// Returns the byte offset just past the `\r\n\r\n` header terminator, or `None`.
915fn 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
1044/// Close a denied TCP connection, answering HTTP clients with 403.
1045///
1046/// TLS first-flights stay silent: injecting plaintext HTTP into a TLS
1047/// stream is worse than a reset, and intercepted HTTPS is handled by
1048/// [`crate::engine::tls::proxy`].
1049pub(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    // Reply only once a complete HTTP/1.x request line identifies the protocol.
1058    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
1086/// Extract the `Host:` header value from an already-buffered HTTP header block.
1087///
1088/// Returns `None` if:
1089/// - The first byte is `0x16` (TLS — not HTTP)
1090/// - The buffer does not yet contain `\r\n\r\n` (headers incomplete)
1091/// - No `Host:` header is present
1092///
1093/// Strips port suffix, lowercases, and trims whitespace. Result is
1094/// ready for byte-equal matching against `SecretEntry::allowed_hosts`.
1095fn extract_http_host(buf: &[u8]) -> Option<String> {
1096    if buf.first() == Some(&0x16) {
1097        return None;
1098    }
1099    // Size the header pool to the buffer rather than a fixed array: a header
1100    // line is at least four bytes (`a:\r\n`), so `len / 4` always covers the
1101    // real header count, and `httparse` never reports `TooManyHeaders` (which
1102    // would make a request with many headers look hostless). The first flight
1103    // is capped at PEEK_BUF_SIZE, so this stays bounded.
1104    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            // Strip port suffix.
1114            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/// Finish classifying the guest's first flight after the upstream socket is
1123/// open, returning the (possibly extended) first-flight buffer and whether it
1124/// is a TLS record.
1125///
1126/// `buf` carries whatever a pre-connect domain-rule peek already captured; when
1127/// it is non-empty the TLS/plain decision is already settled and only header
1128/// top-up runs. `want_headers` is set when at least one secret can be
1129/// substituted over plain HTTP (`SecretsConfig::has_plain_http_candidates`); it
1130/// makes the peek keep reading a non-TLS flight until `\r\n\r\n` so
1131/// [`extract_http_host`] sees a complete header block.
1132///
1133/// Crucially, this relays server→guest while it waits. Server-first protocols
1134/// (SSH, SMTP, databases) send nothing until they have seen the server's
1135/// banner; draining the server side here lets the banner reach the guest
1136/// immediately, so the guest's eventual first flight — not a 5s timeout — is
1137/// what ends the peek.
1138#[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        // Stop as soon as the protocol class is known and — for plain-HTTP
1155        // candidates — a full header block has arrived. Bail the moment a
1156        // non-TLS flight stops looking like an HTTP request so non-HTTP
1157        // protocols (SSH, Postgres) aren't withheld from upstream for the
1158        // whole budget while we wait for a `\r\n\r\n` that never comes.
1159        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 → buffer (not forwarded here; the caller replays it once the
1180            // handler is built, so substitution applies to the first flight too).
1181            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 → guest: relay immediately so a server-first banner is never
1189            // held hostage by the peek.
1190            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
1209/// Buffer a denied plaintext first flight through its HTTP headers, or until
1210/// the shared peek budget/cap is exhausted.
1211///
1212/// Unlike [`peek_for_sni`], this does not return on the first non-TLS chunk:
1213/// a request method may be split across chunks (`GE` then `T / ...`). It stops
1214/// once the headers are complete so the denial can include the Host header,
1215/// or immediately for a conclusively non-HTTP prefix. No upstream connection
1216/// exists on this path.
1217pub(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
1256/// Buffer the first flight until SNI can be extracted, or until one
1257/// of the bail-out conditions hits (channel close, buffer cap,
1258/// timeout). Never errors; non-TLS / slow / malformed input all
1259/// fall through to `None`.
1260///
1261/// On hit, the SNI is canonicalized (lowercase + trim trailing dot)
1262/// for byte-equal matching against rule destinations. The returned
1263/// buffer must be replayed verbatim to upstream before the caller
1264/// starts its relay loop.
1265pub(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                        // First byte of a TLS record is the ContentType;
1283                        // 0x16 is handshake. Anything else can't be a
1284                        // ClientHello, so don't burn the full budget on
1285                        // plain HTTP / SSH / etc.
1286                        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//--------------------------------------------------------------------------------------------------
1307// Tests
1308//--------------------------------------------------------------------------------------------------
1309
1310#[cfg(test)]
1311mod tests {
1312    use super::*;
1313
1314    /// Synthetic TLS ClientHello carrying SNI `example.com`. Bytes
1315    /// borrowed from `tls::sni` test fixtures so the parser sees a
1316    /// well-formed record.
1317    fn synthetic_client_hello(sni: &str) -> Vec<u8> {
1318        // Minimal but valid TLS 1.2 ClientHello with one SNI entry.
1319        // Layout: record header (5) + handshake header (4) + body.
1320        let host_bytes = sni.as_bytes();
1321        let host_len = host_bytes.len() as u16;
1322        let server_name_list_len = 3 + host_len; // type(1) + len(2) + host
1323        let extension_data_len = 2 + server_name_list_len; // list-len(2) + list
1324        let extensions_total = 4 + extension_data_len; // type(2) + len(2) + data
1325
1326        let mut body = Vec::new();
1327        // Client version
1328        body.extend_from_slice(&[0x03, 0x03]);
1329        // Random (32 bytes)
1330        body.extend_from_slice(&[0u8; 32]);
1331        // Session id length + (empty)
1332        body.push(0);
1333        // Cipher suites length + one cipher
1334        body.extend_from_slice(&[0x00, 0x02, 0x00, 0x2f]);
1335        // Compression methods length + null
1336        body.extend_from_slice(&[0x01, 0x00]);
1337        // Extensions length
1338        body.extend_from_slice(&extensions_total.to_be_bytes());
1339        // SNI extension: type 0x0000
1340        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); // host_name type
1344        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); // ClientHello
1350        hs.extend_from_slice(&handshake_len.to_be_bytes()[1..]); // 24-bit length
1351        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]); // Handshake, TLS 1.0
1356        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        // This is the HTTP proxy the guest originally dialed. It is never
1395        // contacted directly; the SOCKS5 request below must carry this address.
1396        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        // HTTP/2 must never receive an HTTP/1.1 response.
1493        assert!(!first_flight_is_http(b"PRI * HTTP/2.0\r\n\r\nSM\r\n\r\n"));
1494        // A complete line is judged by its version, so custom methods pass.
1495        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        // Split before its first CRLF, an SSH banner or SMTP greeting is a
1507        // valid ASCII token but no HTTP method.
1508        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); // close so peek returns even if SNI didn't satisfy
1704
1705        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        // Plaintext HTTP request; not a TLS record so extract_sni returns None.
1723        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        // Hold the sender open but send nothing — peek must time out.
1741        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        // First byte 0x16 keeps the peek collecting past the early
1751        // non-TLS bail. Padding bytes are zero so the SNI parser never
1752        // matches and the loop drives to the size cap.
1753        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        // Plain HTTP request: first byte 'G' (0x47) — clearly not TLS.
1773        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        // 5-second nominal budget; assert we returned in well under
1779        // that — the early-bail must not wait for the full window.
1780        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    //----------------------------------------------------------------------------------------------
1792    // peek_for_sni × evaluate_egress_with_source — combined integration tests
1793    //----------------------------------------------------------------------------------------------
1794
1795    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    /// Over-allow case: cache says IP X is `pypi.org` (allowed); SNI
1835    /// is `evil.com`. SNI must override the cache and deny.
1836    #[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    /// Over-block case: cache says IP X is `ads.example.com` (denied);
1869    /// SNI is `api.example.com`. SNI must override the cache and allow.
1870    #[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    /// Non-TLS first-flight falls back to `CacheOnly`; the cache
1904    /// match decides.
1905    #[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        // Plain HTTP request; not a TLS record.
1917        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    /// SNI matches a `DomainSuffix` rule with a cache binding for the
1944    /// claimed name. Genuine pre-resolved traffic passes.
1945    #[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    /// Spoofed SNI on an IP with no cache binding for any matching
1979    /// name: byte-equality with the suffix passes, but no DNS lookup
1980    /// ever tied a `*.pythonhosted.org` name to the destination, so
1981    /// the AND-check fails and the connection is denied.
1982    #[tokio::test]
1983    async fn integration_sni_denies_domain_suffix_without_cache_binding() {
1984        let shared = SharedState::new(4); // empty cache
1985        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    // ── extract_http_host ──────────────────────────────────────────────────────
2016
2017    #[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        // Far more headers than a small fixed parse array would hold: the Host
2065        // must still be found rather than the request looking hostless.
2066        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    // ── plain-HTTP secret substitution ────────────────────────────────────────
2075
2076    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    /// Relay `chunks` as separate guest reads through a proxy with one domain
2420    /// rule and no secrets. The guest side stays open, so the proxy must close
2421    /// the connection by itself. Returns the bytes upstream received and
2422    /// whether sandbox termination was requested.
2423    ///
2424    /// Both the proxy and the upstream sink are bounded. A regression that
2425    /// refuses the guest *before* dialing leaks the proxy future but never
2426    /// connects the sink, so the sink wait has its own deadline and aborts the
2427    /// task instead of hanging `cargo test` forever.
2428    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    /// A valid h2c first request: preface, empty SETTINGS, then `:method GET`,
2480    /// `:scheme http`, `:path /` and an incrementally indexed `:authority`
2481    /// (index 62).
2482    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    /// The wire bytes for [`h2c_first_request`]: the proxy re-encodes every
2497    /// field as a never-indexed raw literal.
2498    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    /// One HTTP/2 frame.
2523    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        // Prior-knowledge HTTP/2 whose only HEADERS block is `ff`: a truncated
2541        // indexed-field integer. A domain rule alone installs the secrets
2542        // handler here, with no secrets configured.
2543        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        // The whole first flight is rejected, so not even the preface is sent.
2554        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        // `82 86 84` = GET, http, `/`. `41 0f ...` adds `:authority` to the
2564        // guest's dynamic table, and `be` refers to it. The second block ends
2565        // with a truncated integer (`ff`), so the structural check rejects it
2566        // before the decoder runs; the third block is never reached.
2567        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        // The first request is re-encoded as never-indexed raw literals;
2576        // nothing from stream 3 or 5 is forwarded.
2577        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        // The failing block is structurally complete, so the structural check
2584        // accepts it and the decoder runs. `82 86 84` = GET, http, `/`; `be` is
2585        // the authority inserted at index 62 by the first block. `40 03 78 2d
2586        // 61 01 31` is an incremental literal that inserts `x-a: 1` at index 62
2587        // (moving the authority to 63), and the final `80` is an indexed field
2588        // with the invalid index 0. `httlib-hpack` has already applied the
2589        // insertion when it reports the error. The third request would now
2590        // reference the authority through index `bf` (63), but the connection
2591        // must be closed first, so nothing leaks upstream and the sandbox is
2592        // not terminated.
2593        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        // The block `7f c5` (truncated integer) is split across HEADERS and
2615        // CONTINUATION, and every frame is split across guest reads.
2616        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        // `BINARY3` is deliberately a valid HTTP token followed by a space,
2822        // but the control bytes make this an invalid HTTP request line.
2823        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        // Host header split across TCP segments — classify_first_flight must keep
2837        // reading until \r\n\r\n before extract_http_host is called.
2838        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        // A default (require_tls_identity = true) host-bound secret is never substituted over plain
2882        // HTTP. Explicit passthrough allows its placeholder to remain unchanged even when the Host
2883        // arrives in a later segment than the request line.
2884        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        // Body arrives in a separate segment after headers — flows through the relay
3009        // loop, not the peek path. Ensures no bytes are dropped and header substitution
3010        // still happens.
3011        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}