Cluvex
2026-10-08 04:53:07 +00:00
committed by GitHub
parent 7da5dae650
commit 9e55a6ed18
14 changed files with 1126 additions and 50 deletions
+122
View File
@@ -1,7 +1,14 @@
package conf_test
import (
"crypto/ecdsa"
"crypto/ed25519"
"crypto/elliptic"
"crypto/rand"
"crypto/x509"
"encoding/base64"
"encoding/json"
"encoding/pem"
"testing"
"github.com/xtls/xray-core/common/protocol"
@@ -67,6 +74,121 @@ func TestMasqueConfig(t *testing.T) {
}
}
func TestMasqueWarpConfig(t *testing.T) {
creator := func() Buildable {
return new(MasqueConfig)
}
key, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader)
if err != nil {
t.Fatal(err)
}
pkcs8, err := x509.MarshalPKCS8PrivateKey(key)
if err != nil {
t.Fatal(err)
}
sec1, err := x509.MarshalECPrivateKey(key)
if err != nil {
t.Fatal(err)
}
quote := func(s string) string {
b, _ := json.Marshal(s)
return string(b)
}
server, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader)
if err != nil {
t.Fatal(err)
}
publicKey, err := x509.MarshalPKIXPublicKey(&server.PublicKey)
if err != nil {
t.Fatal(err)
}
publicPEM := string(pem.EncodeToMemory(&pem.Block{Type: "PUBLIC KEY", Bytes: publicKey}))
warpInput := func(key string, extra string) string {
return `{` + extra + `"warp": {"privateKey": ` + quote(key) + `, "publicKey": ` + quote(publicPEM) + `, "address": ["172.16.0.2", "2606:4700:110:8a36::2/128"]}}`
}
address := []string{"172.16.0.2/32", "2606:4700:110:8a36::2/128"}
for _, input := range []string{
string(pem.EncodeToMemory(&pem.Block{Type: "PRIVATE KEY", Bytes: pkcs8})),
string(pem.EncodeToMemory(&pem.Block{Type: "EC PRIVATE KEY", Bytes: sec1})),
base64.StdEncoding.EncodeToString(pkcs8),
base64.StdEncoding.EncodeToString(sec1),
} {
runMultiTestCase(t, []TestCase{
{
Input: warpInput(input, ""),
Parser: loadJSON(creator),
Output: &masque.Config{
Host: "cloudflareaccess.com",
Path: "/",
Warp: &masque.Warp{PrivateKey: pkcs8, PublicKey: publicKey, Address: address},
},
},
})
}
runMultiTestCase(t, []TestCase{
{
Input: warpInput(base64.StdEncoding.EncodeToString(sec1), `"host": "example.com", "path": "/warp", `),
Parser: loadJSON(creator),
Output: &masque.Config{
Host: "example.com",
Path: "/warp",
Warp: &masque.Warp{PrivateKey: pkcs8, PublicKey: publicKey, Address: address},
},
},
})
p384, err := ecdsa.GenerateKey(elliptic.P384(), rand.Reader)
if err != nil {
t.Fatal(err)
}
p384DER, err := x509.MarshalPKCS8PrivateKey(p384)
if err != nil {
t.Fatal(err)
}
ed, err := x509.MarshalPKCS8PrivateKey(ed25519.NewKeyFromSeed(make([]byte, ed25519.SeedSize)))
if err != nil {
t.Fatal(err)
}
withAddress := func(address string) string {
return `{"warp": {"privateKey": ` + quote(base64.StdEncoding.EncodeToString(pkcs8)) + `, "publicKey": ` + quote(publicPEM) + `, "address": ` + address + `}}`
}
withPublicKey := func(key string) string {
return `{"warp": {"privateKey": ` + quote(base64.StdEncoding.EncodeToString(pkcs8)) + `, "publicKey": ` + quote(key) + `, "address": ["172.16.0.2"]}}`
}
runMultiTestCase(t, []TestCase{
{
Input: withPublicKey(base64.StdEncoding.EncodeToString(publicKey)),
Parser: loadJSON(creator),
Output: &masque.Config{
Host: "cloudflareaccess.com",
Path: "/",
Warp: &masque.Warp{PrivateKey: pkcs8, PublicKey: publicKey, Address: []string{"172.16.0.2/32"}},
},
},
})
for _, input := range []string{
`{"warp": {}}`,
withAddress(`[]`),
withPublicKey(""),
withPublicKey("not a key"),
withPublicKey(base64.StdEncoding.EncodeToString([]byte("not a key"))),
withPublicKey(base64.StdEncoding.EncodeToString(pkcs8)),
withAddress(`["172.16.0"]`),
withAddress(`["172.16.0.2", "172.16.0.3"]`),
withAddress(`["2606:4700::1", "2606:4700::2/128"]`),
warpInput("not a key", ""),
warpInput(base64.StdEncoding.EncodeToString([]byte("not a key")), ""),
warpInput(base64.StdEncoding.EncodeToString(p384DER), ""),
warpInput(base64.StdEncoding.EncodeToString(ed), ""),
warpInput(base64.StdEncoding.EncodeToString(pkcs8), `"user": "u", "pass": "p", `),
} {
if _, err := loadJSON(creator)(input); err == nil {
t.Errorf("expected an error for %s", input)
}
}
}
func TestMasqueOutboundConfig(t *testing.T) {
build := func(s string) error {
c := new(OutboundDetourConfig)
+113 -4
View File
@@ -1,10 +1,15 @@
package conf
import (
"crypto/ecdsa"
"crypto/elliptic"
"crypto/x509"
"encoding/base64"
"encoding/json"
"encoding/pem"
"maps"
"math/big"
"net/netip"
"net/url"
"sort"
"strconv"
@@ -790,16 +795,49 @@ func (c *HysteriaConfig) Build() (proto.Message, error) {
return config, nil
}
type MasqueWarpConfig struct {
PrivateKey string `json:"privateKey"`
PublicKey string `json:"publicKey"`
Address []string `json:"address"`
}
type MasqueConfig struct {
Host string `json:"host"`
Path string `json:"path"`
User string `json:"user"`
Pass string `json:"pass"`
Headers map[string]string `json:"headers"`
Warp *MasqueWarpConfig `json:"warp"`
}
func (c *MasqueConfig) Build() (proto.Message, error) {
var warp *masque.Warp
host := c.Host
path := c.Path
if c.Warp != nil {
if c.User != "" || c.Pass != "" {
return nil, errors.New(`"user" and "pass" can't be used with "warp"`)
}
key, err := parseWarpPrivateKey(c.Warp.PrivateKey)
if err != nil {
return nil, errors.New(`invalid "privateKey" in "warp"`).Base(err)
}
publicKey, err := parseWarpPublicKey(c.Warp.PublicKey)
if err != nil {
return nil, errors.New(`invalid "publicKey" in "warp"`).Base(err)
}
address, err := parseWarpAddress(c.Warp.Address)
if err != nil {
return nil, err
}
warp = &masque.Warp{PrivateKey: key, PublicKey: publicKey, Address: address}
if host == "" {
host = masque.WarpHost
}
if path == "" {
path = masque.WarpPath
}
}
if path == "" {
path = masque.DefaultPath
}
@@ -811,9 +849,9 @@ func (c *MasqueConfig) Build() (proto.Message, error) {
if !strings.HasPrefix(path, "/") || strings.ContainsAny(path, "{}") {
return nil, errors.New(`invalid "path": `, path, `, only the variables {target} and {ipproto} are supported`)
}
if c.Host != "" {
if u, err := url.Parse("https://" + c.Host); err != nil || u.Host != c.Host {
return nil, errors.New(`invalid "host": `, c.Host)
if host != "" {
if u, err := url.Parse("https://" + host); err != nil || u.Host != host {
return nil, errors.New(`invalid "host": `, host)
}
}
for k, v := range c.Headers {
@@ -841,12 +879,83 @@ func (c *MasqueConfig) Build() (proto.Message, error) {
headers["Authorization"] = "Basic " + base64.StdEncoding.EncodeToString([]byte(c.User+":"+c.Pass))
}
return &masque.Config{
Host: c.Host,
Host: host,
Path: path,
Headers: headers,
Warp: warp,
}, nil
}
func parseWarpAddress(list []string) ([]string, error) {
if len(list) == 0 {
return nil, errors.New(`"address" in "warp" is not set`)
}
var v4, v6 bool
address := make([]string, 0, len(list))
for _, s := range list {
prefix, err := netip.ParsePrefix(s)
if err != nil {
addr, err := netip.ParseAddr(s)
if err != nil {
return nil, errors.New(`invalid "address" in "warp": `, s)
}
prefix = netip.PrefixFrom(addr, addr.BitLen())
}
if prefix.Addr().Is4() && v4 || prefix.Addr().Is6() && v6 {
return nil, errors.New(`"address" in "warp" takes at most one IPv4 and one IPv6 address`)
}
v4 = v4 || prefix.Addr().Is4()
v6 = v6 || prefix.Addr().Is6()
address = append(address, prefix.String())
}
return address, nil
}
func decodeWarpKey(s string) ([]byte, error) {
s = strings.TrimSpace(s)
if s == "" {
return nil, errors.New("empty key")
}
if block, _ := pem.Decode([]byte(s)); block != nil {
return block.Bytes, nil
}
der, err := base64.StdEncoding.DecodeString(s)
if err != nil {
return nil, errors.New("neither PEM nor base64").Base(err)
}
return der, nil
}
func parseWarpPublicKey(s string) ([]byte, error) {
der, err := decodeWarpKey(s)
if err != nil {
return nil, err
}
if _, err := x509.ParsePKIXPublicKey(der); err != nil {
return nil, errors.New("not a PKIX public key").Base(err)
}
return der, nil
}
func parseWarpPrivateKey(s string) ([]byte, error) {
der, err := decodeWarpKey(s)
if err != nil {
return nil, err
}
var key any
key, err = x509.ParsePKCS8PrivateKey(der)
if err != nil {
if key, err = x509.ParseECPrivateKey(der); err != nil {
return nil, errors.New("neither a PKCS #8 nor a SEC 1 private key")
}
}
ecKey, ok := key.(*ecdsa.PrivateKey)
if !ok || ecKey.Curve != elliptic.P256() {
return nil, errors.New("not an ECDSA P-256 key")
}
return x509.MarshalPKCS8PrivateKey(ecKey)
}
func readFileOrString(f string, s []string) ([]byte, error) {
if len(f) > 0 {
return filesystem.ReadCert(f)
+89 -12
View File
@@ -26,6 +26,7 @@ type Config struct {
Host string `protobuf:"bytes,1,opt,name=host,proto3" json:"host,omitempty"`
Path string `protobuf:"bytes,2,opt,name=path,proto3" json:"path,omitempty"`
Headers map[string]string `protobuf:"bytes,3,rep,name=headers,proto3" json:"headers,omitempty" protobuf_key:"bytes,1,opt,name=key" protobuf_val:"bytes,2,opt,name=value"`
Warp *Warp `protobuf:"bytes,4,opt,name=warp,proto3" json:"warp,omitempty"`
unknownFields protoimpl.UnknownFields
sizeCache protoimpl.SizeCache
}
@@ -81,18 +82,92 @@ func (x *Config) GetHeaders() map[string]string {
return nil
}
func (x *Config) GetWarp() *Warp {
if x != nil {
return x.Warp
}
return nil
}
type Warp struct {
state protoimpl.MessageState `protogen:"open.v1"`
PrivateKey []byte `protobuf:"bytes,1,opt,name=private_key,json=privateKey,proto3" json:"private_key,omitempty"`
PublicKey []byte `protobuf:"bytes,2,opt,name=public_key,json=publicKey,proto3" json:"public_key,omitempty"`
Address []string `protobuf:"bytes,3,rep,name=address,proto3" json:"address,omitempty"`
unknownFields protoimpl.UnknownFields
sizeCache protoimpl.SizeCache
}
func (x *Warp) Reset() {
*x = Warp{}
mi := &file_transport_internet_masque_config_proto_msgTypes[1]
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
ms.StoreMessageInfo(mi)
}
func (x *Warp) String() string {
return protoimpl.X.MessageStringOf(x)
}
func (*Warp) ProtoMessage() {}
func (x *Warp) ProtoReflect() protoreflect.Message {
mi := &file_transport_internet_masque_config_proto_msgTypes[1]
if x != nil {
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
if ms.LoadMessageInfo() == nil {
ms.StoreMessageInfo(mi)
}
return ms
}
return mi.MessageOf(x)
}
// Deprecated: Use Warp.ProtoReflect.Descriptor instead.
func (*Warp) Descriptor() ([]byte, []int) {
return file_transport_internet_masque_config_proto_rawDescGZIP(), []int{1}
}
func (x *Warp) GetPrivateKey() []byte {
if x != nil {
return x.PrivateKey
}
return nil
}
func (x *Warp) GetPublicKey() []byte {
if x != nil {
return x.PublicKey
}
return nil
}
func (x *Warp) GetAddress() []string {
if x != nil {
return x.Address
}
return nil
}
var File_transport_internet_masque_config_proto protoreflect.FileDescriptor
const file_transport_internet_masque_config_proto_rawDesc = "" +
"\n" +
"&transport/internet/masque/config.proto\x12\x1exray.transport.internet.masque\"\xbb\x01\n" +
"&transport/internet/masque/config.proto\x12\x1exray.transport.internet.masque\"\xf5\x01\n" +
"\x06Config\x12\x12\n" +
"\x04host\x18\x01 \x01(\tR\x04host\x12\x12\n" +
"\x04path\x18\x02 \x01(\tR\x04path\x12M\n" +
"\aheaders\x18\x03 \x03(\v23.xray.transport.internet.masque.Config.HeadersEntryR\aheaders\x1a:\n" +
"\aheaders\x18\x03 \x03(\v23.xray.transport.internet.masque.Config.HeadersEntryR\aheaders\x128\n" +
"\x04warp\x18\x04 \x01(\v2$.xray.transport.internet.masque.WarpR\x04warp\x1a:\n" +
"\fHeadersEntry\x12\x10\n" +
"\x03key\x18\x01 \x01(\tR\x03key\x12\x14\n" +
"\x05value\x18\x02 \x01(\tR\x05value:\x028\x01B|\n" +
"\x05value\x18\x02 \x01(\tR\x05value:\x028\x01\"`\n" +
"\x04Warp\x12\x1f\n" +
"\vprivate_key\x18\x01 \x01(\fR\n" +
"privateKey\x12\x1d\n" +
"\n" +
"public_key\x18\x02 \x01(\fR\tpublicKey\x12\x18\n" +
"\aaddress\x18\x03 \x03(\tR\aaddressB|\n" +
"\"com.xray.transport.internet.masqueP\x01Z3github.com/xtls/xray-core/transport/internet/masque\xaa\x02\x1eXray.Transport.Internet.Masqueb\x06proto3"
var (
@@ -107,18 +182,20 @@ func file_transport_internet_masque_config_proto_rawDescGZIP() []byte {
return file_transport_internet_masque_config_proto_rawDescData
}
var file_transport_internet_masque_config_proto_msgTypes = make([]protoimpl.MessageInfo, 2)
var file_transport_internet_masque_config_proto_msgTypes = make([]protoimpl.MessageInfo, 3)
var file_transport_internet_masque_config_proto_goTypes = []any{
(*Config)(nil), // 0: xray.transport.internet.masque.Config
nil, // 1: xray.transport.internet.masque.Config.HeadersEntry
(*Warp)(nil), // 1: xray.transport.internet.masque.Warp
nil, // 2: xray.transport.internet.masque.Config.HeadersEntry
}
var file_transport_internet_masque_config_proto_depIdxs = []int32{
1, // 0: xray.transport.internet.masque.Config.headers:type_name -> xray.transport.internet.masque.Config.HeadersEntry
1, // [1:1] is the sub-list for method output_type
1, // [1:1] is the sub-list for method input_type
1, // [1:1] is the sub-list for extension type_name
1, // [1:1] is the sub-list for extension extendee
0, // [0:1] is the sub-list for field type_name
2, // 0: xray.transport.internet.masque.Config.headers:type_name -> xray.transport.internet.masque.Config.HeadersEntry
1, // 1: xray.transport.internet.masque.Config.warp:type_name -> xray.transport.internet.masque.Warp
2, // [2:2] is the sub-list for method output_type
2, // [2:2] is the sub-list for method input_type
2, // [2:2] is the sub-list for extension type_name
2, // [2:2] is the sub-list for extension extendee
0, // [0:2] is the sub-list for field type_name
}
func init() { file_transport_internet_masque_config_proto_init() }
@@ -132,7 +209,7 @@ func file_transport_internet_masque_config_proto_init() {
GoPackagePath: reflect.TypeOf(x{}).PkgPath(),
RawDescriptor: unsafe.Slice(unsafe.StringData(file_transport_internet_masque_config_proto_rawDesc), len(file_transport_internet_masque_config_proto_rawDesc)),
NumEnums: 0,
NumMessages: 2,
NumMessages: 3,
NumExtensions: 0,
NumServices: 0,
},
+7
View File
@@ -10,4 +10,11 @@ message Config {
string host = 1;
string path = 2;
map<string, string> headers = 3;
Warp warp = 4;
}
message Warp {
bytes private_key = 1;
bytes public_key = 2;
repeated string address = 3;
}
+42 -2
View File
@@ -14,16 +14,28 @@ import (
"github.com/apernet/quic-go"
"github.com/apernet/quic-go/http3"
"github.com/apernet/quic-go/quicvarint"
)
const (
cloudflareProtocol = "cf-connect-ip"
SettingDatagramDraft00 uint64 = 0x276
)
type ClientConn struct {
clientConn *http3.ClientConn
quicConn *quic.Conn
}
func NewClientConn(conn *http3.ClientConn) *ClientConn {
return &ClientConn{clientConn: conn}
}
func NewCloudflareClientConn(conn *http3.ClientConn, quicConn *quic.Conn) *ClientConn {
return &ClientConn{clientConn: conn, quicConn: quicConn}
}
func (c *ClientConn) Dial(req *Request) (*Conn, *http.Response, error) {
httpReq := req.httpRequest()
if httpReq.URL == nil {
@@ -32,6 +44,11 @@ func (c *ClientConn) Dial(req *Request) (*Conn, *http.Response, error) {
if httpReq.Host == "" && httpReq.URL.Host == "" {
return nil, nil, errors.New("connect-ip: request needs a host")
}
cloudflare := c.quicConn != nil
if cloudflare {
httpReq = httpReq.Clone(httpReq.Context())
httpReq.Proto = cloudflareProtocol
}
select {
case <-httpReq.Context().Done():
@@ -42,10 +59,11 @@ func (c *ClientConn) Dial(req *Request) (*Conn, *http.Response, error) {
}
settings := c.clientConn.Settings()
if !settings.EnableExtendedConnect {
if !settings.EnableExtendedConnect && !cloudflare {
return nil, nil, errors.New("connect-ip: server didn't enable Extended CONNECT")
}
if !settings.EnableDatagrams {
draftDatagrams := cloudflare && !settings.EnableDatagrams && settings.Other[SettingDatagramDraft00] == 1
if !settings.EnableDatagrams && !draftDatagrams {
return nil, nil, errors.New("connect-ip: server didn't enable datagrams")
}
@@ -71,5 +89,27 @@ func (c *ClientConn) Dial(req *Request) (*Conn, *http.Response, error) {
return nil, rsp, fmt.Errorf("connect-ip: server responded with %d", rsp.StatusCode)
}
keepStream = true
if draftDatagrams {
return newProxiedConn(&draftDatagramStream{RequestStream: rstr, conn: c.quicConn}), rsp, nil
}
return newProxiedConn(rstr), rsp, nil
}
type draftDatagramStream struct {
*http3.RequestStream
conn *quic.Conn
}
func (s *draftDatagramStream) ReceiveDatagram(ctx context.Context) ([]byte, error) {
for {
b, err := s.conn.ReceiveDatagram(ctx)
if err != nil {
return nil, err
}
quarterStreamID, n, err := quicvarint.Parse(b)
if err != nil || quic.StreamID(quarterStreamID*4) != s.StreamID() {
continue
}
return b[n:], nil
}
}
+34 -10
View File
@@ -96,10 +96,21 @@ type Conn struct {
closeResult error
datagramCapsuleOnce sync.Once
bare bool
}
func newProxiedConn(str requestStream) *Conn {
return startProxiedConn(str, false)
}
func newBareProxiedConn(str requestStream) *Conn {
return startProxiedConn(str, true)
}
func startProxiedConn(str requestStream, bare bool) *Conn {
c := &Conn{
bare: bare,
str: str,
writeNotify: make(chan struct{}, 1),
writeDone: make(chan error, 1),
@@ -243,6 +254,12 @@ func (c *Conn) ReceiveAddressAssignment(ctx context.Context) ([]AssignedAddress,
}
}
func (c *Conn) SetAssignedAddresses(prefixes []netip.Prefix) {
c.mu.Lock()
c.assignedAddresses = slices.Clone(prefixes)
c.mu.Unlock()
}
func (c *Conn) ReceiveAddressRequest(ctx context.Context) (*AddressRequest, error) {
var requested *addressRequestCapsule
select {
@@ -491,15 +508,18 @@ func (c *Conn) ReadPacket(b []byte) (int, error) {
return 0, err
}
}
contextID, n, err := quicvarint.Parse(data)
if err != nil {
errors.LogDebugInner(context.Background(), err, "dropping malformed datagram")
continue
packet := data
if !c.bare || len(data) == 0 || data[0] == 0 {
contextID, n, err := quicvarint.Parse(data)
if err != nil {
errors.LogDebugInner(context.Background(), err, "dropping malformed datagram")
continue
}
if contextID != 0 {
continue
}
packet = data[n:]
}
if contextID != 0 {
continue
}
packet := data[n:]
if err := c.handleIncomingProxiedPacket(packet); err != nil {
errors.LogDebugInner(context.Background(), err, "dropping proxied packet")
continue
@@ -644,7 +664,11 @@ func (c *Conn) composeDatagram(b []byte) ([]byte, error) {
}
b[7]--
}
size := len(contextIDZero) + len(b)
contextID := contextIDZero
if c.bare {
contextID = nil
}
size := len(contextID) + len(b)
var data []byte
if c.h3 == nil {
data = make([]byte, 0, quicvarint.Len(uint64(capsuleTypeDatagram))+quicvarint.Len(uint64(size))+size)
@@ -653,7 +677,7 @@ func (c *Conn) composeDatagram(b []byte) ([]byte, error) {
} else {
data = make([]byte, 0, size)
}
data = append(data, contextIDZero...)
data = append(data, contextID...)
data = append(data, b...)
return data, nil
}
+23 -3
View File
@@ -9,22 +9,29 @@ import (
"net"
"net/http"
"os"
"strings"
"sync"
"time"
"github.com/apernet/quic-go"
"github.com/apernet/quic-go/http3"
)
const maxStreamBuffer = 32 << 10
type HTTP2ClientConn struct {
roundTripper http.RoundTripper
cloudflare bool
}
func NewHTTP2ClientConn(rt http.RoundTripper) *HTTP2ClientConn {
return &HTTP2ClientConn{roundTripper: rt}
}
func NewCloudflareHTTP2ClientConn(rt http.RoundTripper) *HTTP2ClientConn {
return &HTTP2ClientConn{roundTripper: rt, cloudflare: true}
}
func (c *HTTP2ClientConn) Dial(req *Request) (*Conn, *http.Response, error) {
httpReq := req.httpRequest()
if httpReq.URL == nil {
@@ -39,7 +46,16 @@ func (c *HTTP2ClientConn) Dial(req *Request) (*Conn, *http.Response, error) {
stop := context.AfterFunc(ctx, cancel)
body := NewStreamBuffer()
r := httpReq.Clone(streamCtx)
r.Header[":protocol"] = []string{requestProtocol}
if c.cloudflare {
r.Header.Del(http3.CapsuleProtocolHeader)
r.Header.Set("Cf-Connect-Proto", cloudflareProtocol)
r.Header.Set("Pq-Enabled", "false")
if _, _, err := net.SplitHostPort(r.Host); err != nil {
r.Host = net.JoinHostPort(strings.Trim(r.Host, "[]"), "443")
}
} else {
r.Header[":protocol"] = []string{requestProtocol}
}
r.Body = body
rsp, err := c.roundTripper.RoundTrip(r)
if !stop() {
@@ -57,12 +73,16 @@ func (c *HTTP2ClientConn) Dial(req *Request) (*Conn, *http.Response, error) {
rsp.Body.Close()
return nil, rsp, fmt.Errorf("connect-ip: server responded with %d", rsp.StatusCode)
}
return newProxiedConn(&http2Stream{
str := &http2Stream{
reader: bufio.NewReader(rsp.Body),
body: body,
rsp: rsp.Body,
cancel: cancel,
}), rsp, nil
}
if c.cloudflare {
return newBareProxiedConn(str), rsp, nil
}
return newProxiedConn(str), rsp, nil
}
type http2Stream struct {
@@ -140,6 +140,87 @@ func TestHTTP2Request(t *testing.T) {
require.Equal(t, maxCapsulePacketSize, conn.MaxPacketSize())
}
func TestCloudflareHTTP2Request(t *testing.T) {
for _, c := range []struct{ host, want string }{
{"cloudflareaccess.com", "cloudflareaccess.com:443"},
{"cloudflareaccess.com:8443", "cloudflareaccess.com:8443"},
{"[2001:db8::1]", "[2001:db8::1]:443"},
} {
requests := make(chan *http.Request, 1)
pr, pw := io.Pipe()
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://"+c.host+"/")
require.NoError(t, err)
conn, _, err := NewCloudflareHTTP2ClientConn(rt).Dial(req)
require.NoError(t, err)
r := <-requests
require.Equal(t, http.MethodConnect, r.Method)
require.Equal(t, c.want, r.Host)
require.Empty(t, r.Header.Values(":protocol"))
require.Empty(t, r.Header.Values("Capsule-Protocol"))
require.Equal(t, cloudflareProtocol, r.Header.Get("Cf-Connect-Proto"))
require.Equal(t, "false", r.Header.Get("Pq-Enabled"))
require.Equal(t, "?1", req.Header().Get("Capsule-Protocol"))
conn.Close()
pw.Close()
}
}
func TestHTTP2BareDatagramCapsules(t *testing.T) {
str, pw := newTestHTTP2Stream()
defer pw.Close()
conn := newBareProxiedConn(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(packet),
capsule(append([]byte{0x02}, packet...)),
capsule(append(bytes.Clone(contextIDZero), packet...)),
} {
if _, err := pw.Write(c); err != nil {
return
}
}
}()
for range 2 {
b := make([]byte, 1500)
n, err := conn.ReadPacket(b)
require.NoError(t, err)
require.Equal(t, packet, b[:n])
}
_, err := conn.WritePacket(slices.Clone(packet))
require.NoError(t, err)
p := http3.NewCapsuleParser(str.body)
var sent []byte
for sent == nil {
typ, cr, err := p.Next()
require.NoError(t, err)
data, err := io.ReadAll(cr)
require.NoError(t, err)
if typ == capsuleTypeDatagram {
sent = data
}
}
require.Len(t, sent, len(packet))
require.Equal(t, packet[8]-1, sent[8])
require.Equal(t, packet[ipv4.HeaderLen:], sent[ipv4.HeaderLen:])
}
func TestHTTP2DialErrors(t *testing.T) {
newReq := func(ctx context.Context) *Request {
req, err := NewRequest(ctx, "https://example.org/connect-ip")
+55 -16
View File
@@ -45,6 +45,9 @@ func Dial(ctx context.Context, dest net.Destination, streamSettings *internet.Me
gotlsConfig := tlsConfig.GetTLSConfig(tls.WithDestination(dest))
gotlsConfig.NextProtos = []string{http3.NextProtoH3}
if err := useWarp(config, gotlsConfig); err != nil {
return nil, err
}
quicParams := streamSettings.QuicParams
if quicParams == nil {
@@ -99,6 +102,9 @@ func Dial(ctx context.Context, dest net.Destination, streamSettings *internet.Me
}
tr := &quic.Transport{Conn: pktConn, DisableGSO: quicParams.DisableGSO}
if config.Warp != nil {
tr.ConnectionIDLength = 20
}
qconn, err := tr.Dial(ctx, udpAddr, gotlsConfig, quicConfig)
if err != nil {
tr.Close()
@@ -118,8 +124,16 @@ func Dial(ctx context.Context, dest net.Destination, streamSettings *internet.Me
return nil, errors.New("unknown congestion control: ", quicParams.Congestion)
}
cc := (&http3.Transport{EnableDatagrams: true, DisableCompression: true}).NewClientConn(qconn)
conn, err := establish(ctx, connectip.NewClientConn(cc), quicConn{qconn}, func() {
h3 := &http3.Transport{EnableDatagrams: true, DisableCompression: true}
if config.Warp != nil {
h3.AdditionalSettings = map[uint64]uint64{connectip.SettingDatagramDraft00: 1}
}
cc := h3.NewClientConn(qconn)
var client tunnelClient = connectip.NewClientConn(cc)
if config.Warp != nil {
client = connectip.NewCloudflareClientConn(cc, qconn)
}
conn, err := establish(ctx, client, quicConn{qconn}, func() {
qconn.CloseWithError(quic.ApplicationErrorCode(http3.ErrCodeRequestCanceled), "")
}, config, authority(config, gotlsConfig.ServerName, dest.Port))
if err != nil {
@@ -136,6 +150,9 @@ func usesHTTP2(config *tls.Config) bool {
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))
if err := useWarp(config, gotlsConfig); err != nil {
return nil, err
}
var conn net.Conn
var err error
@@ -157,7 +174,7 @@ func dialHTTP2(ctx context.Context, dest net.Destination, streamSettings *intern
conn.Close()
return nil, err
}
if protocol := tlsConn.NegotiatedProtocol(); protocol != http2.NextProtoTLS {
if protocol := tlsConn.NegotiatedProtocol(); protocol != http2.NextProtoTLS && (config.Warp == nil || protocol != "") {
conn.Close()
return nil, errors.New("the server negotiated ", strconv.Quote(protocol), " instead of h2")
}
@@ -167,7 +184,11 @@ func dialHTTP2(ctx context.Context, dest net.Destination, streamSettings *intern
conn.Close()
return nil, err
}
mconn, err := establish(ctx, connectip.NewHTTP2ClientConn(cc), cc, func() { cc.Close() }, config, authority(config, gotlsConfig.ServerName, dest.Port))
var client tunnelClient = connectip.NewHTTP2ClientConn(cc)
if config.Warp != nil {
client = connectip.NewCloudflareHTTP2ClientConn(cc)
}
mconn, err := establish(ctx, client, cc, func() { cc.Close() }, config, authority(config, gotlsConfig.ServerName, dest.Port))
if err != nil {
cc.Close()
return nil, err
@@ -221,21 +242,39 @@ func establish(ctx context.Context, client tunnelClient, hconn httpConn, abort f
return nil, errors.New("the tunnel can only carry ", n, "-byte packets, less than ", MinPacketSize)
}
if _, err := ipConn.RequestAddresses([]netip.Prefix{
netip.PrefixFrom(netip.IPv4Unspecified(), 32),
netip.PrefixFrom(netip.IPv6Unspecified(), 128),
}); err != nil {
ipConn.Close()
return nil, err
}
var local []netip.Addr
for len(local) == 0 {
assigned, err := ipConn.ReceiveAddressAssignment(ctx)
if err != nil {
if config.Warp != nil {
prefixes := make([]netip.Prefix, 0, len(config.Warp.Address))
for _, s := range config.Warp.Address {
prefix, err := netip.ParsePrefix(s)
if err != nil {
ipConn.Close()
return nil, errors.New("invalid WARP address ", s).Base(err)
}
prefixes = append(prefixes, prefix)
local = append(local, prefix.Addr())
}
if len(local) == 0 {
ipConn.Close()
return nil, errors.New("no address assigned").Base(err)
return nil, errors.New("WARP needs an address")
}
ipConn.SetAssignedAddresses(prefixes)
} else {
if _, err := ipConn.RequestAddresses([]netip.Prefix{
netip.PrefixFrom(netip.IPv4Unspecified(), 32),
netip.PrefixFrom(netip.IPv6Unspecified(), 128),
}); err != nil {
ipConn.Close()
return nil, err
}
for len(local) == 0 {
assigned, err := ipConn.ReceiveAddressAssignment(ctx)
if err != nil {
ipConn.Close()
return nil, errors.New("no address assigned").Base(err)
}
local = localAddrs(assigned)
}
local = localAddrs(assigned)
}
if !stop() {
ipConn.Close()
+6 -3
View File
@@ -207,11 +207,14 @@ func (c *http2ClientConn) writeHeaders(req *http.Request, maxFrameSize int) erro
if host == "" {
host = req.URL.Host
}
protocol := req.Header.Get(":protocol")
field(":method", req.Method)
field(":authority", host)
field(":scheme", req.URL.Scheme)
field(":path", req.URL.RequestURI())
if protocol := req.Header.Get(":protocol"); protocol != "" {
if req.Method != http.MethodConnect || protocol != "" {
field(":scheme", req.URL.Scheme)
field(":path", req.URL.RequestURI())
}
if protocol != "" {
field(":protocol", protocol)
}
if _, ok := req.Header["User-Agent"]; !ok {
+20
View File
@@ -194,6 +194,26 @@ func TestHTTP2ClientDefaultUserAgent(t *testing.T) {
require.Equal(t, []string{http2DefaultUserAgent}, userAgents)
}
func TestHTTP2ClientClassicConnect(t *testing.T) {
cc, p := newHTTP2Peer(t)
req, err := http.NewRequestWithContext(context.Background(), http.MethodConnect, "https://cloudflareaccess.com:443", nil)
require.NoError(t, err)
req.Header.Set("Cf-Connect-Proto", "cf-connect-ip")
req.Header["User-Agent"] = nil
go cc.RoundTrip(req)
f := p.readFrame()
require.IsType(t, &http2.MetaHeadersFrame{}, f)
var fields []string
for _, hf := range f.(*http2.MetaHeadersFrame).Fields {
fields = append(fields, hf.Name+": "+hf.Value)
}
require.Equal(t, []string{
":method: CONNECT",
":authority: cloudflareaccess.com:443",
"cf-connect-proto: cf-connect-ip",
}, fields)
}
func TestHTTP2ClientNeedsExtendedConnect(t *testing.T) {
cc, _ := newHTTP2Peer(t)
_, err := cc.RoundTrip(connectRequest(t, context.Background(), io.NopCloser(strings.NewReader(""))))
+78
View File
@@ -0,0 +1,78 @@
package masque
import (
"bytes"
"crypto/ecdsa"
"crypto/rand"
gotls "crypto/tls"
"crypto/x509"
"math/big"
"time"
"github.com/xtls/xray-core/common/errors"
)
const (
WarpHost = "cloudflareaccess.com"
WarpPath = "/"
)
func warpCertificate(der []byte) (*gotls.Certificate, error) {
parsed, err := x509.ParsePKCS8PrivateKey(der)
if err != nil {
return nil, errors.New("invalid WARP private key").Base(err)
}
key, ok := parsed.(*ecdsa.PrivateKey)
if !ok {
return nil, errors.New("the WARP private key is not an ECDSA key")
}
serial, err := rand.Int(rand.Reader, new(big.Int).Lsh(big.NewInt(1), 128))
if err != nil {
return nil, err
}
now := time.Now()
template := &x509.Certificate{
SerialNumber: serial,
NotBefore: now.Add(-time.Hour),
NotAfter: now.Add(24 * time.Hour),
}
cert, err := x509.CreateCertificate(rand.Reader, template, template, &key.PublicKey, key)
if err != nil {
return nil, err
}
return &gotls.Certificate{Certificate: [][]byte{cert}, PrivateKey: key}, nil
}
func useWarp(config *Config, tlsConfig *gotls.Config) error {
if config.Warp == nil {
return nil
}
cert, err := warpCertificate(config.Warp.PrivateKey)
if err != nil {
return err
}
tlsConfig.GetClientCertificate = func(*gotls.CertificateRequestInfo) (*gotls.Certificate, error) {
return cert, nil
}
if publicKey := config.Warp.PublicKey; len(publicKey) > 0 {
verify := tlsConfig.VerifyPeerCertificate
tlsConfig.InsecureSkipVerify = true
tlsConfig.VerifyPeerCertificate = func(raw [][]byte, chains [][]*x509.Certificate) error {
if len(raw) == 0 {
return errors.New("the WARP endpoint sent no certificate")
}
leaf, err := x509.ParseCertificate(raw[0])
if err != nil {
return err
}
if !bytes.Equal(leaf.RawSubjectPublicKeyInfo, publicKey) {
return errors.New("the WARP endpoint's key doesn't match \"publicKey\"")
}
if verify != nil {
return verify(raw, chains)
}
return nil
}
}
return nil
}
+439
View File
@@ -0,0 +1,439 @@
package masque
import (
"bytes"
"context"
"crypto/ecdsa"
"crypto/elliptic"
"crypto/rand"
gotls "crypto/tls"
"crypto/x509"
"encoding/binary"
"errors"
"io"
"math/big"
gonet "net"
"net/http"
"net/netip"
"slices"
"strings"
"testing"
"time"
"github.com/apernet/quic-go"
"github.com/apernet/quic-go/http3"
"github.com/apernet/quic-go/quicvarint"
"github.com/xtls/xray-core/common/net"
"github.com/xtls/xray-core/transport/internet"
"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"
)
var (
warpLocal4 = netip.MustParsePrefix("172.16.0.2/32")
warpLocal6 = netip.MustParsePrefix("2606:4700:110:8a36::2/128")
warpRemote = netip.MustParseAddr("1.1.1.1")
)
func newWarpKey(t *testing.T) (*ecdsa.PrivateKey, []byte) {
t.Helper()
key, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader)
if err != nil {
t.Fatal(err)
}
der, err := x509.MarshalPKCS8PrivateKey(key)
if err != nil {
t.Fatal(err)
}
return key, der
}
func warpServerTLS(t *testing.T, client *ecdsa.PublicKey, alpn string) (*gotls.Config, []byte, []byte) {
t.Helper()
key, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader)
if err != nil {
t.Fatal(err)
}
template := &x509.Certificate{
SerialNumber: big.NewInt(1),
DNSNames: []string{"localhost"},
NotBefore: time.Now().Add(-time.Hour),
NotAfter: time.Now().Add(time.Hour),
}
der, err := x509.CreateCertificate(rand.Reader, template, template, &key.PublicKey, key)
if err != nil {
t.Fatal(err)
}
config := &gotls.Config{
Certificates: []gotls.Certificate{{Certificate: [][]byte{der}, PrivateKey: key}},
ClientAuth: gotls.RequireAnyClientCert,
VerifyPeerCertificate: func(raw [][]byte, _ [][]*x509.Certificate) error {
cert, err := x509.ParseCertificate(raw[0])
if err != nil {
return err
}
if pub, ok := cert.PublicKey.(*ecdsa.PublicKey); !ok || !pub.Equal(client) {
return errors.New("unknown client key")
}
return nil
},
}
if alpn != "" {
config.NextProtos = []string{alpn}
}
publicKey, err := x509.MarshalPKIXPublicKey(&key.PublicKey)
if err != nil {
t.Fatal(err)
}
return config, publicKey, tls.GenerateCertHash(der)
}
func warpStreamSettings(key, publicKey []byte, alpn ...string) *internet.MemoryStreamConfig {
return &internet.MemoryStreamConfig{
ProtocolName: protocolName,
ProtocolSettings: &Config{Host: WarpHost, Path: WarpPath, Warp: &Warp{
PrivateKey: key,
PublicKey: publicKey,
Address: []string{warpLocal4.String(), warpLocal6.String()},
}},
SecurityType: "tls",
SecuritySettings: &tls.Config{
ServerName: "consumer-masque.cloudflareclient.com",
NextProtocol: alpn,
},
}
}
func warpPacket(src, dst netip.Addr, payload string) []byte {
b := make([]byte, 20+len(payload))
b[0] = 0x45
binary.BigEndian.PutUint16(b[2:], uint16(len(b)))
b[8] = 64
b[9] = 17
copy(b[12:16], src.AsSlice())
copy(b[16:20], dst.AsSlice())
copy(b[20:], payload)
return b
}
func warpCapsule(typ uint64, value []byte) []byte {
b := quicvarint.Append(nil, typ)
b = quicvarint.Append(b, uint64(len(value)))
return append(b, value...)
}
func readCapsuleTypes(r io.Reader, datagrams chan<- []byte) []uint64 {
var types []uint64
p := http3.NewCapsuleParser(r)
for {
typ, cr, err := p.Next()
if err != nil {
return types
}
types = append(types, uint64(typ))
b, err := io.ReadAll(cr)
if err != nil {
return types
}
if typ == 0 && datagrams != nil {
datagrams <- b
}
}
}
func checkWarpTunnel(t *testing.T, conn stat.Connection, sent <-chan []byte, reply func([]byte)) {
t.Helper()
mconn := conn.(*Conn)
if want := []netip.Addr{warpLocal4.Addr(), warpLocal6.Addr()}; !slices.Equal(mconn.LocalAddrs(), want) {
t.Fatalf("local addresses %v, want %v", mconn.LocalAddrs(), want)
}
out := warpPacket(warpLocal4.Addr(), warpRemote, "ping")
if _, err := conn.Write(slices.Clone(out)); err != nil {
t.Fatal(err)
}
select {
case got := <-sent:
if got[8] != 63 || !bytes.Equal(got[12:], out[12:]) {
t.Fatalf("the proxy got % x, want % x with TTL 63", got, out)
}
case <-time.After(5 * time.Second):
t.Fatal("the proxy got no packet")
}
reply(warpPacket(warpRemote, netip.MustParseAddr("172.16.0.3"), "lost"))
in := warpPacket(warpRemote, warpLocal4.Addr(), "pong")
reply(in)
b := make([]byte, 1500)
done := make(chan error, 1)
var n int
go func() {
var err error
n, err = conn.Read(b)
done <- err
}()
select {
case err := <-done:
if err != nil {
t.Fatal(err)
}
if !bytes.Equal(b[:n], in) {
t.Fatalf("read % x, want % x", b[:n], in)
}
case <-time.After(5 * time.Second):
t.Fatal("no packet came back")
}
}
func TestWarpCertificate(t *testing.T) {
key, der := newWarpKey(t)
cert, err := warpCertificate(der)
if err != nil {
t.Fatal(err)
}
leaf, err := x509.ParseCertificate(cert.Certificate[0])
if err != nil {
t.Fatal(err)
}
if pub, ok := leaf.PublicKey.(*ecdsa.PublicKey); !ok || !pub.Equal(&key.PublicKey) {
t.Fatal("the certificate doesn't carry the WARP key")
}
if err := leaf.CheckSignature(leaf.SignatureAlgorithm, leaf.RawTBSCertificate, leaf.Signature); err != nil {
t.Fatal(err)
}
if now := time.Now(); now.Before(leaf.NotBefore) || now.After(leaf.NotAfter) {
t.Fatal("the certificate isn't valid now")
}
if _, err := warpCertificate([]byte("not a key")); err == nil {
t.Fatal("expected an error for an invalid key")
}
}
func TestDialWarpHTTP3(t *testing.T) {
for _, draft := range []bool{true, false} {
t.Run(map[bool]string{true: "draft datagrams", false: "RFC 9297 datagrams"}[draft], func(t *testing.T) {
key, der := newWarpKey(t)
serverTLS, serverKey, _ := warpServerTLS(t, &key.PublicKey, http3.NextProtoH3)
sent := make(chan []byte, 1)
replies := make(chan []byte, 2)
capsules := make(chan []uint64, 1)
handler := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
settings := w.(http3.Settingser)
<-settings.ReceivedSettings()
switch {
case settings.Settings().Other[connectip.SettingDatagramDraft00] != 1:
t.Error("the client didn't send the draft datagram setting")
case r.Method != http.MethodConnect, r.Proto != "cf-connect-ip", r.Host != WarpHost, r.URL.Path != WarpPath:
t.Errorf("unexpected request %s %s %s%s", r.Method, r.Proto, r.Host, r.URL.Path)
}
w.WriteHeader(http.StatusOK)
w.(http.Flusher).Flush()
str := w.(http3.HTTPStreamer).HTTPStream()
go func() { capsules <- readCapsuleTypes(str, nil) }()
b, err := str.ReceiveDatagram(r.Context())
if err != nil {
t.Error(err)
return
}
if b[0] != 0 {
t.Errorf("datagram context ID %d, want 0", b[0])
}
sent <- b[1:]
for range 2 {
if err := str.SendDatagram(append([]byte{0}, <-replies...)); err != nil {
t.Error(err)
}
}
<-r.Context().Done()
})
udp, err := gonet.ListenUDP("udp4", &gonet.UDPAddr{IP: gonet.IPv4(127, 0, 0, 1)})
if err != nil {
t.Fatal(err)
}
server := &http3.Server{
Handler: handler,
TLSConfig: serverTLS,
QUICConfig: &quic.Config{EnableDatagrams: true},
EnableDatagrams: !draft,
}
if draft {
server.AdditionalSettings = map[uint64]uint64{connectip.SettingDatagramDraft00: 1}
}
go server.Serve(udp)
defer server.Close()
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
defer cancel()
dest := net.UDPDestination(net.LocalHostIP, net.Port(udp.LocalAddr().(*gonet.UDPAddr).Port))
conn, err := Dial(ctx, dest, warpStreamSettings(der, serverKey))
if err != nil {
t.Fatal(err)
}
checkWarpTunnel(t, conn, sent, func(b []byte) { replies <- b })
conn.Close()
select {
case types := <-capsules:
if slices.Contains(types, 2) {
t.Error("the client sent an ADDRESS_REQUEST")
}
case <-time.After(5 * time.Second):
t.Error("the request stream didn't end")
}
})
}
}
func TestDialWarpHTTP2(t *testing.T) {
for _, alpn := range []string{http2.NextProtoTLS, ""} {
t.Run(map[string]string{http2.NextProtoTLS: "h2 ALPN", "": "no ALPN"}[alpn], func(t *testing.T) {
key, der := newWarpKey(t)
serverTLS, serverKey, _ := warpServerTLS(t, &key.PublicKey, alpn)
sent := make(chan []byte, 1)
replies := make(chan []byte, 2)
capsules := make(chan []uint64, 1)
handler := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
switch {
case r.Method != http.MethodConnect, r.ProtoMajor != 2, r.Host != WarpHost+":443":
t.Errorf("unexpected request %s %s %s", r.Method, r.Proto, r.Host)
case r.Header.Get("Cf-Connect-Proto") != "cf-connect-ip", r.Header.Get("Pq-Enabled") != "false":
t.Errorf("unexpected headers %v", r.Header)
case r.Header.Get("Capsule-Protocol") != "", r.Header.Get(":protocol") != "":
t.Errorf("unexpected headers %v", r.Header)
}
w.WriteHeader(http.StatusOK)
w.(http.Flusher).Flush()
datagrams := make(chan []byte, 1)
go func() { capsules <- readCapsuleTypes(r.Body, datagrams) }()
b := <-datagrams
if b[0] != 0x45 {
t.Errorf("the DATAGRAM capsule starts with %#x, want a bare IPv4 packet", b[0])
}
sent <- b
for range 2 {
w.Write(warpCapsule(0, <-replies))
w.(http.Flusher).Flush()
}
<-r.Context().Done()
})
ln, err := gonet.Listen("tcp4", "127.0.0.1:0")
if err != nil {
t.Fatal(err)
}
defer ln.Close()
go func() {
for {
c, err := ln.Accept()
if err != nil {
return
}
go func() {
defer c.Close()
tc := gotls.Server(c, serverTLS)
if err := tc.Handshake(); err != nil {
return
}
(&http2.Server{}).ServeConn(tc, &http2.ServeConnOpts{Handler: handler})
}()
}
}()
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
defer cancel()
dest := net.TCPDestination(net.LocalHostIP, net.Port(ln.Addr().(*gonet.TCPAddr).Port))
conn, err := Dial(ctx, dest, warpStreamSettings(der, serverKey, http2.NextProtoTLS))
if err != nil {
t.Fatal(err)
}
checkWarpTunnel(t, conn, sent, func(b []byte) { replies <- b })
conn.Close()
select {
case types := <-capsules:
if slices.Contains(types, 2) {
t.Error("the client sent an ADDRESS_REQUEST")
}
case <-time.After(5 * time.Second):
t.Error("the request stream didn't end")
}
})
}
}
func TestDialWarpVerification(t *testing.T) {
key, der := newWarpKey(t)
serverTLS, serverKey, pin := warpServerTLS(t, &key.PublicKey, http3.NextProtoH3)
_, otherKey, _ := warpServerTLS(t, &key.PublicKey, http3.NextProtoH3)
udp, err := gonet.ListenUDP("udp4", &gonet.UDPAddr{IP: gonet.IPv4(127, 0, 0, 1)})
if err != nil {
t.Fatal(err)
}
server := &http3.Server{
Handler: http.NotFoundHandler(),
TLSConfig: serverTLS,
QUICConfig: &quic.Config{EnableDatagrams: true},
AdditionalSettings: map[uint64]uint64{connectip.SettingDatagramDraft00: 1},
}
go server.Serve(udp)
defer server.Close()
dest := net.UDPDestination(net.LocalHostIP, net.Port(udp.LocalAddr().(*gonet.UDPAddr).Port))
for _, c := range []struct {
name string
publicKey []byte
pin []byte
want string
}{
{"matching key", serverKey, nil, "404"},
{"matching key and pin", serverKey, pin, "404"},
{"other key", otherKey, nil, `doesn't match "publicKey"`},
{"matching key, other pin", serverKey, make([]byte, 32), "pinnedPeerCertSha256"},
} {
t.Run(c.name, func(t *testing.T) {
settings := warpStreamSettings(der, c.publicKey)
if c.pin != nil {
settings.SecuritySettings.(*tls.Config).PinnedPeerCertSha256 = [][]byte{c.pin}
}
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
defer cancel()
conn, err := Dial(ctx, dest, settings)
if err == nil {
conn.Close()
t.Fatal("expected an error")
}
if !strings.Contains(err.Error(), c.want) {
t.Fatalf("error %q doesn't mention %q", err, c.want)
}
})
}
}
func TestDialWarpRejectedKey(t *testing.T) {
_, der := newWarpKey(t)
other, _ := newWarpKey(t)
serverTLS, serverKey, _ := warpServerTLS(t, &other.PublicKey, http3.NextProtoH3)
udp, err := gonet.ListenUDP("udp4", &gonet.UDPAddr{IP: gonet.IPv4(127, 0, 0, 1)})
if err != nil {
t.Fatal(err)
}
server := &http3.Server{
Handler: http.NotFoundHandler(),
TLSConfig: serverTLS,
QUICConfig: &quic.Config{EnableDatagrams: true},
AdditionalSettings: map[uint64]uint64{connectip.SettingDatagramDraft00: 1},
}
go server.Serve(udp)
defer server.Close()
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
defer cancel()
dest := net.UDPDestination(net.LocalHostIP, net.Port(udp.LocalAddr().(*gonet.UDPAddr).Port))
if conn, err := Dial(ctx, dest, warpStreamSettings(der, serverKey)); err == nil {
conn.Close()
t.Fatal("expected the proxy to reject an unknown key")
}
}
+17
View File
@@ -174,6 +174,23 @@ func copyConfig(c *tls.Config) *utls.Config {
EncryptedClientHelloConfigList: c.EncryptedClientHelloConfigList,
NextProtos: c.NextProtos,
}
if c.GetClientCertificate != nil {
config.GetClientCertificate = func(info *utls.CertificateRequestInfo) (*utls.Certificate, error) {
schemes := make([]tls.SignatureScheme, len(info.SignatureSchemes))
for i, s := range info.SignatureSchemes {
schemes[i] = tls.SignatureScheme(s)
}
cert, err := c.GetClientCertificate(&tls.CertificateRequestInfo{
AcceptableCAs: info.AcceptableCAs,
SignatureSchemes: schemes,
Version: info.Version,
})
if err != nil || cert == nil {
return &utls.Certificate{}, err
}
return &utls.Certificate{Certificate: cert.Certificate, PrivateKey: cert.PrivateKey, Leaf: cert.Leaf}, nil
}
}
return config
}