mirror of
https://github.com/XTLS/Xray-core.git
synced 2026-10-08 14:58:00 +03:00
MASQUE client: Support WARP (#6878)
https://github.com/XTLS/Xray-core/pull/6844#issuecomment-5859092964 https://github.com/XTLS/Xray-core/pull/6862#issuecomment-5897508127 https://github.com/XTLS/Xray-core/pull/6878#issuecomment-6052448898
This commit is contained in:
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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,
|
||||
},
|
||||
|
||||
@@ -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;
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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(""))))
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user