From df261e447907233af19a20ed90c0cc41905e67dd Mon Sep 17 00:00:00 2001 From: Cluvex <125141320+CluvexStudio@users.noreply.github.com> Date: Sat, 26 Sep 2026 06:18:16 +0330 Subject: [PATCH] MASQUE client: Support HTTP/2 (Extended CONNECT, RFC 8441) (#6810) https://github.com/XTLS/Xray-core/pull/6807#issuecomment-5808933074 https://github.com/XTLS/Xray-core/pull/6810#issuecomment-5842441136 --- testing/scenarios/masque_test.go | 205 +++++- transport/internet/masque/conn.go | 22 +- transport/internet/masque/connectip/conn.go | 103 ++- transport/internet/masque/connectip/http2.go | 198 ++++++ .../internet/masque/connectip/http2_test.go | 502 +++++++++++++++ transport/internet/masque/connectip/proxy.go | 21 +- .../internet/masque/connectip/request.go | 16 +- .../internet/masque/connectip/request_test.go | 18 + transport/internet/masque/dialer.go | 73 ++- transport/internet/masque/dialer_test.go | 19 + transport/internet/masque/http2.go | 600 ++++++++++++++++++ transport/internet/masque/http2_test.go | 405 ++++++++++++ 12 files changed, 2144 insertions(+), 38 deletions(-) create mode 100644 transport/internet/masque/connectip/http2.go create mode 100644 transport/internet/masque/connectip/http2_test.go create mode 100644 transport/internet/masque/http2.go create mode 100644 transport/internet/masque/http2_test.go diff --git a/testing/scenarios/masque_test.go b/testing/scenarios/masque_test.go index 1b203a88a..6f0f92d3d 100644 --- a/testing/scenarios/masque_test.go +++ b/testing/scenarios/masque_test.go @@ -1,6 +1,8 @@ package scenarios import ( + "bufio" + "bytes" "context" gotls "crypto/tls" "crypto/x509" @@ -8,12 +10,18 @@ import ( "io" "net/http" "net/netip" + "net/url" + "strconv" + "strings" + "sync" "sync/atomic" "testing" "time" "github.com/apernet/quic-go" "github.com/apernet/quic-go/http3" + "golang.org/x/net/http2" + "golang.org/x/net/http2/hpack" "golang.org/x/sync/errgroup" "gvisor.dev/gvisor/pkg/tcpip" "gvisor.dev/gvisor/pkg/tcpip/adapters/gonet" @@ -52,7 +60,7 @@ const ( masqueAuthorization = "Basic dTpw" ) -func startMasqueServer(t *testing.T) (net.Port, [32]byte) { +func startMasqueServer(t *testing.T, h2 bool) (net.Port, [32]byte) { dev, _, gstack, err := wireguard.CreateNetTUN([]netip.Addr{masqueServerV4, masqueServerV6}, nil, transmasque.MinPacketSize, false) common.Must(err) t.Cleanup(func() { dev.Close() }) @@ -180,6 +188,13 @@ func startMasqueServer(t *testing.T) (net.Port, [32]byte) { Certificates: []gotls.Certificate{{Certificate: [][]byte{certificate.Certificate}, PrivateKey: key}}, NextProtos: []string{http3.NextProtoH3}, } + if h2 { + tlsConfig.NextProtos = []string{http2.NextProtoTLS} + ln := common.Must2(gotls.Listen("tcp", "127.0.0.1:0", tlsConfig)) + t.Cleanup(func() { ln.Close() }) + go serveHTTP2(ln, http.HandlerFunc(handler)) + return net.Port(ln.Addr().(*net.TCPAddr).Port), certHash + } pktConn := common.Must2(net.ListenUDP("udp", &net.UDPAddr{IP: net.LocalHostIP.IP()})) tr := &quic.Transport{Conn: pktConn} ln := common.Must2(tr.ListenEarly(tlsConfig, &quic.Config{EnableDatagrams: true, InitialPacketSize: 1350})) @@ -195,8 +210,182 @@ func startMasqueServer(t *testing.T) (net.Port, [32]byte) { return net.Port(pktConn.LocalAddr().(*net.UDPAddr).Port), certHash } +func serveHTTP2(ln net.Listener, handler http.Handler) { + for { + conn, err := ln.Accept() + if err != nil { + return + } + go serveHTTP2Conn(conn, handler) + } +} + +type http2ServerConn struct { + mu sync.Mutex + fr *http2.Framer + hbuf bytes.Buffer + henc *hpack.Encoder +} + +func (c *http2ServerConn) write(f func(*http2.Framer) error) error { + c.mu.Lock() + defer c.mu.Unlock() + return f(c.fr) +} + +func (c *http2ServerConn) writeHeaders(streamID uint32, status int, header http.Header) error { + c.mu.Lock() + defer c.mu.Unlock() + c.hbuf.Reset() + c.henc.WriteField(hpack.HeaderField{Name: ":status", Value: strconv.Itoa(status)}) + for k, vv := range header { + for _, v := range vv { + c.henc.WriteField(hpack.HeaderField{Name: strings.ToLower(k), Value: v}) + } + } + return c.fr.WriteHeaders(http2.HeadersFrameParam{StreamID: streamID, BlockFragment: c.hbuf.Bytes(), EndHeaders: true}) +} + +func (c *http2ServerConn) writeData(streamID uint32, endStream bool, data []byte) error { + c.mu.Lock() + defer c.mu.Unlock() + for { + n := min(len(data), 16384) + if err := c.fr.WriteData(streamID, endStream && n == len(data), data[:n]); err != nil { + return err + } + if data = data[n:]; len(data) == 0 { + return nil + } + } +} + +func serveHTTP2Conn(conn net.Conn, handler http.Handler) { + defer conn.Close() + br := bufio.NewReader(conn) + preface := make([]byte, len(http2.ClientPreface)) + if _, err := io.ReadFull(br, preface); err != nil || string(preface) != http2.ClientPreface { + return + } + sc := &http2ServerConn{fr: http2.NewFramer(conn, br)} + sc.henc = hpack.NewEncoder(&sc.hbuf) + sc.fr.ReadMetaHeaders = hpack.NewDecoder(4096, nil) + if err := sc.write(func(fr *http2.Framer) error { + if err := fr.WriteSettings( + http2.Setting{ID: http2.SettingEnableConnectProtocol, Val: 1}, + http2.Setting{ID: http2.SettingInitialWindowSize, Val: 1 << 30}, + ); err != nil { + return err + } + return fr.WriteWindowUpdate(0, 1<<30) + }); err != nil { + return + } + + bodies := make(map[uint32]*io.PipeWriter) + defer func() { + for _, body := range bodies { + body.Close() + } + }() + for { + f, err := sc.fr.ReadFrame() + if err != nil { + return + } + switch f := f.(type) { + case *http2.SettingsFrame: + if !f.IsAck() { + err = sc.write((*http2.Framer).WriteSettingsAck) + } + case *http2.PingFrame: + if !f.IsAck() { + err = sc.write(func(fr *http2.Framer) error { return fr.WritePing(true, f.Data) }) + } + case *http2.MetaHeadersFrame: + u, err := url.ParseRequestURI(f.PseudoValue("path")) + if err != nil { + return + } + pr, pw := io.Pipe() + bodies[f.StreamID] = pw + req := &http.Request{ + Method: f.PseudoValue("method"), + URL: u, + Proto: "HTTP/2.0", + ProtoMajor: 2, + Header: http.Header{}, + Host: f.PseudoValue("authority"), + Body: pr, + } + for _, hf := range f.RegularFields() { + req.Header.Add(hf.Name, hf.Value) + } + if protocol := f.PseudoValue("protocol"); protocol != "" { + req.Header.Set(":protocol", protocol) + } + streamID := f.StreamID + w := &http2ResponseWriter{conn: sc, streamID: streamID, header: http.Header{}} + go func() { + handler.ServeHTTP(w, req) + w.WriteHeader(http.StatusOK) + sc.writeData(streamID, true, nil) + }() + case *http2.DataFrame: + if body := bodies[f.StreamID]; body != nil { + if _, err := body.Write(f.Data()); err != nil || f.StreamEnded() { + body.Close() + delete(bodies, f.StreamID) + } + } + case *http2.RSTStreamFrame: + if body := bodies[f.StreamID]; body != nil { + body.CloseWithError(http2.StreamError{StreamID: f.StreamID, Code: f.ErrCode}) + delete(bodies, f.StreamID) + } + } + if err != nil { + return + } + } +} + +type http2ResponseWriter struct { + conn *http2ServerConn + streamID uint32 + header http.Header + wroteHeader bool +} + +func (w *http2ResponseWriter) Header() http.Header { return w.header } + +func (w *http2ResponseWriter) WriteHeader(code int) { + if !w.wroteHeader { + w.wroteHeader = true + w.conn.writeHeaders(w.streamID, code, w.header) + } +} + +func (w *http2ResponseWriter) Write(b []byte) (int, error) { + w.WriteHeader(http.StatusOK) + if err := w.conn.writeData(w.streamID, false, b); err != nil { + return 0, err + } + return len(b), nil +} + +func (w *http2ResponseWriter) Flush() {} + func TestMasque(t *testing.T) { - serverPort, certHash := startMasqueServer(t) + testMasque(t, false) +} + +func TestMasqueHTTP2(t *testing.T) { + testMasque(t, true) +} + +func testMasque(t *testing.T, h2 bool) { + serverPort, certHash := startMasqueServer(t, h2) tcpPort := tcp.PickPort() tcp6Port := tcp.PickPort() @@ -214,6 +403,13 @@ func TestMasque(t *testing.T) { }), } } + tlsConfig := &tls.Config{ + ServerName: "localhost", + PinnedPeerCertSha256: [][]byte{certHash[:]}, + } + if h2 { + tlsConfig.NextProtocol = []string{http2.NextProtoTLS} + } clientConfig := &core.Config{ App: []*serial.TypedMessage{ serial.ToTypedMessage(&log.Config{ @@ -248,10 +444,7 @@ func TestMasque(t *testing.T) { }, SecurityType: serial.GetMessageType(&tls.Config{}), SecuritySettings: []*serial.TypedMessage{ - serial.ToTypedMessage(&tls.Config{ - ServerName: "localhost", - PinnedPeerCertSha256: [][]byte{certHash[:]}, - }), + serial.ToTypedMessage(tlsConfig), }, }, }), diff --git a/transport/internet/masque/conn.go b/transport/internet/masque/conn.go index 6aeaf162b..9ccea3d83 100644 --- a/transport/internet/masque/conn.go +++ b/transport/internet/masque/conn.go @@ -23,9 +23,23 @@ func (e *PacketTooBigError) Error() string { return "packet too big for the tunnel" } +type httpConn interface { + LocalAddr() net.Addr + RemoteAddr() net.Addr + Close() error +} + +type quicConn struct { + *quic.Conn +} + +func (c quicConn) Close() error { + return c.CloseWithError(quic.ApplicationErrorCode(http3.ErrCodeNoError), "") +} + type Conn struct { ipConn *connectip.Conn - quicConn *quic.Conn + httpConn httpConn local []netip.Addr closeOnce sync.Once } @@ -58,17 +72,17 @@ func (c *Conn) Write(b []byte) (int, error) { func (c *Conn) Close() error { c.closeOnce.Do(func() { c.ipConn.Close() - c.quicConn.CloseWithError(quic.ApplicationErrorCode(http3.ErrCodeNoError), "") + c.httpConn.Close() }) return nil } func (c *Conn) LocalAddr() net.Addr { - return c.quicConn.LocalAddr() + return c.httpConn.LocalAddr() } func (c *Conn) RemoteAddr() net.Addr { - return c.quicConn.RemoteAddr() + return c.httpConn.RemoteAddr() } func (c *Conn) SetDeadline(time.Time) error { diff --git a/transport/internet/masque/connectip/conn.go b/transport/internet/masque/connectip/conn.go index 3e5fc6362..510c8e567 100644 --- a/transport/internet/masque/connectip/conn.go +++ b/transport/internet/masque/connectip/conn.go @@ -38,22 +38,30 @@ const ( ipProtoICMPv6 = 58 ) -type http3Stream interface { +type requestStream interface { io.ReadWriteCloser - StreamID() quic.StreamID - ReceiveDatagram(context.Context) ([]byte, error) - SendDatagram([]byte) error CancelRead(quic.StreamErrorCode) CancelWrite(quic.StreamErrorCode) SetWriteDeadline(time.Time) error } +type http3Stream interface { + requestStream + StreamID() quic.StreamID + ReceiveDatagram(context.Context) ([]byte, error) + SendDatagram([]byte) error +} + var ( _ http3Stream = &http3.Stream{} _ http3Stream = &http3.RequestStream{} ) -const maxQueuedCapsules = 128 +const ( + maxQueuedCapsules = 128 + maxQueuedDatagrams = 128 + maxCapsulePacketSize = 1<<16 - 1 +) var errCapsuleLimit = goerrors.New("connect-ip: capsule limit exceeded") @@ -63,7 +71,10 @@ type streamWrite struct { } type Conn struct { - str http3Stream + str requestStream + h3 http3Stream + datagrams chan []byte + writeMu sync.Mutex writeNotify chan struct{} writeDone chan error @@ -87,7 +98,7 @@ type Conn struct { datagramCapsuleOnce sync.Once } -func newProxiedConn(str http3Stream) *Conn { +func newProxiedConn(str requestStream) *Conn { c := &Conn{ str: str, writeNotify: make(chan struct{}, 1), @@ -97,6 +108,9 @@ func newProxiedConn(str http3Stream) *Conn { availableRouteUpdates: make(chan []IPRoute, 1), closeChan: make(chan struct{}), } + if c.h3, _ = str.(http3Stream); c.h3 == nil { + c.datagrams = make(chan []byte, maxQueuedDatagrams) + } go func() { err := c.readFromStream() c.mu.Lock() @@ -382,6 +396,12 @@ func (c *Conn) readFromStream() error { } queueLatest(c.availableRouteUpdates, capsule.IPAddressRanges) case capsuleTypeDatagram: + if c.h3 == nil { + if err := c.queueDatagram(cr); err != nil { + return err + } + continue + } c.datagramCapsuleOnce.Do(func() { errors.LogWarning(context.Background(), "connect-ip: dropping IP packets sent in DATAGRAM capsules, only QUIC DATAGRAM frames are supported") }) @@ -412,7 +432,7 @@ func (c *Conn) writeToStream() error { if w.Fin { return c.str.Close() } - if _, err := c.str.Write(w.Data); err != nil { + if err := c.write(w.Data); err != nil { return err } } @@ -420,6 +440,41 @@ func (c *Conn) writeToStream() error { return c.closeErr } +func (c *Conn) write(b []byte) error { + c.writeMu.Lock() + defer c.writeMu.Unlock() + _, err := c.str.Write(b) + return err +} + +func (c *Conn) queueDatagram(cr http3.CapsuleReader) error { + if cr.Remaining() > int64(len(contextIDZero)+maxCapsulePacketSize) { + errors.LogDebug(context.Background(), "connect-ip: dropping a ", cr.Remaining(), "-byte DATAGRAM capsule") + return cr.Discard() + } + data := make([]byte, cr.Remaining()) + if _, err := io.ReadFull(cr, data); err != nil { + return err + } + select { + case c.datagrams <- data: + case <-c.closeChan: + } + return nil +} + +func (c *Conn) receiveDatagram() ([]byte, error) { + if c.h3 != nil { + return c.h3.ReceiveDatagram(context.Background()) + } + select { + case data := <-c.datagrams: + return data, nil + case <-c.closeChan: + return nil, c.closeErr + } +} + func (c *Conn) ReadPacket(b []byte) (int, error) { for { select { @@ -427,7 +482,7 @@ func (c *Conn) ReadPacket(b []byte) (int, error) { return 0, c.closeErr default: } - data, err := c.str.ReceiveDatagram(context.Background()) + data, err := c.receiveDatagram() if err != nil { select { case <-c.closeChan: @@ -525,7 +580,18 @@ func (c *Conn) WritePacket(b []byte) (icmp []byte, err error) { errors.LogDebugInner(context.Background(), err, "dropping proxied packet (", len(b), " bytes) that can't be proxied") return nil, nil } - if err := c.str.SendDatagram(data); err != nil { + if c.h3 == nil { + if err := c.write(data); err != nil { + select { + case <-c.closeChan: + return nil, c.closeErr + default: + return nil, err + } + } + return nil, nil + } + if err := c.h3.SendDatagram(data); err != nil { if tooLarge, ok := goerrors.AsType[*quic.DatagramTooLargeError](err); ok { icmpPacket, err := composeICMPTooLargePacket(b, int(tooLarge.MaxDatagramPayloadSize)-c.datagramOverhead()) if err != nil { @@ -578,14 +644,22 @@ func (c *Conn) composeDatagram(b []byte) ([]byte, error) { } b[7]-- } - data := make([]byte, 0, len(contextIDZero)+len(b)) + size := len(contextIDZero) + len(b) + var data []byte + if c.h3 == nil { + data = make([]byte, 0, quicvarint.Len(uint64(capsuleTypeDatagram))+quicvarint.Len(uint64(size))+size) + data = quicvarint.Append(data, uint64(capsuleTypeDatagram)) + data = quicvarint.Append(data, uint64(size)) + } else { + data = make([]byte, 0, size) + } data = append(data, contextIDZero...) data = append(data, b...) return data, nil } func (c *Conn) datagramOverhead() int { - return quicvarint.Len(uint64(c.str.StreamID()/4)) + len(contextIDZero) + return quicvarint.Len(uint64(c.h3.StreamID()/4)) + len(contextIDZero) } func (c *Conn) MaxPacketSize() int { @@ -594,7 +668,10 @@ func (c *Conn) MaxPacketSize() int { return 0 default: } - err := c.str.SendDatagram(make([]byte, 1<<16)) + if c.h3 == nil { + return maxCapsulePacketSize + } + err := c.h3.SendDatagram(make([]byte, 1<<16)) tooLarge, ok := goerrors.AsType[*quic.DatagramTooLargeError](err) if !ok { return 0 diff --git a/transport/internet/masque/connectip/http2.go b/transport/internet/masque/connectip/http2.go new file mode 100644 index 000000000..a83f513e5 --- /dev/null +++ b/transport/internet/masque/connectip/http2.go @@ -0,0 +1,198 @@ +package connectip + +import ( + "bufio" + "context" + "errors" + "fmt" + "io" + "net" + "net/http" + "os" + "sync" + "time" + + "github.com/apernet/quic-go" +) + +const maxBufferedRequestBody = 32 << 10 + +type HTTP2ClientConn struct { + roundTripper http.RoundTripper +} + +func NewHTTP2ClientConn(rt http.RoundTripper) *HTTP2ClientConn { + return &HTTP2ClientConn{roundTripper: rt} +} + +func (c *HTTP2ClientConn) Dial(req *Request) (*Conn, *http.Response, error) { + httpReq := req.httpRequest() + if httpReq.URL == nil { + return nil, nil, errors.New("connect-ip: request URL is nil") + } + if httpReq.Host == "" && httpReq.URL.Host == "" { + return nil, nil, errors.New("connect-ip: request needs a host") + } + + ctx := httpReq.Context() + streamCtx, cancel := context.WithCancel(context.WithoutCancel(ctx)) + stop := context.AfterFunc(ctx, cancel) + body := newRequestBody() + r := httpReq.Clone(streamCtx) + r.Header[":protocol"] = []string{requestProtocol} + r.Body = body + rsp, err := c.roundTripper.RoundTrip(r) + if !stop() { + if err == nil { + rsp.Body.Close() + } + err = context.Cause(ctx) + } + if err != nil { + cancel() + return nil, nil, fmt.Errorf("connect-ip: failed to send request: %w", err) + } + if rsp.StatusCode < 200 || rsp.StatusCode > 299 { + cancel() + rsp.Body.Close() + return nil, rsp, fmt.Errorf("connect-ip: server responded with %d", rsp.StatusCode) + } + return newProxiedConn(&http2Stream{ + reader: bufio.NewReader(rsp.Body), + body: body, + rsp: rsp.Body, + cancel: cancel, + }), rsp, nil +} + +type http2Stream struct { + reader *bufio.Reader + body *requestBody + rsp io.Closer + cancel context.CancelFunc +} + +func (s *http2Stream) Read(b []byte) (int, error) { return s.reader.Read(b) } +func (s *http2Stream) ReadByte() (byte, error) { return s.reader.ReadByte() } +func (s *http2Stream) Write(b []byte) (int, error) { return s.body.Write(b) } +func (s *http2Stream) Close() error { return s.body.Close() } +func (s *http2Stream) CancelRead(quic.StreamErrorCode) { s.abort() } +func (s *http2Stream) CancelWrite(quic.StreamErrorCode) { s.abort() } +func (s *http2Stream) SetWriteDeadline(t time.Time) error { return s.body.SetWriteDeadline(t) } + +func (s *http2Stream) abort() { + s.cancel() + s.body.CloseWithError(net.ErrClosed) + s.rsp.Close() +} + +type requestBody struct { + mu sync.Mutex + cond sync.Cond + buf []byte + closed bool + err error + deadline time.Time +} + +func newRequestBody() *requestBody { + b := &requestBody{} + b.cond.L = &b.mu + return b +} + +func (b *requestBody) Read(p []byte) (int, error) { + b.mu.Lock() + defer b.mu.Unlock() + for len(b.buf) == 0 && !b.closed && b.err == nil { + b.cond.Wait() + } + if b.err != nil { + return 0, b.err + } + if len(b.buf) == 0 { + return 0, io.EOF + } + n := copy(p, b.buf) + b.buf = b.buf[:copy(b.buf, b.buf[n:])] + b.cond.Broadcast() + return n, nil +} + +func (b *requestBody) Write(p []byte) (int, error) { + b.mu.Lock() + defer b.mu.Unlock() + for { + switch { + case b.err != nil: + return 0, b.err + case b.closed: + return 0, io.ErrClosedPipe + case !b.deadline.IsZero() && !time.Now().Before(b.deadline): + return 0, os.ErrDeadlineExceeded + case len(b.buf) < maxBufferedRequestBody: + b.buf = append(b.buf, p...) + b.cond.Broadcast() + return len(p), nil + } + b.cond.Wait() + } +} + +func (b *requestBody) Close() error { + b.mu.Lock() + b.closed = true + b.cond.Broadcast() + b.mu.Unlock() + return nil +} + +func (b *requestBody) CloseWithError(err error) { + b.mu.Lock() + if b.err == nil { + b.err = err + b.buf = nil + } + b.cond.Broadcast() + b.mu.Unlock() +} + +func (b *requestBody) SetWriteDeadline(t time.Time) error { + b.mu.Lock() + b.deadline = t + b.cond.Broadcast() + b.mu.Unlock() + if d := time.Until(t); d > 0 { + time.AfterFunc(d, func() { + b.mu.Lock() + b.cond.Broadcast() + b.mu.Unlock() + }) + } + return nil +} + +type http2ResponseStream struct { + reader *bufio.Reader + body io.Closer + w io.Writer + controller *http.ResponseController +} + +func (s *http2ResponseStream) Read(b []byte) (int, error) { return s.reader.Read(b) } +func (s *http2ResponseStream) ReadByte() (byte, error) { return s.reader.ReadByte() } + +func (s *http2ResponseStream) Write(b []byte) (int, error) { + n, err := s.w.Write(b) + if err == nil { + err = s.controller.Flush() + } + return n, err +} + +func (s *http2ResponseStream) Close() error { return nil } +func (s *http2ResponseStream) CancelRead(quic.StreamErrorCode) { s.body.Close() } +func (s *http2ResponseStream) CancelWrite(quic.StreamErrorCode) { s.body.Close() } +func (s *http2ResponseStream) SetWriteDeadline(t time.Time) error { + return s.controller.SetWriteDeadline(t) +} diff --git a/transport/internet/masque/connectip/http2_test.go b/transport/internet/masque/connectip/http2_test.go new file mode 100644 index 000000000..e5c9d3cd0 --- /dev/null +++ b/transport/internet/masque/connectip/http2_test.go @@ -0,0 +1,502 @@ +package connectip + +import ( + "bufio" + "bytes" + "context" + "errors" + "io" + "net" + "net/http" + "net/http/httptest" + "net/netip" + "os" + "slices" + "sync" + "testing" + "time" + + "github.com/apernet/quic-go/http3" + "github.com/apernet/quic-go/quicvarint" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "golang.org/x/net/ipv4" + "golang.org/x/net/ipv6" +) + +type roundTripFunc func(*http.Request) (*http.Response, error) + +func (f roundTripFunc) RoundTrip(r *http.Request) (*http.Response, error) { return f(r) } + +type pipeResponseWriter struct { + *io.PipeWriter + header http.Header + status int + headerOnce sync.Once + headerDone chan struct{} +} + +func (w *pipeResponseWriter) Header() http.Header { return w.header } + +func (w *pipeResponseWriter) WriteHeader(code int) { + w.headerOnce.Do(func() { + w.status = code + close(w.headerDone) + }) +} + +func (w *pipeResponseWriter) Write(b []byte) (int, error) { + w.WriteHeader(http.StatusOK) + return w.PipeWriter.Write(b) +} + +func (w *pipeResponseWriter) Flush() {} + +func http2RoundTripper(handler http.HandlerFunc) http.RoundTripper { + return roundTripFunc(func(r *http.Request) (*http.Response, error) { + pr, pw := io.Pipe() + w := &pipeResponseWriter{PipeWriter: pw, header: http.Header{}, headerDone: make(chan struct{})} + sr := r.Clone(r.Context()) + sr.Proto, sr.ProtoMajor, sr.ProtoMinor = "HTTP/2.0", 2, 0 + go func() { + handler(w, sr) + w.WriteHeader(http.StatusOK) + pw.Close() + }() + <-w.headerDone + return &http.Response{StatusCode: w.status, Header: w.header, Body: pr}, nil + }) +} + +func setupHTTP2Conns(t *testing.T) (client, server *Conn) { + t.Helper() + + serverConns := make(chan *Conn, 1) + rt := http2RoundTripper(func(w http.ResponseWriter, r *http.Request) { + assert.Equal(t, "Bearer token", r.Header.Get("Authorization")) + req, err := ParseProxyRequest(r) + if !assert.NoError(t, err) { + w.WriteHeader(http.StatusBadRequest) + return + } + conn, err := (&Proxy{}).Proxy(w, req) + if !assert.NoError(t, err) { + return + } + serverConns <- conn + <-conn.closeChan + }) + + ctx, cancel := context.WithTimeout(t.Context(), 5*time.Second) + defer cancel() + req, err := NewRequest(ctx, "https://example.org/connect-ip") + require.NoError(t, err) + req.Header().Set("Authorization", "Bearer token") + client, rsp, err := NewHTTP2ClientConn(rt).Dial(req) + require.NoError(t, err) + t.Cleanup(func() { client.Close() }) + require.Equal(t, http.StatusOK, rsp.StatusCode) + require.Equal(t, "?1", rsp.Header.Get("Capsule-Protocol")) + + select { + case <-time.After(5 * time.Second): + t.Fatal("timed out") + case server = <-serverConns: + } + t.Cleanup(func() { server.Close() }) + return client, server +} + +func newTestHTTP2Stream() (*http2Stream, *io.PipeWriter) { + pr, pw := io.Pipe() + return &http2Stream{reader: bufio.NewReader(pr), body: newRequestBody(), rsp: pr, cancel: func() {}}, pw +} + +func TestHTTP2Request(t *testing.T) { + requests := make(chan *http.Request, 1) + pr, pw := io.Pipe() + defer pw.Close() + rt := roundTripFunc(func(r *http.Request) (*http.Response, error) { + requests <- r + return &http.Response{StatusCode: http.StatusOK, Body: pr}, nil + }) + req, err := NewRequest(t.Context(), "https://proxy.example:8443/.well-known/masque/ip/*/*/") + require.NoError(t, err) + req.Header().Set("Authorization", "Bearer token") + conn, _, err := NewHTTP2ClientConn(rt).Dial(req) + require.NoError(t, err) + defer conn.Close() + + r := <-requests + require.Equal(t, http.MethodConnect, r.Method) + require.Equal(t, []string{requestProtocol}, r.Header[":protocol"]) + require.Equal(t, "?1", r.Header.Get("Capsule-Protocol")) + require.Equal(t, "Bearer token", r.Header.Get("Authorization")) + require.Equal(t, "proxy.example:8443", r.Host) + require.Equal(t, "https", r.URL.Scheme) + require.Equal(t, "/.well-known/masque/ip/*/*/", r.URL.Path) + require.NotNil(t, r.Body) + require.Empty(t, req.Header().Values(":protocol")) + require.Equal(t, maxCapsulePacketSize, conn.MaxPacketSize()) +} + +func TestHTTP2DialErrors(t *testing.T) { + newReq := func(ctx context.Context) *Request { + req, err := NewRequest(ctx, "https://example.org/connect-ip") + require.NoError(t, err) + return req + } + + t.Run("status", func(t *testing.T) { + var streamCtx context.Context + rt := roundTripFunc(func(r *http.Request) (*http.Response, error) { + streamCtx = r.Context() + return &http.Response{StatusCode: http.StatusForbidden, Body: io.NopCloser(bytes.NewReader(nil))}, nil + }) + _, rsp, err := NewHTTP2ClientConn(rt).Dial(newReq(t.Context())) + require.EqualError(t, err, "connect-ip: server responded with 403") + require.Equal(t, http.StatusForbidden, rsp.StatusCode) + require.ErrorIs(t, streamCtx.Err(), context.Canceled) + }) + + t.Run("round trip", func(t *testing.T) { + errRoundTrip := errors.New("extended connect not supported by peer") + rt := roundTripFunc(func(*http.Request) (*http.Response, error) { return nil, errRoundTrip }) + _, _, err := NewHTTP2ClientConn(rt).Dial(newReq(t.Context())) + require.ErrorIs(t, err, errRoundTrip) + }) + + t.Run("context", func(t *testing.T) { + rt := roundTripFunc(func(r *http.Request) (*http.Response, error) { + <-r.Context().Done() + return nil, r.Context().Err() + }) + ctx, cancel := context.WithTimeout(t.Context(), 50*time.Millisecond) + defer cancel() + _, _, err := NewHTTP2ClientConn(rt).Dial(newReq(ctx)) + require.ErrorIs(t, err, context.DeadlineExceeded) + }) +} + +func TestHTTP2Packets(t *testing.T) { + client, server := setupHTTP2Conns(t) + clientV4 := netip.MustParseAddr("192.0.2.2") + clientV6 := netip.MustParseAddr("2001:db8::2") + require.NoError(t, server.AssignAddresses([]netip.Prefix{netip.PrefixFrom(clientV4, 32), netip.PrefixFrom(clientV6, 128)})) + require.NoError(t, server.AdvertiseRoute([]IPRoute{ + {StartIP: netip.IPv4Unspecified(), EndIP: netip.MustParseAddr("255.255.255.255")}, + {StartIP: netip.IPv6Unspecified(), EndIP: netip.MustParseAddr("ffff:ffff:ffff:ffff:ffff:ffff:ffff:ffff")}, + })) + ctx, cancel := context.WithTimeout(t.Context(), 5*time.Second) + defer cancel() + _, err := client.ReceiveAddressAssignment(ctx) + require.NoError(t, err) + _, err = client.Routes(ctx) + require.NoError(t, err) + require.Equal(t, maxCapsulePacketSize, client.MaxPacketSize()) + + for _, tc := range []struct { + name string + up []byte + down []byte + ttlOff int + }{ + { + name: "IPv4", + up: ipv4Packet(64, 17, clientV4, testDst4, nil, []byte("foobar")), + down: ipv4Packet(64, 17, testDst4, clientV4, nil, []byte("barfoo")), + ttlOff: 8, + }, + { + name: "IPv6 larger than a QUIC datagram", + up: ipv6Packet(64, 17, clientV6, testDst6, bytes.Repeat([]byte("up"), 4500)), + down: ipv6Packet(64, 17, testDst6, clientV6, bytes.Repeat([]byte("down"), 2250)), + ttlOff: 7, + }, + } { + t.Run(tc.name, func(t *testing.T) { + for _, dir := range []struct { + from, to *Conn + packet []byte + }{ + {client, server, tc.up}, + {server, client, tc.down}, + } { + icmp, err := dir.from.WritePacket(slices.Clone(dir.packet)) + require.NoError(t, err) + require.Nil(t, icmp) + b := make([]byte, 1<<16) + n, err := dir.to.ReadPacket(b) + require.NoError(t, err) + require.Len(t, b[:n], len(dir.packet)) + require.Equal(t, dir.packet[tc.ttlOff]-1, b[tc.ttlOff]) + if tc.ttlOff == 8 { + require.True(t, ipv4ChecksumValid(b[:ipv4.HeaderLen])) + require.Equal(t, dir.packet[ipv4.HeaderLen:], b[ipv4.HeaderLen:n]) + } else { + require.Equal(t, dir.packet[ipv6.HeaderLen:], b[ipv6.HeaderLen:n]) + } + } + }) + } + + t.Run("in order both ways at once", func(t *testing.T) { + const count = 2000 + var wg sync.WaitGroup + for _, dir := range []struct { + from, to *Conn + src, dst netip.Addr + }{ + {client, server, clientV4, testDst4}, + {server, client, testDst4, clientV4}, + } { + wg.Go(func() { + for i := range count { + payload := make([]byte, 1200) + payload[0], payload[1] = byte(i>>8), byte(i) + if _, err := dir.from.WritePacket(ipv4Packet(64, 17, dir.src, dir.dst, nil, payload)); !assert.NoError(t, err) { + return + } + } + }) + wg.Go(func() { + b := make([]byte, 1500) + for i := range count { + n, err := dir.to.ReadPacket(b) + if !assert.NoError(t, err) || !assert.Equal(t, ipv4.HeaderLen+1200, n) { + return + } + if !assert.Equal(t, i, int(b[ipv4.HeaderLen])<<8|int(b[ipv4.HeaderLen+1])) { + return + } + } + }) + } + wg.Wait() + }) +} + +func TestHTTP2AddressRequest(t *testing.T) { + client, server := setupHTTP2Conns(t) + ctx, cancel := context.WithTimeout(t.Context(), 5*time.Second) + defer cancel() + + _, err := client.RequestAddresses([]netip.Prefix{ + netip.PrefixFrom(netip.IPv4Unspecified(), 32), + netip.PrefixFrom(netip.IPv6Unspecified(), 128), + }) + require.NoError(t, err) + req, err := server.ReceiveAddressRequest(ctx) + require.NoError(t, err) + require.Len(t, req.Prefixes, 2) + require.NoError(t, req.Respond([]netip.Prefix{netip.MustParsePrefix("192.0.2.2/32"), {}}, nil)) + + assigned, err := client.ReceiveAddressAssignment(ctx) + require.NoError(t, err) + require.Len(t, assigned, 2) + require.Equal(t, netip.MustParsePrefix("192.0.2.2/32"), assigned[0].IPPrefix) + require.True(t, assigned[1].Rejected()) +} + +func TestHTTP2Closing(t *testing.T) { + for _, side := range []string{"client", "proxy"} { + t.Run(side, func(t *testing.T) { + client, server := setupHTTP2Conns(t) + closing, peer := client, server + if side == "proxy" { + closing, peer = server, client + } + + require.NoError(t, closing.Close()) + _, err := closing.ReadPacket(make([]byte, 1500)) + require.ErrorIs(t, err, net.ErrClosed) + _, err = closing.WritePacket(ipv4Packet(64, 17, testSrc4, testDst4, nil, nil)) + require.ErrorIs(t, err, net.ErrClosed) + + ctx, cancel := context.WithTimeout(t.Context(), 5*time.Second) + defer cancel() + _, err = peer.Routes(ctx) + require.ErrorIs(t, err, net.ErrClosed) + var closeErr *CloseError + require.ErrorAs(t, err, &closeErr) + require.True(t, closeErr.Remote) + _, err = peer.ReadPacket(make([]byte, 1500)) + require.ErrorIs(t, err, net.ErrClosed) + }) + } +} + +func TestHTTP2CloseUnblocksWrites(t *testing.T) { + str, pw := newTestHTTP2Stream() + defer pw.Close() + conn := newProxiedConn(str) + + writeErr := make(chan error, 1) + go func() { + for { + if _, err := conn.WritePacket(ipv4Packet(64, 17, testSrc4, testDst4, nil, make([]byte, 1000))); err != nil { + writeErr <- err + return + } + } + }() + require.Eventually(t, func() bool { + str.body.mu.Lock() + defer str.body.mu.Unlock() + return len(str.body.buf) >= maxBufferedRequestBody + }, 5*time.Second, time.Millisecond) + + closed := make(chan error, 1) + go func() { closed <- conn.Close() }() + select { + case err := <-closed: + require.NoError(t, err) + case <-time.After(5 * time.Second): + t.Fatal("Close blocked on a stalled stream") + } + require.ErrorIs(t, <-writeErr, net.ErrClosed) +} + +func TestHTTP2DatagramCapsules(t *testing.T) { + str, pw := newTestHTTP2Stream() + defer pw.Close() + conn := newProxiedConn(str) + t.Cleanup(func() { conn.Close() }) + require.NoError(t, conn.AdvertiseRoute([]IPRoute{ + {StartIP: netip.IPv4Unspecified(), EndIP: netip.MustParseAddr("255.255.255.255")}, + })) + + capsule := func(payload []byte) []byte { + b := quicvarint.Append(nil, uint64(capsuleTypeDatagram)) + b = quicvarint.Append(b, uint64(len(payload))) + return append(b, payload...) + } + packet := ipv4Packet(64, 17, testSrc4, testDst4, nil, []byte("foobar")) + go func() { + for _, c := range [][]byte{ + capsule(nil), + capsule([]byte{0x40}), + capsule(append([]byte{0x02}, packet...)), + capsule(append(bytes.Clone(contextIDZero), make([]byte, maxCapsulePacketSize+1)...)), + capsule(append(bytes.Clone(contextIDZero), packet...)), + } { + if _, err := pw.Write(c); err != nil { + return + } + } + }() + b := make([]byte, 1500) + n, err := conn.ReadPacket(b) + require.NoError(t, err) + require.Equal(t, packet, b[:n]) +} + +func TestHTTP2WritesDatagramCapsules(t *testing.T) { + str, pw := newTestHTTP2Stream() + defer pw.Close() + conn := newProxiedConn(str) + t.Cleanup(func() { conn.Close() }) + + packet := ipv4Packet(64, 17, testSrc4, testDst4, nil, []byte("foobar")) + _, err := conn.WritePacket(slices.Clone(packet)) + require.NoError(t, err) + + p := http3.NewCapsuleParser(str.body) + typ, cr, err := p.Next() + require.NoError(t, err) + require.Equal(t, capsuleTypeDatagram, typ) + data, err := io.ReadAll(cr) + require.NoError(t, err) + require.Equal(t, contextIDZero, data[:len(contextIDZero)]) + sent := data[len(contextIDZero):] + require.Len(t, sent, len(packet)) + require.Equal(t, packet[8]-1, sent[8]) + require.Equal(t, packet[ipv4.HeaderLen:], sent[ipv4.HeaderLen:]) +} + +func TestRequestBody(t *testing.T) { + t.Run("coalesces writes", func(t *testing.T) { + b := newRequestBody() + for _, s := range []string{"foo", "bar", "baz"} { + _, err := b.Write([]byte(s)) + require.NoError(t, err) + } + p := make([]byte, 16) + n, err := b.Read(p) + require.NoError(t, err) + require.Equal(t, "foobarbaz", string(p[:n])) + }) + + t.Run("blocks writes while full", func(t *testing.T) { + b := newRequestBody() + _, err := b.Write(make([]byte, maxBufferedRequestBody)) + require.NoError(t, err) + written := make(chan struct{}) + go func() { + b.Write([]byte("x")) + close(written) + }() + select { + case <-written: + t.Fatal("write did not block") + case <-time.After(50 * time.Millisecond): + } + _, err = b.Read(make([]byte, maxBufferedRequestBody)) + require.NoError(t, err) + select { + case <-written: + case <-time.After(time.Second): + t.Fatal("write stayed blocked") + } + }) + + t.Run("close", func(t *testing.T) { + b := newRequestBody() + _, err := b.Write([]byte("foo")) + require.NoError(t, err) + require.NoError(t, b.Close()) + _, err = b.Write([]byte("bar")) + require.ErrorIs(t, err, io.ErrClosedPipe) + data, err := io.ReadAll(b) + require.NoError(t, err) + require.Equal(t, "foo", string(data)) + }) + + t.Run("write deadline", func(t *testing.T) { + b := newRequestBody() + _, err := b.Write(make([]byte, maxBufferedRequestBody)) + require.NoError(t, err) + writeErr := make(chan error, 1) + go func() { + _, err := b.Write([]byte("x")) + writeErr <- err + }() + require.NoError(t, b.SetWriteDeadline(time.Now().Add(50*time.Millisecond))) + select { + case err := <-writeErr: + require.ErrorIs(t, err, os.ErrDeadlineExceeded) + case <-time.After(time.Second): + t.Fatal("write deadline did not unblock the write") + } + require.NoError(t, b.Close()) + data, err := io.ReadAll(b) + require.NoError(t, err) + require.Len(t, data, maxBufferedRequestBody) + }) + + t.Run("close with error", func(t *testing.T) { + b := newRequestBody() + _, err := b.Write([]byte("foo")) + require.NoError(t, err) + b.CloseWithError(net.ErrClosed) + _, err = b.Read(make([]byte, 16)) + require.ErrorIs(t, err, net.ErrClosed) + _, err = b.Write([]byte("bar")) + require.ErrorIs(t, err, net.ErrClosed) + }) +} + +func TestProxyNeedsAnHTTPStream(t *testing.T) { + _, err := (&Proxy{}).Proxy(httptest.NewRecorder(), &ProxyRequest{}) + require.EqualError(t, err, "connect-ip: response writer is neither an HTTP/3 nor an HTTP/2 stream") +} diff --git a/transport/internet/masque/connectip/proxy.go b/transport/internet/masque/connectip/proxy.go index 323d666df..ab0686ab5 100644 --- a/transport/internet/masque/connectip/proxy.go +++ b/transport/internet/masque/connectip/proxy.go @@ -7,6 +7,7 @@ package connectip import ( + "bufio" "errors" "net/http" @@ -18,13 +19,25 @@ var contextIDZero = quicvarint.Append([]byte{}, 0) type Proxy struct{} -func (s *Proxy) Proxy(w http.ResponseWriter, _ *ProxyRequest) (*Conn, error) { +func (s *Proxy) Proxy(w http.ResponseWriter, r *ProxyRequest) (*Conn, error) { streamer, ok := w.(http3.HTTPStreamer) - if !ok { - return nil, errors.New("connect-ip: response writer is not an HTTP/3 stream") + if !ok && (r == nil || r.body == nil) { + return nil, errors.New("connect-ip: response writer is neither an HTTP/3 nor an HTTP/2 stream") } w.Header().Set(http3.CapsuleProtocolHeader, capsuleProtocolHeaderValue) w.WriteHeader(http.StatusOK) - return newProxiedConn(streamer.HTTPStream()), nil + if ok { + return newProxiedConn(streamer.HTTPStream()), nil + } + controller := http.NewResponseController(w) + if err := controller.Flush(); err != nil { + return nil, err + } + return newProxiedConn(&http2ResponseStream{ + reader: bufio.NewReader(r.body), + body: r.body, + w: w, + controller: controller, + }), nil } diff --git a/transport/internet/masque/connectip/request.go b/transport/internet/masque/connectip/request.go index 4d2c3d6ea..9091caf4e 100644 --- a/transport/internet/masque/connectip/request.go +++ b/transport/internet/masque/connectip/request.go @@ -10,6 +10,7 @@ import ( "context" "errors" "fmt" + "io" "net/http" "strings" @@ -45,7 +46,9 @@ func (r *Request) Header() http.Header { return r.req.Header } func (r *Request) httpRequest() *http.Request { return r.req } -type ProxyRequest struct{} +type ProxyRequest struct { + body io.ReadCloser +} type ProxyRequestParseError struct { HTTPStatus int @@ -62,10 +65,14 @@ func ParseProxyRequest(r *http.Request) (*ProxyRequest, error) { Err: fmt.Errorf("expected CONNECT request, got %s", r.Method), } } - if r.Proto != requestProtocol { + protocol := r.Proto + if r.ProtoMajor == 2 { + protocol = r.Header.Get(":protocol") + } + if protocol != requestProtocol { return nil, &ProxyRequestParseError{ HTTPStatus: http.StatusNotImplemented, - Err: fmt.Errorf("unexpected protocol: %s", r.Proto), + Err: fmt.Errorf("unexpected protocol: %s", protocol), } } capsuleHeaderValues, ok := r.Header[http3.CapsuleProtocolHeader] @@ -82,6 +89,9 @@ func ParseProxyRequest(r *http.Request) (*ProxyRequest, error) { } } + if r.ProtoMajor == 2 { + return &ProxyRequest{body: r.Body}, nil + } return &ProxyRequest{}, nil } diff --git a/transport/internet/masque/connectip/request_test.go b/transport/internet/masque/connectip/request_test.go index 61c8a9c43..223b8a3ec 100644 --- a/transport/internet/masque/connectip/request_test.go +++ b/transport/internet/masque/connectip/request_test.go @@ -71,6 +71,24 @@ func TestProxyRequestParsing(t *testing.T) { require.Equal(t, http.StatusNotImplemented, err.(*ProxyRequestParseError).HTTPStatus) }) + t.Run("HTTP/2", func(t *testing.T) { + req := newRequest("https://localhost:1234/masque/ip") + req.Proto, req.ProtoMajor = "HTTP/2.0", 2 + req.Header.Set(":protocol", requestProtocol) + r, err := ParseProxyRequest(req) + require.NoError(t, err) + require.Equal(t, &ProxyRequest{body: req.Body}, r) + }) + + t.Run("wrong protocol over HTTP/2", func(t *testing.T) { + req := newRequest("https://localhost:1234/masque") + req.Proto, req.ProtoMajor = "HTTP/2.0", 2 + req.Header.Set(":protocol", "websocket") + _, err := ParseProxyRequest(req) + require.EqualError(t, err, "unexpected protocol: websocket") + require.Equal(t, http.StatusNotImplemented, err.(*ProxyRequestParseError).HTTPStatus) + }) + t.Run("wrong request method", func(t *testing.T) { req := newRequest("https://localhost:1234/masque") req.Method = http.MethodHead diff --git a/transport/internet/masque/dialer.go b/transport/internet/masque/dialer.go index 65c75cc81..aff454308 100644 --- a/transport/internet/masque/dialer.go +++ b/transport/internet/masque/dialer.go @@ -2,9 +2,12 @@ package masque import ( "context" + "net/http" "net/netip" "reflect" "runtime" + "slices" + "strconv" "strings" "time" @@ -22,6 +25,7 @@ import ( "github.com/xtls/xray-core/transport/internet/masque/connectip" "github.com/xtls/xray-core/transport/internet/stat" "github.com/xtls/xray-core/transport/internet/tls" + "golang.org/x/net/http2" ) const ( @@ -35,6 +39,9 @@ func Dial(ctx context.Context, dest net.Destination, streamSettings *internet.Me return nil, errors.New("tls config is nil") } config := streamSettings.ProtocolSettings.(*Config) + if usesHTTP2(tlsConfig) { + return dialHTTP2(ctx, dest, streamSettings, tlsConfig, config) + } dest.Network = net.Network_UDP gotlsConfig := tlsConfig.GetTLSConfig(tls.WithDestination(dest)) @@ -112,7 +119,10 @@ func Dial(ctx context.Context, dest net.Destination, streamSettings *internet.Me return nil, errors.New("unknown congestion control: ", quicParams.Congestion) } - conn, err := establish(ctx, qconn, config, authority(config, gotlsConfig.ServerName, dest.Port)) + cc := (&http3.Transport{EnableDatagrams: true, DisableCompression: true}).NewClientConn(qconn) + conn, err := establish(ctx, connectip.NewClientConn(cc), quicConn{qconn}, func() { + qconn.CloseWithError(quic.ApplicationErrorCode(http3.ErrCodeRequestCanceled), "") + }, config, authority(config, gotlsConfig.ServerName, dest.Port)) if err != nil { qconn.CloseWithError(quic.ApplicationErrorCode(http3.ErrCodeNoError), "") return nil, err @@ -120,10 +130,58 @@ func Dial(ctx context.Context, dest net.Destination, streamSettings *internet.Me return conn, nil } -func establish(ctx context.Context, qconn *quic.Conn, config *Config, host string) (*Conn, error) { - stop := context.AfterFunc(ctx, func() { - qconn.CloseWithError(quic.ApplicationErrorCode(http3.ErrCodeRequestCanceled), "") - }) +func usesHTTP2(config *tls.Config) bool { + return slices.Contains(config.NextProtocol, http2.NextProtoTLS) && !slices.Contains(config.NextProtocol, http3.NextProtoH3) +} + +func dialHTTP2(ctx context.Context, dest net.Destination, streamSettings *internet.MemoryStreamConfig, tlsConfig *tls.Config, config *Config) (stat.Connection, error) { + dest.Network = net.Network_TCP + gotlsConfig := tlsConfig.GetTLSConfig(tls.WithDestination(dest)) + + var conn net.Conn + var err error + if streamSettings.FinalMask != nil { + conn, err = streamSettings.FinalMask.DialTCP(ctx, dest) + } else { + conn, err = internet.DialSystem(ctx, dest, streamSettings.SocketSettings) + } + if err != nil { + return nil, errors.New("failed to dial to dest").Base(err) + } + if fingerprint := tls.GetFingerprint(tlsConfig.Fingerprint); fingerprint != nil { + conn = tls.UClient(conn, gotlsConfig, fingerprint) + } else { + conn = tls.Client(conn, gotlsConfig) + } + tlsConn := conn.(tls.Interface) + if err := tlsConn.HandshakeContext(ctx); err != nil { + conn.Close() + return nil, err + } + if protocol := tlsConn.NegotiatedProtocol(); protocol != http2.NextProtoTLS { + conn.Close() + return nil, errors.New("the server negotiated ", strconv.Quote(protocol), " instead of h2") + } + + cc, err := newHTTP2ClientConn(conn) + if err != nil { + conn.Close() + return nil, err + } + mconn, err := establish(ctx, connectip.NewHTTP2ClientConn(cc), cc, func() { cc.Close() }, config, authority(config, gotlsConfig.ServerName, dest.Port)) + if err != nil { + cc.Close() + return nil, err + } + return mconn, nil +} + +type tunnelClient interface { + Dial(*connectip.Request) (*connectip.Conn, *http.Response, error) +} + +func establish(ctx context.Context, client tunnelClient, hconn httpConn, abort func(), config *Config, host string) (*Conn, error) { + stop := context.AfterFunc(ctx, abort) defer stop() req, err := connectip.NewRequest(ctx, "https://"+host+config.Path) @@ -151,8 +209,7 @@ func establish(ctx context.Context, qconn *quic.Conn, config *Config, host strin header.Del("User-Agent") } - cc := (&http3.Transport{EnableDatagrams: true, DisableCompression: true}).NewClientConn(qconn) - ipConn, _, err := connectip.NewClientConn(cc).Dial(req) + ipConn, _, err := client.Dial(req) if err != nil { if ctx.Err() != nil { err = context.Cause(ctx) @@ -188,7 +245,7 @@ func establish(ctx context.Context, qconn *quic.Conn, config *Config, host strin conn := &Conn{ ipConn: ipConn, - quicConn: qconn, + httpConn: hconn, local: local, } go conn.serveAddressAssignments() diff --git a/transport/internet/masque/dialer_test.go b/transport/internet/masque/dialer_test.go index 20dd15e18..ddeef00b2 100644 --- a/transport/internet/masque/dialer_test.go +++ b/transport/internet/masque/dialer_test.go @@ -7,8 +7,27 @@ import ( "github.com/xtls/xray-core/common/net" "github.com/xtls/xray-core/transport/internet/masque/connectip" + "github.com/xtls/xray-core/transport/internet/tls" ) +func TestUsesHTTP2(t *testing.T) { + for _, c := range []struct { + alpn []string + want bool + }{ + {alpn: nil, want: false}, + {alpn: []string{"h3"}, want: false}, + {alpn: []string{"h2"}, want: true}, + {alpn: []string{"h2", "http/1.1"}, want: true}, + {alpn: []string{"h3", "h2"}, want: false}, + {alpn: []string{"http/1.1"}, want: false}, + } { + if got := usesHTTP2(&tls.Config{NextProtocol: c.alpn}); got != c.want { + t.Errorf("usesHTTP2(%q) = %v, want %v", c.alpn, got, c.want) + } + } +} + func TestAuthority(t *testing.T) { for _, c := range []struct { host, serverName string diff --git a/transport/internet/masque/http2.go b/transport/internet/masque/http2.go new file mode 100644 index 000000000..4665aeaf4 --- /dev/null +++ b/transport/internet/masque/http2.go @@ -0,0 +1,600 @@ +package masque + +import ( + "bufio" + "bytes" + "context" + go_errors "errors" + "io" + "maps" + "net" + "net/http" + "slices" + "strconv" + "strings" + "sync" + "sync/atomic" + "time" + + "github.com/xtls/xray-core/common/errors" + "golang.org/x/net/http2" + "golang.org/x/net/http2/hpack" +) + +const ( + http2StreamID = 1 + http2DefaultWindow = 65535 + http2DefaultFrameSize = 16 << 10 + http2HeaderTableSize = 64 << 10 + http2StreamWindow = 6 << 20 + http2ConnectionWindow = 15 << 20 + http2MaxHeaderListSize = 256 << 10 + http2WindowUpdateSize = 1 << 20 + http2KeepAlivePeriod = 10 * time.Second + http2IdleTimeout = 30 * time.Second + http2DefaultUserAgent = "Go-http-client/2.0" +) + +var ( + errHTTP2StreamUsed = go_errors.New("http2: the connection carries a single stream") + errHTTP2NoExtendedConnect = go_errors.New("http2: the server did not enable extended CONNECT") + errHTTP2BodyClosed = go_errors.New("http2: response body closed") + errHTTP2IdleTimeout = go_errors.New("http2: no frame received within the idle timeout") +) + +type http2ClientConn struct { + conn net.Conn + + wmu sync.Mutex + bw *bufio.Writer + fr *http2.Framer + hbuf bytes.Buffer + henc *hpack.Encoder + + lastFrame atomic.Int64 + settings chan struct{} + responses chan *http.Response + aborted chan struct{} + done chan struct{} + + mu sync.Mutex + cond sync.Cond + err error + gotSettings bool + extendedConnect bool + maxFrameSize uint32 + initialWindow int64 + connSendWindow int64 + streamSendWindow int64 + connRecvWindow int64 + streamRecvWindow int64 + streamOpen bool + gotResponse bool + sentEnd bool + recvEnd bool + streamErr error + reqBody io.Closer + recv bytes.Buffer + recvErr error + recvUnacked int64 +} + +func newHTTP2ClientConn(conn net.Conn) (*http2ClientConn, error) { + c := &http2ClientConn{ + conn: conn, + bw: bufio.NewWriter(conn), + settings: make(chan struct{}), + responses: make(chan *http.Response, 1), + aborted: make(chan struct{}), + done: make(chan struct{}), + maxFrameSize: http2DefaultFrameSize, + initialWindow: http2DefaultWindow, + connSendWindow: http2DefaultWindow, + connRecvWindow: http2ConnectionWindow, + streamRecvWindow: http2StreamWindow, + } + c.cond.L = &c.mu + c.fr = http2.NewFramer(c.bw, bufio.NewReader(conn)) + c.fr.SetMaxReadFrameSize(http2DefaultFrameSize) + c.henc = hpack.NewEncoder(&c.hbuf) + c.henc.SetMaxDynamicTableSizeLimit(0) + c.fr.ReadMetaHeaders = hpack.NewDecoder(http2HeaderTableSize, nil) + c.fr.MaxHeaderListSize = http2MaxHeaderListSize + c.lastFrame.Store(time.Now().UnixNano()) + + if err := c.write(func(fr *http2.Framer) error { + if _, err := c.bw.WriteString(http2.ClientPreface); err != nil { + return err + } + if err := fr.WriteSettings( + http2.Setting{ID: http2.SettingHeaderTableSize, Val: http2HeaderTableSize}, + http2.Setting{ID: http2.SettingEnablePush, Val: 0}, + http2.Setting{ID: http2.SettingInitialWindowSize, Val: http2StreamWindow}, + http2.Setting{ID: http2.SettingMaxHeaderListSize, Val: http2MaxHeaderListSize}, + ); err != nil { + return err + } + return fr.WriteWindowUpdate(0, http2ConnectionWindow-http2DefaultWindow) + }); err != nil { + return nil, err + } + go c.readLoop() + go c.keepAlive() + return c, nil +} + +func (c *http2ClientConn) LocalAddr() net.Addr { + return c.conn.LocalAddr() +} + +func (c *http2ClientConn) RemoteAddr() net.Addr { + return c.conn.RemoteAddr() +} + +func (c *http2ClientConn) Close() error { + c.fail(net.ErrClosed) + return nil +} + +func (c *http2ClientConn) RoundTrip(req *http.Request) (*http.Response, error) { + rsp, err := c.roundTrip(req) + if err != nil && req.Body != nil { + req.Body.Close() + } + return rsp, err +} + +func (c *http2ClientConn) roundTrip(req *http.Request) (*http.Response, error) { + ctx := req.Context() + select { + case <-c.settings: + case <-c.done: + return nil, c.connErr() + case <-ctx.Done(): + return nil, context.Cause(ctx) + } + + c.mu.Lock() + switch { + case c.err != nil: + err := c.err + c.mu.Unlock() + return nil, err + case c.streamOpen: + c.mu.Unlock() + return nil, errHTTP2StreamUsed + case req.Header.Get(":protocol") != "" && !c.extendedConnect: + c.mu.Unlock() + return nil, errHTTP2NoExtendedConnect + } + c.streamOpen = true + c.streamSendWindow = c.initialWindow + c.reqBody = req.Body + maxFrameSize := int(c.maxFrameSize) + c.mu.Unlock() + + if err := c.writeHeaders(req, maxFrameSize); err != nil { + c.fail(err) + return nil, err + } + if req.Body != nil { + go c.writeBody(req.Body) + } else { + c.endStream() + } + context.AfterFunc(ctx, func() { c.abortStream(context.Cause(ctx), true) }) + + select { + case rsp := <-c.responses: + return rsp, nil + case <-c.aborted: + c.mu.Lock() + err := c.streamErr + c.mu.Unlock() + return nil, err + } +} + +func (c *http2ClientConn) writeHeaders(req *http.Request, maxFrameSize int) error { + c.wmu.Lock() + defer c.wmu.Unlock() + + c.hbuf.Reset() + field := func(name, value string) { + c.henc.WriteField(hpack.HeaderField{Name: name, Value: value}) + } + host := req.Host + if host == "" { + host = req.URL.Host + } + field(":method", req.Method) + field(":authority", host) + field(":scheme", req.URL.Scheme) + field(":path", req.URL.RequestURI()) + if protocol := req.Header.Get(":protocol"); protocol != "" { + field(":protocol", protocol) + } + if _, ok := req.Header["User-Agent"]; !ok { + field("user-agent", http2DefaultUserAgent) + } + for _, k := range slices.Sorted(maps.Keys(req.Header)) { + name := strings.ToLower(k) + switch name { + case ":protocol", "host", "connection", "proxy-connection", "keep-alive", "transfer-encoding", "upgrade", "content-length": + continue + } + for _, v := range req.Header[k] { + if name == "user-agent" && v == "" { + continue + } + field(name, v) + } + } + + block := c.hbuf.Bytes() + for first := true; first || len(block) > 0; first = false { + chunk := block[:min(len(block), maxFrameSize)] + block = block[len(chunk):] + var err error + if first { + err = c.fr.WriteHeaders(http2.HeadersFrameParam{StreamID: http2StreamID, BlockFragment: chunk, EndHeaders: len(block) == 0}) + } else { + err = c.fr.WriteContinuation(http2StreamID, len(block) == 0, chunk) + } + if err != nil { + return err + } + } + return c.bw.Flush() +} + +func (c *http2ClientConn) writeBody(body io.ReadCloser) { + defer body.Close() + buf := make([]byte, http2DefaultFrameSize) + for { + n, err := body.Read(buf) + for data := buf[:n]; len(data) > 0; { + allowed, err := c.awaitSendWindow(len(data)) + if err != nil { + return + } + if err := c.write(func(fr *http2.Framer) error { + return fr.WriteData(http2StreamID, false, data[:allowed]) + }); err != nil { + c.fail(err) + return + } + data = data[allowed:] + } + if err == io.EOF { + c.endStream() + return + } + if err != nil { + c.abortStream(err, true) + return + } + } +} + +func (c *http2ClientConn) awaitSendWindow(n int) (int, error) { + c.mu.Lock() + defer c.mu.Unlock() + for { + if c.streamErr != nil { + return 0, c.streamErr + } + if window := min(c.connSendWindow, c.streamSendWindow); window > 0 { + n = int(min(int64(n), window, int64(c.maxFrameSize))) + c.connSendWindow -= int64(n) + c.streamSendWindow -= int64(n) + return n, nil + } + c.cond.Wait() + } +} + +func (c *http2ClientConn) endStream() { + c.mu.Lock() + if c.streamErr != nil || c.sentEnd { + c.mu.Unlock() + return + } + c.sentEnd = true + c.mu.Unlock() + if err := c.write(func(fr *http2.Framer) error { + return fr.WriteData(http2StreamID, true, nil) + }); err != nil { + c.fail(err) + } +} + +func (c *http2ClientConn) write(f func(*http2.Framer) error) error { + c.wmu.Lock() + defer c.wmu.Unlock() + if err := f(c.fr); err != nil { + return err + } + return c.bw.Flush() +} + +func (c *http2ClientConn) connErr() error { + c.mu.Lock() + defer c.mu.Unlock() + return c.err +} + +func (c *http2ClientConn) fail(err error) { + c.mu.Lock() + if c.err == nil { + c.err = err + } + c.mu.Unlock() + c.abortStream(err, false) + c.conn.Close() +} + +func (c *http2ClientConn) abortStream(err error, reset bool) { + c.mu.Lock() + if c.streamErr != nil { + c.mu.Unlock() + return + } + c.streamErr = err + if c.recvErr == nil { + c.recvErr = err + } + reset = reset && c.streamOpen && !(c.sentEnd && c.recvEnd) + body := c.reqBody + close(c.aborted) + c.cond.Broadcast() + c.mu.Unlock() + + if body != nil { + body.Close() + } + if reset { + go c.write(func(fr *http2.Framer) error { + return fr.WriteRSTStream(http2StreamID, http2.ErrCodeCancel) + }) + } +} + +func (c *http2ClientConn) keepAlive() { + ticker := time.NewTicker(http2KeepAlivePeriod) + defer ticker.Stop() + for { + select { + case <-c.done: + return + case <-ticker.C: + } + idle := time.Since(time.Unix(0, c.lastFrame.Load())) + if idle >= http2IdleTimeout { + c.fail(errHTTP2IdleTimeout) + return + } + if idle >= http2KeepAlivePeriod { + go c.write(func(fr *http2.Framer) error { + return fr.WritePing(false, [8]byte{}) + }) + } + } +} + +func (c *http2ClientConn) readLoop() { + defer close(c.done) + for { + f, err := c.fr.ReadFrame() + if err != nil { + var streamErr http2.StreamError + if go_errors.As(err, &streamErr) && streamErr.StreamID == http2StreamID { + c.abortStream(streamErr, true) + continue + } + c.fail(err) + return + } + c.lastFrame.Store(time.Now().UnixNano()) + if err := c.handleFrame(f); err != nil { + c.fail(err) + return + } + } +} + +func (c *http2ClientConn) handleFrame(f http2.Frame) error { + switch f := f.(type) { + case *http2.SettingsFrame: + if f.IsAck() { + return nil + } + if err := c.applySettings(f); err != nil { + return err + } + return c.write((*http2.Framer).WriteSettingsAck) + case *http2.PingFrame: + if f.IsAck() { + return nil + } + return c.write(func(fr *http2.Framer) error { + return fr.WritePing(true, f.Data) + }) + case *http2.WindowUpdateFrame: + c.mu.Lock() + switch f.StreamID { + case 0: + c.connSendWindow += int64(f.Increment) + case http2StreamID: + c.streamSendWindow += int64(f.Increment) + } + c.cond.Broadcast() + c.mu.Unlock() + case *http2.MetaHeadersFrame: + if f.StreamID == http2StreamID { + c.handleHeaders(f) + } + case *http2.DataFrame: + return c.handleData(f) + case *http2.RSTStreamFrame: + if f.StreamID == http2StreamID { + c.abortStream(http2.StreamError{StreamID: f.StreamID, Code: f.ErrCode}, false) + } + case *http2.GoAwayFrame: + if f.ErrCode != http2.ErrCodeNo || f.LastStreamID < http2StreamID { + return errors.New("http2: the server sent GOAWAY (", f.ErrCode, ")") + } + case *http2.PushPromiseFrame: + return http2.ConnectionError(http2.ErrCodeProtocol) + } + return nil +} + +func (c *http2ClientConn) applySettings(f *http2.SettingsFrame) error { + c.mu.Lock() + defer c.mu.Unlock() + if err := f.ForeachSetting(func(s http2.Setting) error { + if err := s.Valid(); err != nil { + return err + } + switch s.ID { + case http2.SettingMaxFrameSize: + c.maxFrameSize = s.Val + case http2.SettingInitialWindowSize: + c.streamSendWindow += int64(s.Val) - c.initialWindow + c.initialWindow = int64(s.Val) + case http2.SettingEnableConnectProtocol: + if !c.gotSettings { + c.extendedConnect = s.Val == 1 + } + } + return nil + }); err != nil { + return err + } + if !c.gotSettings { + c.gotSettings = true + close(c.settings) + } + c.cond.Broadcast() + return nil +} + +func (c *http2ClientConn) handleHeaders(f *http2.MetaHeadersFrame) { + c.mu.Lock() + gotResponse := c.gotResponse + c.mu.Unlock() + if !gotResponse { + status, err := strconv.Atoi(f.PseudoValue("status")) + if err != nil || status < 100 || status > 999 { + c.abortStream(errors.New("http2: invalid response status ", strconv.Quote(f.PseudoValue("status"))), true) + return + } + if status < 200 { + return + } + header := make(http.Header) + for _, hf := range f.RegularFields() { + header.Add(hf.Name, hf.Value) + } + c.mu.Lock() + c.gotResponse = true + c.mu.Unlock() + c.responses <- &http.Response{ + Status: strconv.Itoa(status) + " " + http.StatusText(status), + StatusCode: status, + Proto: "HTTP/2.0", + ProtoMajor: 2, + Header: header, + Body: &http2ResponseBody{c}, + ContentLength: -1, + } + } + if f.StreamEnded() { + c.mu.Lock() + c.recvEnd = true + if c.recvErr == nil { + c.recvErr = io.EOF + } + c.cond.Broadcast() + c.mu.Unlock() + } +} + +func (c *http2ClientConn) handleData(f *http2.DataFrame) error { + size := int64(f.Length) + c.mu.Lock() + c.connRecvWindow -= size + if c.connRecvWindow < 0 { + c.mu.Unlock() + return http2.ConnectionError(http2.ErrCodeFlowControl) + } + if f.StreamID != http2StreamID || c.recvErr != nil { + c.connRecvWindow += size + c.mu.Unlock() + if size == 0 { + return nil + } + return c.write(func(fr *http2.Framer) error { + return fr.WriteWindowUpdate(0, uint32(size)) + }) + } + c.streamRecvWindow -= size + if c.streamRecvWindow < 0 { + c.mu.Unlock() + return http2.ConnectionError(http2.ErrCodeFlowControl) + } + c.recv.Write(f.Data()) + c.recvUnacked += size - int64(len(f.Data())) + if f.StreamEnded() { + c.recvEnd = true + c.recvErr = io.EOF + } + c.cond.Broadcast() + c.mu.Unlock() + return nil +} + +type http2ResponseBody struct { + c *http2ClientConn +} + +func (b *http2ResponseBody) Read(p []byte) (int, error) { + c := b.c + c.mu.Lock() + for c.recv.Len() == 0 && c.recvErr == nil { + c.cond.Wait() + } + if c.recv.Len() == 0 { + err := c.recvErr + c.mu.Unlock() + return 0, err + } + n, _ := c.recv.Read(p) + c.recvUnacked += int64(n) + var update int64 + if c.recvUnacked >= http2WindowUpdateSize && !c.recvEnd { + update = c.recvUnacked + c.recvUnacked = 0 + c.connRecvWindow += update + c.streamRecvWindow += update + } + c.mu.Unlock() + + if update > 0 { + if err := c.write(func(fr *http2.Framer) error { + if err := fr.WriteWindowUpdate(0, uint32(update)); err != nil { + return err + } + return fr.WriteWindowUpdate(http2StreamID, uint32(update)) + }); err != nil { + c.fail(err) + } + } + return n, nil +} + +func (b *http2ResponseBody) Close() error { + b.c.abortStream(errHTTP2BodyClosed, true) + return nil +} diff --git a/transport/internet/masque/http2_test.go b/transport/internet/masque/http2_test.go new file mode 100644 index 000000000..b2d73b71d --- /dev/null +++ b/transport/internet/masque/http2_test.go @@ -0,0 +1,405 @@ +package masque + +import ( + "bytes" + "context" + "io" + "net" + "net/http" + "strings" + "testing" + "time" + + "github.com/stretchr/testify/require" + "golang.org/x/net/http2" + "golang.org/x/net/http2/hpack" +) + +type http2Peer struct { + t *testing.T + conn net.Conn + fr *http2.Framer + hbuf bytes.Buffer + henc *hpack.Encoder +} + +func tcpPipe(t *testing.T) (net.Conn, net.Conn) { + t.Helper() + ln, err := net.Listen("tcp", "127.0.0.1:0") + require.NoError(t, err) + defer ln.Close() + accepted := make(chan net.Conn, 1) + go func() { + conn, _ := ln.Accept() + accepted <- conn + }() + client, err := net.Dial("tcp", ln.Addr().String()) + require.NoError(t, err) + server := <-accepted + require.NotNil(t, server) + return client, server +} + +func newHTTP2Peer(t *testing.T, settings ...http2.Setting) (*http2ClientConn, *http2Peer) { + t.Helper() + client, server := tcpPipe(t) + p := &http2Peer{t: t, conn: server, fr: http2.NewFramer(server, server)} + p.henc = hpack.NewEncoder(&p.hbuf) + p.fr.ReadMetaHeaders = hpack.NewDecoder(4096, nil) + t.Cleanup(func() { server.Close() }) + + ccErr := make(chan error, 1) + var cc *http2ClientConn + go func() { + var err error + cc, err = newHTTP2ClientConn(client) + ccErr <- err + }() + preface := make([]byte, len(http2.ClientPreface)) + _, err := io.ReadFull(server, preface) + require.NoError(t, err) + require.Equal(t, http2.ClientPreface, string(preface)) + + f := p.readFrame() + require.IsType(t, &http2.SettingsFrame{}, f) + var got []http2.Setting + f.(*http2.SettingsFrame).ForeachSetting(func(s http2.Setting) error { + got = append(got, s) + return nil + }) + require.Equal(t, []http2.Setting{ + {ID: http2.SettingHeaderTableSize, Val: http2HeaderTableSize}, + {ID: http2.SettingEnablePush, Val: 0}, + {ID: http2.SettingInitialWindowSize, Val: http2StreamWindow}, + {ID: http2.SettingMaxHeaderListSize, Val: http2MaxHeaderListSize}, + }, got) + f = p.readFrame() + require.IsType(t, &http2.WindowUpdateFrame{}, f) + require.Equal(t, uint32(0), f.Header().StreamID) + require.Equal(t, uint32(http2ConnectionWindow-http2DefaultWindow), f.(*http2.WindowUpdateFrame).Increment) + require.NoError(t, <-ccErr) + t.Cleanup(func() { cc.Close() }) + + require.NoError(t, p.fr.WriteSettings(settings...)) + f = p.readFrame() + require.IsType(t, &http2.SettingsFrame{}, f) + require.True(t, f.(*http2.SettingsFrame).IsAck()) + return cc, p +} + +func (p *http2Peer) readFrame() http2.Frame { + p.t.Helper() + p.conn.SetReadDeadline(time.Now().Add(5 * time.Second)) + f, err := p.fr.ReadFrame() + require.NoError(p.t, err) + return f +} + +func (p *http2Peer) writeHeaders(endStream bool, fields ...string) { + p.t.Helper() + p.hbuf.Reset() + for i := 0; i < len(fields); i += 2 { + require.NoError(p.t, p.henc.WriteField(hpack.HeaderField{Name: fields[i], Value: fields[i+1]})) + } + require.NoError(p.t, p.fr.WriteHeaders(http2.HeadersFrameParam{ + StreamID: http2StreamID, + BlockFragment: p.hbuf.Bytes(), + EndHeaders: true, + EndStream: endStream, + })) +} + +func connectRequest(t *testing.T, ctx context.Context, body io.ReadCloser) *http.Request { + req, err := http.NewRequestWithContext(ctx, http.MethodConnect, "https://proxy.example/.well-known/masque/ip/*/*/", body) + require.NoError(t, err) + req.Header[":protocol"] = []string{"connect-ip"} + req.Header.Set("Capsule-Protocol", "?1") + req.Header.Set("Authorization", "Basic dTpw") + req.Header["User-Agent"] = nil + return req +} + +func TestHTTP2ClientRequest(t *testing.T) { + cc, p := newHTTP2Peer(t, http2.Setting{ID: http2.SettingEnableConnectProtocol, Val: 1}) + pr, pw := io.Pipe() + + type result struct { + rsp *http.Response + err error + } + results := make(chan result, 1) + go func() { + rsp, err := cc.RoundTrip(connectRequest(t, context.Background(), pr)) + results <- result{rsp, err} + }() + + f := p.readFrame() + require.IsType(t, &http2.MetaHeadersFrame{}, f) + headers := f.(*http2.MetaHeadersFrame) + require.False(t, headers.StreamEnded()) + var fields []string + for _, hf := range headers.Fields { + fields = append(fields, hf.Name+": "+hf.Value) + } + require.Equal(t, []string{ + ":method: CONNECT", + ":authority: proxy.example", + ":scheme: https", + ":path: /.well-known/masque/ip/*/*/", + ":protocol: connect-ip", + "authorization: Basic dTpw", + "capsule-protocol: ?1", + }, fields) + + p.writeHeaders(false, ":status", "200", "capsule-protocol", "?1") + r := <-results + require.NoError(t, r.err) + require.Equal(t, http.StatusOK, r.rsp.StatusCode) + require.Equal(t, "?1", r.rsp.Header.Get("Capsule-Protocol")) + + go pw.Write([]byte("ping")) + f = p.readFrame() + require.IsType(t, &http2.DataFrame{}, f) + require.Equal(t, "ping", string(f.(*http2.DataFrame).Data())) + + require.NoError(t, p.fr.WriteData(http2StreamID, false, []byte("pong"))) + b := make([]byte, 16) + n, err := r.rsp.Body.Read(b) + require.NoError(t, err) + require.Equal(t, "pong", string(b[:n])) + + require.NoError(t, pw.Close()) + f = p.readFrame() + require.IsType(t, &http2.DataFrame{}, f) + require.True(t, f.(*http2.DataFrame).StreamEnded()) + + require.NoError(t, p.fr.WriteData(http2StreamID, true, nil)) + _, err = r.rsp.Body.Read(b) + require.ErrorIs(t, err, io.EOF) +} + +func TestHTTP2ClientDefaultUserAgent(t *testing.T) { + cc, p := newHTTP2Peer(t, http2.Setting{ID: http2.SettingEnableConnectProtocol, Val: 1}) + req := connectRequest(t, context.Background(), nil) + delete(req.Header, "User-Agent") + go cc.RoundTrip(req) + f := p.readFrame() + require.IsType(t, &http2.MetaHeadersFrame{}, f) + var userAgents []string + for _, hf := range f.(*http2.MetaHeadersFrame).Fields { + if hf.Name == "user-agent" { + userAgents = append(userAgents, hf.Value) + } + } + require.Equal(t, []string{http2DefaultUserAgent}, userAgents) +} + +func TestHTTP2ClientNeedsExtendedConnect(t *testing.T) { + cc, _ := newHTTP2Peer(t) + _, err := cc.RoundTrip(connectRequest(t, context.Background(), io.NopCloser(strings.NewReader("")))) + require.ErrorIs(t, err, errHTTP2NoExtendedConnect) +} + +func TestHTTP2ClientSingleStream(t *testing.T) { + cc, p := newHTTP2Peer(t, http2.Setting{ID: http2.SettingEnableConnectProtocol, Val: 1}) + go cc.RoundTrip(connectRequest(t, context.Background(), nil)) + p.readFrame() + _, err := cc.RoundTrip(connectRequest(t, context.Background(), nil)) + require.ErrorIs(t, err, errHTTP2StreamUsed) +} + +func TestHTTP2ClientFlowControl(t *testing.T) { + cc, p := newHTTP2Peer(t, + http2.Setting{ID: http2.SettingEnableConnectProtocol, Val: 1}, + http2.Setting{ID: http2.SettingInitialWindowSize, Val: 10}, + ) + pr, pw := io.Pipe() + go cc.RoundTrip(connectRequest(t, context.Background(), pr)) + require.IsType(t, &http2.MetaHeadersFrame{}, p.readFrame()) + + go pw.Write([]byte("0123456789abcdef")) + f := p.readFrame() + require.Equal(t, "0123456789", string(f.(*http2.DataFrame).Data())) + + require.NoError(t, p.fr.WriteWindowUpdate(http2StreamID, 4)) + f = p.readFrame() + require.Equal(t, "abcd", string(f.(*http2.DataFrame).Data())) + + require.NoError(t, p.fr.WriteSettings(http2.Setting{ID: http2.SettingInitialWindowSize, Val: 12})) + var acked bool + var data string + for range 2 { + switch f := p.readFrame().(type) { + case *http2.SettingsFrame: + acked = f.IsAck() + case *http2.DataFrame: + data = string(f.Data()) + } + } + require.True(t, acked) + require.Equal(t, "ef", data) +} + +func TestHTTP2ClientReceiveWindow(t *testing.T) { + cc, p := newHTTP2Peer(t, http2.Setting{ID: http2.SettingEnableConnectProtocol, Val: 1}) + rsps := make(chan *http.Response, 1) + go func() { + rsp, err := cc.RoundTrip(connectRequest(t, context.Background(), nil)) + if err == nil { + rsps <- rsp + } + }() + require.IsType(t, &http2.MetaHeadersFrame{}, p.readFrame()) + require.True(t, p.readFrame().(*http2.DataFrame).StreamEnded()) + p.writeHeaders(false, ":status", "200") + rsp := <-rsps + + chunk := bytes.Repeat([]byte("x"), http2DefaultFrameSize) + sent := 0 + go func() { + for sent+len(chunk) <= http2WindowUpdateSize { + if p.fr.WriteData(http2StreamID, false, chunk) != nil { + return + } + sent += len(chunk) + } + }() + _, err := io.CopyN(io.Discard, rsp.Body, http2WindowUpdateSize) + require.NoError(t, err) + for _, id := range []uint32{0, http2StreamID} { + f := p.readFrame() + require.IsType(t, &http2.WindowUpdateFrame{}, f) + require.Equal(t, id, f.Header().StreamID) + require.Equal(t, uint32(http2WindowUpdateSize), f.(*http2.WindowUpdateFrame).Increment) + } +} + +func TestHTTP2ClientRejectsOverflow(t *testing.T) { + cc, p := newHTTP2Peer(t, http2.Setting{ID: http2.SettingEnableConnectProtocol, Val: 1}) + rsps := make(chan *http.Response, 1) + go func() { + rsp, err := cc.RoundTrip(connectRequest(t, context.Background(), nil)) + if err == nil { + rsps <- rsp + } + }() + p.readFrame() + p.readFrame() + p.writeHeaders(false, ":status", "200") + rsp := <-rsps + + chunk := make([]byte, http2DefaultFrameSize) + go func() { + for range http2StreamWindow/len(chunk) + 1 { + if p.fr.WriteData(http2StreamID, false, chunk) != nil { + return + } + } + }() + select { + case <-cc.done: + case <-time.After(5 * time.Second): + t.Fatal("the connection outlived a flow control violation") + } + require.ErrorIs(t, cc.connErr(), http2.ConnectionError(http2.ErrCodeFlowControl)) + _, err := io.Copy(io.Discard, rsp.Body) + require.ErrorIs(t, err, http2.ConnectionError(http2.ErrCodeFlowControl)) +} + +func TestHTTP2ClientRejectsOversizedFrames(t *testing.T) { + cc, p := newHTTP2Peer(t) + require.NoError(t, p.fr.WritePing(false, [8]byte{})) + require.True(t, p.readFrame().(*http2.PingFrame).IsAck()) + + p.fr.AllowIllegalWrites = true + require.NoError(t, p.fr.WriteData(http2StreamID, false, make([]byte, http2DefaultFrameSize+1))) + select { + case <-cc.done: + case <-time.After(5 * time.Second): + t.Fatal("the connection accepted a frame larger than it allows") + } + require.ErrorIs(t, cc.connErr(), http2.ErrFrameTooLarge) +} + +func TestHTTP2ClientStatus(t *testing.T) { + cc, p := newHTTP2Peer(t, http2.Setting{ID: http2.SettingEnableConnectProtocol, Val: 1}) + rsps := make(chan *http.Response, 1) + go func() { + rsp, err := cc.RoundTrip(connectRequest(t, context.Background(), nil)) + if err == nil { + rsps <- rsp + } + }() + p.readFrame() + p.readFrame() + p.writeHeaders(false, ":status", "100") + p.writeHeaders(true, ":status", "407", "proxy-authenticate", "Basic") + rsp := <-rsps + require.Equal(t, http.StatusProxyAuthRequired, rsp.StatusCode) + require.Equal(t, "Basic", rsp.Header.Get("Proxy-Authenticate")) + _, err := rsp.Body.Read(make([]byte, 1)) + require.ErrorIs(t, err, io.EOF) +} + +func TestHTTP2ClientReset(t *testing.T) { + t.Run("by the server", func(t *testing.T) { + cc, p := newHTTP2Peer(t, http2.Setting{ID: http2.SettingEnableConnectProtocol, Val: 1}) + errs := make(chan error, 1) + go func() { + _, err := cc.RoundTrip(connectRequest(t, context.Background(), nil)) + errs <- err + }() + p.readFrame() + p.readFrame() + require.NoError(t, p.fr.WriteRSTStream(http2StreamID, http2.ErrCodeRefusedStream)) + require.Equal(t, http2.StreamError{StreamID: http2StreamID, Code: http2.ErrCodeRefusedStream}, <-errs) + }) + + t.Run("by the context", func(t *testing.T) { + cc, p := newHTTP2Peer(t, http2.Setting{ID: http2.SettingEnableConnectProtocol, Val: 1}) + ctx, cancel := context.WithCancel(context.Background()) + pr, pw := io.Pipe() + defer pw.Close() + rsps := make(chan *http.Response, 1) + go func() { + rsp, err := cc.RoundTrip(connectRequest(t, ctx, pr)) + if err == nil { + rsps <- rsp + } + }() + p.readFrame() + p.writeHeaders(false, ":status", "200") + rsp := <-rsps + cancel() + f := p.readFrame() + require.IsType(t, &http2.RSTStreamFrame{}, f) + require.Equal(t, http2.ErrCodeCancel, f.(*http2.RSTStreamFrame).ErrCode) + _, err := rsp.Body.Read(make([]byte, 1)) + require.ErrorIs(t, err, context.Canceled) + _, err = pw.Write([]byte("x")) + require.ErrorIs(t, err, io.ErrClosedPipe) + }) + + t.Run("by GOAWAY", func(t *testing.T) { + cc, p := newHTTP2Peer(t, http2.Setting{ID: http2.SettingEnableConnectProtocol, Val: 1}) + errs := make(chan error, 1) + go func() { + _, err := cc.RoundTrip(connectRequest(t, context.Background(), nil)) + errs <- err + }() + p.readFrame() + p.readFrame() + require.NoError(t, p.fr.WriteGoAway(0, http2.ErrCodeNo, nil)) + require.ErrorContains(t, <-errs, "GOAWAY") + }) +} + +func TestHTTP2ClientAnswersPings(t *testing.T) { + _, p := newHTTP2Peer(t) + data := [8]byte{1, 2, 3, 4, 5, 6, 7, 8} + require.NoError(t, p.fr.WritePing(false, data)) + f := p.readFrame() + require.IsType(t, &http2.PingFrame{}, f) + require.True(t, f.(*http2.PingFrame).IsAck()) + require.Equal(t, data, f.(*http2.PingFrame).Data) +}