diff --git a/DECISIONS.md b/DECISIONS.md index 150e22a..cf17045 100644 --- a/DECISIONS.md +++ b/DECISIONS.md @@ -77,3 +77,4 @@ - 2026-02-28 m+git@andri.dk — `replayConn` wraps a `net.Conn` with a `bufio.Reader` to replay peeked ClientHello bytes during the TLS handshake. The `extractSNI` function peeks at the TLS record without consuming bytes; `replayConn.Read` reads from the buffered reader first, then from the underlying connection. This avoids copying or buffering the entire ClientHello separately. - 2026-02-28 m+git@andri.dk — Transparent HTTPS portal handler now passes `proxy.api` to `portalHandler`. Without it, any `/api/*` request on the transparent-mode portal would nil-pointer panic because the `api` field was zero-valued. - 2026-02-28 m+git@andri.dk — Transparent HTTP forwarding now handles WebSocket upgrades. `forwardHTTPUpgrade` mirrors `proxyHandler.handleHTTPUpgrade` — re-adds hop-by-hop upgrade headers, hijacks both sides on 101, and does bidirectional copy. Without this, WebSocket connections through the transparent HTTP proxy would fail with a `RoundTrip` error. +- 2026-02-28 m+git@andri.dk — Fixed `singleConnListener` race condition. The old implementation returned an error immediately from the second `Accept`, causing `http.Server.Serve` to return before the in-flight request handler finished writing the response. The fix wraps the connection in `notifyCloseConn` which signals a channel on close; the second `Accept` blocks on that channel so `Serve` doesn't exit prematurely. Also added `IdleTimeout: 5s` to the portal's `http.Server` so connections don't block forever after the last response. diff --git a/transparent.go b/transparent.go index 7e44bc1..8336efd 100644 --- a/transparent.go +++ b/transparent.go @@ -509,9 +509,15 @@ func handleTransparentPortalTLS(conn net.Conn, br *bufio.Reader, proxy *proxyHan conn.SetDeadline(time.Time{}) - // Serve portal pages over the TLS connection + // Serve portal pages over the TLS connection. + // singleConnListener blocks the second Accept until the wrapped + // connection is closed (client disconnect or idle timeout), so + // Serve doesn't return before the response is fully written. portalH := &portalHandler{proxy: proxy, api: proxy.api} - server := http.Server{Handler: portalH} + server := http.Server{ + Handler: portalH, + IdleTimeout: 5 * time.Second, + } serverConn := &singleConnListener{conn: tlsConn} server.Serve(serverConn) } @@ -581,26 +587,44 @@ func (c *replayConn) Read(p []byte) (int, error) { } // singleConnListener is a net.Listener that serves exactly one connection. -// Used to serve HTTP over an already-established TLS connection. +// The first Accept returns a tracked wrapper; the second Accept blocks +// until that connection is closed, then returns an error so Serve exits +// cleanly after the request is fully handled. type singleConnListener struct { conn net.Conn once sync.Once + done chan struct{} } func (l *singleConnListener) Accept() (net.Conn, error) { - var conn net.Conn + var first bool l.once.Do(func() { - conn = l.conn + first = true + l.done = make(chan struct{}) }) - if conn != nil { - return conn, nil + if first { + return ¬ifyCloseConn{Conn: l.conn, done: l.done}, nil } + <-l.done return nil, errors.New("listener closed") } func (l *singleConnListener) Close() error { return nil } func (l *singleConnListener) Addr() net.Addr { return l.conn.LocalAddr() } +// notifyCloseConn wraps a net.Conn and signals a channel when closed. +type notifyCloseConn struct { + net.Conn + done chan struct{} + closeOnce sync.Once +} + +func (c *notifyCloseConn) Close() error { + err := c.Conn.Close() + c.closeOnce.Do(func() { close(c.done) }) + return err +} + // startTransparentHTTP starts the HTTP server for transparent proxy mode. // It intercepts plain HTTP traffic, serves captive portal for untrusted // clients, and forwards requests upstream for trusted clients. diff --git a/transparent_test.go b/transparent_test.go index 43c44ca..84cbead 100644 --- a/transparent_test.go +++ b/transparent_test.go @@ -554,6 +554,248 @@ func TestTransparentHTTPPortalAccess(t *testing.T) { } } +// --- Transparent HTTPS portal access --- + +func TestTransparentHTTPSPortalAccess(t *testing.T) { + env := startTransparentTestEnv(t, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + t.Error("upstream should not be reached for portal access") + }), nil) + env.serveTransparentHTTPS(t) + + // Connect with SNI matching the portal hostname — the proxy should + // serve the management portal instead of proxying to upstream. + conn, err := net.DialTimeout("tcp", env.proxyListener.Addr().String(), 2*time.Second) + if err != nil { + t.Fatalf("dial proxy: %v", err) + } + defer conn.Close() + + tlsConn := tls.Client(conn, &tls.Config{ + ServerName: env.portalHost, + RootCAs: env.caPool, + }) + if err := tlsConn.Handshake(); err != nil { + t.Fatalf("TLS handshake: %v", err) + } + defer tlsConn.Close() + + // Request the portal index page + req, _ := http.NewRequest("GET", "https://"+env.portalHost+"/", nil) + if err := req.Write(tlsConn); err != nil { + t.Fatalf("write request: %v", err) + } + + resp, err := http.ReadResponse(bufio.NewReader(tlsConn), req) + if err != nil { + t.Fatalf("read response: %v", err) + } + defer resp.Body.Close() + + if resp.StatusCode != http.StatusOK { + t.Errorf("status = %d, want %d", resp.StatusCode, http.StatusOK) + } + body, _ := io.ReadAll(resp.Body) + if !strings.Contains(string(body), "ublproxy") { + t.Errorf("portal page should contain 'ublproxy', got %q", string(body[:min(len(body), 200)])) + } +} + +// --- URL-level blocking in transparent mode --- + +func TestTransparentHTTPSBlocksByURLPattern(t *testing.T) { + var requestPaths []string + + rs := blocklist.NewRuleSet() + rs.AddRule("/ads/tracking.js") + + env := startTransparentTestEnv(t, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + requestPaths = append(requestPaths, r.URL.Path) + w.WriteHeader(http.StatusOK) + w.Write([]byte("upstream response")) + }), rs) + + sniHost := "urlblock.example.com" + upstreamAddr := env.upstream.Listener.Addr().String() + _, upstreamPort, _ := net.SplitHostPort(upstreamAddr) + + env.proxy.transport.DialContext = func(_ context.Context, network, addr string) (net.Conn, error) { + h, _, _ := net.SplitHostPort(addr) + if h == sniHost { + return net.DialTimeout(network, upstreamAddr, 5*time.Second) + } + return net.DialTimeout(network, addr, 5*time.Second) + } + + env.serveTransparentHTTPS(t) + + conn, err := net.DialTimeout("tcp", env.proxyListener.Addr().String(), 2*time.Second) + if err != nil { + t.Fatalf("dial proxy: %v", err) + } + defer conn.Close() + + tlsConn := tls.Client(conn, &tls.Config{ + ServerName: sniHost, + RootCAs: env.caPool, + }) + defer tlsConn.Close() + + hostPort := net.JoinHostPort(sniHost, upstreamPort) + br := bufio.NewReader(tlsConn) + + // Non-matching path should pass through + req1, _ := http.NewRequest("GET", fmt.Sprintf("https://%s/page.html", hostPort), nil) + req1.Host = hostPort + if err := req1.Write(tlsConn); err != nil { + t.Fatalf("write request 1: %v", err) + } + resp1, err := http.ReadResponse(br, req1) + if err != nil { + t.Fatalf("read response 1: %v", err) + } + io.Copy(io.Discard, resp1.Body) + resp1.Body.Close() + + if resp1.StatusCode != http.StatusOK { + t.Errorf("page.html: status = %d, want %d", resp1.StatusCode, http.StatusOK) + } + + // Matching URL pattern should be blocked inside the MITM tunnel + req2, _ := http.NewRequest("GET", fmt.Sprintf("https://%s/ads/tracking.js", hostPort), nil) + req2.Host = hostPort + if err := req2.Write(tlsConn); err != nil { + t.Fatalf("write request 2: %v", err) + } + resp2, err := http.ReadResponse(br, req2) + if err != nil { + t.Fatalf("read response 2: %v", err) + } + io.Copy(io.Discard, resp2.Body) + resp2.Body.Close() + + if resp2.StatusCode != http.StatusNoContent { + t.Errorf("tracking.js: status = %d, want %d", resp2.StatusCode, http.StatusNoContent) + } + + // Only the allowed request should have reached upstream + if len(requestPaths) != 1 || requestPaths[0] != "/page.html" { + t.Errorf("upstream received %v, want [/page.html]", requestPaths) + } +} + +func TestTransparentHTTPBlocksByURLPattern(t *testing.T) { + var requestPaths []string + + upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + requestPaths = append(requestPaths, r.URL.Path) + w.WriteHeader(http.StatusOK) + w.Write([]byte("upstream response")) + })) + defer upstream.Close() + + rs := blocklist.NewRuleSet() + rs.AddRule("/ads/tracking.js") + + env := startTransparentHTTPTestEnv(t) + env.proxy.baselineRules.Store(rs) + + upstreamHost := strings.TrimPrefix(upstream.URL, "http://") + + // Non-matching path should pass through + req1, _ := http.NewRequest("GET", env.server.URL+"/page.html", nil) + req1.Host = upstreamHost + + resp1, err := http.DefaultClient.Do(req1) + if err != nil { + t.Fatalf("GET /page.html: %v", err) + } + io.Copy(io.Discard, resp1.Body) + resp1.Body.Close() + + if resp1.StatusCode != http.StatusOK { + t.Errorf("page.html: status = %d, want %d", resp1.StatusCode, http.StatusOK) + } + + // Matching URL pattern should be blocked + req2, _ := http.NewRequest("GET", env.server.URL+"/ads/tracking.js", nil) + req2.Host = upstreamHost + + resp2, err := http.DefaultClient.Do(req2) + if err != nil { + t.Fatalf("GET /ads/tracking.js: %v", err) + } + io.Copy(io.Discard, resp2.Body) + resp2.Body.Close() + + if resp2.StatusCode != http.StatusNoContent { + t.Errorf("tracking.js: status = %d, want %d", resp2.StatusCode, http.StatusNoContent) + } + + // Only the allowed request should have reached upstream + if len(requestPaths) != 1 || requestPaths[0] != "/page.html" { + t.Errorf("upstream received %v, want [/page.html]", requestPaths) + } +} + +// --- WebSocket in transparent HTTP mode --- + +func TestTransparentHTTPWebSocketUpgrade(t *testing.T) { + upstream := httptest.NewServer(http.HandlerFunc(wsEchoHandler)) + defer upstream.Close() + + env := startTransparentHTTPTestEnv(t) + upstreamHost := strings.TrimPrefix(upstream.URL, "http://") + + // Connect raw TCP to the transparent HTTP test server and perform + // a WebSocket upgrade with Host header pointing at the upstream. + serverHost := strings.TrimPrefix(env.server.URL, "http://") + conn, err := net.DialTimeout("tcp", serverHost, 2*time.Second) + if err != nil { + t.Fatalf("dial: %v", err) + } + defer conn.Close() + conn.SetDeadline(time.Now().Add(5 * time.Second)) + + // Send WebSocket upgrade request + wsKey := "dGhlIHNhbXBsZSBub25jZQ==" + reqStr := fmt.Sprintf( + "GET /ws HTTP/1.1\r\n"+ + "Host: %s\r\n"+ + "Upgrade: websocket\r\n"+ + "Connection: Upgrade\r\n"+ + "Sec-WebSocket-Key: %s\r\n"+ + "Sec-WebSocket-Version: 13\r\n"+ + "\r\n", + upstreamHost, wsKey) + if _, err := conn.Write([]byte(reqStr)); err != nil { + t.Fatalf("write upgrade: %v", err) + } + + // Read the 101 response + br := bufio.NewReader(conn) + resp, err := http.ReadResponse(br, nil) + if err != nil { + t.Fatalf("read response: %v", err) + } + if resp.StatusCode != http.StatusSwitchingProtocols { + t.Fatalf("status = %d, want %d", resp.StatusCode, http.StatusSwitchingProtocols) + } + + // Send a masked WebSocket text frame and read the echo back + payload := []byte("hello transparent websocket") + if err := writeMaskedWSFrame(conn, payload); err != nil { + t.Fatalf("write frame: %v", err) + } + + frame, err := readWSFrame(br) + if err != nil { + t.Fatalf("read frame: %v", err) + } + if string(frame.payload) != string(payload) { + t.Errorf("frame payload = %q, want %q", frame.payload, payload) + } +} + // --- CA trust tracker tests --- func TestCATrustTracker(t *testing.T) {