mirror of
https://github.com/XTLS/Xray-core.git
synced 2026-09-26 19:20:06 +00:00
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
This commit is contained in:
@@ -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),
|
||||
},
|
||||
},
|
||||
}),
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
Reference in New Issue
Block a user