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