diff --git a/README.md b/README.md index 5d4cba8..63e477e 100644 --- a/README.md +++ b/README.md @@ -21,6 +21,10 @@ AutoCAR 在本地提供 SOCKS5、HTTP 和 HTTPS Proxy,在远端解析并连接 HTTP/3/UDP,UDP 不可用时让新 TCP 流继续走 HTTPS/HTTP/2/TCP;SOCKS5 UDP 使用 H3 RFC 9298 CONNECT-UDP,不跨入 H2 fallback。 +源代码版本还支持固定上游网站的 HTTP/1.1 WebSocket,并在服务端关闭时回收 +升级连接;它仍是网站流量,不是新的代理隧道。限制见 +[WebSocket 网站兼容性](docs/WEB_COVER.md#website-websocket-support-in-source-builds)。 + v1.0.1 是功能增强与问题修复版本。默认仍为 `native` 服务端与 `auto` 客户端; **Web-cover 是需要显式开启的实验性功能**,发布不代表其被动抗识别能力已经验证。 变更、升级与限制见 [v1.0.1 发布说明](docs/releases/v1.0.1.md)。 diff --git a/docs/ARCHITECTURE.md b/docs/ARCHITECTURE.md index 922c86b..90a5354 100644 --- a/docs/ARCHITECTURE.md +++ b/docs/ARCHITECTURE.md @@ -99,12 +99,22 @@ connection selected after GOAWAY run the full bootstrap again. QUIC migration or NAT rebinding that remains the same `*quic.Conn` retains authentication state. The cover is either a local static directory or a reverse proxy to one fixed, -operator-authorized HTTP(S) origin. TCP cover responses advertise the bound H3 -service through `Alt-Svc`. Neither the cover nor an unauthenticated probe sees +operator-authorized HTTP(S) origin. Ordinary combined H1/H2/H3 cover responses +advertise the bound H3 service through `Alt-Svc`; raw upgraded 101 responses do +not carry that override guarantee. Neither cover nor an unauthenticated probe sees the `autocar/2` ALPN or native binary request header. This reduces active-probe exposure but does not prove browser-indistinguishable passive behavior; see [WEB_COVER.md](WEB_COVER.md). +Source builds additionally permit validated H1.1 WebSocket upgrades to that +same fixed website. Request-local handshake state prevents an optional +transport's response metadata from choosing eligibility; legal duplex bodies +retain optional half-close capability and close once on errors. A TCP-side +physical-connection owner outlives net/http's hijack bookkeeping, so server +shutdown cancels requests and closes upgraded raw sockets without cancelling +shared destination dialers or cover transports. See the WebSocket section of +[WEB_COVER.md](WEB_COVER.md#website-websocket-support-in-source-builds). + The public H1/H2 listener accepts TLS 1.2 and TLS 1.3 for ordinary website compatibility, while AutoCAR H2 clients and authenticated H2 tunnels require TLS 1.3. The handler rejects TLS 1.2 from the tunnel path before ticket diff --git a/docs/DEPLOYMENT.md b/docs/DEPLOYMENT.md index 469fe69..88f5ee2 100644 --- a/docs/DEPLOYMENT.md +++ b/docs/DEPLOYMENT.md @@ -148,6 +148,15 @@ headers, and preserves the request path and query. Treat that origin as an Internet-facing application; do not point it at metadata or control-plane services. +Source builds after v1.0.1 can forward narrowly validated H1.1 WebSocket GET +upgrades to this same fixed origin; no new flag or arbitrary Upgrade proxy is +introduced. Authorization and nominated hop fields are still stripped. +Printable, parseable malformed upgrades remain ordinary scrubbed website requests; +Go can reject invalid raw/header characters earlier. Invalid upstream +101 responses become generic 502. Shutdown aborts owned upgraded sockets. +See [WebSocket boundaries](WEB_COVER.md#website-websocket-support-in-source-builds) +for handshake, application-policy and close semantics. + `--listen` binds H3/UDP and `--tcp-listen` binds HTTPS/H1/H2. They must use the same numeric port; if `--tcp-listen` is omitted it inherits `--listen`. Open both TCP and UDP in the deployment firewall. `--disable-tcp-fallback` is invalid in diff --git a/docs/PROTOCOL.md b/docs/PROTOCOL.md index c466cb5..9a49a20 100644 --- a/docs/PROTOCOL.md +++ b/docs/PROTOCOL.md @@ -170,6 +170,15 @@ HTTP stream; there is no AutoCAR binary stream preface. H3 additionally recognizes an authenticated Extended CONNECT with `:protocol=connect-udp` as described below. +Source builds after v1.0.1 support the configured website's narrowly validated +H1.1 WebSocket GET/101 exchange. It remains fixed-origin cover traffic, never +an authenticated H1 tunnel, arbitrary Upgrade, h2c or H2/H3 WebSocket extended +CONNECT. Ordinary combined H3 cover responses also use the bound Alt-Svc +policy; authenticated writers and standalone H3 policy are unchanged. +See [website WebSocket support](WEB_COVER.md#website-websocket-support-in-source-builds) +for exact validation and abort/cleanup limits; raw upgraded 101 advertisement +is not covered by the ordinary-response override. + An authenticated request that exceeds stream admission receives generic HTTP `503`; an allowed request whose destination cannot be opened receives generic HTTP `502`. These responses are available only after a valid credential has diff --git a/docs/WEB_COVER.md b/docs/WEB_COVER.md index ebe198a..15c90c6 100644 --- a/docs/WEB_COVER.md +++ b/docs/WEB_COVER.md @@ -40,6 +40,55 @@ HTTP/2 CONNECT presented to the public origin is always treated as cover, including when it carries an otherwise valid ticket: the handler removes `Proxy-Authorization` before delegation and never dials its authority. +### Website WebSocket support in source builds + +Source builds after v1.0.1 also forward valid HTTP/1.1 WebSocket upgrades to +the configured `--cover-upstream` origin. This is website traffic, never an +AutoCAR tunnel or a requester-selected upstream. It works on the public TCP +listener with TLS 1.2 or 1.3; static cover and H2/H3 extended CONNECT behavior +are unchanged. An HTTPS website's ordinary H2 connection can remain reusable +while its WebSocket handshake uses H1. + +The initial upgrade allowlist is deliberately narrow: body-free H1.1 GET, +one WebSocket Upgrade value, one valid Connection token list containing +Upgrade, version 13, and one canonical base64 key decoding to 16 bytes. +Ambiguous/duplicate fields, other upgrade protocols, body/transfer coding, +Connection close, and nominations of required or negotiation handshake fields +are not upgraded. Parseable, printable malformed handshakes continue as +scrubbed ordinary requests to the same fixed website. Go's existing HTTP +parser or reverse proxy can reject invalid raw/header characters before this +rewrite policy runs; those inputs do not promise an origin request, and the +proxy's early error remains generic 502. Connection-nominated fields and +authorization headers are removed; only the validated +`Connection: Upgrade` / `Upgrade: websocket` +pair is restored. Origin, cookies and ordinary safe negotiation fields remain +end-to-end; the website is responsible for its own access and Origin policy. +An application requiring forwarded Authorization headers remains incompatible +with the cover's intentional credential-stripping policy. + +A final 101 must match the validated request and its key, contain the correct +accept value, and supply a duplex body. Unexpected, mismatched or non-duplex +101 responses produce the same generic 502 and close their upstream body. +This follows the [RFC 6455 opening-handshake mechanism](https://www.rfc-editor.org/rfc/rfc6455.html#section-4), +with the additional body-free and single-field restrictions described above. +After the handshake, bytes are relayed without interpreting application frames; +subprotocol and extension negotiation remain the website/client's policy. + +Ordinary-response scrubbing can only use nominations still exposed by the +upstream transport. Go's native response parser removes the entire Connection +field when it sees `close`, so additional nominations in that same field are +not available to this handler. This existing parser boundary is unchanged; +explicit authorization-header stripping does not depend on those nominations. + +Client disconnection releases the upgraded connection's admission slot. +Server Close or Serve-context cancellation aborts owned physical TCP sockets, +including hijacked upgrades, and cancels their request contexts. This is an +abort operation, not a graceful WebSocket close-frame exchange, and does not +close the shared destination dialer or upstream transport. Ordinary public +responses retain the bound Alt-Svc policy described above; an actual hijacked +101 is written by the reverse proxy's raw upgrade path and is not promised +the same bound-header override. + H2 preserves a client upload half-close while the destination's reply drains. When the destination itself reaches EOF, H2 finishes that CONNECT response and stops any remaining upload on that stream; the HTTP handler interface cannot diff --git a/internal/cover/handler.go b/internal/cover/handler.go index 8048ca8..a2d2a4b 100644 --- a/internal/cover/handler.go +++ b/internal/cover/handler.go @@ -80,11 +80,38 @@ func NewReverseProxyHandler(origin *url.URL, transport http.RoundTripper) (http. proxy := &httputil.ReverseProxy{ Transport: &informationalHeaderTransport{base: transport}, Rewrite: func(request *httputil.ProxyRequest) { + upgrade := websocketRequestEligibility(request.In) request.SetURL(target) request.Out.Host = target.Host + // ReverseProxy already removes nominated fields before Rewrite, + // then restores its generic Upgrade pair. Retain the original + // nomination boundary, and restore only our validated H1 WebSocket. + removeConnectionNominatedHeaders(request.Out.Header, request.In.Header) removeUnsafeHeaders(request.Out.Header) + if upgrade.eligible { + request.Out.Header.Set("Connection", "Upgrade") + request.Out.Header.Set("Upgrade", "websocket") + } + request.Out = withWebsocketRequestEligibility(request.Out, upgrade) }, ModifyResponse: func(response *http.Response) error { + if response.StatusCode == http.StatusSwitchingProtocols { + if response.Body == nil { + // ReverseProxy closes Body unconditionally on hook failure. + response.Body = http.NoBody + } + if err := validateWebsocketResponse(response); err != nil { + return err + } + if err := ownWebsocketResponse(response); err != nil { + return err + } + removeUnsafeHeaders(response.Header) + removeUnsafeHeaders(response.Trailer) + response.Header.Set("Connection", "Upgrade") + response.Header.Set("Upgrade", "websocket") + return nil // Preserve duplex I/O and optional CloseWrite. + } removeUnsafeHeaders(response.Header) removeUnsafeHeaders(response.Trailer) // An upgraded body is duplex, not an HTTP message with trailers. @@ -93,7 +120,8 @@ func NewReverseProxyHandler(origin *url.URL, transport http.RoundTripper) (http. } return nil }, - ErrorHandler: func(w http.ResponseWriter, _ *http.Request, _ error) { + ErrorHandler: func(w http.ResponseWriter, request *http.Request, _ error) { + closeWebsocketResponse(request) http.Error(w, http.StatusText(http.StatusBadGateway), http.StatusBadGateway) }, ErrorLog: log.New(io.Discard, "", 0), @@ -111,6 +139,9 @@ type informationalHeaderTransport struct { } func (t *informationalHeaderTransport) RoundTrip(request *http.Request) (*http.Response, error) { + // Capture the immutable Rewrite result before an optional custom transport + // sees the request. Its response.Request is not evidence of eligibility. + trusted := websocketResponseRequest(request) trace := &httptrace.ClientTrace{ Got1xxResponse: func(_ int, header textproto.MIMEHeader) error { removeUnsafeHeaders(http.Header(header)) @@ -118,7 +149,11 @@ func (t *informationalHeaderTransport) RoundTrip(request *http.Request) (*http.R }, } request = request.WithContext(httptrace.WithClientTrace(request.Context(), trace)) - return t.base.RoundTrip(request) + response, err := t.base.RoundTrip(request) + if response != nil && response.StatusCode == http.StatusSwitchingProtocols { + response.Request = trusted + } + return response, err } // responseTrailerBody filters fields that a transport discovers only at EOF @@ -187,15 +222,63 @@ func normalizeOrigin(origin *url.URL) (*url.URL, error) { } func removeUnsafeHeaders(header http.Header) { - for _, value := range header.Values("Connection") { + removeConnectionNominatedHeaders(header, header) + for _, name := range hopByHopHeaders { + deleteHeaderFold(header, name) + } + deleteHeaderFold(header, "Authorization") +} + +func removeConnectionNominatedHeaders(header, connectionSource http.Header) { + var nominations map[string]struct{} + var scratch [64]byte + folded := scratch[:0] + for _, value := range headerValuesFold(connectionSource, "Connection") { for token := range strings.SplitSeq(value, ",") { - if name := strings.TrimSpace(token); name != "" { - header.Del(name) + name := strings.TrimSpace(token) + if !httpToken(name) { + continue + } + folded = foldASCIIHeaderName(folded, name) + if _, exists := nominations[string(folded)]; !exists { + if nominations == nil { + nominations = make(map[string]struct{}) + } + // Only a new nomination owns a copied key. Repeated tokens + // reuse scratch; they never rescan the destination header. + nominations[string(folded)] = struct{}{} } } } - for _, name := range hopByHopHeaders { - header.Del(name) + if len(nominations) == 0 { + return + } + // Collect first: header and connectionSource may be the same map, and + // Connection itself may be nominated without hiding later nominations. + for field := range header { + if !httpToken(field) { + continue + } + folded = foldASCIIHeaderName(folded, field) + if _, nominated := nominations[string(folded)]; nominated { + delete(header, field) + } + } +} + +// Call only for validated ASCII HTTP tokens. The scratch buffer is local to +// one filtering call; map lookups need no separately retained folded string. +func foldASCIIHeaderName(buffer []byte, name string) []byte { + if cap(buffer) < len(name) { + buffer = make([]byte, len(name)) + } + buffer = buffer[:len(name)] + for index := range name { + char := name[index] + if char >= 'A' && char <= 'Z' { + char += 'a' - 'A' + } + buffer[index] = char } - header.Del("Authorization") + return buffer } diff --git a/internal/cover/websocket.go b/internal/cover/websocket.go new file mode 100644 index 0000000..3a5b443 --- /dev/null +++ b/internal/cover/websocket.go @@ -0,0 +1,218 @@ +package cover + +import ( + "context" + "crypto/sha1" + "encoding/base64" + "errors" + "io" + "net/http" + "strings" + "sync" +) + +const websocketGUID = "258EAFA5-E914-47DA-95CA-C5AB0DC85B11" + +type websocketRequestContextKey struct{} +type websocketResponseContextKey struct{} +type websocketCleanupContextKey struct{} + +// A value, not a shared pointer: each request keeps its own original key and +// eligibility even when concurrent custom RoundTrippers replace response.Request. +type websocketRequestInfo struct { + eligible bool + key string +} + +func websocketRequestEligibility(request *http.Request) websocketRequestInfo { + if request == nil || request.Method != http.MethodGet || request.ProtoMajor != 1 || request.ProtoMinor != 1 || + request.ContentLength != 0 || request.Body != nil && request.Body != http.NoBody || + len(request.TransferEncoding) != 0 || len(request.Trailer) != 0 || len(headerValuesFold(request.Header, "Transfer-Encoding")) != 0 || + !websocketUpgradeHeaders(request.Header) { + return websocketRequestInfo{} + } + version, ok := singleHeaderValue(request.Header, "Sec-WebSocket-Version") + if !ok || version != "13" { + return websocketRequestInfo{} + } + key, ok := singleHeaderValue(request.Header, "Sec-WebSocket-Key") + if !ok { + return websocketRequestInfo{} + } + decoded, err := base64.StdEncoding.Strict().DecodeString(key) + if err != nil || len(decoded) != 16 || base64.StdEncoding.EncodeToString(decoded) != key { + return websocketRequestInfo{} + } + return websocketRequestInfo{eligible: true, key: key} +} + +func withWebsocketRequestEligibility(request *http.Request, info websocketRequestInfo) *http.Request { + ctx := context.WithValue(request.Context(), websocketRequestContextKey{}, info) + ctx = context.WithValue(ctx, websocketCleanupContextKey{}, &websocketCleanupState{}) + return request.WithContext(ctx) +} + +func websocketResponseRequest(request *http.Request) *http.Request { + info, _ := request.Context().Value(websocketRequestContextKey{}).(websocketRequestInfo) + return request.WithContext(context.WithValue(request.Context(), websocketResponseContextKey{}, info)) +} + +func validateWebsocketResponse(response *http.Response) error { + invalid := errors.New("invalid WebSocket upgrade response") + if response.Request == nil || response.ProtoMajor != 1 || response.ProtoMinor != 1 || + !websocketUpgradeHeaders(response.Header) { + return invalid + } + info, ok := response.Request.Context().Value(websocketResponseContextKey{}).(websocketRequestInfo) + if !ok || !info.eligible { + return invalid + } + accept, ok := singleHeaderValue(response.Header, "Sec-WebSocket-Accept") + // SHA-1 is the RFC 6455 handshake transform, not authentication or a MAC. + want := sha1.Sum([]byte(info.key + websocketGUID)) + if !ok || accept != base64.StdEncoding.EncodeToString(want[:]) { + return invalid + } + if _, ok := response.Body.(io.ReadWriteCloser); !ok { + return invalid + } + return nil +} + +// ReverseProxy closes invalid 101 bodies itself, but its early unsupported- +// Hijacker branch does not close an already validated body. Track only those +// accepted bodies, without changing the shared upstream transport's lifetime. +type websocketCleanupState struct { + mu sync.Mutex + body io.ReadWriteCloser +} + +func ownWebsocketResponse(response *http.Response) error { + state, _ := response.Request.Context().Value(websocketCleanupContextKey{}).(*websocketCleanupState) + if state == nil { + return errors.New("missing WebSocket response owner") + } + body := response.Body.(io.ReadWriteCloser) // validated before ownership + guard := &websocketDuplexBody{ReadWriteCloser: body} + var protected io.ReadWriteCloser = guard + if halfClose, ok := body.(interface{ CloseWrite() error }); ok { + protected = &websocketHalfCloseBody{websocketDuplexBody: guard, halfClose: halfClose} + } + response.Body = protected + state.mu.Lock() + state.body = protected + state.mu.Unlock() + return nil +} + +func closeWebsocketResponse(request *http.Request) { + if request == nil { + return + } + state, _ := request.Context().Value(websocketCleanupContextKey{}).(*websocketCleanupState) + if state == nil { + return + } + state.mu.Lock() + body := state.body + state.mu.Unlock() + if body != nil { + _ = body.Close() + } +} + +type websocketDuplexBody struct { + io.ReadWriteCloser + once sync.Once + closeErr error +} + +func (b *websocketDuplexBody) Close() error { + b.once.Do(func() { b.closeErr = b.ReadWriteCloser.Close() }) + return b.closeErr +} + +// Do not advertise CloseWrite when the original duplex body lacks it. +type websocketHalfCloseBody struct { + *websocketDuplexBody + halfClose interface{ CloseWrite() error } +} + +func (b *websocketHalfCloseBody) CloseWrite() error { return b.halfClose.CloseWrite() } + +// Only the WebSocket Upgrade pair may survive hop removal. Other valid tokens +// remain nominations to scrub; mandatory/negotiation fields may not be revived. +func websocketUpgradeHeaders(header http.Header) bool { + upgrade, ok := singleHeaderValue(header, "Upgrade") + if !ok || !httpToken(upgrade) || !strings.EqualFold(upgrade, "websocket") { + return false + } + connection, ok := singleHeaderValue(header, "Connection") + if !ok { + return false + } + seen := make(map[string]bool) + for part := range strings.SplitSeq(connection, ",") { + part = strings.Trim(part, " \t") + token := strings.ToLower(part) + if !httpToken(part) || seen[token] || token == "close" || websocketHandshakeHeader(token) { + return false + } + seen[token] = true + } + return seen["upgrade"] +} + +func websocketHandshakeHeader(name string) bool { + switch name { + case "sec-websocket-key", "sec-websocket-version", "sec-websocket-accept", "sec-websocket-protocol", "sec-websocket-extensions": + return true + default: + return false + } +} + +func httpToken(value string) bool { + if value == "" { + return false + } + for index := 0; index < len(value); index++ { + char := value[index] + if char >= 'a' && char <= 'z' || char >= 'A' && char <= 'Z' || char >= '0' && char <= '9' || + strings.ContainsRune("!#$%&'*+-.^_`|~", rune(char)) { + continue + } + return false + } + return true +} + +func singleHeaderValue(header http.Header, name string) (string, bool) { + values := headerValuesFold(header, name) + if len(values) != 1 { + return "", false + } + value := strings.Trim(values[0], " \t") + return value, value != "" +} + +// Native net/http fields are canonical, but an optional custom RoundTripper +// need not canonicalize its map. HTTP field names are ASCII tokens: Unicode +// case-fold aliases must not qualify when Header.Write would drop the field. +func headerValuesFold(header http.Header, name string) []string { + var values []string + for field, entries := range header { + if httpToken(field) && httpToken(name) && strings.EqualFold(field, name) { + values = append(values, entries...) + } + } + return values +} + +func deleteHeaderFold(header http.Header, name string) { + for field := range header { + if httpToken(field) && httpToken(name) && strings.EqualFold(field, name) { + delete(header, field) + } + } +} diff --git a/internal/cover/websocket_integration_test.go b/internal/cover/websocket_integration_test.go new file mode 100644 index 0000000..b3cc3e8 --- /dev/null +++ b/internal/cover/websocket_integration_test.go @@ -0,0 +1,1284 @@ +package cover + +import ( + "bufio" + "context" + "crypto/sha1" + "encoding/base64" + "errors" + "fmt" + "io" + "log" + "net" + "net/http" + "net/http/httptest" + "net/url" + "reflect" + "strings" + "sync" + "sync/atomic" + "syscall" + "testing" + "time" +) + +const webSocketIntegrationKey = "dGhlIHNhbXBsZSBub25jZQ==" + +// These tests use ordinary Go HTTP/1 sockets and complete, small RFC 6455 +// frames. They do not claim browser identity or full WebSocket conformance. +func TestReverseProxyWebSocketOnWire(t *testing.T) { + for _, caseFold := range []bool{false, true} { + t.Run(fmt.Sprintf("case_fold_%t", caseFold), func(t *testing.T) { + origin := newWebSocketIntegrationOrigin(t, 1, func(conn net.Conn, reader *bufio.Reader, request *http.Request) error { + if err := webSocketIntegrationWriteUpgrade(conn, request, nil); err != nil { + return err + } + if err := webSocketIntegrationWriteFrame(conn, "origin greeting", false); err != nil { + return err + } + for _, want := range []string{"client payload one", "client payload two"} { + payload, err := webSocketIntegrationReadFrame(reader, true) + if err != nil || payload != want { + return fmt.Errorf("masked client frame = %q/%v, want %q", payload, err, want) + } + if err := webSocketIntegrationWriteFrame(conn, "echo:"+payload, false); err != nil { + return err + } + } + _, err := reader.ReadByte() + if !webSocketIntegrationPeerClosed(err) { + return fmt.Errorf("origin peer after client close = %v", err) + } + return nil + }) + front := webSocketIntegrationProxy(t, origin.url, nil) + request := webSocketIntegrationRequest("GET", "HTTP/1.1") + request.Header.Set("Connection", "keep-alive, Upgrade, X-Request-Hop, Authorization, Proxy-Authorization") + request.Header.Set("X-Request-Hop", "fictional request hop") + request.Header.Set("Keep-Alive", "timeout=5") + request.Header.Set("Proxy-Connection", "keep-alive") + request.Header.Set("Authorization", "Bearer fictional origin credential") + request.Header.Set("Proxy-Authorization", "Basic fictional proxy credential") + request.Header.Set("Origin", "https://ordinary.example") + request.Header.Set("Cookie", "ordinary=one; other=two") + request.Header.Set("Sec-WebSocket-Protocol", "chat, superchat") + request.Header.Set("Sec-WebSocket-Extensions", "permessage-deflate; client_max_window_bits") + request.Header.Set("X-End-To-End", "retained request") + if caseFold { + request.Header.Set("Upgrade", "\tWebSocket ") + request.Header.Set("Connection", "Keep-Alive, uPgRaDe, X-Request-Hop, Authorization, Proxy-Authorization") + request.Header.Set("Sec-WebSocket-Key", "\t"+webSocketIntegrationKey+" ") + request.Header.Set("Sec-WebSocket-Version", "\t13 ") + } + conn, reader := webSocketIntegrationDial(t, front.addr) + if err := webSocketIntegrationWriteRequest(conn, request); err != nil { + t.Fatal(err) + } + response := webSocketIntegrationReadResponse(t, reader, request) + if response.StatusCode != http.StatusSwitchingProtocols || response.Proto != "HTTP/1.1" { + t.Fatalf("opening response = %s %d, want HTTP/1.1 101", response.Proto, response.StatusCode) + } + webSocketIntegrationAssertUpgrade(t, response.Header, webSocketIntegrationKey) + for name, want := range map[string]string{ + "Sec-WebSocket-Protocol": "chat", "Sec-WebSocket-Extensions": "permessage-deflate", + "Set-Cookie": "website=retained; HttpOnly", "X-End-To-End": "retained response", + } { + if got := response.Header.Get(name); got != want { + t.Errorf("ordinary upgraded %s = %q, want %q", name, got, want) + } + } + payload, err := webSocketIntegrationReadFrame(reader, false) + if err != nil || payload != "origin greeting" { + t.Fatalf("complete unmasked origin frame = %q/%v", payload, err) + } + for _, payload := range []string{"client payload one", "client payload two"} { + if err := webSocketIntegrationWriteFrame(conn, payload, true); err != nil { + t.Fatal(err) + } + got, err := webSocketIntegrationReadFrame(reader, false) + if err != nil || got != "echo:"+payload { + t.Fatalf("complete unmasked echo = %q/%v, want %q", got, err, "echo:"+payload) + } + } + _ = conn.Close() + got := webSocketIntegrationTake(t, origin.requests, "fixed-origin request") + webSocketIntegrationAssertFixedOrigin(t, got, origin.url) + if got.Header.Get("Connection") != "Upgrade" || got.Header.Get("Upgrade") != "websocket" { + t.Errorf("upstream normalized pair = %v", got.Header) + } + webSocketIntegrationAssertScrubbed(t, got.Header, "X-Request-Hop") + for name, want := range map[string]string{ + "Origin": "https://ordinary.example", "Cookie": "ordinary=one; other=two", + "Sec-WebSocket-Protocol": "chat, superchat", "Sec-WebSocket-Extensions": "permessage-deflate; client_max_window_bits", + "Sec-WebSocket-Key": webSocketIntegrationKey, "Sec-WebSocket-Version": "13", "X-End-To-End": "retained request", + } { + if got.Header.Get(name) != want { + t.Errorf("upstream ordinary %s = %q, want %q", name, got.Header.Get(name), want) + } + } + webSocketIntegrationAssertOriginFinished(t, origin) + front.waitHandlers(t) + }) + } + t.Run("https_origin_warm_h2_then_h1_upgrade", func(t *testing.T) { + requests := make(chan *http.Request, 3) + result := make(chan error, 1) + joined := make(chan struct{}) + var mu sync.Mutex + var hijacked net.Conn + closing := false + var started atomic.Bool + origin := httptest.NewUnstartedServer(http.HandlerFunc(func(w http.ResponseWriter, request *http.Request) { + requests <- request.Clone(context.Background()) + if request.Header.Get("Upgrade") == "" { + _, _ = io.WriteString(w, "warm fixed website") + return + } + started.Store(true) + defer close(joined) + conn, rw, err := w.(http.Hijacker).Hijack() + if err != nil { + result <- err + return + } + mu.Lock() + if closing { + mu.Unlock() + _ = conn.Close() + result <- net.ErrClosed + return + } + hijacked = conn + mu.Unlock() + defer conn.Close() + _ = conn.SetDeadline(time.Now().Add(3 * time.Second)) + err = webSocketIntegrationWriteUpgrade(conn, request, nil) + if err == nil { + err = webSocketIntegrationWriteFrame(conn, "secure greeting", false) + } + if err == nil { + var payload string + payload, err = webSocketIntegrationReadFrame(rw, true) + if err == nil && payload != "secure payload" { + err = fmt.Errorf("secure masked payload = %q", payload) + } + if err == nil { + err = webSocketIntegrationWriteFrame(conn, "secure echo:"+payload, false) + } + } + if err == nil { + _, err = rw.ReadByte() + if webSocketIntegrationPeerClosed(err) { + err = nil + } + } + result <- err + })) + origin.Config.ReadHeaderTimeout = time.Second + origin.Config.ReadTimeout = 3 * time.Second + origin.Config.WriteTimeout = 3 * time.Second + origin.EnableHTTP2 = true + origin.StartTLS() + t.Cleanup(func() { + mu.Lock() + closing = true + conn := hijacked + mu.Unlock() + if conn != nil { + _ = conn.Close() + } + origin.Close() + if started.Load() { + webSocketIntegrationJoin(t, joined, "HTTPS origin duplex handler cleanup") + } + }) + target, err := url.Parse(origin.URL + "/base?operator=one") + if err != nil { + t.Fatal(err) + } + upstream := origin.Client().Transport.(*http.Transport).Clone() + upstream.Proxy = nil + upstream.ForceAttemptHTTP2 = true + upstream.ResponseHeaderTimeout = time.Second + upstream.DialContext = func(ctx context.Context, network, address string) (net.Conn, error) { + if network != "tcp" || address != target.Host { + return nil, errors.New("HTTPS fixture refuses requester-controlled destination") + } + return (&net.Dialer{Timeout: time.Second}).DialContext(ctx, network, address) + } + t.Cleanup(upstream.CloseIdleConnections) + front := webSocketIntegrationProxy(t, target, upstream) + clientTransport := &http.Transport{Proxy: nil} + t.Cleanup(clientTransport.CloseIdleConnections) + client := &http.Client{Transport: clientTransport, Timeout: 2 * time.Second} + ordinary := func() { + t.Helper() + response, err := client.Get("http://" + front.addr + "/socket?client=one") + if err != nil { + t.Fatal(err) + } + body, err := io.ReadAll(response.Body) + _ = response.Body.Close() + if err != nil || response.StatusCode != 200 || string(body) != "warm fixed website" { + t.Fatalf("HTTPS ordinary response = %d/%q/%v", response.StatusCode, body, err) + } + got := webSocketIntegrationTake(t, requests, "HTTPS ordinary origin request") + webSocketIntegrationAssertFixedOrigin(t, got, target) + if got.ProtoMajor != 2 { + t.Fatalf("ordinary HTTPS origin protocol = %s, want actual H2", got.Proto) + } + } + ordinary() // Keep the same upstream transport and its warm H2 pool. + request := webSocketIntegrationRequest("GET", "HTTP/1.1") + conn, reader := webSocketIntegrationDial(t, front.addr) + if err := webSocketIntegrationWriteRequest(conn, request); err != nil { + t.Fatal(err) + } + response := webSocketIntegrationReadResponse(t, reader, request) + if response.StatusCode != 101 { + t.Fatalf("HTTPS fixed-origin opening response = %d, want 101", response.StatusCode) + } + webSocketIntegrationAssertUpgrade(t, response.Header, webSocketIntegrationKey) + if payload, err := webSocketIntegrationReadFrame(reader, false); err != nil || payload != "secure greeting" { + t.Fatalf("secure unmasked greeting = %q/%v", payload, err) + } + if err := webSocketIntegrationWriteFrame(conn, "secure payload", true); err != nil { + t.Fatal(err) + } + if payload, err := webSocketIntegrationReadFrame(reader, false); err != nil || payload != "secure echo:secure payload" { + t.Fatalf("secure unmasked echo = %q/%v", payload, err) + } + got := webSocketIntegrationTake(t, requests, "HTTPS WebSocket origin request") + webSocketIntegrationAssertFixedOrigin(t, got, target) + if got.Proto != "HTTP/1.1" { + t.Errorf("HTTPS WebSocket origin protocol = %s, want actual HTTP/1.1", got.Proto) + } + _ = conn.Close() + if err := webSocketIntegrationTake(t, result, "HTTPS origin result before cleanup"); err != nil { + t.Error(err) + } + webSocketIntegrationJoin(t, joined, "HTTPS origin duplex handler before cleanup") + front.waitHandlers(t) + ordinary() // WebSocket did not disable H2 for subsequent ordinary work. + }) +} + +func TestReverseProxyWebSocketInvalidRequestsRemainOrdinary(t *testing.T) { + tests := []struct { + name string + edit func(*http.Request) + }{ + {"post", func(r *http.Request) { r.Method = http.MethodPost }}, + {"head", func(r *http.Request) { r.Method = http.MethodHead }}, + {"http_1_0", func(r *http.Request) { r.Proto, r.ProtoMajor, r.ProtoMinor = "HTTP/1.0", 1, 0 }}, + {"missing_upgrade", func(r *http.Request) { r.Header.Del("Upgrade") }}, + {"h2c", func(r *http.Request) { r.Header.Set("Upgrade", "h2c") }}, + {"other_upgrade", func(r *http.Request) { r.Header.Set("Upgrade", "other") }}, + {"upgrade_list", func(r *http.Request) { r.Header.Set("Upgrade", "websocket, h2c") }}, + {"duplicate_upgrade_fields", func(r *http.Request) { r.Header.Add("Upgrade", "websocket") }}, + {"missing_connection", func(r *http.Request) { r.Header.Del("Connection") }}, + {"connection_without_upgrade", func(r *http.Request) { r.Header.Set("Connection", "keep-alive") }}, + {"duplicate_connection_fields", func(r *http.Request) { r.Header.Add("Connection", "Upgrade") }}, + {"duplicate_connection_upgrade", func(r *http.Request) { r.Header.Set("Connection", "Upgrade, upgrade") }}, + {"empty_connection_token", func(r *http.Request) { r.Header.Set("Connection", "Upgrade,,keep-alive") }}, + {"connection_close", func(r *http.Request) { r.Header.Set("Connection", "Upgrade, close") }}, + {"bad_connection_token", func(r *http.Request) { r.Header.Set("Connection", "Upgrade, bad(token") }}, + {"missing_key", func(r *http.Request) { r.Header.Del("Sec-WebSocket-Key") }}, + {"duplicate_key", func(r *http.Request) { r.Header.Add("Sec-WebSocket-Key", webSocketIntegrationKey) }}, + {"malformed_key", func(r *http.Request) { r.Header.Set("Sec-WebSocket-Key", "not-base64") }}, + {"short_key", func(r *http.Request) { + r.Header.Set("Sec-WebSocket-Key", base64.StdEncoding.EncodeToString([]byte("too short"))) + }}, + {"noncanonical_key", func(r *http.Request) { r.Header.Set("Sec-WebSocket-Key", "dGhlIHNhbXBsZSBub25jZR==") }}, + {"unicode_key_whitespace", func(r *http.Request) { r.Header.Set("Sec-WebSocket-Key", "\u00a0"+webSocketIntegrationKey) }}, + {"missing_version", func(r *http.Request) { r.Header.Del("Sec-WebSocket-Version") }}, + {"wrong_version", func(r *http.Request) { r.Header.Set("Sec-WebSocket-Version", "12") }}, + {"duplicate_version", func(r *http.Request) { r.Header.Add("Sec-WebSocket-Version", "13") }}, + {"body", func(r *http.Request) { r.Body, r.ContentLength = io.NopCloser(strings.NewReader("x")), 1 }}, + {"chunked_body", func(r *http.Request) { + r.Body, r.ContentLength, r.TransferEncoding = io.NopCloser(strings.NewReader("x")), -1, []string{"chunked"} + }}, + {"nominated_key", func(r *http.Request) { r.Header.Set("Connection", "Upgrade, Sec-WebSocket-Key") }}, + {"nominated_version", func(r *http.Request) { r.Header.Set("Connection", "Upgrade, Sec-WebSocket-Version") }}, + {"nominated_protocol", func(r *http.Request) { + r.Header.Set("Connection", "Upgrade, Sec-WebSocket-Protocol") + r.Header.Set("Sec-WebSocket-Protocol", "chat") + }}, + {"nominated_extensions", func(r *http.Request) { + r.Header.Set("Connection", "Upgrade, Sec-WebSocket-Extensions") + r.Header.Set("Sec-WebSocket-Extensions", "permessage-deflate") + }}, + } + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + origin := newWebSocketIntegrationOrigin(t, 1, func(conn net.Conn, _ *bufio.Reader, _ *http.Request) error { + _, err := io.WriteString(conn, "HTTP/1.1 200 OK\r\nContent-Length: 16\r\nConnection: close\r\nX-Website: ordinary\r\n\r\nordinary website") + return err + }) + front := webSocketIntegrationProxy(t, origin.url, nil) + request := webSocketIntegrationRequest("GET", "HTTP/1.1") + request.Header.Set("Authorization", "fictional auth") + request.Header.Set("Proxy-Authorization", "fictional proxy auth") + request.Header.Set("X-End-To-End", "retained") + test.edit(request) + conn, reader := webSocketIntegrationDial(t, front.addr) + if err := webSocketIntegrationWriteRequest(conn, request); err != nil { + t.Fatal(err) + } + response := webSocketIntegrationReadResponse(t, reader, request) + body, err := io.ReadAll(response.Body) + _ = response.Body.Close() + if err != nil { + t.Fatal(err) + } + wantBody := "ordinary website" + if request.Method == http.MethodHead { + wantBody = "" + } + if response.StatusCode != http.StatusOK || string(body) != wantBody || response.Header.Get("X-Website") != "ordinary" { + t.Fatalf("ineligible request result = %d/%q/%v, want fixed website 200/%q", response.StatusCode, body, response.Header, wantBody) + } + got := webSocketIntegrationTake(t, origin.requests, "ordinary fixed-origin request") + webSocketIntegrationAssertFixedOrigin(t, got, origin.url) + if got.Method != request.Method || got.Header.Get("Connection") != "" || got.Header.Get("Upgrade") != "" { + t.Errorf("ineligible upstream method/upgrade = %s/%v", got.Method, got.Header) + } + webSocketIntegrationAssertScrubbed(t, got.Header) + if got.Header.Get("X-End-To-End") != "retained" { + t.Error("ordinary request header lost on ineligible upgrade") + } + for _, nominated := range []string{"Sec-WebSocket-Key", "Sec-WebSocket-Version", "Sec-WebSocket-Protocol", "Sec-WebSocket-Extensions"} { + if strings.Contains(strings.ToLower(request.Header.Get("Connection")), strings.ToLower(nominated)) && got.Header.Get(nominated) != "" { + t.Errorf("Connection-nominated %s was resurrected: %v", nominated, got.Header) + } + } + _ = conn.Close() + webSocketIntegrationAssertOriginFinished(t, origin) + front.waitHandlers(t) + }) + } + t.Run("non_ascii_upgrade_native_rejection", func(t *testing.T) { + // This obs-text field is HTTP-parseable, but the standard ReverseProxy + // rejects a non-printable upgrade type before its Rewrite callback. + // Unlike the printable malformed cases above, no origin call is made. + origin := newWebSocketIntegrationOrigin(t, 0, func(net.Conn, *bufio.Reader, *http.Request) error { + return errors.New("unexpected non-ASCII upgrade origin call") + }) + upstream := webSocketIntegrationNewTransport(t, origin.url) + var calls atomic.Int32 + front := webSocketIntegrationProxy(t, origin.url, webSocketIntegrationRoundTripper(func(r *http.Request) (*http.Response, error) { + calls.Add(1) + return upstream.RoundTrip(r) + })) + request := webSocketIntegrationRequest("GET", "HTTP/1.1") + request.Header.Set("Upgrade", "webs\u00f6cket") + conn, reader := webSocketIntegrationDial(t, front.addr) + if err := webSocketIntegrationWriteRequest(conn, request); err != nil { + t.Fatal(err) + } + webSocketIntegrationAssertGeneric502(t, webSocketIntegrationReadResponse(t, reader, request)) + front.waitHandlers(t) + if calls.Load() != 0 { + t.Errorf("native non-ASCII upgrade rejection made %d origin transport calls", calls.Load()) + } + webSocketIntegrationJoin(t, origin.joined, "unused fixed-origin worker") + }) +} + +func TestReverseProxyWebSocketRejectsUnsafeSwitchingResponses(t *testing.T) { + tests := []struct { + name string + unsolicited bool + edit func(http.Header) + }{ + {name: "unsolicited", unsolicited: true}, + {name: "wrong_accept", edit: func(h http.Header) { h.Set("Sec-WebSocket-Accept", "fictional private origin proof") }}, + {name: "missing_accept", edit: func(h http.Header) { h.Del("Sec-WebSocket-Accept") }}, + {name: "duplicate_accept", edit: func(h http.Header) { h.Add("Sec-WebSocket-Accept", h.Get("Sec-WebSocket-Accept")) }}, + {name: "wrong_upgrade", edit: func(h http.Header) { h.Set("Upgrade", "h2c") }}, + {name: "upgrade_list", edit: func(h http.Header) { h.Set("Upgrade", "websocket, h2c") }}, + {name: "duplicate_upgrade", edit: func(h http.Header) { h.Add("Upgrade", "websocket") }}, + {name: "missing_upgrade", edit: func(h http.Header) { h.Del("Upgrade") }}, + {name: "missing_connection", edit: func(h http.Header) { h.Del("Connection") }}, + {name: "duplicate_connection", edit: func(h http.Header) { h.Add("Connection", "Upgrade") }}, + {name: "connection_close", edit: func(h http.Header) { h.Set("Connection", "Upgrade, close") }}, + {name: "connection_duplicate_token", edit: func(h http.Header) { h.Set("Connection", "Upgrade, upgrade") }}, + {name: "connection_empty_token", edit: func(h http.Header) { h.Set("Connection", "Upgrade,") }}, + {name: "connection_bad_token", edit: func(h http.Header) { h.Set("Connection", "Upgrade, bad(token") }}, + {name: "nominated_accept", edit: func(h http.Header) { h.Set("Connection", "Upgrade, Sec-WebSocket-Accept") }}, + {name: "nominated_protocol", edit: func(h http.Header) { h.Set("Connection", "Upgrade, Sec-WebSocket-Protocol") }}, + {name: "nominated_extensions", edit: func(h http.Header) { h.Set("Connection", "Upgrade, Sec-WebSocket-Extensions") }}, + } + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + origin := newWebSocketIntegrationOrigin(t, 1, func(conn net.Conn, reader *bufio.Reader, request *http.Request) error { + if err := webSocketIntegrationWriteUpgrade(conn, request, test.edit); err != nil { + return err + } + _, err := reader.ReadByte() + if !webSocketIntegrationPeerClosed(err) { + return fmt.Errorf("rejected 101 origin peer closure before cleanup = %v", err) + } + return nil + }) + upstream := webSocketIntegrationNewTransport(t, origin.url) + closed := make(chan struct{}) + front := webSocketIntegrationProxy(t, origin.url, webSocketIntegrationRoundTripper(func(r *http.Request) (*http.Response, error) { + response, err := upstream.RoundTrip(r) + if err != nil { + return nil, err + } + tracked := &webSocketIntegrationCloseBody{ReadCloser: response.Body, closed: closed} + if writer, ok := response.Body.(io.Writer); ok { + response.Body = &webSocketIntegrationDuplexBody{webSocketIntegrationCloseBody: tracked, writer: writer} + } else { + response.Body = tracked + } + return response, nil + })) + request := webSocketIntegrationRequest("GET", "HTTP/1.1") + if test.unsolicited { + request.Header.Del("Connection") + request.Header.Del("Upgrade") + } + conn, reader := webSocketIntegrationDial(t, front.addr) + if err := webSocketIntegrationWriteRequest(conn, request); err != nil { + t.Fatal(err) + } + response := webSocketIntegrationReadResponse(t, reader, request) + webSocketIntegrationAssertGeneric502(t, response) + // Both oracles happen before any explicit test client/origin close. + webSocketIntegrationJoin(t, closed, "original upstream Body.Close before fixture cleanup") + webSocketIntegrationAssertOriginFinished(t, origin) + front.waitHandlers(t) + _ = conn.Close() + }) + } + for _, test := range []struct { + name string + edit func(*http.Response) + }{ + {"not_duplex", func(r *http.Response) {}}, + {"nil_body", func(r *http.Response) { r.Body = nil }}, + {"http_2_response", func(r *http.Response) { r.Proto, r.ProtoMajor, r.ProtoMinor = "HTTP/2.0", 2, 0 }}, + {"unspecified_protocol", func(r *http.Response) { r.Proto, r.ProtoMajor, r.ProtoMinor = "", 0, 0 }}, + {"case_alias_duplicate_accept", func(r *http.Response) { + r.Header["sec-websocket-accept"] = []string{r.Header.Get("Sec-WebSocket-Accept")} + }}, + {"case_alias_duplicate_connection", func(r *http.Response) { r.Header["connection"] = []string{"Upgrade"} }}, + {"unicode_alias_accept", func(r *http.Response) { + accept := r.Header.Get("Sec-WebSocket-Accept") + r.Header.Del("Sec-WebSocket-Accept") + // Unicode case folding is not HTTP field-name equivalence. Go's + // wire serializer drops this invalid non-ASCII name, so accepting + // it as the mandatory proof could send an unverified 101 on wire. + r.Header["\u017fec-WebSocket-Accept"] = []string{accept} + }}, + } { + t.Run("custom_transport/"+test.name, func(t *testing.T) { + body := &webSocketIntegrationCloseBody{ReadCloser: io.NopCloser(strings.NewReader("fictional private upstream body")), closed: make(chan struct{})} + var originalBodyPresent bool + front := webSocketIntegrationProxy(t, &url.URL{Scheme: "http", Host: "fixed-origin.invalid"}, webSocketIntegrationRoundTripper(func(r *http.Request) (*http.Response, error) { + response := webSocketIntegrationResponse(r.Header.Get("Sec-WebSocket-Key"), body) + if test.name != "not_duplex" && test.name != "nil_body" { + response.Body = &webSocketIntegrationDuplexBody{webSocketIntegrationCloseBody: body, writer: io.Discard} + } + test.edit(response) + originalBodyPresent = response.Body != nil + return response, nil + })) + request := webSocketIntegrationRequest("GET", "HTTP/1.1") + conn, reader := webSocketIntegrationDial(t, front.addr) + if err := webSocketIntegrationWriteRequest(conn, request); err != nil { + t.Fatal(err) + } + response := webSocketIntegrationReadResponse(t, reader, request) + webSocketIntegrationAssertGeneric502(t, response) + front.waitHandlers(t) + if originalBodyPresent { + webSocketIntegrationJoin(t, body.closed, "rejected custom upstream Body.Close before fixture cleanup") + if body.closes.Load() != 1 { + t.Errorf("rejected custom upstream Close calls = %d, want 1", body.closes.Load()) + } + if body.reads.Load() != 0 { + t.Errorf("rejected upstream body was read %d times", body.reads.Load()) + } + } + _ = conn.Close() + }) + } +} + +func TestReverseProxyWebSocketTransportRequestAssociation(t *testing.T) { + for _, failure := range []string{"recorder_without_hijack", "hijack_not_supported", "hijack_other_error"} { + t.Run(failure, func(t *testing.T) { + body := &webSocketIntegrationCloseBody{ReadCloser: io.NopCloser(strings.NewReader("fictional private duplex bytes")), closed: make(chan struct{})} + handler, err := NewReverseProxyHandler(&url.URL{Scheme: "http", Host: "fixed-origin.invalid"}, webSocketIntegrationRoundTripper(func(*http.Request) (*http.Response, error) { + return webSocketIntegrationResponse(webSocketIntegrationKey, &webSocketIntegrationDuplexBody{webSocketIntegrationCloseBody: body, writer: io.Discard}), nil + })) + if err != nil { + t.Fatal(err) + } + recorder := httptest.NewRecorder() + var writer http.ResponseWriter = recorder + if failure != "recorder_without_hijack" { + hijackErr := errors.New("fictional downstream hijack failure") + if failure == "hijack_not_supported" { + hijackErr = http.ErrNotSupported + } + writer = &webSocketIntegrationHijackFailure{ResponseRecorder: recorder, err: hijackErr} + } + request := httptest.NewRequest("GET", "http://requester-target.invalid/socket", nil) + request.Header = webSocketIntegrationRequest("GET", "HTTP/1.1").Header + handler.ServeHTTP(writer, request) + if recorder.Code != 502 || recorder.Body.String() != "Bad Gateway\n" { + t.Errorf("failed downstream hijack = %d/%q, want generic 502", recorder.Code, recorder.Body.String()) + } + webSocketIntegrationJoin(t, body.closed, "legal upstream 101 Close after failed Hijack before cleanup") + if body.closes.Load() != 1 || body.reads.Load() != 0 { + t.Errorf("failed Hijack original body calls = Close %d / Read %d, want 1 / 0", body.closes.Load(), body.reads.Load()) + } + }) + } + for _, responseRequest := range []string{"nil", "forged", "lowercase_fields"} { + for _, eligible := range []bool{true, false} { + t.Run(fmt.Sprintf("response_request_%s/eligible_%t", responseRequest, eligible), func(t *testing.T) { + body, pipe := newWebSocketIntegrationPipe(t) + front := webSocketIntegrationProxy(t, &url.URL{Scheme: "http", Host: "fixed-origin.invalid", Path: "/base", RawQuery: "operator=one"}, webSocketIntegrationRoundTripper(func(r *http.Request) (*http.Response, error) { + if r.URL.Host != "fixed-origin.invalid" || r.Host != "fixed-origin.invalid" || r.URL.Path != "/base/socket" || r.URL.RawQuery != "operator=one&client=one" { + return nil, fmt.Errorf("custom transport received requester destination: %s / %s", r.URL, r.Host) + } + response := webSocketIntegrationResponse(webSocketIntegrationKey, body) + if responseRequest == "forged" { + response.Request = webSocketIntegrationRequest("GET", "HTTP/1.1") + response.Request.Header.Set("Sec-WebSocket-Key", base64.StdEncoding.EncodeToString([]byte("different-key-16"))) + } + if responseRequest == "lowercase_fields" { + lowercase := make(http.Header) + for name, values := range response.Header { + lowercase[strings.ToLower(name)] = append([]string(nil), values...) + } + response.Header = lowercase + } + return response, nil + })) + request := webSocketIntegrationRequest("GET", "HTTP/1.1") + if !eligible { + request.Header.Del("Sec-WebSocket-Version") + } + conn, reader := webSocketIntegrationDial(t, front.addr) + if err := webSocketIntegrationWriteRequest(conn, request); err != nil { + t.Fatal(err) + } + response := webSocketIntegrationReadResponse(t, reader, request) + if eligible { + if response.StatusCode != http.StatusSwitchingProtocols { + t.Fatalf("trusted eligible request with %s response.Request = %d, want 101", responseRequest, response.StatusCode) + } + webSocketIntegrationAssertUpgrade(t, response.Header, webSocketIntegrationKey) + webSocketIntegrationExchangePipe(t, conn, reader) + _ = conn.Close() + if err := webSocketIntegrationTake(t, pipe.result, "custom transport duplex result"); err != nil { + t.Error(err) + } + } else { + webSocketIntegrationAssertGeneric502(t, response) + webSocketIntegrationJoin(t, body.closed, "ineligible forged/nil response upstream close before cleanup") + _ = conn.Close() + } + webSocketIntegrationJoin(t, pipe.joined, "custom transport duplex worker before cleanup") + front.waitHandlers(t) + }) + } + } + t.Run("public_synthetic_request_boundaries", func(t *testing.T) { + for _, test := range []struct { + name string + edit func(*http.Request) + }{ + {"http_2_request", func(r *http.Request) { r.Proto, r.ProtoMajor, r.ProtoMinor = "HTTP/2.0", 2, 0 }}, + {"empty_non_nil_body", func(r *http.Request) { r.Body = io.NopCloser(strings.NewReader("")) }}, + {"trailer_without_body", func(r *http.Request) { r.Trailer = http.Header{"X-End": nil} }}, + {"transfer_encoding_header", func(r *http.Request) { r.Header.Set("Transfer-Encoding", "chunked") }}, + {"case_alias_duplicate_key", func(r *http.Request) { r.Header["sec-websocket-key"] = []string{webSocketIntegrationKey} }}, + {"unicode_alias_key", func(r *http.Request) { + r.Header.Del("Sec-WebSocket-Key") + r.Header["\u017fec-WebSocket-Key"] = []string{webSocketIntegrationKey} + }}, + } { + t.Run(test.name, func(t *testing.T) { + var calls int + handler, err := NewReverseProxyHandler(&url.URL{Scheme: "http", Host: "fixed-origin.invalid"}, webSocketIntegrationRoundTripper(func(r *http.Request) (*http.Response, error) { + calls++ + if r.Header.Get("Upgrade") != "" || r.Header.Get("Connection") != "" || r.URL.Host != "fixed-origin.invalid" || r.Host != "fixed-origin.invalid" { + t.Errorf("synthetic ineligible request did not remain scrubbed fixed-origin: %v/%s/%s", r.Header, r.URL, r.Host) + } + return &http.Response{StatusCode: 200, Header: make(http.Header), Body: io.NopCloser(strings.NewReader("ordinary"))}, nil + })) + if err != nil { + t.Fatal(err) + } + request := httptest.NewRequest("GET", "http://requester-target.invalid/socket", nil) + for key, values := range webSocketIntegrationRequest("GET", "HTTP/1.1").Header { + request.Header[key] = values + } + test.edit(request) + response := httptest.NewRecorder() + handler.ServeHTTP(response, request) + if response.Code != 200 || response.Body.String() != "ordinary" || calls != 1 { + t.Errorf("synthetic ordinary fixed-origin result = %d/%q, calls %d", response.Code, response.Body, calls) + } + }) + } + }) + t.Run("concurrent_marker_isolation", func(t *testing.T) { + var calls atomic.Int32 + ready := make(chan struct{}) + keys := map[string]string{"eligible-a": webSocketIntegrationKey, "eligible-b": "AAAAAAAAAAAAAAAAAAAAAA==", "invalid": webSocketIntegrationKey} + bodies := make(map[string]*webSocketIntegrationDuplexBody) + pipes := make(map[string]*webSocketIntegrationPipe) + for _, marker := range []string{"eligible-a", "eligible-b", "invalid"} { + bodies[marker], pipes[marker] = newWebSocketIntegrationPipe(t) + } + front := webSocketIntegrationProxy(t, &url.URL{Scheme: "http", Host: "fixed-origin.invalid"}, webSocketIntegrationRoundTripper(func(r *http.Request) (*http.Response, error) { + if calls.Add(1) == 3 { + close(ready) + } + select { + case <-ready: + case <-r.Context().Done(): + return nil, context.Cause(r.Context()) + case <-time.After(time.Second): + return nil, errors.New("concurrent request barrier timed out") + } + marker := r.Header.Get("X-Marker") + response := webSocketIntegrationResponse(keys[marker], bodies[marker]) + // Every malicious association claims a valid request with another key. + response.Request = webSocketIntegrationRequest("GET", "HTTP/1.1") + response.Request.Header.Set("Sec-WebSocket-Key", "AAAAAAAAAAAAAAAAAAAAAA==") + return response, nil + })) + type result struct { + marker string + status int + accept string + body string + err error + } + results := make(chan result, 3) + joined := make(chan struct{}) + var workers sync.WaitGroup + for _, marker := range []string{"eligible-a", "eligible-b", "invalid"} { + conn, reader := webSocketIntegrationDial(t, front.addr) + workerJoined := make(chan struct{}) + t.Cleanup(func() { + _ = conn.Close() + webSocketIntegrationJoin(t, workerJoined, "independent concurrent client worker cleanup") + }) + workers.Add(1) + go func() { + defer close(workerJoined) + defer workers.Done() + defer conn.Close() + request := webSocketIntegrationRequest("GET", "HTTP/1.1") + request.Header.Set("X-Marker", marker) + request.Header.Set("Sec-WebSocket-Key", keys[marker]) + if marker == "invalid" { + request.Header.Del("Sec-WebSocket-Version") + } + got := result{marker: marker} + if got.err = webSocketIntegrationWriteRequest(conn, request); got.err == nil { + var response *http.Response + response, got.err = http.ReadResponse(reader, request) + if got.err == nil { + got.status = response.StatusCode + got.accept = response.Header.Get("Sec-WebSocket-Accept") + if response.StatusCode == 101 { + got.err = webSocketIntegrationExchangePipeError(conn, reader) + } else { + var body []byte + body, got.err = io.ReadAll(response.Body) + _ = response.Body.Close() + got.body = string(body) + } + } + } + results <- got + }() + } + go func() { defer close(joined); workers.Wait() }() + for index := 0; index < 3; index++ { + got := webSocketIntegrationTake(t, results, "concurrent marker result") + wantStatus := 101 + if got.marker == "invalid" { + wantStatus = 502 + } + if got.err != nil || got.status != wantStatus || got.marker == "invalid" && got.body != "Bad Gateway\n" { + t.Errorf("concurrent %s = %d/%q/%v, want %d with request-local eligibility", got.marker, got.status, got.body, got.err, wantStatus) + } + if got.marker != "invalid" && got.accept != webSocketIntegrationAccept(keys[got.marker]) { + t.Errorf("concurrent %s accept = %q, want proof bound to its own key", got.marker, got.accept) + } + } + webSocketIntegrationJoin(t, joined, "concurrent clients before fixture cleanup") + for _, marker := range []string{"eligible-a", "eligible-b", "invalid"} { + webSocketIntegrationJoin(t, bodies[marker].closed, marker+" original body close") + webSocketIntegrationJoin(t, pipes[marker].joined, marker+" pipe worker") + } + front.waitHandlers(t) + }) +} + +func TestReverseProxyWebSocketOrdinaryResponseControls(t *testing.T) { + for _, status := range []int{http.StatusOK, http.StatusUpgradeRequired} { + t.Run(fmt.Sprintf("status_%d", status), func(t *testing.T) { + origin := newWebSocketIntegrationOrigin(t, 1, func(conn net.Conn, _ *bufio.Reader, _ *http.Request) error { + for index, code := range []int{103, 102, 103} { + if _, err := fmt.Fprintf(conn, "HTTP/1.1 %d Information\r\nConnection: X-Info-Hop\r\nX-Info-Hop: fictional\r\nAuthorization: fictional\r\nProxy-Authorization: fictional\r\nLink: ; rel=preload\r\n\r\n", code, index); err != nil { + return err + } + } + // Do not combine this nomination oracle with Connection: close: + // Go's native response parser deletes the entire Connection field + // for that token before RoundTrip exposes it to the proxy. The + // owned origin still physically closes after the complete message. + _, err := fmt.Fprintf(conn, "HTTP/1.1 %d %s\r\nConnection: X-Final-Hop\r\nUpgrade: websocket\r\nX-Final-Hop: fictional\r\nAuthorization: fictional initial\r\nProxy-Authorization: fictional initial\r\nX-Website: retained\r\nTransfer-Encoding: chunked\r\nTrailer: X-End, Authorization, Proxy-Authorization, Connection, X-Late-Hop\r\n\r\n1\r\nx\r\n0\r\nAuthorization: fictional late\r\nProxy-Authorization: fictional late\r\nConnection: X-Late-Hop\r\nX-Late-Hop: fictional late\r\nX-End: retained trailer\r\n\r\n", status, http.StatusText(status)) + return err + }) + front := webSocketIntegrationProxy(t, origin.url, nil) + request := webSocketIntegrationRequest("GET", "HTTP/1.1") + conn, reader := webSocketIntegrationDial(t, front.addr) + if err := webSocketIntegrationWriteRequest(conn, request); err != nil { + t.Fatal(err) + } + for index, want := range []int{103, 102, 103} { + response := webSocketIntegrationReadResponse(t, reader, request) + if response.StatusCode != want || response.Header.Get("Link") != fmt.Sprintf("; rel=preload", index) { + t.Errorf("ordinary info response %d = %d/%v", index, response.StatusCode, response.Header) + } + webSocketIntegrationAssertScrubbed(t, response.Header, "Connection", "Upgrade", "X-Info-Hop") + } + response := webSocketIntegrationReadResponse(t, reader, request) + body, err := io.ReadAll(response.Body) + _ = response.Body.Close() + if err != nil || response.StatusCode != status || string(body) != "x" || response.Header.Get("X-Website") != "retained" || response.Trailer.Get("X-End") != "retained trailer" { + t.Errorf("ordinary non-101 final = %d/%q/%v, headers=%v trailers=%v", response.StatusCode, body, err, response.Header, response.Trailer) + } + // Wire framing may legitimately add Transfer-Encoding or Trailer; + // credential and Connection-nominated values may never reappear. + for _, fields := range []http.Header{response.Header, response.Trailer} { + for _, name := range []string{"Authorization", "Proxy-Authorization", "X-Final-Hop", "X-Late-Hop", "Upgrade"} { + if len(fields.Values(name)) != 0 { + t.Errorf("unsafe ordinary final/trailer %s = %v", name, fields.Values(name)) + } + } + } + _ = conn.Close() + webSocketIntegrationAssertOriginFinished(t, origin) + front.waitHandlers(t) + }) + } +} + +type webSocketIntegrationCloseBody struct { + io.ReadCloser + closed chan struct{} + once sync.Once + reads atomic.Int64 + closes atomic.Int64 +} + +func (body *webSocketIntegrationCloseBody) Read(p []byte) (int, error) { + body.reads.Add(1) + return body.ReadCloser.Read(p) +} + +func (body *webSocketIntegrationCloseBody) Close() error { + body.closes.Add(1) + err := body.ReadCloser.Close() + body.once.Do(func() { close(body.closed) }) + return err +} + +type webSocketIntegrationHijackFailure struct { + *httptest.ResponseRecorder + err error +} + +func (writer *webSocketIntegrationHijackFailure) Hijack() (net.Conn, *bufio.ReadWriter, error) { + return nil, nil, writer.err +} + +type webSocketIntegrationDuplexBody struct { + *webSocketIntegrationCloseBody + writer io.Writer +} + +func (body *webSocketIntegrationDuplexBody) Write(p []byte) (int, error) { + return body.writer.Write(p) +} + +func webSocketIntegrationResponse(key string, body io.ReadCloser) *http.Response { + return &http.Response{StatusCode: 101, Proto: "HTTP/1.1", ProtoMajor: 1, ProtoMinor: 1, + Header: webSocketIntegrationUpgradeHeader(key), Body: body} +} + +func webSocketIntegrationAssertGeneric502(t *testing.T, response *http.Response) { + t.Helper() + if response.StatusCode != http.StatusBadGateway { + t.Fatalf("invalid upstream 101 result = %d/%v, want generic 502", response.StatusCode, response.Header) + } + body, err := io.ReadAll(response.Body) + _ = response.Body.Close() + if err != nil || string(body) != "Bad Gateway\n" { + t.Fatalf("generic 502 body = %q/%v", body, err) + } + for _, name := range []string{"Connection", "Upgrade", "Sec-WebSocket-Accept", "X-End-To-End", "Authorization", "Proxy-Authorization", "Proxy-Authenticate"} { + if response.Header.Get(name) != "" { + t.Errorf("upstream/private field %s escaped generic rejection: %v", name, response.Header) + } + } +} + +type webSocketIntegrationPipe struct { + result chan error + joined chan struct{} +} + +func newWebSocketIntegrationPipe(t *testing.T) (*webSocketIntegrationDuplexBody, *webSocketIntegrationPipe) { + t.Helper() + local, peer := net.Pipe() + _ = local.SetDeadline(time.Now().Add(3 * time.Second)) + _ = peer.SetDeadline(time.Now().Add(3 * time.Second)) + body := &webSocketIntegrationDuplexBody{webSocketIntegrationCloseBody: &webSocketIntegrationCloseBody{ReadCloser: local, closed: make(chan struct{})}, writer: local} + pipe := &webSocketIntegrationPipe{result: make(chan error, 1), joined: make(chan struct{})} + go func() { + defer close(pipe.joined) + defer peer.Close() + reader := bufio.NewReader(peer) + err := webSocketIntegrationWriteFrame(peer, "custom greeting", false) + if err == nil { + var payload string + payload, err = webSocketIntegrationReadFrame(reader, true) + if err == nil && payload != "custom payload" { + err = fmt.Errorf("custom masked payload = %q", payload) + } + if err == nil { + err = webSocketIntegrationWriteFrame(peer, "custom echo:"+payload, false) + } + if err == nil { + _, err = reader.ReadByte() + if webSocketIntegrationPeerClosed(err) { + err = nil + } + } + } + pipe.result <- err + }() + t.Cleanup(func() { + _ = local.Close() + _ = peer.Close() + webSocketIntegrationJoin(t, pipe.joined, "custom transport pipe cleanup worker") + }) + return body, pipe +} + +func webSocketIntegrationExchangePipe(t *testing.T, conn net.Conn, reader *bufio.Reader) { + t.Helper() + if err := webSocketIntegrationExchangePipeError(conn, reader); err != nil { + t.Fatal(err) + } +} + +func webSocketIntegrationExchangePipeError(conn net.Conn, reader *bufio.Reader) error { + payload, err := webSocketIntegrationReadFrame(reader, false) + if err != nil || payload != "custom greeting" { + return fmt.Errorf("custom unmasked greeting = %q/%v", payload, err) + } + if err := webSocketIntegrationWriteFrame(conn, "custom payload", true); err != nil { + return err + } + payload, err = webSocketIntegrationReadFrame(reader, false) + if err != nil || payload != "custom echo:custom payload" { + return fmt.Errorf("custom unmasked echo = %q/%v", payload, err) + } + return nil +} + +type webSocketIntegrationRoundTripper func(*http.Request) (*http.Response, error) + +func (f webSocketIntegrationRoundTripper) RoundTrip(r *http.Request) (*http.Response, error) { + return f(r) +} + +type webSocketIntegrationFront struct { + addr string + server *http.Server + listener net.Listener + mu sync.Mutex + conns map[net.Conn]struct{} + closing bool + handlers sync.WaitGroup + joined chan struct{} + serveErr chan error +} + +func webSocketIntegrationProxy(t *testing.T, origin *url.URL, transport http.RoundTripper) *webSocketIntegrationFront { + t.Helper() + if transport == nil { + transport = webSocketIntegrationNewTransport(t, origin) + } + handler, err := NewReverseProxyHandler(origin, transport) + if err != nil { + t.Fatal(err) + } + listener, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatal(err) + } + front := &webSocketIntegrationFront{addr: listener.Addr().String(), listener: listener, conns: make(map[net.Conn]struct{}), joined: make(chan struct{}), serveErr: make(chan error, 1)} + front.server = &http.Server{ReadHeaderTimeout: time.Second, ReadTimeout: 3 * time.Second, WriteTimeout: 3 * time.Second, IdleTimeout: time.Second, ErrorLog: log.New(io.Discard, "", 0)} + front.server.ConnState = func(conn net.Conn, state http.ConnState) { + front.mu.Lock() + defer front.mu.Unlock() + if state == http.StateNew { + front.conns[conn] = struct{}{} + } else if state == http.StateClosed { + delete(front.conns, conn) + } + } + front.server.Handler = http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + front.mu.Lock() + if front.closing { + front.mu.Unlock() + http.Error(w, "closing", http.StatusServiceUnavailable) + return + } + front.handlers.Add(1) + front.mu.Unlock() + defer front.handlers.Done() + handler.ServeHTTP(w, r) + }) + go func() { + defer close(front.joined) + front.serveErr <- front.server.Serve(listener) + }() + t.Cleanup(func() { + front.mu.Lock() + front.closing = true + var sockets []net.Conn + for conn := range front.conns { + sockets = append(sockets, conn) + } + front.mu.Unlock() + _ = listener.Close() + _ = front.server.Close() + for _, conn := range sockets { + _ = conn.Close() + } + webSocketIntegrationJoin(t, front.joined, "front Serve worker") + front.waitHandlers(t) + }) + return front +} + +func webSocketIntegrationNewTransport(t *testing.T, origin *url.URL) *http.Transport { + t.Helper() + transport := &http.Transport{Proxy: nil, DialContext: func(ctx context.Context, network, address string) (net.Conn, error) { + if network != "tcp" || address != origin.Host { + return nil, errors.New("test refuses requester-controlled destination") + } + return (&net.Dialer{Timeout: time.Second}).DialContext(ctx, network, address) + }, ResponseHeaderTimeout: time.Second, IdleConnTimeout: time.Second} + t.Cleanup(transport.CloseIdleConnections) + return transport +} + +func (front *webSocketIntegrationFront) waitHandlers(t *testing.T) { + t.Helper() + joined := make(chan struct{}) + go func() { defer close(joined); front.handlers.Wait() }() + webSocketIntegrationJoin(t, joined, "front cover handlers") +} + +type webSocketIntegrationOrigin struct { + url *url.URL + requests chan *http.Request + results chan error + joined chan struct{} +} + +func newWebSocketIntegrationOrigin(t *testing.T, attempts int, serve func(net.Conn, *bufio.Reader, *http.Request) error) *webSocketIntegrationOrigin { + t.Helper() + listener, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatal(err) + } + _ = listener.(*net.TCPListener).SetDeadline(time.Now().Add(3 * time.Second)) + originURL, err := url.Parse("http://" + listener.Addr().String() + "/base?operator=one") + if err != nil { + t.Fatal(err) + } + origin := &webSocketIntegrationOrigin{url: originURL, requests: make(chan *http.Request, attempts), results: make(chan error, attempts), joined: make(chan struct{})} + var mu sync.Mutex + var sockets []net.Conn + closing := false + go func() { + defer close(origin.joined) + var workers sync.WaitGroup + defer workers.Wait() + for attempt := 0; attempt < attempts; attempt++ { + conn, err := listener.Accept() + if err != nil { + origin.results <- err + return + } + mu.Lock() + if closing { + mu.Unlock() + _ = conn.Close() + origin.results <- net.ErrClosed + return + } + sockets = append(sockets, conn) + mu.Unlock() + workers.Add(1) + go func() { + defer workers.Done() + defer conn.Close() + _ = conn.SetDeadline(time.Now().Add(3 * time.Second)) + reader := bufio.NewReader(conn) + request, err := http.ReadRequest(reader) + if err != nil { + origin.results <- err + return + } + if _, err := io.Copy(io.Discard, request.Body); err != nil { + origin.results <- err + return + } + _ = request.Body.Close() + origin.requests <- request.Clone(context.Background()) + origin.results <- serve(conn, reader, request) + }() + } + }() + t.Cleanup(func() { + mu.Lock() + closing = true + owned := append([]net.Conn(nil), sockets...) + mu.Unlock() + _ = listener.Close() + for _, conn := range owned { + _ = conn.Close() + } + webSocketIntegrationJoin(t, origin.joined, "origin accept/frame workers") + }) + return origin +} + +func webSocketIntegrationRequest(method, protocol string) *http.Request { + request := &http.Request{Method: method, Proto: protocol, ProtoMajor: 1, ProtoMinor: 1, + URL: &url.URL{Scheme: "http", Host: "requester-target.invalid:9876", Path: "/socket", RawQuery: "client=one"}, + Host: "requester-host.invalid:5432", Header: make(http.Header)} + request.Header.Set("Connection", "Upgrade") + request.Header.Set("Upgrade", "websocket") + request.Header.Set("Sec-WebSocket-Key", webSocketIntegrationKey) + request.Header.Set("Sec-WebSocket-Version", "13") + request.Header.Set("Sec-WebSocket-Protocol", "chat, superchat") + request.Header.Set("Sec-WebSocket-Extensions", "permessage-deflate; client_max_window_bits") + return request +} + +func webSocketIntegrationWriteRequest(w io.Writer, request *http.Request) error { + var body []byte + if request.Body != nil { + var err error + body, err = io.ReadAll(request.Body) + _ = request.Body.Close() + if err != nil { + return err + } + } + if _, err := fmt.Fprintf(w, "%s %s %s\r\nHost: %s\r\n", request.Method, request.URL.String(), request.Proto, request.Host); err != nil { + return err + } + if err := request.Header.Write(w); err != nil { + return err + } + if len(request.TransferEncoding) != 0 { + if _, err := fmt.Fprintf(w, "Transfer-Encoding: chunked\r\n\r\n%x\r\n%s\r\n0\r\n\r\n", len(body), body); err != nil { + return err + } + return nil + } + if request.ContentLength > 0 { + if _, err := fmt.Fprintf(w, "Content-Length: %d\r\n", request.ContentLength); err != nil { + return err + } + } + _, err := fmt.Fprintf(w, "\r\n%s", body) + return err +} + +func webSocketIntegrationDial(t *testing.T, address string) (net.Conn, *bufio.Reader) { + t.Helper() + conn, err := net.DialTimeout("tcp", address, time.Second) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = conn.Close() }) + _ = conn.SetDeadline(time.Now().Add(2 * time.Second)) + return conn, bufio.NewReader(conn) +} + +func webSocketIntegrationReadResponse(t *testing.T, reader *bufio.Reader, request *http.Request) *http.Response { + t.Helper() + response, err := http.ReadResponse(reader, request) + if err != nil { + t.Fatal(err) + } + return response +} + +func webSocketIntegrationAccept(key string) string { + digest := sha1.Sum([]byte(strings.Trim(key, " \t") + "258EAFA5-E914-47DA-95CA-C5AB0DC85B11")) + return base64.StdEncoding.EncodeToString(digest[:]) +} + +func webSocketIntegrationUpgradeHeader(key string) http.Header { + return http.Header{ + "Connection": {"Upgrade, X-Origin-Hop, Authorization, Proxy-Authorization"}, "Upgrade": {"websocket"}, + "Sec-Websocket-Accept": {webSocketIntegrationAccept(key)}, "Sec-Websocket-Protocol": {"chat"}, + "Sec-Websocket-Extensions": {"permessage-deflate"}, "Set-Cookie": {"website=retained; HttpOnly"}, + "X-End-To-End": {"retained response"}, "X-Origin-Hop": {"fictional response hop"}, + "Authorization": {"fictional origin credential"}, "Proxy-Authorization": {"fictional proxy credential"}, + "Proxy-Authenticate": {"Basic realm=fictional"}, "Keep-Alive": {"timeout=5"}, "Proxy-Connection": {"keep-alive"}, + } +} + +func webSocketIntegrationWriteUpgrade(conn net.Conn, request *http.Request, edit func(http.Header)) error { + header := webSocketIntegrationUpgradeHeader(request.Header.Get("Sec-WebSocket-Key")) + if edit != nil { + edit(header) + } + if _, err := io.WriteString(conn, "HTTP/1.1 101 Switching Protocols\r\n"); err != nil { + return err + } + if err := header.Write(conn); err != nil { + return err + } + _, err := io.WriteString(conn, "\r\n") + return err +} + +func webSocketIntegrationWriteFrame(w io.Writer, payload string, masked bool) error { + if len(payload) > 125 { + return errors.New("fixture frame exceeds small-frame boundary") + } + frame := []byte{0x81, byte(len(payload))} + data := []byte(payload) + if masked { + frame[1] |= 0x80 + mask := []byte{0x13, 0x37, 0x42, 0x81} + frame = append(frame, mask...) + for index := range data { + data[index] ^= mask[index%4] + } + } + frame = append(frame, data...) + for len(frame) != 0 { + n, err := w.Write(frame) + if err != nil { + return err + } + if n == 0 { + return io.ErrShortWrite + } + frame = frame[n:] + } + return nil +} + +func webSocketIntegrationReadFrame(r io.Reader, wantMasked bool) (string, error) { + var header [2]byte + if _, err := io.ReadFull(r, header[:]); err != nil { + return "", err + } + if header[0] != 0x81 || (header[1]&0x80 != 0) != wantMasked || header[1]&0x7f > 125 { + return "", fmt.Errorf("unexpected complete frame header %x, want masked=%t", header, wantMasked) + } + var mask [4]byte + if wantMasked { + if _, err := io.ReadFull(r, mask[:]); err != nil { + return "", err + } + } + payload := make([]byte, int(header[1]&0x7f)) + if _, err := io.ReadFull(r, payload); err != nil { + return "", err + } + if wantMasked { + for index := range payload { + payload[index] ^= mask[index%4] + } + } + return string(payload), nil +} + +func webSocketIntegrationAssertFixedOrigin(t *testing.T, request *http.Request, target *url.URL) { + t.Helper() + if request.Host != target.Host || request.URL.Path != "/base/socket" || request.URL.RawQuery != "operator=one&client=one" { + t.Errorf("fixed upstream Host/path/query = %q/%q/%q, want %q/base/socket?operator=one&client=one", request.Host, request.URL.Path, request.URL.RawQuery, target.Host) + } +} + +func webSocketIntegrationAssertUpgrade(t *testing.T, header http.Header, key string) { + t.Helper() + if !reflect.DeepEqual(header.Values("Connection"), []string{"Upgrade"}) || !reflect.DeepEqual(header.Values("Upgrade"), []string{"websocket"}) || !reflect.DeepEqual(header.Values("Sec-WebSocket-Accept"), []string{webSocketIntegrationAccept(key)}) { + t.Errorf("normalized verified upgrade response = %v", header) + } + webSocketIntegrationAssertScrubbed(t, header, "X-Origin-Hop") +} + +func webSocketIntegrationAssertScrubbed(t *testing.T, header http.Header, extra ...string) { + t.Helper() + for _, name := range append([]string{"Authorization", "Proxy-Authorization", "Proxy-Authenticate", "Proxy-Connection", "Keep-Alive", "Te", "Trailer", "Transfer-Encoding"}, extra...) { + if values := header.Values(name); len(values) != 0 { + t.Errorf("unsafe %s survived: %v", name, values) + } + } +} + +func webSocketIntegrationPeerClosed(err error) bool { + return errors.Is(err, io.EOF) || errors.Is(err, syscall.ECONNRESET) +} + +func webSocketIntegrationTake[T any](t *testing.T, channel <-chan T, name string) T { + t.Helper() + select { + case value := <-channel: + return value + case <-time.After(2 * time.Second): + t.Fatalf("timed out waiting for %s", name) + var zero T + return zero + } +} + +func webSocketIntegrationJoin(t *testing.T, joined <-chan struct{}, name string) { + t.Helper() + select { + case <-joined: + case <-time.After(2 * time.Second): + t.Errorf("%s did not independently join", name) + } +} + +func webSocketIntegrationAssertOriginFinished(t *testing.T, origin *webSocketIntegrationOrigin) { + t.Helper() + if err := webSocketIntegrationTake(t, origin.results, "origin result before fixture cleanup"); err != nil { + t.Errorf("origin worker result = %v", err) + } + webSocketIntegrationJoin(t, origin.joined, "origin before fixture cleanup") +} diff --git a/internal/cover/websocket_test.go b/internal/cover/websocket_test.go new file mode 100644 index 0000000..d7e38bc --- /dev/null +++ b/internal/cover/websocket_test.go @@ -0,0 +1,539 @@ +package cover + +import ( + "bufio" + "bytes" + "crypto/sha1" + "encoding/base64" + "errors" + "fmt" + "io" + "net" + "net/http" + "net/http/httptest" + "net/http/httputil" + "net/url" + "reflect" + "strings" + "sync" + "sync/atomic" + "testing" +) + +const websocketUnitKey = "dGhlIHNhbXBsZSBub25jZQ==" + +func websocketUnitRequest() *http.Request { + request := httptest.NewRequest(http.MethodGet, "http://visitor.invalid/socket", nil) + request.Header.Set("Connection", "Upgrade") + request.Header.Set("Upgrade", "websocket") + request.Header.Set("Sec-WebSocket-Version", "13") + request.Header.Set("Sec-WebSocket-Key", websocketUnitKey) + return request +} + +func TestWebsocketRequestEligibilityStrictBoundary(t *testing.T) { + set := func(name, value string) func(*http.Request) { + return func(r *http.Request) { r.Header.Set(name, value) } + } + tests := []struct { + name string + change func(*http.Request) + want bool + }{ + {"valid", func(*http.Request) {}, true}, + {"nil_body", func(r *http.Request) { r.Body = nil }, true}, + {"upgrade_case", set("Upgrade", "WebSocket"), true}, + {"key_ows", set("Sec-WebSocket-Key", "\t "+websocketUnitKey+" \t"), true}, + {"version_ows", set("Sec-WebSocket-Version", " 13\t"), true}, + {"safe_nominations", set("Connection", "keep-alive, Upgrade, Authorization, Proxy-Authorization, X-Private"), true}, + {"h10", func(r *http.Request) { r.ProtoMinor = 0 }, false}, + {"h12", func(r *http.Request) { r.ProtoMinor = 2 }, false}, + {"h2", func(r *http.Request) { r.ProtoMajor, r.ProtoMinor = 2, 0 }, false}, + {"h3", func(r *http.Request) { r.ProtoMajor, r.ProtoMinor = 3, 0 }, false}, + {"post", func(r *http.Request) { r.Method = http.MethodPost }, false}, + {"connect", func(r *http.Request) { r.Method = http.MethodConnect }, false}, + {"body_even_zero_length", func(r *http.Request) { r.Body = io.NopCloser(strings.NewReader("")) }, false}, + {"positive_length", func(r *http.Request) { r.ContentLength = 1 }, false}, + {"unknown_length", func(r *http.Request) { r.ContentLength = -1 }, false}, + {"transfer_field", func(r *http.Request) { r.TransferEncoding = []string{"chunked"} }, false}, + {"transfer_header", set("Transfer-Encoding", "chunked"), false}, + {"trailer", func(r *http.Request) { r.Trailer = http.Header{"X-Late": {"value"}} }, false}, + {"missing_upgrade", func(r *http.Request) { r.Header.Del("Upgrade") }, false}, + {"h2c", set("Upgrade", "h2c"), false}, + {"upgrade_list", set("Upgrade", "websocket, h2c"), false}, + {"duplicate_upgrade", func(r *http.Request) { r.Header.Add("Upgrade", "websocket") }, false}, + {"unicode_upgrade_fold", set("Upgrade", "webſocket"), false}, + {"missing_connection", func(r *http.Request) { r.Header.Del("Connection") }, false}, + {"connection_without_upgrade", set("Connection", "keep-alive"), false}, + {"duplicate_connection", func(r *http.Request) { r.Header.Add("Connection", "Upgrade") }, false}, + {"duplicate_token", set("Connection", "Upgrade, uPgRaDe"), false}, + {"duplicate_other_token", set("Connection", "Upgrade, X-Private, x-private"), false}, + {"close", set("Connection", "Upgrade, close"), false}, + {"empty_token", set("Connection", "Upgrade,,X-Private"), false}, + {"invalid_token", set("Connection", "Upgrade, X Private"), false}, + {"unicode_token_fold", set("Connection", "Upgrade, X-Key"), false}, + {"missing_version", func(r *http.Request) { r.Header.Del("Sec-WebSocket-Version") }, false}, + {"version12", set("Sec-WebSocket-Version", "12"), false}, + {"version_list", set("Sec-WebSocket-Version", "13, 12"), false}, + {"duplicate_version", func(r *http.Request) { r.Header.Add("Sec-WebSocket-Version", "13") }, false}, + {"missing_key", func(r *http.Request) { r.Header.Del("Sec-WebSocket-Key") }, false}, + {"key_bad_encoding", set("Sec-WebSocket-Key", "not-base64"), false}, + {"key15", set("Sec-WebSocket-Key", base64.StdEncoding.EncodeToString(make([]byte, 15))), false}, + {"key17", set("Sec-WebSocket-Key", base64.StdEncoding.EncodeToString(make([]byte, 17))), false}, + {"noncanonical_padding_bits", set("Sec-WebSocket-Key", "AAAAAAAAAAAAAAAAAAAAAB=="), false}, + {"key_embedded_newline", set("Sec-WebSocket-Key", websocketUnitKey[:8]+"\n"+websocketUnitKey[8:]), false}, + {"key_unicode_ows", set("Sec-WebSocket-Key", "\u00a0"+websocketUnitKey), false}, + {"duplicate_key", func(r *http.Request) { r.Header.Add("Sec-WebSocket-Key", websocketUnitKey) }, false}, + {"case_alias_duplicate_key", func(r *http.Request) { r.Header["sec-websocket-key"] = []string{websocketUnitKey} }, false}, + {"unicode_key_field_alias", func(r *http.Request) { + r.Header.Del("Sec-WebSocket-Key") + r.Header["Sec-WebSocket-Key"] = []string{websocketUnitKey} + }, false}, + } + for _, name := range []string{"Sec-WebSocket-Key", "Sec-WebSocket-Version", "Sec-WebSocket-Accept", "Sec-WebSocket-Protocol", "Sec-WebSocket-Extensions"} { + tests = append(tests, struct { + name string + change func(*http.Request) + want bool + }{"nominated_" + name, set("Connection", "Upgrade, "+name), false}) + } + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + request := websocketUnitRequest() + test.change(request) + info := websocketRequestEligibility(request) + if info.eligible != test.want || test.want && info.key != websocketUnitKey { + t.Fatalf("eligibility = %+v, want eligible=%t", info, test.want) + } + }) + } + if websocketRequestEligibility(nil).eligible { + t.Fatal("nil request was eligible") + } +} + +type websocketUnitDuplex struct { + closes atomic.Int64 + writes atomic.Int64 + halfCloses atomic.Int64 +} + +func (*websocketUnitDuplex) Read([]byte) (int, error) { return 0, io.EOF } +func (b *websocketUnitDuplex) Write(p []byte) (int, error) { + b.writes.Add(int64(len(p))) + return len(p), nil +} +func (b *websocketUnitDuplex) Close() error { b.closes.Add(1); return nil } +func (b *websocketUnitDuplex) CloseWrite() error { b.halfCloses.Add(1); return nil } + +func websocketUnitAccept(key string) string { + sum := sha1.Sum([]byte(key + websocketGUID)) + return base64.StdEncoding.EncodeToString(sum[:]) +} + +func websocketUnitResponse(request *http.Request, body io.ReadCloser) *http.Response { + return &http.Response{ + StatusCode: http.StatusSwitchingProtocols, Proto: "HTTP/1.1", ProtoMajor: 1, ProtoMinor: 1, + Header: http.Header{"Connection": {"Upgrade"}, "Upgrade": {"websocket"}, "Sec-Websocket-Accept": {websocketUnitAccept(websocketUnitKey)}}, + Body: body, Request: request, + } +} + +func TestWebsocketResponseValidationAndScrubPreservesDuplex(t *testing.T) { + request := websocketUnitRequest() + proxyHandler, err := NewReverseProxyHandler(&url.URL{Scheme: "http", Host: "owned.invalid"}, nil) + if err != nil { + t.Fatal(err) + } + // Public constructor intentionally still returns the native ReverseProxy. + proxy, ok := proxyHandler.(*httputil.ReverseProxy) + if !ok { + t.Fatalf("handler type = %T", proxyHandler) + } + set := func(name, value string) func(*http.Response) { + return func(r *http.Response) { r.Header.Set(name, value) } + } + tests := []struct { + name string + change func(*http.Response) + want bool + }{ + {"valid", func(*http.Response) {}, true}, + {"safe_nominations", set("Connection", "keep-alive, Upgrade, X-Private, Authorization, Proxy-Authorization"), true}, + {"without_closewrite", func(r *http.Response) { + r.Body = struct{ io.ReadWriteCloser }{r.Body.(io.ReadWriteCloser)} + }, true}, + {"missing_request", func(r *http.Response) { r.Request = nil }, false}, + {"untrusted_request", func(r *http.Response) { r.Request = websocketUnitRequest() }, false}, + {"h10", func(r *http.Response) { r.ProtoMinor = 0 }, false}, + {"h2", func(r *http.Response) { r.ProtoMajor = 2 }, false}, + {"nil_body", func(r *http.Response) { r.Body = nil }, false}, + {"read_only_body", func(r *http.Response) { r.Body = io.NopCloser(strings.NewReader("ordinary")) }, false}, + {"missing_accept", func(r *http.Response) { r.Header.Del("Sec-WebSocket-Accept") }, false}, + {"wrong_accept", set("Sec-WebSocket-Accept", "private origin failure"), false}, + {"duplicate_accept", func(r *http.Response) { r.Header.Add("Sec-WebSocket-Accept", websocketUnitAccept(websocketUnitKey)) }, false}, + {"accept_case_alias", func(r *http.Response) { + r.Header["sec-websocket-accept"] = []string{websocketUnitAccept(websocketUnitKey)} + }, false}, + {"unicode_accept_field_alias", func(r *http.Response) { + r.Header.Del("Sec-WebSocket-Accept") + r.Header["Sec-WebsocKet-Accept"] = []string{websocketUnitAccept(websocketUnitKey)} + }, false}, + {"unicode_long_s_accept_field_alias", func(r *http.Response) { + r.Header.Del("Sec-WebSocket-Accept") + r.Header["ſec-WebSocket-Accept"] = []string{websocketUnitAccept(websocketUnitKey)} + }, false}, + {"missing_upgrade", func(r *http.Response) { r.Header.Del("Upgrade") }, false}, + {"unrelated_upgrade", set("Upgrade", "h2c"), false}, + {"upgrade_list", set("Upgrade", "websocket, h2c"), false}, + {"missing_connection", func(r *http.Response) { r.Header.Del("Connection") }, false}, + {"close", set("Connection", "Upgrade, close"), false}, + {"duplicate_upgrade_token", set("Connection", "Upgrade, upgrade"), false}, + {"empty_token", set("Connection", "Upgrade,"), false}, + {"nominated_accept", set("Connection", "Upgrade, Sec-WebSocket-Accept"), false}, + {"nominated_protocol", set("Connection", "Upgrade, Sec-WebSocket-Protocol"), false}, + {"nominated_extensions", set("Connection", "Upgrade, Sec-WebSocket-Extensions"), false}, + } + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + trusted := websocketResponseRequest(withWebsocketRequestEligibility(request, websocketRequestEligibility(request))) + body := &websocketUnitDuplex{} + response := websocketUnitResponse(trusted, body) + response.Header["authorization"] = []string{"private credential"} + response.Header["proxy-authorization"] = []string{"private proxy credential"} + response.Header["x-private"] = []string{"private hop"} + response.Header.Set("Sec-WebSocket-Protocol", "chat") + response.Trailer = http.Header{"connection": {"X-Late, Authorization"}, "x-late": {"private trailer"}, "authorization": {"private credential"}, "X-End-To-End": {"kept"}} + test.change(response) + err := proxy.ModifyResponse(response) + if (err == nil) != test.want { + t.Fatalf("ModifyResponse = %v, want valid=%t", err, test.want) + } + if !test.want { + return // ReverseProxy, not ModifyResponse itself, owns error Close. + } + duplex, ok := response.Body.(io.ReadWriteCloser) + if !ok { + t.Fatal("duplex capability was hidden") + } + if n, err := duplex.Write([]byte("guarded")); err != nil || n != 7 || body.writes.Load() != 7 { + t.Fatal("duplex Write delegation changed") + } + if n, err := duplex.Read(make([]byte, 1)); n != 0 || err != io.EOF { + t.Fatal("duplex Read delegation changed") + } + halfClose, ok := response.Body.(interface{ CloseWrite() error }) + if ok != (test.name != "without_closewrite") { + t.Fatal("CloseWrite capability was lost or falsely advertised") + } + if ok { + if err := halfClose.CloseWrite(); err != nil || body.halfCloses.Load() != 1 { + t.Fatal("CloseWrite delegation changed") + } + } + if response.Header.Get("Connection") != "Upgrade" || response.Header.Get("Upgrade") != "websocket" || response.Header.Get("Sec-WebSocket-Protocol") != "chat" { + t.Fatalf("normalized/safe headers = %v", response.Header) + } + for _, name := range []string{"Authorization", "Proxy-Authorization"} { + if len(headerValuesFold(response.Header, name)) != 0 { + t.Fatalf("%s leaked: %v", name, response.Header) + } + } + if test.name == "safe_nominations" && len(headerValuesFold(response.Header, "X-Private")) != 0 { + t.Fatal("nominated case alias leaked") + } + if len(response.Trailer) != 1 || response.Trailer.Get("X-End-To-End") != "kept" { + t.Fatalf("unsafe trailer survived: %v", response.Trailer) + } + _ = duplex.Close() + closeWebsocketResponse(trusted) + _ = duplex.Close() + if body.closes.Load() != 1 { + t.Fatal("accepted body Close was not once-only") + } + }) + } +} + +type websocketUnitHijackErrorWriter struct { + *httptest.ResponseRecorder + err error +} + +func (w *websocketUnitHijackErrorWriter) Hijack() (net.Conn, *bufio.ReadWriter, error) { + return nil, nil, w.err +} + +func TestWebsocketAccepted101OwnsPreHijackFailures(t *testing.T) { + for _, mode := range []string{"non_hijacker", "unsupported_hijack", "other_hijack_error"} { + t.Run(mode, func(t *testing.T) { + body := &websocketUnitDuplex{} + response := websocketUnitResponse(nil, body) + handler, err := NewReverseProxyHandler(&url.URL{Scheme: "http", Host: "owned.invalid"}, roundTripFunc(func(*http.Request) (*http.Response, error) { return response, nil })) + if err != nil { + t.Fatal(err) + } + recorder := httptest.NewRecorder() + var writer http.ResponseWriter = recorder + if mode == "unsupported_hijack" { + writer = &websocketUnitHijackErrorWriter{recorder, http.ErrNotSupported} + } else if mode == "other_hijack_error" { + writer = &websocketUnitHijackErrorWriter{recorder, errors.New("private hijack error")} + } + handler.ServeHTTP(writer, websocketUnitRequest()) + // Before explicit cleanup: even the early unsupported-Hijack path + // must close the accepted original backend exactly once. + if recorder.Code != http.StatusBadGateway || recorder.Body.String() != "Bad Gateway\n" || body.closes.Load() != 1 { + t.Fatalf("status/body/actual Close=%d/%q/%d", recorder.Code, recorder.Body.String(), body.closes.Load()) + } + // Generic Hijack errors install the standard cancellation closer; + // concurrent callers still close the original backend only once. + var wg sync.WaitGroup + for range 12 { + wg.Add(1) + go func() { defer wg.Done(); _ = response.Body.Close() }() + } + wg.Wait() + if body.closes.Load() != 1 { + t.Fatal("cancellation/error cleanup double-closed the original body") + } + }) + } +} + +type websocketUnitReadBody struct{ closes atomic.Int64 } + +func (*websocketUnitReadBody) Read([]byte) (int, error) { return 0, io.EOF } +func (b *websocketUnitReadBody) Close() error { b.closes.Add(1); return nil } + +func TestWebsocketRejected101HasGenericErrorAndOwnsBodyClose(t *testing.T) { + for _, test := range []string{"unexpected", "wrong_accept", "nonduplex", "nil_body"} { + t.Run(test, func(t *testing.T) { + duplex := &websocketUnitDuplex{} + readOnly := &websocketUnitReadBody{} + response := websocketUnitResponse(nil, duplex) + response.Header.Set("Authorization", "private origin credential") + response.Header.Set("X-Private", "private origin details") + request := websocketUnitRequest() + switch test { + case "unexpected": + request.Header.Del("Upgrade") + case "wrong_accept": + response.Header.Set("Sec-WebSocket-Accept", "private wrong proof") + case "nonduplex": + response.Body = readOnly + case "nil_body": + response.Body = nil + } + handler, err := NewReverseProxyHandler(&url.URL{Scheme: "http", Host: "owned.invalid"}, roundTripFunc(func(*http.Request) (*http.Response, error) { return response, nil })) + if err != nil { + t.Fatal(err) + } + writer := httptest.NewRecorder() + handler.ServeHTTP(writer, request) + if writer.Code != http.StatusBadGateway || writer.Body.String() != "Bad Gateway\n" || writer.Header().Get("Authorization") != "" || writer.Header().Get("X-Private") != "" { + t.Fatalf("private rejection leaked or status changed: %d %v %q", writer.Code, writer.Header(), writer.Body.String()) + } + if test == "nonduplex" && readOnly.closes.Load() != 1 || test != "nonduplex" && test != "nil_body" && duplex.closes.Load() != 1 { + t.Fatalf("upstream Close count duplex/read-only=%d/%d", duplex.closes.Load(), readOnly.closes.Load()) + } + }) + } +} + +func TestWebsocketTransportResponseProvenanceIsRequestLocal(t *testing.T) { + const attempts = 24 + var wg sync.WaitGroup + errorsCh := make(chan string, attempts) + base := roundTripFunc(func(request *http.Request) (*http.Response, error) { + // A custom transport's Request pointer must never establish eligibility. + forged := withWebsocketRequestEligibility(websocketUnitRequest(), websocketRequestInfo{eligible: true, key: websocketUnitKey}) + response := websocketUnitResponse(forged, &websocketUnitDuplex{}) + response.Header.Set("Sec-WebSocket-Accept", websocketUnitAccept(request.Header.Get("Sec-WebSocket-Key"))) + return response, nil + }) + transport := &informationalHeaderTransport{base: base} + for index := range attempts { + wg.Add(1) + go func() { + defer wg.Done() + request := websocketUnitRequest() + keyBytes := make([]byte, 16) + keyBytes[0] = byte(index) + request.Header.Set("Sec-WebSocket-Key", base64.StdEncoding.EncodeToString(keyBytes)) + if index%2 != 0 { + request.Method = http.MethodPost + } + request = withWebsocketRequestEligibility(request, websocketRequestEligibility(request)) + response, err := transport.RoundTrip(request) + if err != nil || response == nil { + errorsCh <- "transport failed" + return + } + valid := validateWebsocketResponse(response) == nil + if valid != (index%2 == 0) { + errorsCh <- "another request or forged Request changed eligibility" + } + }() + } + wg.Wait() + close(errorsCh) + for message := range errorsCh { + t.Error(message) + } + ordinaryRequest := websocketUnitRequest() + ordinary := &http.Response{StatusCode: http.StatusOK, Request: ordinaryRequest, Body: http.NoBody} + transport.base = roundTripFunc(func(*http.Request) (*http.Response, error) { return ordinary, nil }) + response, err := transport.RoundTrip(ordinaryRequest) + if err != nil || response.Request != ordinaryRequest { + t.Fatal("ordinary response.Request contract changed") + } +} + +func TestWebsocketHeaderLookupUsesASCIIFieldNames(t *testing.T) { + for _, alias := range []string{"Sec-WebsocKet-Accept", "ſec-WebSocket-Accept"} { + t.Run(alias, func(t *testing.T) { + if !strings.EqualFold(alias, "Sec-WebSocket-Accept") { + t.Fatal("fixture does not exercise Unicode case folding") + } + header := http.Header{alias: {"invalid-field-proof"}, "X-Ordinary": {"kept"}} + var wire bytes.Buffer + if err := header.Write(&wire); err != nil { + t.Fatal(err) + } + if strings.Contains(wire.String(), "invalid-field-proof") || !strings.Contains(wire.String(), "X-Ordinary: kept") { + t.Fatalf("native HTTP field-name behavior = %q", wire.String()) + } + if _, ok := singleHeaderValue(header, "Sec-WebSocket-Accept"); ok { + t.Fatal("a field discarded by native HTTP writing established handshake proof") + } + header["sec-websocket-accept"] = []string{"ascii-field-proof"} + if value, ok := singleHeaderValue(header, "Sec-WebSocket-Accept"); !ok || value != "ascii-field-proof" { + t.Fatal("valid ASCII field alias stopped qualifying") + } + deleteHeaderFold(header, "Sec-WebSocket-Accept") + if _, exists := header["sec-websocket-accept"]; exists { + t.Fatal("valid ASCII alias was not removed") + } + if _, exists := header[alias]; !exists { + t.Fatal("deletion treated an invalid Unicode field as an ASCII alias") + } + }) + } +} + +func websocketUnitNominationHeaders(fields int) http.Header { + header := make(http.Header, fields) + for index := range fields { + header[fmt.Sprintf("X-%04d", index)] = []string{"a"} + } + return header +} + +func websocketUnitNominationList(tokens int, distinct bool, name string) string { + if !distinct { + return strings.Repeat(name+",", tokens-1) + name + } + names := make([]string, tokens) + for index := range tokens { + names[index] = fmt.Sprintf("m%03x", index) + } + return strings.Join(names, ",") +} + +func TestWebsocketConnectionNominationFilterBounded(t *testing.T) { + for _, test := range []struct { + name string + tokens int + distinct bool + }{ + {"repeated_unmatched", 1600, false}, + {"distinct_unmatched", 1024, true}, + } { + t.Run(test.name, func(t *testing.T) { + header := websocketUnitNominationHeaders(800) + before := header.Clone() + source := http.Header{"Connection": {websocketUnitNominationList(test.tokens, test.distinct, "MiSsInG")}} + sourceBefore := source.Clone() + removeConnectionNominatedHeaders(header, source) + if !reflect.DeepEqual(header, before) || !reflect.DeepEqual(source, sourceBefore) { + t.Fatal("unmatched nominations mutated ordinary fields or a separate source") + } + }) + } + t.Run("ascii_case_aliases_and_trimspace", func(t *testing.T) { + header := http.Header{ + "X-Private": {"one"}, "x-private": {"two"}, "AUTHORIZATION": {"credential"}, + "proxy-authorization": {"proxy credential"}, "Sec-Websocket-Key": {websocketUnitKey}, + "X-Ordinary": {"kept"}, "X-Key": {"invalid name"}, "X-Key": {"ordinary"}, + } + source := http.Header{"connection": {"\u00a0 X-private \u00a0, Authorization, PROXY-AUTHORIZATION, Sec-WebSocket-Key, X-Key,,bad(token"}} + removeConnectionNominatedHeaders(header, source) + want := http.Header{"X-Ordinary": {"kept"}, "X-Key": {"invalid name"}, "X-Key": {"ordinary"}} + if !reflect.DeepEqual(header, want) { + t.Fatalf("ASCII aliases/TrimSpace/invalid Unicode behavior = %v, want %v", header, want) + } + }) + t.Run("same_map_collects_before_deletion", func(t *testing.T) { + header := http.Header{ + "Connection": {"Connection, X-Private"}, "connection": {"Authorization, X-Later"}, + "x-private": {"private"}, "AUTHORIZATION": {"credential"}, "X-Later": {"late"}, "X-Ordinary": {"kept"}, + } + removeConnectionNominatedHeaders(header, header) + if !reflect.DeepEqual(header, http.Header{"X-Ordinary": {"kept"}}) { + t.Fatalf("self-nomination hid later source values: %v", header) + } + }) + t.Run("full_scrub_with_large_unmatched_list", func(t *testing.T) { + header := websocketUnitNominationHeaders(800) + before := header.Clone() + header["connection"] = []string{websocketUnitNominationList(1600, false, "ABSENT")} + header["authorization"] = []string{"credential"} + header["PROXY-AUTHORIZATION"] = []string{"proxy credential"} + header["Upgrade"] = []string{"websocket"} + removeUnsafeHeaders(header) + if !reflect.DeepEqual(header, before) { + t.Fatal("unmatched nominations bypassed credential/hop scrubbing or removed ordinary fields") + } + }) + t.Run("long_names_keep_owned_folded_keys", func(t *testing.T) { + name := "X-" + strings.Repeat("A", 96) + header := http.Header{name: {"private"}, "X-Other": {"kept"}} + source := http.Header{"Connection": {name + ",x-other"}} + removeConnectionNominatedHeaders(header, source) + if len(header) != 0 { + t.Fatal("resizing or reusing fold scratch changed a stored nomination") + } + }) +} + +// Local helper microbenchmark only: no listener, request, transport or network +// traffic. Header/list setup and semantic checks are outside measured loops. +func BenchmarkWebsocketConnectionNominationFilter(b *testing.B) { + for _, test := range []struct { + name string + fields int + tokens int + distinct bool + token string + }{ + {"repeated_absent_32x1", 32, 1, false, "q"}, + {"repeated_absent_800x1600", 800, 1600, false, "q"}, + {"repeated_mixed_case_800x1600", 800, 1600, false, "MiSsInG"}, + {"distinct_absent_800x1024", 800, 1024, true, ""}, + } { + b.Run(test.name, func(b *testing.B) { + header := websocketUnitNominationHeaders(test.fields) + source := http.Header{"Connection": {websocketUnitNominationList(test.tokens, test.distinct, test.token)}} + b.ReportAllocs() + b.ResetTimer() + for range b.N { + removeConnectionNominatedHeaders(header, source) + } + b.StopTimer() + if len(header) != test.fields { + b.Fatal("benchmark's absent nominations changed header fields") + } + }) + } +} diff --git a/internal/tunnel/web_h2.go b/internal/tunnel/web_h2.go index fdc9bd5..e23359a 100644 --- a/internal/tunnel/web_h2.go +++ b/internal/tunnel/web_h2.go @@ -59,6 +59,7 @@ type WebH2ServerConfig struct { type WebH2Server struct { listener net.Listener server *http.Server + tcpOwner *webTCPConnectionOwner serveMu sync.Mutex serving bool @@ -163,9 +164,12 @@ func listenWebH2WithCore(config WebH2ServerConfig, core *serverCore, auth *webAu if err != nil { return nil, fmt.Errorf("tunnel: listen web-cover TCP: %w", err) } + owner := newWebTCPConnectionOwner() + httpServer.BaseContext = func(net.Listener) context.Context { return owner.ctx } return &WebH2Server{ - listener: tls.NewListener(&webAdmissionListener{Listener: raw, admission: connectionAdmission}, tlsConfig), + listener: tls.NewListener(&webAdmissionListener{Listener: raw, admission: connectionAdmission, owner: owner}, tlsConfig), server: httpServer, + tcpOwner: owner, }, nil } @@ -196,9 +200,11 @@ func (s *WebH2Server) Serve(ctx context.Context) error { return fmt.Errorf("tunnel: serve web-cover HTTP/2: %w", err) } -// Close stops the listener and aborts active HTTP streams. +// Close stops the listener and aborts active HTTP streams and hijacked cover +// connections. It does not close a shared destination dialer or cover transport. func (s *WebH2Server) Close() error { s.closeOnce.Do(func() { + ownerErr := s.tcpOwner.close() serverErr := s.server.Close() listenerErr := s.listener.Close() if errors.Is(serverErr, http.ErrServerClosed) || errors.Is(serverErr, net.ErrClosed) { @@ -207,7 +213,7 @@ func (s *WebH2Server) Close() error { if errors.Is(listenerErr, net.ErrClosed) { listenerErr = nil } - s.closeErr = errors.Join(serverErr, listenerErr) + s.closeErr = errors.Join(ownerErr, serverErr, listenerErr) }) return s.closeErr } diff --git a/internal/tunnel/web_limits.go b/internal/tunnel/web_limits.go index 6d4bc19..8d01e32 100644 --- a/internal/tunnel/web_limits.go +++ b/internal/tunnel/web_limits.go @@ -68,6 +68,7 @@ func (a *webConnectionAdmission) acquire(remote net.Addr) (func(), bool) { type webAdmissionListener struct { net.Listener admission *webConnectionAdmission + owner *webTCPConnectionOwner } func (l *webAdmissionListener) Accept() (net.Conn, error) { @@ -81,19 +82,32 @@ func (l *webAdmissionListener) Accept() (net.Conn, error) { _ = conn.Close() continue } - return &webAdmissionConn{Conn: conn, release: release}, nil + owned := &webAdmissionConn{Conn: conn, release: release, owner: l.owner} + if l.owner != nil && !l.owner.register(owned) { + _ = owned.Close() + return nil, net.ErrClosed + } + return owned, nil } } type webAdmissionConn struct { net.Conn - release func() + release func() + owner *webTCPConnectionOwner + once sync.Once + closeErr error } func (c *webAdmissionConn) Close() error { - err := c.Conn.Close() - c.release() - return err + c.once.Do(func() { + c.closeErr = c.Conn.Close() + c.release() + if c.owner != nil { + c.owner.remove(c) + } + }) + return c.closeErr } // webAdmissionQUICListener adapts a quic.Listener to http3.QUICListener while diff --git a/internal/tunnel/web_tcp_owner.go b/internal/tunnel/web_tcp_owner.go new file mode 100644 index 0000000..2dafc3a --- /dev/null +++ b/internal/tunnel/web_tcp_owner.go @@ -0,0 +1,59 @@ +package tunnel + +import ( + "context" + "errors" + "sync" +) + +// webTCPConnectionOwner keeps accepted physical sockets owned after net/http +// relinquishes a hijacked connection. It never reads application bytes and is +// independent of the shared destination dialer and cover transport. +type webTCPConnectionOwner struct { + ctx context.Context + cancel context.CancelFunc + mu sync.Mutex + closed bool + conns map[*webAdmissionConn]struct{} +} + +func newWebTCPConnectionOwner() *webTCPConnectionOwner { + ctx, cancel := context.WithCancel(context.Background()) + return &webTCPConnectionOwner{ctx: ctx, cancel: cancel, conns: make(map[*webAdmissionConn]struct{})} +} + +func (o *webTCPConnectionOwner) register(conn *webAdmissionConn) bool { + o.mu.Lock() + defer o.mu.Unlock() + if o.closed { + return false + } + o.conns[conn] = struct{}{} + return true +} + +func (o *webTCPConnectionOwner) remove(conn *webAdmissionConn) { + o.mu.Lock() + delete(o.conns, conn) + o.mu.Unlock() +} + +func (o *webTCPConnectionOwner) close() error { + o.mu.Lock() + o.closed = true + conns := make([]*webAdmissionConn, 0, len(o.conns)) + for conn := range o.conns { + conns = append(conns, conn) + } + o.mu.Unlock() + // Cancellation can invoke transport callbacks; Close can unblock handlers + // which unregister. Neither operation may execute under the registry lock. + o.cancel() + var closeErrors []error + for _, conn := range conns { + if err := conn.Close(); err != nil { + closeErrors = append(closeErrors, err) + } + } + return errors.Join(closeErrors...) +} diff --git a/internal/tunnel/web_tcp_owner_test.go b/internal/tunnel/web_tcp_owner_test.go new file mode 100644 index 0000000..b73d8ee --- /dev/null +++ b/internal/tunnel/web_tcp_owner_test.go @@ -0,0 +1,202 @@ +package tunnel + +import ( + "errors" + "io" + "net" + "sync" + "sync/atomic" + "syscall" + "testing" + "time" +) + +type webTCPOwnerObservedConn struct { + net.Conn + closes atomic.Int64 + check func() +} + +func (c *webTCPOwnerObservedConn) Close() error { + c.closes.Add(1) + if c.check != nil { + c.check() + } + return c.Conn.Close() +} + +func webTCPOwnerRealPair(t *testing.T) (*webTCPOwnerObservedConn, net.Conn) { + t.Helper() + listener, err := net.ListenTCP("tcp", &net.TCPAddr{IP: net.ParseIP("127.0.0.1")}) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = listener.Close() }) + peer, err := net.DialTimeout("tcp", listener.Addr().String(), time.Second) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = peer.Close() }) + _ = listener.SetDeadline(time.Now().Add(time.Second)) + raw, err := listener.AcceptTCP() + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = raw.Close() }) + return &webTCPOwnerObservedConn{Conn: raw}, peer +} + +func TestWebTCPConnectionOwnerClosesOnceOutsideLock(t *testing.T) { + owner := newWebTCPConnectionOwner() + t.Cleanup(func() { _ = owner.close() }) + admission, err := newWebConnectionAdmission(1, 1) + if err != nil { + t.Fatal(err) + } + raw, peer := webTCPOwnerRealPair(t) + var lockHeld atomic.Bool + raw.check = func() { + if !owner.mu.TryLock() { + lockHeld.Store(true) + return + } + owner.mu.Unlock() + } + // The cancellation callback executes synchronously, making the lock + // assertion independent of scheduler timing of context.AfterFunc. + cancel := owner.cancel + owner.cancel = func() { + if !owner.mu.TryLock() { + lockHeld.Store(true) + } else { + owner.mu.Unlock() + } + cancel() + } + release, ok := admission.acquire(raw.RemoteAddr()) + if !ok { + t.Fatal("real TCP admission failed") + } + conn := &webAdmissionConn{Conn: raw, release: release, owner: owner} + if !owner.register(conn) { + t.Fatal("live owner rejected registration") + } + // First exercise owner shutdown without lock contention, then duplicate + // normal/owner Close concurrently. Every call must retain the first result. + if err := owner.close(); err != nil { + t.Fatal(err) + } + if lockHeld.Load() { + t.Error("owner canceled or closed the raw socket under its registry lock") + } + var workers sync.WaitGroup + for range 8 { + workers.Add(1) + go func() { defer workers.Done(); _ = conn.Close(); _ = owner.close() }() + } + joined := make(chan struct{}) + go func() { workers.Wait(); close(joined) }() + websocketLifecycleJoin(t, "concurrent duplicate owner/connection closes", joined) + if raw.closes.Load() != 1 || len(admission.slots) != 0 || owner.ctx.Err() == nil { + t.Errorf("raw closes=%d slots=%d owner context=%v", raw.closes.Load(), len(admission.slots), owner.ctx.Err()) + } + owner.mu.Lock() + registered := len(owner.conns) + owner.mu.Unlock() + if registered != 0 { + t.Errorf("closed socket retained %d registry entries", registered) + } + _ = peer.SetReadDeadline(time.Now().Add(time.Second)) + if _, err := peer.Read(make([]byte, 1)); !errors.Is(err, io.EOF) && !errors.Is(err, syscall.ECONNRESET) { + t.Errorf("real peer did not observe close before cleanup: %v", err) + } +} + +type webTCPOwnerDelayedListener struct { + net.Listener + entered chan struct{} + release <-chan struct{} + raw chan *webTCPOwnerObservedConn +} + +func (l *webTCPOwnerDelayedListener) Accept() (net.Conn, error) { + conn, err := l.Listener.Accept() + if err != nil { + return nil, err + } + raw := &webTCPOwnerObservedConn{Conn: conn} + l.raw <- raw + close(l.entered) + <-l.release + return raw, nil +} + +func TestWebTCPConnectionOwnerRejectsLateAcceptedRegistration(t *testing.T) { + owner := newWebTCPConnectionOwner() + admission, err := newWebConnectionAdmission(1, 1) + if err != nil { + t.Fatal(err) + } + listener, err := net.ListenTCP("tcp", &net.TCPAddr{IP: net.ParseIP("127.0.0.1")}) + if err != nil { + t.Fatal(err) + } + entered, release, joined := make(chan struct{}), make(chan struct{}), make(chan struct{}) + var once sync.Once + ungate := func() { once.Do(func() { close(release) }) } + observed := make(chan *webTCPOwnerObservedConn, 1) + acceptResult := make(chan error, 1) + delayed := &webTCPOwnerDelayedListener{Listener: listener, entered: entered, release: release, raw: observed} + owned := &webAdmissionListener{Listener: delayed, admission: admission, owner: owner} + t.Cleanup(func() { + ungate() + _ = listener.Close() + _ = owner.close() + websocketLifecycleJoin(t, "late raw Accept cleanup", joined) + }) + go func() { + defer close(joined) + conn, err := owned.Accept() + if conn != nil { + _ = conn.Close() + acceptResult <- errors.New("closed owner delivered a late accepted connection") + return + } + acceptResult <- err + }() + peer, err := net.DialTimeout("tcp", listener.Addr().String(), time.Second) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = peer.Close() }) + select { + case <-entered: + case <-time.After(time.Second): + t.Fatal("real raw Accept did not reach pre-registration gate") + } + raw := <-observed // Buffered publication happens before entered closes. + t.Cleanup(func() { _ = raw.Conn.Close() }) + if len(admission.slots) != 0 { + t.Fatal("gate did not precede admission and ownership registration") + } + if err := owner.close(); err != nil { + t.Fatal(err) + } + ungate() + websocketLifecycleJoin(t, "late raw Accept", joined) + select { + case err := <-acceptResult: + if !errors.Is(err, net.ErrClosed) { + t.Errorf("late Accept error=%v, want net.ErrClosed", err) + } + default: + t.Error("late Accept did not publish result") + } + if raw.closes.Load() != 1 || len(admission.slots) != 0 { + t.Errorf("late raw closes=%d slots=%d", raw.closes.Load(), len(admission.slots)) + } + _ = peer.SetReadDeadline(time.Now().Add(time.Second)) + if _, err := peer.Read(make([]byte, 1)); !errors.Is(err, io.EOF) && !errors.Is(err, syscall.ECONNRESET) { + t.Errorf("late peer did not observe physical close before cleanup: %v", err) + } +} diff --git a/internal/tunnel/websocket_combined_test.go b/internal/tunnel/websocket_combined_test.go new file mode 100644 index 0000000..912356f --- /dev/null +++ b/internal/tunnel/websocket_combined_test.go @@ -0,0 +1,299 @@ +package tunnel + +import ( + "bufio" + "context" + "crypto/tls" + "errors" + "fmt" + "io" + "net" + "net/http" + "net/url" + "sync" + "sync/atomic" + "syscall" + "testing" + "time" + + "github.com/cppla/autocar/internal/cover" + "github.com/cppla/autocar/internal/transport" +) + +// Exercise the actual cover constructor through the combined server, including +// ResponseController/Unwrap, across both public TLS versions. Closing the +// server aborts upgrades; it does not negotiate a WebSocket close handshake. +func TestWebSocketCoverCombinedTLSAndOwnerClose(t *testing.T) { + for _, version := range []uint16{tls.VersionTLS12, tls.VersionTLS13} { + for _, stop := range []string{"close", "serve-cancel"} { + t.Run(fmt.Sprintf("tls%d/%s", version, stop), func(t *testing.T) { + runCombinedWebSocketOwnerClose(t, version, stop) + }) + } + } +} + +func runCombinedWebSocketOwnerClose(t *testing.T, version uint16, stop string) { + t.Helper() + origin, err := net.ListenTCP("tcp", &net.TCPAddr{IP: net.ParseIP("127.0.0.1")}) + if err != nil { + t.Fatal(err) + } + var originMu sync.Mutex + var ownedOrigin net.Conn + originClosed := false + originResult := make(chan error, 1) + originJoined := make(chan struct{}) + t.Cleanup(func() { + _ = origin.Close() + originMu.Lock() + originClosed = true + if ownedOrigin != nil { + _ = ownedOrigin.Close() + } + originMu.Unlock() + select { + case <-originJoined: + case <-time.After(2 * time.Second): + t.Error("origin Accept/frame worker did not join after independent cleanup") + } + }) + go func() { + defer close(originJoined) + _ = origin.SetDeadline(time.Now().Add(3 * time.Second)) + raw, err := origin.AcceptTCP() + if err != nil { + originResult <- err + return + } + defer raw.Close() + originMu.Lock() + if originClosed { + originMu.Unlock() + originResult <- raw.Close() + return + } + ownedOrigin = raw + originMu.Unlock() + _ = raw.SetDeadline(time.Now().Add(3 * time.Second)) + reader := bufio.NewReader(raw) + request, err := http.ReadRequest(reader) + if err != nil { + originResult <- err + return + } + _ = request.Body.Close() + if request.Method != http.MethodGet || request.Header.Get("Upgrade") != "websocket" || request.Header.Get("Connection") != "Upgrade" || request.Header.Get("Sec-WebSocket-Key") != "dGhlIHNhbXBsZSBub25jZQ==" || request.Host != origin.Addr().String() || request.URL.Path != "/base/socket" || request.URL.RawQuery != "origin=1&client=2" || request.Header.Get("Origin") != "https://visitor.invalid" || request.Header.Get("Cookie") != "ordinary=yes" || request.Header.Get("Authorization") != "" || request.Header.Get("Proxy-Authorization") != "" || request.Header.Get("X-Remove") != "" { + originResult <- errors.New("owned origin did not receive the expected standard handshake") + return + } + if _, err := io.WriteString(raw, "HTTP/1.1 101 Switching Protocols\r\nUpgrade: websocket\r\nConnection: Upgrade, X-Origin-Hop\r\nSec-WebSocket-Accept: s3pPLMBiTxaQ9kYGzzhZRbK+xOo=\r\nX-Site: ordinary\r\nX-Origin-Hop: do-not-forward\r\nProxy-Authenticate: private\r\nAuthorization: private\r\n\r\n"); err != nil { + originResult <- err + return + } + for range 2 { + frame := make([]byte, 9) + if _, err := io.ReadFull(reader, frame); err != nil { + originResult <- err + return + } + if frame[0] != 0x81 || frame[1] != 0x83 { + originResult <- errors.New("unexpected real masked WebSocket frame") + return + } + payload := make([]byte, 3) + for i := range payload { + payload[i] = frame[i+6] ^ frame[2+i%4] + } + if string(payload) != "hey" { + originResult <- errors.New("real masked WebSocket payload was corrupted") + return + } + if _, err := raw.Write(append([]byte{0x81, 0x03}, payload...)); err != nil { + originResult <- err + return + } + } + _, err = reader.ReadByte() + originResult <- err + }() + + upstream := &http.Transport{Proxy: nil} + t.Cleanup(upstream.CloseIdleConnections) + originURL, err := url.Parse("http://" + origin.Addr().String() + "/base?origin=1") + if err != nil { + t.Fatal(err) + } + proxy, err := cover.NewReverseProxyHandler(originURL, upstream) + if err != nil { + t.Fatal(err) + } + coverDone := make(chan struct{}) + coverContext := make(chan context.Context, 1) + var coverStarted atomic.Bool + website := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + coverStarted.Store(true) + defer close(coverDone) + coverContext <- r.Context() + proxy.ServeHTTP(w, r) + }) + serverTLS, clientTLS := testTLSConfigs(t) + var destinationCalls atomic.Int64 + server, err := ListenWeb(WebServerConfig{ + TCPAddress: "127.0.0.1:0", UDPAddress: "127.0.0.1:0", Token: webTestToken, + TLSConfig: serverTLS, Cover: website, MaxConnections: 1, MaxClientConnections: 1, + Dialer: transport.DialFunc(func(context.Context, string, string) (net.Conn, error) { + destinationCalls.Add(1) + return nil, errors.New("combined WebSocket test forbids every tunnel destination") + }), + }) + if err != nil { + t.Fatal(err) + } + serveResult := make(chan error, 1) + serveJoined := make(chan struct{}) + serveCtx, cancelServe := context.WithCancel(context.Background()) + defer cancelServe() + go func() { defer close(serveJoined); serveResult <- server.Serve(serveCtx) }() + var client *tls.Conn + var closeJoined chan struct{} + t.Cleanup(func() { + if client != nil { + _ = client.Close() + } + _ = server.Close() + workers := map[string]<-chan struct{}{"WebServer Serve worker": serveJoined} + if closeJoined != nil { + workers["WebServer stop worker"] = closeJoined + } + if coverStarted.Load() { + workers["hijacked cover handler"] = coverDone + } + for name, done := range workers { + select { + case <-done: + case <-time.After(2 * time.Second): + t.Errorf("%s did not join after independent cleanup", name) + } + } + }) + clientTLS = clientTLS.Clone() + clientTLS.NextProtos = []string{webHTTP11ALPN} + clientTLS.MinVersion, clientTLS.MaxVersion = version, version + ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second) + defer cancel() + raw, err := (&net.Dialer{}).DialContext(ctx, "tcp", server.TCPAddr().String()) + if err != nil { + t.Fatal(err) + } + client = tls.Client(raw, clientTLS) + if err := client.HandshakeContext(ctx); err != nil { + t.Fatal(err) + } + if state := client.ConnectionState(); state.NegotiatedProtocol != webHTTP11ALPN || state.Version != version || len(state.VerifiedChains) == 0 { + t.Fatal("combined WebSocket test did not use verified actual HTTP/1.1 TLS") + } + _ = client.SetDeadline(time.Now().Add(2 * time.Second)) + if _, err := io.WriteString(client, "GET /socket?client=2 HTTP/1.1\r\nHost: requester.invalid\r\nUpgrade: websocket\r\nConnection: Upgrade, X-Remove, Proxy-Authorization\r\nSec-WebSocket-Version: 13\r\nSec-WebSocket-Key: dGhlIHNhbXBsZSBub25jZQ==\r\nOrigin: https://visitor.invalid\r\nCookie: ordinary=yes\r\nAuthorization: Basic fixture-secret\r\nProxy-Authorization: Bearer fixture-ticket\r\nX-Remove: do-not-forward\r\n\r\n"); err != nil { + t.Fatal(err) + } + reader := bufio.NewReader(client) + response, err := http.ReadResponse(reader, &http.Request{Method: http.MethodGet}) + if err != nil || response.StatusCode != http.StatusSwitchingProtocols || response.Header.Get("Upgrade") != "websocket" || response.Header.Get("Sec-WebSocket-Accept") != "s3pPLMBiTxaQ9kYGzzhZRbK+xOo=" { + t.Fatalf("actual public WebSocket handshake status/error=%v/%v", response, err) + } + for _, header := range []string{"X-Origin-Hop", "Authorization", "Proxy-Authenticate", "Proxy-Authorization", webAuthResponseHeader} { + if response.Header.Get(header) != "" { + t.Fatalf("private/hop header %s reached the upgraded client", header) + } + } + if response.Header.Get("Connection") != "Upgrade" || response.Header.Get("X-Site") != "ordinary" { + t.Fatal("safe upgraded response fields changed") + } + frame := []byte{0x81, 0x83, 1, 2, 3, 4, 'h' ^ 1, 'e' ^ 2, 'y' ^ 3} + echo := func() { + t.Helper() + if _, err := client.Write(frame); err != nil { + t.Fatal(err) + } + got := make([]byte, 5) + if _, err := io.ReadFull(reader, got); err != nil || string(got) != string([]byte{0x81, 0x03, 'h', 'e', 'y'}) { + t.Fatalf("complete real frame echo=%x err=%v", got, err) + } + if reader.Buffered() != 0 { + t.Fatal("combined WebSocket test left unread frame bytes") + } + } + echo() + echo() + var requestContext context.Context + select { + case requestContext = <-coverContext: + case <-time.After(time.Second): + t.Fatal("cover did not publish its actual request context") + } + closed := make(chan error, 1) + closeJoined = make(chan struct{}) + go func() { + defer close(closeJoined) + if stop == "serve-cancel" { + cancelServe() + closed <- nil + } else { + closed <- server.Close() + } + }() + select { + case err := <-closed: + if err != nil { + t.Fatal(err) + } + case <-time.After(2 * time.Second): + t.Fatal("WebServer.Close did not return") + } + select { + case <-closeJoined: + case <-time.After(time.Second): + t.Fatal("WebServer.Close worker did not join") + } + select { + case err := <-serveResult: + if err != nil { + t.Fatal(err) + } + case <-time.After(time.Second): + t.Fatal("WebServer.Serve did not return after Close") + } + _ = client.SetReadDeadline(time.Now().Add(2 * time.Second)) + _, readErr := reader.ReadByte() + var timeout net.Error + if readErr == nil || (errors.As(readErr, &timeout) && timeout.Timeout()) { + t.Fatalf("owner stop left the upgraded client alive: %v", readErr) + } + select { + case <-requestContext.Done(): + case <-time.After(time.Second): + t.Fatal("owner stop did not cancel the actual cover request") + } + if destinationCalls.Load() != 0 { + t.Fatal("ordinary H1 upgrade invoked tunnel destination dialer") + } + if got := len(server.h3.listener.admission.slots); got != 0 { + t.Fatalf("owner stop retained %d shared connection admission slots", got) + } + select { + case err := <-originResult: + if !errors.Is(err, io.EOF) && !errors.Is(err, syscall.ECONNRESET) { + t.Fatalf("owner stop did not close the real origin peer before cleanup: %v", err) + } + case <-time.After(2 * time.Second): + t.Fatal("origin worker did not complete after owner stop, before cleanup") + } + for name, done := range map[string]<-chan struct{}{"origin worker": originJoined, "Serve worker": serveJoined, "cover handler": coverDone} { + select { + case <-done: + case <-time.After(2 * time.Second): + t.Fatalf("%s did not strictly join", name) + } + } +} diff --git a/internal/tunnel/websocket_lifecycle_test.go b/internal/tunnel/websocket_lifecycle_test.go new file mode 100644 index 0000000..09e64b1 --- /dev/null +++ b/internal/tunnel/websocket_lifecycle_test.go @@ -0,0 +1,590 @@ +package tunnel + +import ( + "bufio" + "context" + "crypto/tls" + "errors" + "fmt" + "io" + "net" + "net/http" + "net/http/httputil" + "net/url" + "sync" + "sync/atomic" + "syscall" + "testing" + "time" + + "github.com/cppla/autocar/internal/transport" +) + +// These lifecycle tests inject a standard fixed-loopback ReverseProxy. They +// isolate ownership of real hijacked TLS sockets from cover-handshake policy. +type websocketLifecycleOrigin struct { + listener *net.TCPListener + mu sync.Mutex + closed bool + conn *net.TCPConn + done chan struct{} + result chan error + flood chan struct{} + floodOnce sync.Once +} + +func newWebsocketLifecycleOrigin(t *testing.T, flood bool) *websocketLifecycleOrigin { + t.Helper() + listener, err := net.ListenTCP("tcp", &net.TCPAddr{IP: net.ParseIP("127.0.0.1")}) + if err != nil { + t.Fatal(err) + } + o := &websocketLifecycleOrigin{listener: listener, done: make(chan struct{}), result: make(chan error, 1)} + if flood { + o.flood = make(chan struct{}) + } + t.Cleanup(func() { + o.releaseFlood() + _ = listener.Close() + o.mu.Lock() + o.closed = true + conn := o.conn + o.mu.Unlock() + if conn != nil { + _ = conn.Close() + } + websocketLifecycleJoin(t, "owned origin worker cleanup", o.done) + }) + go func() { + defer close(o.done) + o.result <- o.run() + }() + return o +} + +func (o *websocketLifecycleOrigin) releaseFlood() { + if o.flood != nil { + o.floodOnce.Do(func() { close(o.flood) }) + } +} + +func (o *websocketLifecycleOrigin) run() error { + _ = o.listener.SetDeadline(time.Now().Add(4 * time.Second)) + conn, err := o.listener.AcceptTCP() + if err != nil { + return err + } + defer conn.Close() + o.mu.Lock() + if o.closed { + o.mu.Unlock() + return net.ErrClosed + } + o.conn = conn + o.mu.Unlock() + _ = conn.SetDeadline(time.Now().Add(4 * time.Second)) + reader := bufio.NewReader(conn) + request, err := http.ReadRequest(reader) + if err != nil { + return err + } + _ = request.Body.Close() + if request.Method != http.MethodGet || request.Header.Get("Upgrade") != "websocket" || request.Header.Get("Sec-WebSocket-Key") != "dGhlIHNhbXBsZSBub25jZQ==" { + return errors.New("unexpected owned WebSocket origin handshake") + } + if _, err := io.WriteString(conn, "HTTP/1.1 101 Switching Protocols\r\nUpgrade: websocket\r\nConnection: Upgrade\r\nSec-WebSocket-Accept: s3pPLMBiTxaQ9kYGzzhZRbK+xOo=\r\n\r\n"); err != nil { + return err + } + if o.flood != nil { + select { + case <-o.flood: + case <-time.After(2 * time.Second): + return errors.New("origin flood was not released") + } + // One actual unmasked 8-MiB binary frame exceeds loopback socket + // buffering. No synthetic blocking is inserted in the copy path. + if _, err := conn.Write([]byte{0x82, 127, 0, 0, 0, 0, 0, 128, 0, 0}); err != nil { + return err + } + payload := make([]byte, 8<<20) + if _, err := conn.Write(payload); err != nil { + return err + } + } + for { + frame := make([]byte, 9) + if _, err := io.ReadFull(reader, frame); err != nil { + return err + } + if frame[0] != 0x81 || frame[1] != 0x83 || frame[6]^frame[2] != 'h' || frame[7]^frame[3] != 'e' || frame[8]^frame[4] != 'y' { + return errors.New("owned origin received a corrupted masked frame") + } + if _, err := conn.Write([]byte{0x81, 3, 'h', 'e', 'y'}); err != nil { + return err + } + } +} + +type websocketLifecycleFixture struct { + server *WebH2Server + clientTLS *tls.Config + admission *webConnectionAdmission + origin *websocketLifecycleOrigin + serveCancel context.CancelFunc + serveDone chan struct{} + serveResult chan error + handlerDone chan struct{} + requestCtx chan context.Context + started atomic.Bool + targetCalls atomic.Int64 +} + +func newWebsocketLifecycleFixture(t *testing.T, flood bool, wrap func(http.ResponseWriter) http.ResponseWriter, beforeServe func(*WebH2Server)) *websocketLifecycleFixture { + t.Helper() + f := &websocketLifecycleFixture{origin: newWebsocketLifecycleOrigin(t, flood), serveDone: make(chan struct{}), serveResult: make(chan error, 1), handlerDone: make(chan struct{}), requestCtx: make(chan context.Context, 1)} + upstream := &http.Transport{Proxy: nil} + t.Cleanup(upstream.CloseIdleConnections) + originURL := &url.URL{Scheme: "http", Host: f.origin.listener.Addr().String()} + proxy := &httputil.ReverseProxy{Transport: upstream, Rewrite: func(p *httputil.ProxyRequest) { + p.SetURL(originURL) + p.Out.Host = originURL.Host + }} + website := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.URL.Path == "/ordinary" { + _, _ = io.WriteString(w, "ordinary healthy") + return + } + f.started.Store(true) + defer close(f.handlerDone) + f.requestCtx <- r.Context() + if wrap != nil { + w = wrap(w) + } + proxy.ServeHTTP(w, r) + }) + var err error + f.admission, err = newWebConnectionAdmission(1, 1) + if err != nil { + t.Fatal(err) + } + serverTLS, clientTLS := testTLSConfigs(t) + f.clientTLS = clientTLS.Clone() + f.clientTLS.NextProtos = []string{webHTTP11ALPN} + // Preserve the combined server's actual ResponseController/Unwrap path. + f.server, err = ListenWebH2(WebH2ServerConfig{ + Address: "127.0.0.1:0", Token: webTestToken, TLSConfig: serverTLS, + Cover: &webAltSvcCover{next: website, value: `h3=":443"; ma=60`}, connectionAdmission: f.admission, + Dialer: transport.DialFunc(func(context.Context, string, string) (net.Conn, error) { + f.targetCalls.Add(1) + return nil, errors.New("WebSocket lifecycle test forbids tunnel destination dialing") + }), + }) + if err != nil { + t.Fatal(err) + } + ctx, cancel := context.WithCancel(context.Background()) + f.serveCancel = cancel + t.Cleanup(func() { + cancel() + _ = f.server.Close() + websocketLifecycleJoin(t, "Serve cleanup", f.serveDone) + if f.started.Load() { + websocketLifecycleJoin(t, "cover handler cleanup", f.handlerDone) + } + }) + if beforeServe != nil { + beforeServe(f.server) + } + go func() { + defer close(f.serveDone) + f.serveResult <- f.server.Serve(ctx) + }() + return f +} + +func (f *websocketLifecycleFixture) dial(t *testing.T) *tls.Conn { + t.Helper() + ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second) + defer cancel() + raw, err := (&net.Dialer{}).DialContext(ctx, "tcp", f.server.Addr().String()) + if err != nil { + t.Fatal(err) + } + if tcp, ok := raw.(*net.TCPConn); ok { + _ = tcp.SetReadBuffer(1024) + } + client := tls.Client(raw, f.clientTLS.Clone()) + t.Cleanup(func() { _ = client.Close() }) + if err := client.HandshakeContext(ctx); err != nil { + t.Fatal(err) + } + state := client.ConnectionState() + if state.NegotiatedProtocol != webHTTP11ALPN || len(state.VerifiedChains) == 0 { + t.Fatal("expected verified actual HTTP/1.1 TLS") + } + _ = client.SetDeadline(time.Now().Add(2 * time.Second)) + return client +} + +func (f *websocketLifecycleFixture) request(t *testing.T, client *tls.Conn) { + t.Helper() + if _, err := fmt.Fprintf(client, "GET /socket HTTP/1.1\r\nHost: %s\r\nUpgrade: websocket\r\nConnection: Upgrade\r\nSec-WebSocket-Version: 13\r\nSec-WebSocket-Key: dGhlIHNhbXBsZSBub25jZQ==\r\n\r\n", f.server.Addr()); err != nil { + t.Fatal(err) + } +} + +func (f *websocketLifecycleFixture) open(t *testing.T) (*tls.Conn, *bufio.Reader, context.Context) { + t.Helper() + client := f.dial(t) + f.request(t, client) + reader := bufio.NewReader(client) + response, err := http.ReadResponse(reader, &http.Request{Method: http.MethodGet}) + if err != nil || response.StatusCode != http.StatusSwitchingProtocols || response.Header.Get("Upgrade") != "websocket" || response.Header.Get("Sec-WebSocket-Accept") != "s3pPLMBiTxaQ9kYGzzhZRbK+xOo=" { + t.Fatalf("real public handshake response=%v err=%v", response, err) + } + select { + case ctx := <-f.requestCtx: + if ctx.Err() != nil || len(f.admission.slots) != 1 { + t.Fatal("live hijacked connection lost its context or admission slot") + } + return client, reader, ctx + case <-time.After(time.Second): + t.Fatal("actual cover request context was not published") + return nil, nil, nil + } +} + +func websocketLifecycleEcho(t *testing.T, client *tls.Conn, reader *bufio.Reader) { + t.Helper() + if _, err := client.Write([]byte{0x81, 0x83, 1, 2, 3, 4, 'h' ^ 1, 'e' ^ 2, 'y' ^ 3}); err != nil { + t.Fatal(err) + } + got := make([]byte, 5) + if _, err := io.ReadFull(reader, got); err != nil || string(got) != string([]byte{0x81, 3, 'h', 'e', 'y'}) { + t.Fatalf("complete real WebSocket frame=%x err=%v", got, err) + } + if reader.Buffered() != 0 { + t.Fatal("echo left unread frame bytes") + } +} + +func websocketLifecycleJoin(t *testing.T, name string, done <-chan struct{}) { + t.Helper() + select { + case <-done: + case <-time.After(2 * time.Second): + t.Errorf("%s did not join", name) + } +} + +func (f *websocketLifecycleFixture) assertStopped(t *testing.T, requestCtx context.Context) { + t.Helper() + websocketLifecycleJoin(t, "Serve", f.serveDone) + select { + case err := <-f.serveResult: + if err != nil { + t.Errorf("Serve: %v", err) + } + default: + t.Error("Serve did not publish its result") + } + f.assertConnectionReleased(t, requestCtx) +} + +func (f *websocketLifecycleFixture) assertConnectionReleased(t *testing.T, requestCtx context.Context) { + t.Helper() + websocketLifecycleJoin(t, "hijacked cover handler", f.handlerDone) + websocketLifecycleJoin(t, "owned origin frame worker", f.origin.done) + if requestCtx.Err() == nil { + t.Error("hijacked request context is still live") + } + if slots := len(f.admission.slots); slots != 0 { + t.Errorf("closed physical socket retained %d admission slots", slots) + } + if f.targetCalls.Load() != 0 { + t.Error("H1 cover invoked tunnel destination dialer") + } + select { + case err := <-f.origin.result: + if !errors.Is(err, io.EOF) && !errors.Is(err, syscall.ECONNRESET) && !errors.Is(err, syscall.EPIPE) { + t.Errorf("owned upstream did not observe peer closure: %v", err) + } + default: + t.Error("origin result missing") + } +} + +func TestWebSocketLifecycleIdleCloseAndServeCancellation(t *testing.T) { + for _, cancelServe := range []bool{false, true} { + name := "explicit_close" + if cancelServe { + name = "serve_context" + } + t.Run(name, func(t *testing.T) { + f := newWebsocketLifecycleFixture(t, false, nil, nil) + client, reader, requestCtx := f.open(t) + websocketLifecycleEcho(t, client, reader) + if cancelServe { + f.serveCancel() + } else if err := f.server.Close(); err != nil { + t.Fatal(err) + } + _ = client.SetReadDeadline(time.Now().Add(time.Second)) + if _, err := reader.ReadByte(); !errors.Is(err, io.EOF) && !errors.Is(err, syscall.ECONNRESET) { + t.Errorf("idle downstream was not closed by server shutdown: %v", err) + } + f.assertStopped(t, requestCtx) + }) + } +} + +func TestWebSocketLifecycleClientDisconnectReleasesCapacity(t *testing.T) { + f := newWebsocketLifecycleFixture(t, false, nil, nil) + client, reader, requestCtx := f.open(t) + websocketLifecycleEcho(t, client, reader) + _ = client.Close() + f.assertConnectionReleased(t, requestCtx) + // This occurs before server/fixture cleanup: one source with a one-slot + // allowance must be able to open a second ordinary physical connection. + ordinary := f.dial(t) + if _, err := fmt.Fprintf(ordinary, "GET /ordinary HTTP/1.1\r\nHost: %s\r\nConnection: close\r\n\r\n", f.server.Addr()); err != nil { + t.Fatal(err) + } + response, err := http.ReadResponse(bufio.NewReader(ordinary), &http.Request{Method: http.MethodGet}) + if err != nil { + t.Fatal(err) + } + body, err := io.ReadAll(response.Body) + _ = response.Body.Close() + if err != nil || response.StatusCode != 200 || string(body) != "ordinary healthy" { + t.Fatalf("second ordinary request status=%d body=%q err=%v", response.StatusCode, body, err) + } + _ = ordinary.Close() +} + +type websocketLifecycleHijackGate struct { + http.ResponseWriter + entered chan struct{} + release <-chan struct{} +} + +func (w *websocketLifecycleHijackGate) Hijack() (net.Conn, *bufio.ReadWriter, error) { + close(w.entered) + <-w.release // Test cleanup always releases this gate independently. + return http.NewResponseController(w.ResponseWriter).Hijack() +} + +func TestWebSocketLifecycleCloseBeforeHijack(t *testing.T) { + entered, release := make(chan struct{}), make(chan struct{}) + var once sync.Once + ungate := func() { once.Do(func() { close(release) }) } + f := newWebsocketLifecycleFixture(t, false, func(w http.ResponseWriter) http.ResponseWriter { + return &websocketLifecycleHijackGate{ResponseWriter: w, entered: entered, release: release} + }, nil) + // Registered after fixture cleanup, so gate release runs first on failure. + t.Cleanup(ungate) + client := f.dial(t) + f.request(t, client) + select { + case <-entered: + case <-time.After(time.Second): + t.Fatal("real ReverseProxy did not reach Hijack") + } + var requestCtx context.Context + select { + case requestCtx = <-f.requestCtx: + case <-time.After(time.Second): + t.Fatal("request context missing") + } + if err := f.server.Close(); err != nil { + t.Fatal(err) + } + if requestCtx.Err() == nil || len(f.admission.slots) != 0 { + t.Error("Close did not cancel/release the physical connection before delayed Hijack") + } + ungate() + f.assertStopped(t, requestCtx) +} + +type websocketLifecycleAcceptGate struct { + net.Listener + entered chan struct{} + release <-chan struct{} +} + +func (l *websocketLifecycleAcceptGate) Accept() (net.Conn, error) { + conn, err := l.Listener.Accept() + if err == nil { + close(l.entered) + <-l.release // Actual admission already ran; net/http has not seen it. + } + return conn, err +} + +func TestWebSocketLifecycleCloseDuringAcceptedSocketDelivery(t *testing.T) { + entered, release := make(chan struct{}), make(chan struct{}) + var once sync.Once + ungate := func() { once.Do(func() { close(release) }) } + f := newWebsocketLifecycleFixture(t, false, nil, func(s *WebH2Server) { + s.listener = &websocketLifecycleAcceptGate{Listener: s.listener, entered: entered, release: release} + }) + t.Cleanup(ungate) + client, err := net.DialTimeout("tcp", f.server.Addr().String(), time.Second) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = client.Close() }) + select { + case <-entered: + case <-time.After(time.Second): + t.Fatal("actual accepted socket was not gated") + } + if len(f.admission.slots) != 1 { + t.Fatal("test did not pause after real connection admission") + } + closeResult, closeDone := make(chan error, 1), make(chan struct{}) + go func() { defer close(closeDone); closeResult <- f.server.Close() }() + t.Cleanup(func() { ungate(); websocketLifecycleJoin(t, "delayed Accept Close cleanup", closeDone) }) + _ = client.SetReadDeadline(time.Now().Add(time.Second)) + if _, err := client.Read(make([]byte, 1)); !errors.Is(err, io.EOF) && !errors.Is(err, syscall.ECONNRESET) { + t.Errorf("Close did not close not-yet-delivered physical socket: %v", err) + } + ungate() + websocketLifecycleJoin(t, "Close with delayed Accept return", closeDone) + select { + case err := <-closeResult: + if err != nil { + t.Error(err) + } + default: + t.Error("Close result missing") + } + // Remote EOF can precede the local raw Close return and its admission + // release callback. Assert callback completion only after Close joins; + // physical peer closure was independently checked before ungating Accept. + if len(f.admission.slots) != 0 { + t.Error("accepted socket retained its admission slot after Close joined") + } + websocketLifecycleJoin(t, "Serve with delayed Accept return", f.serveDone) +} + +type websocketLifecycleWriteCall struct { + done chan struct{} + err error +} + +type websocketLifecycleWriteWitness struct { + mu sync.Mutex + latest *websocketLifecycleWriteCall + changed chan struct{} +} + +func (w *websocketLifecycleWriteWitness) observe(call *websocketLifecycleWriteCall) { + w.mu.Lock() + w.latest = call + w.mu.Unlock() + // Coalesce notifications, not evidence: the most recent actual call is + // retained even when a burst fills the notification channel. + select { + case w.changed <- struct{}{}: + default: + } +} + +func (w *websocketLifecycleWriteWitness) current() *websocketLifecycleWriteCall { + w.mu.Lock() + defer w.mu.Unlock() + return w.latest +} + +type websocketLifecycleObservedConn struct { + net.Conn + writes *websocketLifecycleWriteWitness +} + +func (c *websocketLifecycleObservedConn) Write(p []byte) (int, error) { + call := &websocketLifecycleWriteCall{done: make(chan struct{})} + c.writes.observe(call) + n, err := c.Conn.Write(p) + call.err = err + close(call.done) + return n, err +} + +type websocketLifecycleWriteObserver struct { + http.ResponseWriter + writes *websocketLifecycleWriteWitness +} + +func (w *websocketLifecycleWriteObserver) Hijack() (net.Conn, *bufio.ReadWriter, error) { + conn, buffered, err := http.NewResponseController(w.ResponseWriter).Hijack() + if err != nil { + return conn, buffered, err + } + return &websocketLifecycleObservedConn{Conn: conn, writes: w.writes}, buffered, nil +} + +func TestWebSocketLifecycleCloseUnblocksRealDownstreamWrite(t *testing.T) { + writes := &websocketLifecycleWriteWitness{changed: make(chan struct{}, 1)} + f := newWebsocketLifecycleFixture(t, true, func(w http.ResponseWriter) http.ResponseWriter { + return &websocketLifecycleWriteObserver{ResponseWriter: w, writes: writes} + }, nil) + client, _, requestCtx := f.open(t) + // Deliberately no client reader after the 101. Observe an actual TLS Write + // that has not completed for 75 ms, not a fixture-created blocking gate. + f.origin.releaseFlood() + deadline := time.NewTimer(2 * time.Second) + defer deadline.Stop() + var blocked *websocketLifecycleWriteCall +findBlocked: + for { + if call := writes.current(); call != nil { + window := time.NewTimer(75 * time.Millisecond) + select { + case <-call.done: + window.Stop() + case <-window.C: + select { + case <-call.done: + default: + blocked = call + break findBlocked + } + case <-deadline.C: + window.Stop() + t.Fatal("did not establish actual downstream write backpressure") + } + } + select { + case <-writes.changed: + case <-deadline.C: + t.Fatal("no actual downstream write blocked") + } + } + if requestCtx.Err() != nil { + t.Fatal("request was canceled before the shutdown oracle") + } + select { + case <-blocked.done: + t.Fatal("observed actual Write finished before the Close oracle") + default: + } + if err := f.server.Close(); err != nil { + t.Fatal(err) + } + websocketLifecycleJoin(t, "actual downstream Write", blocked.done) + select { + case <-blocked.done: + if blocked.err == nil { + t.Error("blocked actual Write unexpectedly succeeded after physical close") + } + var timeout net.Error + if errors.As(blocked.err, &timeout) && timeout.Timeout() { + t.Errorf("blocked Write ended by a timeout, not physical closure: %v", blocked.err) + } + default: + } + f.assertStopped(t, requestCtx) + // Explicit client close is cleanup, and occurs after all shutdown oracles. + _ = client.Close() +}