Cluvex
2026-09-26 03:49:37 +00:00
committed by RPRX
parent 61cad5ec8b
commit df261e4479
12 changed files with 2144 additions and 38 deletions
+199 -6
View File
@@ -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),
},
},
}),
+18 -4
View File
@@ -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 {
+90 -13
View File
@@ -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
@@ -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)
}
@@ -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")
}
+17 -4
View File
@@ -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
}
+13 -3
View File
@@ -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
}
@@ -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
+65 -8
View File
@@ -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()
+19
View File
@@ -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
+600
View File
@@ -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
}
+405
View File
@@ -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)
}