mirror of
https://github.com/XTLS/Xray-core.git
synced 2026-10-08 22:59:51 +03:00
Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
adff74795e | ||
|
|
8989adfd3f | ||
|
|
1a60fc78ff | ||
|
|
2eedef51a2 | ||
|
|
37184948d4 | ||
|
|
9e55a6ed18 | ||
|
|
7da5dae650 | ||
|
|
a3c6cc6aa8 | ||
|
|
1e08d48d88 | ||
|
|
5121c28b85 | ||
|
|
92f7e72490 | ||
|
|
eddad37c69 | ||
|
|
00cde72a9b | ||
|
|
08775afd65 | ||
|
|
e0bae21201 | ||
|
|
9d9a7a1c00 |
@@ -314,12 +314,10 @@ func (m *ClientWorker) Dispatch(ctx context.Context, link *transport.Link) bool
|
||||
}
|
||||
|
||||
sm := m.sessionManager
|
||||
s := sm.Allocate(&m.strategy)
|
||||
s := sm.Allocate(&m.strategy, link.Reader, link.Writer)
|
||||
if s == nil {
|
||||
return false
|
||||
}
|
||||
s.input = link.Reader
|
||||
s.output = link.Writer
|
||||
go fetchInput(ctx, s, m.link.Writer)
|
||||
if _, ok := link.Reader.(*pipe.Reader); !ok {
|
||||
select {
|
||||
|
||||
@@ -51,7 +51,7 @@ func (m *SessionManager) Count() int {
|
||||
return int(m.count)
|
||||
}
|
||||
|
||||
func (m *SessionManager) Allocate(Strategy *ClientStrategy) *Session {
|
||||
func (m *SessionManager) Allocate(Strategy *ClientStrategy, input buf.Reader, output buf.Writer) *Session {
|
||||
m.Lock()
|
||||
defer m.Unlock()
|
||||
|
||||
@@ -64,6 +64,8 @@ func (m *SessionManager) Allocate(Strategy *ClientStrategy) *Session {
|
||||
|
||||
m.count++
|
||||
s := &Session{
|
||||
input: input,
|
||||
output: output,
|
||||
ID: m.count,
|
||||
parent: m,
|
||||
done: done.New(),
|
||||
|
||||
@@ -9,7 +9,7 @@ import (
|
||||
func TestSessionManagerAdd(t *testing.T) {
|
||||
m := NewSessionManager()
|
||||
|
||||
s := m.Allocate(&ClientStrategy{})
|
||||
s := m.Allocate(&ClientStrategy{}, nil, nil)
|
||||
if s.ID != 1 {
|
||||
t.Error("id: ", s.ID)
|
||||
}
|
||||
@@ -17,7 +17,7 @@ func TestSessionManagerAdd(t *testing.T) {
|
||||
t.Error("size: ", m.Size())
|
||||
}
|
||||
|
||||
s = m.Allocate(&ClientStrategy{})
|
||||
s = m.Allocate(&ClientStrategy{}, nil, nil)
|
||||
if s.ID != 2 {
|
||||
t.Error("id: ", s.ID)
|
||||
}
|
||||
@@ -39,7 +39,7 @@ func TestSessionManagerAdd(t *testing.T) {
|
||||
|
||||
func TestSessionManagerClose(t *testing.T) {
|
||||
m := NewSessionManager()
|
||||
s := m.Allocate(&ClientStrategy{})
|
||||
s := m.Allocate(&ClientStrategy{}, nil, nil)
|
||||
|
||||
if m.CloseIfNoSessionAndIdle(m.Size(), m.Count()) {
|
||||
t.Error("able to close")
|
||||
|
||||
@@ -146,6 +146,10 @@ func SniffQUIC(b []byte) (*SniffHeader, error) {
|
||||
}
|
||||
|
||||
restPayload := b[hdrLen+int(packetLen):]
|
||||
// cachedReader can concatenate zero-padded UDP datagrams.
|
||||
for len(restPayload) > 0 && restPayload[0] == 0 {
|
||||
restPayload = restPayload[1:]
|
||||
}
|
||||
if !isQUICInitial { // Skip this packet if it's not initial packet
|
||||
b = restPayload
|
||||
continue
|
||||
|
||||
File diff suppressed because one or more lines are too long
@@ -62,7 +62,7 @@ func (s *Service) Cleanup() error {
|
||||
}
|
||||
|
||||
for name, subs := range s.subs {
|
||||
newSub := make([]*Subscriber, 0, len(s.subs))
|
||||
newSub := make([]*Subscriber, 0, len(subs))
|
||||
for _, sub := range subs {
|
||||
if !sub.IsClosed() {
|
||||
newSub = append(newSub, sub)
|
||||
|
||||
@@ -55,6 +55,7 @@ require (
|
||||
github.com/vishvananda/netns v0.0.5 // indirect
|
||||
github.com/wlynxg/anet v0.0.5 // indirect
|
||||
go.yaml.in/yaml/v3 v3.0.5 // indirect
|
||||
golang.org/x/crypto/x509roots/fallback v0.0.0-20261005185213-c3db4df58582 // indirect
|
||||
golang.org/x/text v0.42.0 // indirect
|
||||
golang.org/x/time v0.14.0 // indirect
|
||||
golang.org/x/tools v0.49.0 // indirect
|
||||
|
||||
@@ -91,6 +91,8 @@ golang.org/x/crypto v0.0.0-20190308221718-c2843e01d9a2/go.mod h1:djNgcEr1/C05ACk
|
||||
golang.org/x/crypto v0.0.0-20191011191535-87dc89f01550/go.mod h1:yigFU9vqHzYiE8UmvKecakEJjdnWj3jj499lnFckfCI=
|
||||
golang.org/x/crypto v0.57.0 h1:3ZVCjf8Ggz7zneR/EHRVx68Ctf+2pmIMP2UFhh9cC6M=
|
||||
golang.org/x/crypto v0.57.0/go.mod h1:Fdz0i5U6CoizGwLda9DttjSk6qlZo25zYNtR+ycvuZA=
|
||||
golang.org/x/crypto/x509roots/fallback v0.0.0-20261005185213-c3db4df58582 h1:wjDBrGbfLuifgrVLFEWUBJYAfh5Q1wkMc5t0FY0tCbs=
|
||||
golang.org/x/crypto/x509roots/fallback v0.0.0-20261005185213-c3db4df58582/go.mod h1:HPze8vhfG6fO06AM+VSvxRm4E3+5Yk375mgrJ5M2z1E=
|
||||
golang.org/x/exp v0.0.0-20240506185415-9bf2ced13842 h1:vr/HnozRka3pE4EsMEg1lgkXJkTFJCVUX+S/ZT6wYzM=
|
||||
golang.org/x/exp v0.0.0-20240506185415-9bf2ced13842/go.mod h1:XtvwrStGgqGPLc4cjQfWqZHG1YFdYs6swckp8vpsjnc=
|
||||
golang.org/x/lint v0.0.0-20200302205851-738671d3881b/go.mod h1:3xt1FjdF8hUf6vQPIChWIBhFzV8gjjsPE/fR3IyQdNY=
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -15,7 +15,6 @@ import (
|
||||
googleuuid "github.com/google/uuid"
|
||||
"github.com/xtls/xray-core/common/errors"
|
||||
"github.com/xtls/xray-core/common/net"
|
||||
"github.com/xtls/xray-core/common/serial"
|
||||
"github.com/xtls/xray-core/transport/internet/finalmask/fragment"
|
||||
"github.com/xtls/xray-core/transport/internet/finalmask/header/custom"
|
||||
"github.com/xtls/xray-core/transport/internet/finalmask/mkcp/aes128gcm"
|
||||
@@ -792,37 +791,15 @@ func (c *Sudoku) Build() (proto.Message, error) {
|
||||
}
|
||||
|
||||
type XDNSDomain struct {
|
||||
Name string `json:"name"`
|
||||
Names []string `json:"names"`
|
||||
LenLimit int32 `json:"lenLimit"`
|
||||
LabelLimit int32 `json:"labelLimit"`
|
||||
Types []int32 `json:"types"`
|
||||
Edns0 int32 `json:"edns0"`
|
||||
}
|
||||
|
||||
type XDNSResolverTCP struct {
|
||||
Addr string `json:"addr"`
|
||||
}
|
||||
|
||||
func (c *XDNSResolverTCP) Build() (proto.Message, error) {
|
||||
return &xdns.TCPResolverProto{Addr: c.Addr}, nil
|
||||
}
|
||||
|
||||
type XDNSResolverUDP struct {
|
||||
Addr string `json:"addr"`
|
||||
}
|
||||
|
||||
func (c *XDNSResolverUDP) Build() (proto.Message, error) {
|
||||
return &xdns.UDPResolverProto{Addr: c.Addr}, nil
|
||||
}
|
||||
|
||||
var xdnsLoader = NewJSONConfigLoader(ConfigCreatorCache{
|
||||
"tcp": func() interface{} { return new(XDNSResolverTCP) },
|
||||
"udp": func() interface{} { return new(XDNSResolverUDP) },
|
||||
}, "type", "settings")
|
||||
|
||||
type XDNSResolver struct {
|
||||
Type string `json:"type"`
|
||||
Settings json.RawMessage `json:"settings"`
|
||||
Addrs []string `json:"addrs"`
|
||||
}
|
||||
|
||||
type XDNS struct {
|
||||
@@ -833,7 +810,7 @@ type XDNS struct {
|
||||
|
||||
func (c *XDNS) Build() (proto.Message, error) {
|
||||
var domains []*xdns.DomainProto
|
||||
var resolvers []*serial.TypedMessage
|
||||
var resolvers []*xdns.ResolverProto
|
||||
for i := range c.Domains {
|
||||
if c.Domains[i].LenLimit == 0 {
|
||||
c.Domains[i].LenLimit = 255
|
||||
@@ -841,33 +818,46 @@ func (c *XDNS) Build() (proto.Message, error) {
|
||||
if c.Domains[i].LabelLimit == 0 {
|
||||
c.Domains[i].LabelLimit = 63
|
||||
}
|
||||
types := make([]uint16, 0, len(c.Domains[i].Types))
|
||||
for j := range c.Domains[i].Types {
|
||||
types = append(types, uint16(c.Domains[i].Types[j]))
|
||||
}
|
||||
domain, err := xdns.NewDomain(c.Domains[i].Name, int(c.Domains[i].LenLimit), int(c.Domains[i].LabelLimit), types, uint16(c.Domains[i].Edns0))
|
||||
for j := range c.Domains[i].Names {
|
||||
domain, err := xdns.NewDomain(c.Domains[i].Names[j], int(c.Domains[i].LenLimit), int(c.Domains[i].LabelLimit), []uint16{1, 5, 16, 28}, uint16(c.Domains[i].Edns0))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
errors.LogInfo(context.Background(), domain.Show())
|
||||
domains = append(domains, &xdns.DomainProto{
|
||||
Name: c.Domains[i].Name,
|
||||
Name: c.Domains[i].Names[j],
|
||||
LenLimit: c.Domains[i].LenLimit,
|
||||
LabelLimit: c.Domains[i].LabelLimit,
|
||||
Types: c.Domains[i].Types,
|
||||
Edns0: c.Domains[i].Edns0,
|
||||
})
|
||||
}
|
||||
}
|
||||
for i := range c.Resolvers {
|
||||
config, err := xdnsLoader.LoadWithID(c.Resolvers[i].Settings, c.Resolvers[i].Type)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
for j := range c.Resolvers[i].Addrs {
|
||||
var u *url.URL
|
||||
var e error
|
||||
if !strings.Contains(c.Resolvers[i].Addrs[j], "://") {
|
||||
u, e = url.Parse("udp://" + c.Resolvers[i].Addrs[j])
|
||||
} else {
|
||||
u, e = url.Parse(c.Resolvers[i].Addrs[j])
|
||||
}
|
||||
pm, err := config.(interface{ Build() (proto.Message, error) }).Build()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
if e != nil {
|
||||
return nil, e
|
||||
}
|
||||
switch u.Scheme {
|
||||
case "tcp", "udp":
|
||||
default:
|
||||
return nil, errors.New("invalid protocol")
|
||||
}
|
||||
var host, port string
|
||||
host = u.Hostname()
|
||||
port = u.Port()
|
||||
if port == "" {
|
||||
port = "53"
|
||||
}
|
||||
resolvers = append(resolvers, &xdns.ResolverProto{Type: u.Scheme, Addr: net.JoinHostPort(host, port)})
|
||||
}
|
||||
resolvers = append(resolvers, serial.ToTypedMessage(pm))
|
||||
}
|
||||
if c.ExtraPoll < 0 || c.ExtraPoll > 3 {
|
||||
return nil, errors.New("c.ExtraPoll < 0 || c.ExtraPoll > 3")
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -316,6 +316,7 @@ type TLSConfig struct {
|
||||
ECHServerKeys string `json:"echServerKeys"`
|
||||
ECHConfigList string `json:"echConfigList"`
|
||||
ECHSocketSettings *SocketConfig `json:"echSockopt"`
|
||||
UseSystemCA bool `json:"useSystemCA"`
|
||||
}
|
||||
|
||||
// Build implements Buildable.
|
||||
@@ -403,6 +404,7 @@ func (c *TLSConfig) Build() (proto.Message, error) {
|
||||
}
|
||||
config.EchSocketSettings = ss
|
||||
}
|
||||
config.UseSystemCa = c.UseSystemCA
|
||||
|
||||
return config, nil
|
||||
}
|
||||
|
||||
@@ -213,6 +213,8 @@ If the filters cannot be added, Xray does not start. They are removed when Xray
|
||||
|
||||
`autoSystemWfpBlockLeak` (Windows only) is empty by default, as the filters break some setups: with `"dns"`, a local DNS resolver other programs use (e.g. on `127.0.0.1:53`), the DNS of another VPN on its own interface, virtual machines whose NAT resolves names on the host, or signing in to a captive portal; with `"misconfigtun"`, IPv4 or IPv6 on the local network while no route of that version leads to the TUN. Without the filters, DNS may leak as described above. To keep an IP version out of the TUN on purpose while still blocking DNS leaks, use only `["dns"]`.
|
||||
|
||||
`autoOutboundsInterface` (the default with `autoSystemRoutingTable`) keeps Xray's own connections out of the TUN by binding them to another interface, which Windows only honors while that interface has weak host send and forwarding off for the IP versions routed to the TUN. Otherwise, Windows sends them into the TUN, from that interface's address, and they stall. While the TUN runs, Xray therefore turns weak host send off on that interface, and on again when it stops or another interface takes over. Forwarding cannot be turned off this way, as Mobile Hotspot and Internet Connection Sharing need it, so a warning is logged while it is on. Having the hotspot share the TUN instead of that interface (Settings, Mobile hotspot, Share my internet connection from) moves forwarding to the TUN, where it does no harm, and sends the hotspot's devices through Xray as well.
|
||||
|
||||
You can give the adapter ip address manually, you can live Windows to give it autogenerated ip address (which take few seconds), it doesn't matter, the traffic going _through_ the interface will be forwarded into the app for proxying. \
|
||||
Minimal configuration that will work for local machine is routing passing the traffic on-link through the interface.
|
||||
You will need the interface id for that, unfortunately it is going to change with every Xray start due to implementation ambiguity between Xray and wintun driver.
|
||||
|
||||
@@ -101,7 +101,7 @@ func (t *stackGVisor) Start() error {
|
||||
// Use custom UDP packet handler, instead of strict gVisor forwarder, for FullCone NAT support
|
||||
udpForwarder := newUdpConnectionHandler(t.handler.HandleConnection, t.writeRawUDPPacket)
|
||||
ipStack.SetTransportProtocolHandler(udp.ProtocolNumber, func(id stack.TransportEndpointID, pkt *stack.PacketBuffer) bool {
|
||||
data := pkt.Clone().Data().AsRange().ToSlice()
|
||||
data := pkt.Data().AsRange().ToSlice()
|
||||
// if len(data) == 0 {
|
||||
// return false
|
||||
// }
|
||||
|
||||
@@ -155,12 +155,20 @@ func NewTun(options *Config) (Tun, error) {
|
||||
fdStr := platform.NewEnvFlag(platform.TunFdKey).GetValue(func() string { return "" })
|
||||
if fdStr != "" {
|
||||
// iOS: use provided fd from NetworkExtension
|
||||
fd, err := strconv.Atoi(fdStr)
|
||||
providedFd, err := strconv.Atoi(fdStr)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// duplicate NetworkExtension fd so Xray can close its own handle
|
||||
// without closing the original.
|
||||
fd, err := unix.FcntlInt(uintptr(providedFd), unix.F_DUPFD_CLOEXEC, 0)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if err = unix.SetNonblock(fd, true); err != nil {
|
||||
_ = unix.Close(fd)
|
||||
return nil, err
|
||||
}
|
||||
|
||||
@@ -232,12 +240,8 @@ func (t *DarwinTun) Close() error {
|
||||
t.waitKq.close()
|
||||
}
|
||||
routeErr := t.unsetSystemRoutes()
|
||||
if t.ownsFd {
|
||||
return xerrors.Combine(routeErr, t.tunFile.Close())
|
||||
}
|
||||
// iOS: don't close the fd, it's owned by NetworkExtension
|
||||
return routeErr
|
||||
}
|
||||
|
||||
func (t *DarwinTun) monitorRouteChanges() {
|
||||
buffer := make([]byte, 64*1024)
|
||||
|
||||
@@ -46,6 +46,7 @@ type WindowsTun struct {
|
||||
luid winipcfg.LUID
|
||||
cbr winipcfg.ChangeCallback
|
||||
cbi winipcfg.ChangeCallback
|
||||
guard outboundGuard
|
||||
wfp windows.Handle
|
||||
resolver *savedResolver
|
||||
skipStop chan struct{}
|
||||
@@ -178,6 +179,11 @@ startOver:
|
||||
}
|
||||
ipif, err := t.luid.IPInterface(family)
|
||||
if err != nil {
|
||||
// With IPv6 disabled system-wide (DisabledComponents), the adapter has no
|
||||
// IPv6 interface at all. Skip the family unless the config asks for it.
|
||||
if err == windows.ERROR_NOT_FOUND && family == windows.AF_INET6 && !address6 && !route6 {
|
||||
continue
|
||||
}
|
||||
return err
|
||||
}
|
||||
ipif.RouterDiscoveryBehavior = winipcfg.RouterDiscoveryDisabled
|
||||
@@ -292,10 +298,21 @@ startOver:
|
||||
}
|
||||
|
||||
if updater != nil {
|
||||
// Xray's own connections have to stay out of the IP versions routed
|
||||
// to the TUN, which needs Windows to honor the binding to updater's
|
||||
// interface.
|
||||
if route4 {
|
||||
t.guard.families = append(t.guard.families, windows.AF_INET)
|
||||
}
|
||||
if route6 {
|
||||
t.guard.families = append(t.guard.families, windows.AF_INET6)
|
||||
}
|
||||
t.guard.check()
|
||||
// Only a registered callback goes into the fields: a nil pointer in
|
||||
// them would not compare equal to nil in Close.
|
||||
cbr, err := winipcfg.RegisterRouteChangeCallback(func(notificationType winipcfg.MibNotificationType, route *winipcfg.MibIPforwardRow2) {
|
||||
updater.Update()
|
||||
t.guard.check()
|
||||
})
|
||||
if err != nil {
|
||||
return err
|
||||
@@ -303,6 +320,7 @@ startOver:
|
||||
t.cbr = cbr
|
||||
cbi, err := winipcfg.RegisterInterfaceChangeCallback(func(notificationType winipcfg.MibNotificationType, iface *winipcfg.MibIPInterfaceRow) {
|
||||
updater.Update()
|
||||
t.guard.check()
|
||||
})
|
||||
if err != nil {
|
||||
return err
|
||||
@@ -326,6 +344,7 @@ func (t *WindowsTun) Close() error {
|
||||
if t.cbi != nil {
|
||||
t.cbi.Unregister()
|
||||
}
|
||||
t.guard.restore()
|
||||
if t.luid != 0 {
|
||||
t.luid.FlushRoutes(windows.AF_INET)
|
||||
t.luid.FlushIPAddresses(windows.AF_INET)
|
||||
|
||||
@@ -0,0 +1,120 @@
|
||||
//go:build windows
|
||||
|
||||
package tun
|
||||
|
||||
import (
|
||||
"context"
|
||||
"slices"
|
||||
"strings"
|
||||
"sync"
|
||||
|
||||
"github.com/xtls/xray-core/common/errors"
|
||||
"golang.org/x/sys/windows"
|
||||
"golang.zx2c4.com/wireguard/windows/tunnel/winipcfg"
|
||||
)
|
||||
|
||||
// outboundGuard keeps Windows to the binding of autoOutboundsInterface, which
|
||||
// keeps Xray's own connections out of the TUN. With weak host send or
|
||||
// forwarding on for an IP version on the bound interface, Windows sends them
|
||||
// where the routes lead, into the TUN, from that interface's address, and
|
||||
// drops what comes back to that address through the TUN, so they stall.
|
||||
//
|
||||
// For the IP versions routed to the TUN, weak host send is turned off on the
|
||||
// bound interface while the TUN runs, and turned on again when the TUN stops
|
||||
// or another interface takes over. Forwarding is what Mobile Hotspot and
|
||||
// Internet Connection Sharing need, so it is only reported.
|
||||
type outboundGuard struct {
|
||||
sync.Mutex
|
||||
families []winipcfg.AddressFamily
|
||||
luid winipcfg.LUID // of the interface last checked
|
||||
name string // of that interface
|
||||
turnedOff []winipcfg.AddressFamily // where weak host send was turned off on it
|
||||
forwarding bool // whether forwarding was on there
|
||||
stopped bool
|
||||
}
|
||||
|
||||
// check turns weak host send off on the bound interface, and warns when
|
||||
// forwarding comes on there, but not again while it stays on.
|
||||
func (g *outboundGuard) check() {
|
||||
g.Lock()
|
||||
defer g.Unlock()
|
||||
if g.stopped {
|
||||
return
|
||||
}
|
||||
var luid winipcfg.LUID
|
||||
var name string
|
||||
if iface := updater.Get(); iface != nil {
|
||||
luid, _ = winipcfg.LUIDFromIndex(uint32(iface.Index))
|
||||
name = iface.Name
|
||||
}
|
||||
if luid != g.luid {
|
||||
g.restoreLocked()
|
||||
g.luid, g.name = luid, name
|
||||
g.forwarding = false // to warn about the new interface as well
|
||||
}
|
||||
if luid == 0 {
|
||||
return
|
||||
}
|
||||
var forwarding []string
|
||||
for _, family := range g.families {
|
||||
row, err := luid.IPInterface(family)
|
||||
if err != nil {
|
||||
continue // the interface lacks that IP version
|
||||
}
|
||||
if row.ForwardingEnabled {
|
||||
forwarding = append(forwarding, familyName(family))
|
||||
}
|
||||
if !row.WeakHostSend {
|
||||
continue
|
||||
}
|
||||
if err := setWeakHostSend(row, false); err != nil {
|
||||
errors.LogWarningInner(context.Background(), err, "[tun] unable to turn weak host send off for ", familyName(family), " on ", name)
|
||||
continue
|
||||
}
|
||||
if !slices.Contains(g.turnedOff, family) {
|
||||
g.turnedOff = append(g.turnedOff, family)
|
||||
errors.LogInfo(context.Background(), "[tun] weak host send turned off for ", familyName(family), " on ", name, " while the TUN runs, as Windows would ignore autoOutboundsInterface")
|
||||
}
|
||||
}
|
||||
wasOn := g.forwarding
|
||||
g.forwarding = len(forwarding) > 0
|
||||
if g.forwarding && !wasOn {
|
||||
errors.LogWarning(context.Background(), "[tun] forwarding is on for ", strings.Join(forwarding, " and "), " on ", name, " (Mobile Hotspot and Internet Connection Sharing turn it on), so Windows ignores autoOutboundsInterface there, and Xray's own connections go into the TUN and stall: turn the hotspot off, or have it share the TUN instead of ", name)
|
||||
}
|
||||
}
|
||||
|
||||
// restore turns weak host send on again where check turned it off, for good.
|
||||
func (g *outboundGuard) restore() {
|
||||
g.Lock()
|
||||
defer g.Unlock()
|
||||
g.restoreLocked()
|
||||
g.stopped = true
|
||||
}
|
||||
|
||||
func (g *outboundGuard) restoreLocked() {
|
||||
for _, family := range g.turnedOff {
|
||||
row, err := g.luid.IPInterface(family)
|
||||
if err == nil {
|
||||
err = setWeakHostSend(row, true)
|
||||
}
|
||||
if err != nil {
|
||||
errors.LogWarningInner(context.Background(), err, "[tun] unable to turn weak host send on again for ", familyName(family), " on ", g.name)
|
||||
}
|
||||
}
|
||||
g.turnedOff = nil
|
||||
}
|
||||
|
||||
func setWeakHostSend(row *winipcfg.MibIPInterfaceRow, on bool) error {
|
||||
row.WeakHostSend = on
|
||||
if row.Family == windows.AF_INET {
|
||||
row.SitePrefixLength = 0 // as SetIpInterfaceEntry requires for IPv4
|
||||
}
|
||||
return row.Set()
|
||||
}
|
||||
|
||||
func familyName(family winipcfg.AddressFamily) string {
|
||||
if family == windows.AF_INET {
|
||||
return "IPv4"
|
||||
}
|
||||
return "IPv6"
|
||||
}
|
||||
@@ -0,0 +1,91 @@
|
||||
package wireguard
|
||||
|
||||
import (
|
||||
"context"
|
||||
|
||||
"github.com/xtls/xray-core/common/errors"
|
||||
tunicmp "github.com/xtls/xray-core/proxy/tun/icmp"
|
||||
"gvisor.dev/gvisor/pkg/buffer"
|
||||
"gvisor.dev/gvisor/pkg/tcpip"
|
||||
"gvisor.dev/gvisor/pkg/tcpip/header"
|
||||
"gvisor.dev/gvisor/pkg/tcpip/stack"
|
||||
"gvisor.dev/gvisor/pkg/tcpip/transport/icmp"
|
||||
)
|
||||
|
||||
// CreateICMPEchoResponder answers ICMP echo requests from peers locally, the way
|
||||
// the TUN inbound does: ICMP is not proxied, but ping and connectivity checks
|
||||
// through the tunnel get a reply instead of timing out.
|
||||
//
|
||||
// In promiscuous mode gVisor skips its own IPv4 echo reply for addresses that are
|
||||
// not assigned to the NIC and leaves it to a custom handler; IPv6 is registered
|
||||
// too so both families behave the same.
|
||||
func CreateICMPEchoResponder(gstack *stack.Stack) {
|
||||
gstack.SetTransportProtocolHandler(icmp.ProtocolNumber4, func(id stack.TransportEndpointID, pkt *stack.PacketBuffer) bool {
|
||||
return handleICMPEcho(gstack, header.IPv4ProtocolNumber, id, pkt)
|
||||
})
|
||||
gstack.SetTransportProtocolHandler(icmp.ProtocolNumber6, func(id stack.TransportEndpointID, pkt *stack.PacketBuffer) bool {
|
||||
return handleICMPEcho(gstack, header.IPv6ProtocolNumber, id, pkt)
|
||||
})
|
||||
}
|
||||
|
||||
func handleICMPEcho(gstack *stack.Stack, netProto tcpip.NetworkProtocolNumber, id stack.TransportEndpointID, pkt *stack.PacketBuffer) bool {
|
||||
srcIP := id.RemoteAddress
|
||||
dstIP := id.LocalAddress
|
||||
if srcIP.Len() == 0 || dstIP.Len() == 0 {
|
||||
return true
|
||||
}
|
||||
|
||||
headerBytes := pkt.TransportHeader().Slice()
|
||||
payloadBytes := pkt.Data().AsRange().ToSlice()
|
||||
message := make([]byte, len(headerBytes)+len(payloadBytes))
|
||||
copy(message, headerBytes)
|
||||
copy(message[len(headerBytes):], payloadBytes)
|
||||
|
||||
if _, _, ok := tunicmp.ParseEchoRequest(netProto, message); !ok {
|
||||
return true
|
||||
}
|
||||
|
||||
reply, err := tunicmp.BuildLocalEchoReply(netProto, message, dstIP, srcIP)
|
||||
if err != nil {
|
||||
errors.LogInfoInner(context.Background(), err, "failed to build local icmp echo reply")
|
||||
return true
|
||||
}
|
||||
if err := writeRawICMPPacket(gstack, netProto, reply, dstIP, srcIP); err != nil {
|
||||
errors.LogInfoInner(context.Background(), err, "failed to write local icmp echo reply")
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
func writeRawICMPPacket(gstack *stack.Stack, netProto tcpip.NetworkProtocolNumber, message []byte, srcIP, dstIP tcpip.Address) error {
|
||||
pkt := stack.NewPacketBuffer(stack.PacketBufferOptions{
|
||||
ReserveHeaderBytes: header.IPv6MinimumSize,
|
||||
Payload: buffer.MakeWithData(message),
|
||||
})
|
||||
defer pkt.DecRef()
|
||||
|
||||
if netProto == header.IPv4ProtocolNumber {
|
||||
ipHdr := header.IPv4(pkt.NetworkHeader().Push(header.IPv4MinimumSize))
|
||||
ipHdr.Encode(&header.IPv4Fields{
|
||||
TotalLength: uint16(header.IPv4MinimumSize + len(message)),
|
||||
TTL: 64,
|
||||
Protocol: uint8(header.ICMPv4ProtocolNumber),
|
||||
SrcAddr: srcIP,
|
||||
DstAddr: dstIP,
|
||||
})
|
||||
ipHdr.SetChecksum(^ipHdr.CalculateChecksum())
|
||||
} else {
|
||||
ipHdr := header.IPv6(pkt.NetworkHeader().Push(header.IPv6MinimumSize))
|
||||
ipHdr.Encode(&header.IPv6Fields{
|
||||
PayloadLength: uint16(len(message)),
|
||||
TransportProtocol: header.ICMPv6ProtocolNumber,
|
||||
HopLimit: 64,
|
||||
SrcAddr: srcIP,
|
||||
DstAddr: dstIP,
|
||||
})
|
||||
}
|
||||
|
||||
if err := gstack.WriteRawPacket(1, netProto, buffer.MakeWithView(pkt.ToView())); err != nil {
|
||||
return errors.New("failed to write raw icmp packet back to stack ", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,176 @@
|
||||
package wireguard
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"net/netip"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/xtls/xray-core/common/net"
|
||||
"gvisor.dev/gvisor/pkg/tcpip"
|
||||
"gvisor.dev/gvisor/pkg/tcpip/checksum"
|
||||
"gvisor.dev/gvisor/pkg/tcpip/header"
|
||||
)
|
||||
|
||||
func newICMPTestStack(t *testing.T) *netTun {
|
||||
t.Helper()
|
||||
dev, _, gstack, err := CreateNetTUN([]netip.Addr{
|
||||
netip.MustParseAddr("10.66.0.1"),
|
||||
netip.MustParseAddr("fd00::1"),
|
||||
}, nil, 1420, false)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
t.Cleanup(func() { dev.Close() })
|
||||
CreateForwarder(gstack, func(conn net.Conn, dest net.Destination) { conn.Close() })
|
||||
CreateICMPEchoResponder(gstack)
|
||||
return dev.(*netTun)
|
||||
}
|
||||
|
||||
// startReader must run before the request is written: the stack may answer
|
||||
// synchronously inside Write, and netTun hands packets over an unbuffered channel.
|
||||
func startReader(dev *netTun) <-chan []byte {
|
||||
got := make(chan []byte, 1)
|
||||
go func() {
|
||||
buf := make([]byte, 2048)
|
||||
sizes := make([]int, 1)
|
||||
if _, err := dev.Read([][]byte{buf}, sizes, 0); err == nil {
|
||||
got <- buf[:sizes[0]]
|
||||
}
|
||||
}()
|
||||
return got
|
||||
}
|
||||
|
||||
func awaitPacket(t *testing.T, got <-chan []byte) []byte {
|
||||
t.Helper()
|
||||
select {
|
||||
case p := <-got:
|
||||
return p
|
||||
case <-time.After(2 * time.Second):
|
||||
t.Fatal("no echo reply from the stack")
|
||||
return nil
|
||||
}
|
||||
}
|
||||
|
||||
func TestICMPv4EchoReply(t *testing.T) {
|
||||
dev := newICMPTestStack(t)
|
||||
src := tcpip.AddrFrom4([4]byte{10, 66, 0, 2})
|
||||
dst := tcpip.AddrFrom4([4]byte{1, 1, 1, 1})
|
||||
payload := []byte("xray wireguard ping")
|
||||
|
||||
icmpMsg := make([]byte, header.ICMPv4MinimumSize+len(payload))
|
||||
req := header.ICMPv4(icmpMsg)
|
||||
req.SetType(header.ICMPv4Echo)
|
||||
req.SetIdent(0x1234)
|
||||
req.SetSequence(7)
|
||||
copy(req.Payload(), payload)
|
||||
req.SetChecksum(header.ICMPv4Checksum(req[:header.ICMPv4MinimumSize], checksum.Checksum(payload, 0)))
|
||||
|
||||
pkt := make([]byte, header.IPv4MinimumSize+len(icmpMsg))
|
||||
ip := header.IPv4(pkt)
|
||||
ip.Encode(&header.IPv4Fields{
|
||||
TotalLength: uint16(len(pkt)),
|
||||
TTL: 64,
|
||||
Protocol: uint8(header.ICMPv4ProtocolNumber),
|
||||
SrcAddr: src,
|
||||
DstAddr: dst,
|
||||
})
|
||||
ip.SetChecksum(^ip.CalculateChecksum())
|
||||
copy(pkt[header.IPv4MinimumSize:], icmpMsg)
|
||||
|
||||
got := startReader(dev)
|
||||
if _, err := dev.Write([][]byte{pkt}, 0); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
reply := header.IPv4(awaitPacket(t, got))
|
||||
if !reply.IsValid(len(reply)) {
|
||||
t.Fatal("invalid ipv4 reply")
|
||||
}
|
||||
if reply.SourceAddress() != dst || reply.DestinationAddress() != src {
|
||||
t.Fatalf("reply addresses %v -> %v, want %v -> %v", reply.SourceAddress(), reply.DestinationAddress(), dst, src)
|
||||
}
|
||||
if reply.TransportProtocol() != header.ICMPv4ProtocolNumber {
|
||||
t.Fatalf("reply protocol %v, want icmpv4", reply.TransportProtocol())
|
||||
}
|
||||
echo := header.ICMPv4(reply.Payload())
|
||||
if echo.Type() != header.ICMPv4EchoReply {
|
||||
t.Fatalf("reply type %v, want echo reply", echo.Type())
|
||||
}
|
||||
if echo.Ident() != 0x1234 || echo.Sequence() != 7 {
|
||||
t.Fatalf("reply ident/seq %#x/%d, want 0x1234/7", echo.Ident(), echo.Sequence())
|
||||
}
|
||||
if !bytes.Equal(echo.Payload(), payload) {
|
||||
t.Fatalf("reply payload %q, want %q", echo.Payload(), payload)
|
||||
}
|
||||
if checksum.Checksum(echo, 0) != 0xffff {
|
||||
t.Fatal("bad icmpv4 checksum")
|
||||
}
|
||||
}
|
||||
|
||||
func TestICMPv6EchoReply(t *testing.T) {
|
||||
dev := newICMPTestStack(t)
|
||||
src := tcpip.AddrFrom16([16]byte{0xfd, 15: 2})
|
||||
dst := tcpip.AddrFrom16([16]byte{0x26, 0x06, 0x47, 0x00, 0x47, 0x00, 15: 0x11})
|
||||
payload := []byte("xray wireguard ping6")
|
||||
|
||||
icmpMsg := make([]byte, header.ICMPv6MinimumSize+len(payload))
|
||||
req := header.ICMPv6(icmpMsg)
|
||||
req.SetType(header.ICMPv6EchoRequest)
|
||||
req.SetIdent(0x4321)
|
||||
req.SetSequence(9)
|
||||
copy(req.Payload(), payload)
|
||||
req.SetChecksum(header.ICMPv6Checksum(header.ICMPv6ChecksumParams{
|
||||
Header: req[:header.ICMPv6MinimumSize],
|
||||
Src: src,
|
||||
Dst: dst,
|
||||
PayloadCsum: checksum.Checksum(payload, 0),
|
||||
PayloadLen: len(payload),
|
||||
}))
|
||||
|
||||
pkt := make([]byte, header.IPv6MinimumSize+len(icmpMsg))
|
||||
ip := header.IPv6(pkt)
|
||||
ip.Encode(&header.IPv6Fields{
|
||||
PayloadLength: uint16(len(icmpMsg)),
|
||||
TransportProtocol: header.ICMPv6ProtocolNumber,
|
||||
HopLimit: 64,
|
||||
SrcAddr: src,
|
||||
DstAddr: dst,
|
||||
})
|
||||
copy(pkt[header.IPv6MinimumSize:], icmpMsg)
|
||||
|
||||
got := startReader(dev)
|
||||
if _, err := dev.Write([][]byte{pkt}, 0); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
reply := header.IPv6(awaitPacket(t, got))
|
||||
if !reply.IsValid(len(reply)) {
|
||||
t.Fatal("invalid ipv6 reply")
|
||||
}
|
||||
if reply.SourceAddress() != dst || reply.DestinationAddress() != src {
|
||||
t.Fatalf("reply addresses %v -> %v, want %v -> %v", reply.SourceAddress(), reply.DestinationAddress(), dst, src)
|
||||
}
|
||||
echo := header.ICMPv6(reply.Payload())
|
||||
if echo.Type() != header.ICMPv6EchoReply {
|
||||
t.Fatalf("reply type %v, want echo reply", echo.Type())
|
||||
}
|
||||
if echo.Ident() != 0x4321 || echo.Sequence() != 9 {
|
||||
t.Fatalf("reply ident/seq %#x/%d, want 0x4321/9", echo.Ident(), echo.Sequence())
|
||||
}
|
||||
if !bytes.Equal(echo.Payload(), payload) {
|
||||
t.Fatalf("reply payload %q, want %q", echo.Payload(), payload)
|
||||
}
|
||||
zeroed := header.ICMPv6(append([]byte(nil), echo[:header.ICMPv6MinimumSize]...))
|
||||
zeroed.SetChecksum(0)
|
||||
want := header.ICMPv6Checksum(header.ICMPv6ChecksumParams{
|
||||
Header: zeroed,
|
||||
Src: dst,
|
||||
Dst: src,
|
||||
PayloadCsum: checksum.Checksum(echo.Payload(), 0),
|
||||
PayloadLen: len(echo.Payload()),
|
||||
})
|
||||
if echo.Checksum() != want {
|
||||
t.Fatalf("icmpv6 checksum %#x, want %#x", echo.Checksum(), want)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,80 @@
|
||||
package wireguard
|
||||
|
||||
import (
|
||||
"runtime"
|
||||
"testing"
|
||||
)
|
||||
|
||||
const benchBatch = 64
|
||||
|
||||
// Raw cost of queueing and draining a small burst, as one flow's reader does.
|
||||
func BenchmarkQueueBurstChan(b *testing.B) {
|
||||
ch := make(chan *packet, udpQueueLimit)
|
||||
p := &packet{}
|
||||
b.ReportAllocs()
|
||||
for i := 0; i < b.N; i++ {
|
||||
for j := 0; j < benchBatch; j++ {
|
||||
ch <- p
|
||||
}
|
||||
for j := 0; j < benchBatch; j++ {
|
||||
<-ch
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func BenchmarkQueueBurstPacketQueue(b *testing.B) {
|
||||
q := newPacketQueue(udpQueueLimit)
|
||||
p := &packet{}
|
||||
b.ReportAllocs()
|
||||
for i := 0; i < b.N; i++ {
|
||||
for j := 0; j < benchBatch; j++ {
|
||||
q.push(p)
|
||||
}
|
||||
for j := 0; j < benchBatch; j++ {
|
||||
q.pop()
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Producer and consumer on different goroutines; the producer yields when the
|
||||
// queue is full instead of spinning, like a blocking channel send would.
|
||||
func BenchmarkQueueStreamChan(b *testing.B) {
|
||||
ch := make(chan *packet, udpQueueLimit)
|
||||
p := &packet{}
|
||||
done := make(chan struct{})
|
||||
go func() {
|
||||
for range ch {
|
||||
}
|
||||
close(done)
|
||||
}()
|
||||
b.ReportAllocs()
|
||||
b.ResetTimer()
|
||||
for i := 0; i < b.N; i++ {
|
||||
ch <- p
|
||||
}
|
||||
close(ch)
|
||||
<-done
|
||||
}
|
||||
|
||||
func BenchmarkQueueStreamPacketQueue(b *testing.B) {
|
||||
q := newPacketQueue(udpQueueLimit)
|
||||
p := &packet{}
|
||||
done := make(chan struct{})
|
||||
go func() {
|
||||
for {
|
||||
if _, ok := q.pop(); !ok {
|
||||
break
|
||||
}
|
||||
}
|
||||
close(done)
|
||||
}()
|
||||
b.ReportAllocs()
|
||||
b.ResetTimer()
|
||||
for i := 0; i < b.N; i++ {
|
||||
for !q.push(p) {
|
||||
runtime.Gosched()
|
||||
}
|
||||
}
|
||||
q.close()
|
||||
<-done
|
||||
}
|
||||
@@ -134,6 +134,7 @@ func NewServer(ctx context.Context, conf *DeviceConfig) (*Server, error) {
|
||||
}
|
||||
// Install the stack's protocol handlers before the device can deliver packets to it (Start -> dev.Up).
|
||||
CreateForwarder(stack, s.HandleConnection)
|
||||
CreateICMPEchoResponder(stack)
|
||||
return s, nil
|
||||
}
|
||||
|
||||
|
||||
+73
-18
@@ -85,7 +85,7 @@ func CreateForwarder(gstack *stack.Stack, handler func(conn net.Conn, dest net.D
|
||||
}
|
||||
|
||||
gstack.SetTransportProtocolHandler(udp.ProtocolNumber, func(id stack.TransportEndpointID, pkt *stack.PacketBuffer) bool {
|
||||
data := pkt.Clone().Data().AsRange().ToSlice()
|
||||
data := pkt.Data().AsRange().ToSlice()
|
||||
// if len(data) == 0 {
|
||||
// return false
|
||||
// }
|
||||
@@ -112,12 +112,7 @@ func (m *udpManager) feed(src net.Destination, dst net.Destination, data []byte)
|
||||
m.mutex.RLock()
|
||||
uc, ok := m.m[src.NetAddr()]
|
||||
if ok {
|
||||
select {
|
||||
case uc.queue <- &packet{
|
||||
p: data,
|
||||
dest: &dst,
|
||||
}:
|
||||
default:
|
||||
if !uc.queue.push(&packet{p: data, dest: &dst}) {
|
||||
errors.LogDebug(context.Background(), "drop udp with size ", len(data), " to ", dst.NetAddr(), " original ", uc.dst.NetAddr(), " > queue full")
|
||||
}
|
||||
m.mutex.RUnlock()
|
||||
@@ -131,7 +126,7 @@ func (m *udpManager) feed(src net.Destination, dst net.Destination, data []byte)
|
||||
uc, ok = m.m[src.NetAddr()]
|
||||
if !ok {
|
||||
uc = &udpConn{
|
||||
queue: make(chan *packet, 1024),
|
||||
queue: newPacketQueue(udpQueueLimit),
|
||||
src: src,
|
||||
dst: dst,
|
||||
}
|
||||
@@ -145,12 +140,7 @@ func (m *udpManager) feed(src net.Destination, dst net.Destination, data []byte)
|
||||
go m.handler(uc, dst)
|
||||
}
|
||||
|
||||
select {
|
||||
case uc.queue <- &packet{
|
||||
p: data,
|
||||
dest: &dst,
|
||||
}:
|
||||
default:
|
||||
if !uc.queue.push(&packet{p: data, dest: &dst}) {
|
||||
errors.LogDebug(context.Background(), "drop udp with size ", len(data), " to ", dst.NetAddr(), " original ", uc.dst.NetAddr(), " > queue full 2")
|
||||
}
|
||||
}
|
||||
@@ -158,7 +148,7 @@ func (m *udpManager) feed(src net.Destination, dst net.Destination, data []byte)
|
||||
func (m *udpManager) close(uc *udpConn) {
|
||||
if !uc.closed {
|
||||
uc.closed = true
|
||||
close(uc.queue)
|
||||
uc.queue.close()
|
||||
delete(m.m, uc.src.NetAddr())
|
||||
}
|
||||
}
|
||||
@@ -232,7 +222,7 @@ type packet struct {
|
||||
}
|
||||
|
||||
type udpConn struct {
|
||||
queue chan *packet
|
||||
queue *packetQueue
|
||||
src net.Destination
|
||||
dst net.Destination
|
||||
writeFunc func(payload []byte, src net.Destination, dst net.Destination) error
|
||||
@@ -242,7 +232,7 @@ type udpConn struct {
|
||||
|
||||
func (c *udpConn) ReadMultiBuffer() (buf.MultiBuffer, error) {
|
||||
for {
|
||||
q, ok := <-c.queue
|
||||
q, ok := c.queue.pop()
|
||||
if !ok {
|
||||
return nil, io.EOF
|
||||
}
|
||||
@@ -261,7 +251,7 @@ func (c *udpConn) ReadMultiBuffer() (buf.MultiBuffer, error) {
|
||||
}
|
||||
|
||||
func (c *udpConn) Read(p []byte) (int, error) {
|
||||
q, ok := <-c.queue
|
||||
q, ok := c.queue.pop()
|
||||
if !ok {
|
||||
return 0, io.EOF
|
||||
}
|
||||
@@ -324,3 +314,68 @@ func (c *udpConn) SetReadDeadline(t time.Time) error {
|
||||
func (c *udpConn) SetWriteDeadline(t time.Time) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
// udpQueueLimit bounds the packets waiting for one UDP flow; more are dropped.
|
||||
const udpQueueLimit = 1024
|
||||
|
||||
// packetQueue holds the packets waiting for one UDP flow. Unlike a buffered
|
||||
// channel of the same bound it only allocates for packets actually queued, so
|
||||
// the many idle flows kept until the idle timeout cost next to nothing.
|
||||
type packetQueue struct {
|
||||
mu sync.Mutex
|
||||
items []*packet
|
||||
limit int
|
||||
notify chan struct{}
|
||||
closed bool
|
||||
}
|
||||
|
||||
func newPacketQueue(limit int) *packetQueue {
|
||||
return &packetQueue{limit: limit, notify: make(chan struct{}, 1)}
|
||||
}
|
||||
|
||||
// push queues p and reports whether it was accepted.
|
||||
func (q *packetQueue) push(p *packet) bool {
|
||||
q.mu.Lock()
|
||||
defer q.mu.Unlock()
|
||||
if q.closed || len(q.items) >= q.limit {
|
||||
return false
|
||||
}
|
||||
q.items = append(q.items, p)
|
||||
select {
|
||||
case q.notify <- struct{}{}:
|
||||
default:
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
// pop blocks until a packet is queued or the queue is closed and drained.
|
||||
func (q *packetQueue) pop() (*packet, bool) {
|
||||
for {
|
||||
q.mu.Lock()
|
||||
if len(q.items) > 0 {
|
||||
p := q.items[0]
|
||||
q.items[0] = nil
|
||||
q.items = q.items[1:]
|
||||
if len(q.items) == 0 {
|
||||
q.items = nil
|
||||
}
|
||||
q.mu.Unlock()
|
||||
return p, true
|
||||
}
|
||||
if q.closed {
|
||||
q.mu.Unlock()
|
||||
return nil, false
|
||||
}
|
||||
q.mu.Unlock()
|
||||
<-q.notify
|
||||
}
|
||||
}
|
||||
|
||||
func (q *packetQueue) close() {
|
||||
q.mu.Lock()
|
||||
defer q.mu.Unlock()
|
||||
if !q.closed {
|
||||
q.closed = true
|
||||
close(q.notify)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -9,7 +9,7 @@ import (
|
||||
"net"
|
||||
"net/netip"
|
||||
"os"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"syscall"
|
||||
|
||||
"golang.org/x/sys/unix"
|
||||
@@ -20,21 +20,10 @@ import (
|
||||
"golang.zx2c4.com/wireguard/tun"
|
||||
)
|
||||
|
||||
var (
|
||||
tableIndex int = 10230
|
||||
mu sync.Mutex
|
||||
)
|
||||
var tableIndex atomic.Uint32
|
||||
|
||||
func allocateIPv6TableIndex() int {
|
||||
mu.Lock()
|
||||
defer mu.Unlock()
|
||||
|
||||
if tableIndex > 10230 {
|
||||
errors.LogInfo(context.Background(), "allocate new ipv6 table index: ", tableIndex)
|
||||
}
|
||||
currentIndex := tableIndex
|
||||
tableIndex++
|
||||
return currentIndex
|
||||
func init() {
|
||||
tableIndex.Store(10230)
|
||||
}
|
||||
|
||||
type kernelTun struct {
|
||||
@@ -111,17 +100,23 @@ func createKernelTun(localAddresses, dnsServers []netip.Addr, mtu int) (tdev tun
|
||||
}
|
||||
}
|
||||
|
||||
ipv6TableIndex := allocateIPv6TableIndex()
|
||||
var ipv6TableIndex int
|
||||
if v6 != nil {
|
||||
r := &netlink.Route{Table: ipv6TableIndex}
|
||||
r := &netlink.Route{}
|
||||
for {
|
||||
ipv6TableIndex = int(tableIndex.Add(1)) - 1
|
||||
r.Table = ipv6TableIndex
|
||||
routeList, fErr := netlink.RouteListFiltered(netlink.FAMILY_V6, r, netlink.RT_FILTER_TABLE)
|
||||
if len(routeList) == 0 || fErr != nil {
|
||||
if fErr != nil {
|
||||
return nil, nil, errors.New("failed to pre check routes for table: ", ipv6TableIndex).Base(fErr)
|
||||
}
|
||||
if len(routeList) == 0 {
|
||||
errors.LogInfo(context.Background(), "allocate new ipv6 table index: ", ipv6TableIndex)
|
||||
break
|
||||
}
|
||||
ipv6TableIndex--
|
||||
if ipv6TableIndex < 0 {
|
||||
return nil, nil, fmt.Errorf("failed to find available ipv6 table index")
|
||||
// to prevent infinite loop
|
||||
if ipv6TableIndex > 65535 {
|
||||
return nil, nil, errors.New("failed to find available ipv6 table index")
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,91 @@
|
||||
package wireguard
|
||||
|
||||
import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/xtls/xray-core/common/net"
|
||||
)
|
||||
|
||||
// BenchmarkUDPManagerNewSession measures what one new UDP flow costs the
|
||||
// inbound while it stays open: QUIC and DNS open many short flows, and each
|
||||
// one lives until the connection idle timeout.
|
||||
func BenchmarkUDPManagerNewSession(b *testing.B) {
|
||||
m := &udpManager{
|
||||
handler: func(conn net.Conn, dest net.Destination) {},
|
||||
m: make(map[string]*udpConn),
|
||||
}
|
||||
dst := net.UDPDestination(net.ParseAddress("1.1.1.1"), 443)
|
||||
payload := make([]byte, 1200)
|
||||
b.ReportAllocs()
|
||||
b.ResetTimer()
|
||||
for i := 0; i < b.N; i++ {
|
||||
src := net.UDPDestination(net.IPAddress([]byte{10, byte(i >> 16), byte(i >> 8), byte(i)}), net.Port(1024+i%60000))
|
||||
m.feed(src, dst, payload)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPacketQueueOrderAndClose(t *testing.T) {
|
||||
q := newPacketQueue(udpQueueLimit)
|
||||
for i := 0; i < 3; i++ {
|
||||
if !q.push(&packet{p: []byte{byte(i)}}) {
|
||||
t.Fatalf("push %d rejected", i)
|
||||
}
|
||||
}
|
||||
for i := 0; i < 3; i++ {
|
||||
p, ok := q.pop()
|
||||
if !ok || p.p[0] != byte(i) {
|
||||
t.Fatalf("pop %d: got %v, %v", i, p, ok)
|
||||
}
|
||||
}
|
||||
q.close()
|
||||
if _, ok := q.pop(); ok {
|
||||
t.Fatal("pop after close returned a packet")
|
||||
}
|
||||
if q.push(&packet{}) {
|
||||
t.Fatal("push after close accepted")
|
||||
}
|
||||
}
|
||||
|
||||
func TestPacketQueueLimit(t *testing.T) {
|
||||
q := newPacketQueue(udpQueueLimit)
|
||||
for i := 0; i < udpQueueLimit; i++ {
|
||||
if !q.push(&packet{}) {
|
||||
t.Fatalf("push %d rejected below the limit", i)
|
||||
}
|
||||
}
|
||||
if q.push(&packet{}) {
|
||||
t.Fatal("push above the limit accepted")
|
||||
}
|
||||
}
|
||||
|
||||
func TestPacketQueueCloseUnblocksReader(t *testing.T) {
|
||||
q := newPacketQueue(udpQueueLimit)
|
||||
done := make(chan bool)
|
||||
go func() {
|
||||
_, ok := q.pop()
|
||||
done <- ok
|
||||
}()
|
||||
q.close()
|
||||
select {
|
||||
case ok := <-done:
|
||||
if ok {
|
||||
t.Fatal("blocked pop returned a packet after close")
|
||||
}
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("close did not wake the reader")
|
||||
}
|
||||
}
|
||||
|
||||
func TestPacketQueueDropsDrainedStorage(t *testing.T) {
|
||||
q := newPacketQueue(udpQueueLimit)
|
||||
for i := 0; i < 100; i++ {
|
||||
q.push(&packet{})
|
||||
}
|
||||
for i := 0; i < 100; i++ {
|
||||
q.pop()
|
||||
}
|
||||
if q.items != nil {
|
||||
t.Fatalf("drained queue still holds %d slots", cap(q.items))
|
||||
}
|
||||
}
|
||||
@@ -115,6 +115,9 @@ func TestDokodemoTCP(t *testing.T) {
|
||||
defer CloseServer(server)
|
||||
break
|
||||
}
|
||||
if server != nil {
|
||||
CloseServer(server)
|
||||
}
|
||||
retry++
|
||||
if retry > 5 {
|
||||
t.Fatal("All attempts failed to start client")
|
||||
@@ -209,6 +212,9 @@ func TestDokodemoUDP(t *testing.T) {
|
||||
defer CloseServer(server)
|
||||
break
|
||||
}
|
||||
if server != nil {
|
||||
CloseServer(server)
|
||||
}
|
||||
retry++
|
||||
if retry > 5 {
|
||||
t.Fatal("All attempts failed to start client")
|
||||
|
||||
@@ -227,6 +227,9 @@ func TestSocksBridageUDP(t *testing.T) {
|
||||
defer CloseServer(server)
|
||||
break
|
||||
}
|
||||
if server != nil {
|
||||
CloseServer(server)
|
||||
}
|
||||
retry++
|
||||
if retry > 5 {
|
||||
t.Fatal("All attempts failed to start server")
|
||||
@@ -342,6 +345,9 @@ func TestSocksBridageUDPWithRouting(t *testing.T) {
|
||||
defer CloseServer(server)
|
||||
break
|
||||
}
|
||||
if server != nil {
|
||||
CloseServer(server)
|
||||
}
|
||||
retry++
|
||||
if retry > 5 {
|
||||
t.Fatal("All attempts failed to start server")
|
||||
|
||||
@@ -47,7 +47,6 @@ type xdnsClient struct {
|
||||
resolverIndex atomic.Uint32
|
||||
|
||||
readCh chan packet
|
||||
sendCh chan []byte
|
||||
poolCh chan struct{}
|
||||
closeCh chan struct{}
|
||||
wg sync.WaitGroup
|
||||
@@ -70,6 +69,9 @@ func NewClient(c *Config, dialer *finalmask.Dialer) (net.PacketConn, error) {
|
||||
for j := range c.Domains[i].Types {
|
||||
types = append(types, uint16(c.Domains[i].Types[j]))
|
||||
}
|
||||
if len(types) == 0 {
|
||||
types = []uint16{16}
|
||||
}
|
||||
domain, err := NewDomain(c.Domains[i].Name, int(c.Domains[i].LenLimit), int(c.Domains[i].LabelLimit), types, uint16(c.Domains[i].Edns0))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
@@ -80,6 +82,9 @@ func NewClient(c *Config, dialer *finalmask.Dialer) (net.PacketConn, error) {
|
||||
for i := range c.Resolvers {
|
||||
resolver, err := NewResolver(c.Resolvers[i], dialer)
|
||||
if err != nil {
|
||||
for _, resolver := range resolvers {
|
||||
resolver.Close()
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
resolvers = append(resolvers, resolver)
|
||||
@@ -95,7 +100,6 @@ func NewClient(c *Config, dialer *finalmask.Dialer) (net.PacketConn, error) {
|
||||
resolverSends: make([]atomic.Uint32, len(c.Resolvers)),
|
||||
|
||||
readCh: make(chan packet),
|
||||
sendCh: make(chan []byte, 16),
|
||||
poolCh: make(chan struct{}, pollLimit),
|
||||
closeCh: make(chan struct{}),
|
||||
}
|
||||
@@ -112,6 +116,107 @@ func (c *xdnsClient) closed() bool {
|
||||
}
|
||||
}
|
||||
|
||||
func (c *xdnsClient) send(p []byte) {
|
||||
domain := c.domains[mrand.Intn(len(c.domains))]
|
||||
qtype := domain.types[mrand.Intn(len(domain.types))]
|
||||
|
||||
var buf [512]byte
|
||||
var data [255]byte
|
||||
|
||||
send := func(p []byte) {
|
||||
msg := dnsmessage.Message{
|
||||
Header: dnsmessage.Header{
|
||||
RecursionDesired: true,
|
||||
},
|
||||
Questions: []dnsmessage.Question{
|
||||
{
|
||||
Name: domain.Encode(p),
|
||||
Type: dnsmessage.Type(qtype),
|
||||
Class: dnsmessage.ClassINET,
|
||||
},
|
||||
},
|
||||
}
|
||||
if domain.edns0 > 0 {
|
||||
msg.Additionals = []dnsmessage.Resource{
|
||||
{
|
||||
Header: dnsmessage.ResourceHeader{
|
||||
Name: dnsmessage.MustNewName("."),
|
||||
Type: dnsmessage.TypeOPT,
|
||||
Class: dnsmessage.Class(domain.edns0),
|
||||
TTL: 0,
|
||||
},
|
||||
Body: &dnsmessage.OPTResource{},
|
||||
},
|
||||
}
|
||||
}
|
||||
pack := common.Must2(msg.AppendPack(buf[:0]))
|
||||
common.Must2(rand.Read(pack[:2]))
|
||||
|
||||
index := c.resolverIndex.Load()
|
||||
cur := c.resolverSends[index].Add(1)
|
||||
i := index
|
||||
for {
|
||||
i++
|
||||
if i == uint32(len(c.resolvers)) {
|
||||
i = 0
|
||||
}
|
||||
if i == index {
|
||||
break
|
||||
}
|
||||
if cur > c.resolverSends[i].Load() {
|
||||
break
|
||||
}
|
||||
}
|
||||
c.resolverIndex.Store(i)
|
||||
c.resolvers[index].Send(pack)
|
||||
}
|
||||
|
||||
if len(p) == 0 {
|
||||
copy(data[:], c.clientID[:])
|
||||
data[0] |= TypeMap[qtype]
|
||||
data[8] = 8
|
||||
common.Must2(rand.Read(data[9:17]))
|
||||
send(data[:17])
|
||||
return
|
||||
}
|
||||
|
||||
if len(p) <= domain.cap-12 {
|
||||
copy(data[:], c.clientID[:])
|
||||
data[0] |= TypeMap[qtype]
|
||||
data[8] = 3
|
||||
common.Must2(rand.Read(data[9:12]))
|
||||
copy(data[12:], p)
|
||||
send(data[:12+len(p)])
|
||||
return
|
||||
}
|
||||
|
||||
if len(p) <= 255*(domain.cap-15) {
|
||||
copy(data[:], c.clientID[:])
|
||||
data[0] |= TypeMap[qtype]
|
||||
data[8] = 3 | 0xC0
|
||||
common.Must2(rand.Read(data[9:12]))
|
||||
|
||||
fragID := byte(c.fragID.Add(1))
|
||||
fragN := len(p) / (domain.cap - 15)
|
||||
if len(p)%(domain.cap-15) > 0 {
|
||||
fragN++
|
||||
}
|
||||
|
||||
for i := range fragN {
|
||||
data[12] = fragID
|
||||
data[13] = byte(i)
|
||||
data[14] = byte(fragN)
|
||||
size := min(len(p), domain.cap-15)
|
||||
copy(data[15:], p[:size])
|
||||
send(data[:15+size])
|
||||
p = p[size:]
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
errors.LogError(context.Background(), "err size ", len(p))
|
||||
}
|
||||
|
||||
func (c *xdnsClient) read(buf []byte, addr net.Addr) bool {
|
||||
msg := dnsmessage.Message{}
|
||||
if err := msg.Unpack(buf); err != nil {
|
||||
@@ -187,11 +292,10 @@ func (c *xdnsClient) run() {
|
||||
}
|
||||
|
||||
c.wg.Add(1)
|
||||
go c.send()
|
||||
go c.poll()
|
||||
|
||||
c.wg.Wait()
|
||||
close(c.readCh)
|
||||
close(c.sendCh)
|
||||
close(c.poolCh)
|
||||
}
|
||||
|
||||
@@ -218,152 +322,36 @@ func (c *xdnsClient) recv(i int) {
|
||||
}
|
||||
}
|
||||
|
||||
func (c *xdnsClient) send() {
|
||||
func (c *xdnsClient) poll() {
|
||||
defer c.wg.Done()
|
||||
|
||||
var buf [512]byte
|
||||
var data [255]byte
|
||||
|
||||
sendMsg := func(p []byte, domain *Domain, qtype uint16) {
|
||||
msg := dnsmessage.Message{
|
||||
Header: dnsmessage.Header{
|
||||
RecursionDesired: true,
|
||||
},
|
||||
Questions: []dnsmessage.Question{
|
||||
{
|
||||
Name: domain.Encode(p),
|
||||
Type: dnsmessage.Type(qtype),
|
||||
Class: dnsmessage.ClassINET,
|
||||
},
|
||||
},
|
||||
select {
|
||||
case <-c.closeCh:
|
||||
case <-c.poolCh:
|
||||
}
|
||||
if domain.edns0 > 0 {
|
||||
msg.Additionals = []dnsmessage.Resource{
|
||||
{
|
||||
Header: dnsmessage.ResourceHeader{
|
||||
Name: dnsmessage.MustNewName("."),
|
||||
Type: dnsmessage.TypeOPT,
|
||||
Class: dnsmessage.Class(domain.edns0),
|
||||
TTL: 0,
|
||||
},
|
||||
Body: &dnsmessage.OPTResource{},
|
||||
},
|
||||
}
|
||||
}
|
||||
pack := common.Must2(msg.AppendPack(buf[:0]))
|
||||
common.Must2(rand.Read(pack[:2]))
|
||||
|
||||
index := c.resolverIndex.Load()
|
||||
cur := c.resolverSends[index].Add(1)
|
||||
i := index
|
||||
for {
|
||||
i++
|
||||
if i == uint32(len(c.resolvers)) {
|
||||
i = 0
|
||||
}
|
||||
if i == index {
|
||||
break
|
||||
}
|
||||
if cur > c.resolverSends[i].Load() {
|
||||
break
|
||||
}
|
||||
}
|
||||
c.resolverIndex.Store(i)
|
||||
c.resolvers[index].Send(pack)
|
||||
}
|
||||
|
||||
send := func(p []byte) {
|
||||
domain := c.domains[mrand.Intn(len(c.domains))]
|
||||
qtype := domain.types[mrand.Intn(len(domain.types))]
|
||||
|
||||
if len(p) == 0 {
|
||||
copy(data[:], c.clientID[:])
|
||||
data[0] |= TypeMap[qtype]
|
||||
data[8] = 8
|
||||
common.Must2(rand.Read(data[9:17]))
|
||||
sendMsg(data[:17], domain, qtype)
|
||||
return
|
||||
}
|
||||
|
||||
if len(p) <= domain.cap-12 {
|
||||
copy(data[:], c.clientID[:])
|
||||
data[0] |= TypeMap[qtype]
|
||||
data[8] = 3
|
||||
common.Must2(rand.Read(data[9:12]))
|
||||
copy(data[12:], p)
|
||||
sendMsg(data[:12+len(p)], domain, qtype)
|
||||
return
|
||||
}
|
||||
|
||||
if len(p) <= 255*(domain.cap-15) {
|
||||
copy(data[:], c.clientID[:])
|
||||
data[0] |= TypeMap[qtype]
|
||||
data[8] = 3 | 0xC0
|
||||
common.Must2(rand.Read(data[9:12]))
|
||||
|
||||
fragID := byte(c.fragID.Add(1))
|
||||
fragN := len(p) / (domain.cap - 15)
|
||||
if len(p)%(domain.cap-15) > 0 {
|
||||
fragN++
|
||||
}
|
||||
|
||||
for i := range fragN {
|
||||
data[12] = fragID
|
||||
data[13] = byte(i)
|
||||
data[14] = byte(fragN)
|
||||
size := min(len(p), domain.cap-15)
|
||||
copy(data[15:], p[:size])
|
||||
sendMsg(data[:15+size], domain, qtype)
|
||||
p = p[size:]
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
errors.LogError(context.Background(), "err size ", len(p))
|
||||
}
|
||||
|
||||
ticker := time.NewTicker(initPollDelay)
|
||||
defer ticker.Stop()
|
||||
delay := initPollDelay
|
||||
p := []byte(nil)
|
||||
timeout := false
|
||||
ticker := time.NewTicker(delay)
|
||||
defer ticker.Stop()
|
||||
for {
|
||||
select {
|
||||
case <-c.closeCh:
|
||||
return
|
||||
default:
|
||||
select {
|
||||
case <-c.closeCh:
|
||||
return
|
||||
case p = <-c.sendCh:
|
||||
case <-c.poolCh:
|
||||
delay = initPollDelay
|
||||
case <-ticker.C:
|
||||
timeout = true
|
||||
}
|
||||
}
|
||||
|
||||
if len(p) > 0 {
|
||||
select {
|
||||
case <-c.poolCh:
|
||||
default:
|
||||
}
|
||||
}
|
||||
|
||||
send(p)
|
||||
for range c.extraPoll {
|
||||
send(nil)
|
||||
}
|
||||
|
||||
if timeout {
|
||||
delay *= pollDelayMultiplier
|
||||
if delay > maxPollDelay {
|
||||
delay = maxPollDelay
|
||||
}
|
||||
timeout = false
|
||||
} else {
|
||||
delay = initPollDelay
|
||||
}
|
||||
if c.closed() {
|
||||
return
|
||||
}
|
||||
ticker.Reset(delay)
|
||||
c.send(nil)
|
||||
for range c.extraPoll {
|
||||
c.send(nil)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -385,11 +373,9 @@ func (c *xdnsClient) WriteTo(p []byte, addr net.Addr) (n int, err error) {
|
||||
errors.LogError(context.Background(), "err size ", len(p))
|
||||
return 0, errors.New("err size")
|
||||
}
|
||||
b := make([]byte, len(p))
|
||||
copy(b, p)
|
||||
select {
|
||||
case c.sendCh <- b:
|
||||
default:
|
||||
c.send(p)
|
||||
for range c.extraPoll {
|
||||
c.send(nil)
|
||||
}
|
||||
return len(p), nil
|
||||
}
|
||||
|
||||
@@ -7,7 +7,7 @@
|
||||
package xdns
|
||||
|
||||
import (
|
||||
serial "github.com/xtls/xray-core/common/serial"
|
||||
_ "github.com/xtls/xray-core/common/serial"
|
||||
protoreflect "google.golang.org/protobuf/reflect/protoreflect"
|
||||
protoimpl "google.golang.org/protobuf/runtime/protoimpl"
|
||||
reflect "reflect"
|
||||
@@ -98,10 +98,62 @@ func (x *DomainProto) GetEdns0() int32 {
|
||||
return 0
|
||||
}
|
||||
|
||||
type ResolverProto struct {
|
||||
state protoimpl.MessageState `protogen:"open.v1"`
|
||||
Type string `protobuf:"bytes,1,opt,name=type,proto3" json:"type,omitempty"`
|
||||
Addr string `protobuf:"bytes,2,opt,name=addr,proto3" json:"addr,omitempty"`
|
||||
unknownFields protoimpl.UnknownFields
|
||||
sizeCache protoimpl.SizeCache
|
||||
}
|
||||
|
||||
func (x *ResolverProto) Reset() {
|
||||
*x = ResolverProto{}
|
||||
mi := &file_transport_internet_finalmask_xdns_config_proto_msgTypes[1]
|
||||
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
|
||||
ms.StoreMessageInfo(mi)
|
||||
}
|
||||
|
||||
func (x *ResolverProto) String() string {
|
||||
return protoimpl.X.MessageStringOf(x)
|
||||
}
|
||||
|
||||
func (*ResolverProto) ProtoMessage() {}
|
||||
|
||||
func (x *ResolverProto) ProtoReflect() protoreflect.Message {
|
||||
mi := &file_transport_internet_finalmask_xdns_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 ResolverProto.ProtoReflect.Descriptor instead.
|
||||
func (*ResolverProto) Descriptor() ([]byte, []int) {
|
||||
return file_transport_internet_finalmask_xdns_config_proto_rawDescGZIP(), []int{1}
|
||||
}
|
||||
|
||||
func (x *ResolverProto) GetType() string {
|
||||
if x != nil {
|
||||
return x.Type
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func (x *ResolverProto) GetAddr() string {
|
||||
if x != nil {
|
||||
return x.Addr
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
type Config struct {
|
||||
state protoimpl.MessageState `protogen:"open.v1"`
|
||||
Domains []*DomainProto `protobuf:"bytes,1,rep,name=domains,proto3" json:"domains,omitempty"`
|
||||
Resolvers []*serial.TypedMessage `protobuf:"bytes,2,rep,name=resolvers,proto3" json:"resolvers,omitempty"`
|
||||
Resolvers []*ResolverProto `protobuf:"bytes,2,rep,name=resolvers,proto3" json:"resolvers,omitempty"`
|
||||
ExtraPoll int32 `protobuf:"varint,3,opt,name=extra_poll,json=extraPoll,proto3" json:"extra_poll,omitempty"`
|
||||
unknownFields protoimpl.UnknownFields
|
||||
sizeCache protoimpl.SizeCache
|
||||
@@ -109,7 +161,7 @@ type Config struct {
|
||||
|
||||
func (x *Config) Reset() {
|
||||
*x = Config{}
|
||||
mi := &file_transport_internet_finalmask_xdns_config_proto_msgTypes[1]
|
||||
mi := &file_transport_internet_finalmask_xdns_config_proto_msgTypes[2]
|
||||
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
|
||||
ms.StoreMessageInfo(mi)
|
||||
}
|
||||
@@ -121,7 +173,7 @@ func (x *Config) String() string {
|
||||
func (*Config) ProtoMessage() {}
|
||||
|
||||
func (x *Config) ProtoReflect() protoreflect.Message {
|
||||
mi := &file_transport_internet_finalmask_xdns_config_proto_msgTypes[1]
|
||||
mi := &file_transport_internet_finalmask_xdns_config_proto_msgTypes[2]
|
||||
if x != nil {
|
||||
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
|
||||
if ms.LoadMessageInfo() == nil {
|
||||
@@ -134,7 +186,7 @@ func (x *Config) ProtoReflect() protoreflect.Message {
|
||||
|
||||
// Deprecated: Use Config.ProtoReflect.Descriptor instead.
|
||||
func (*Config) Descriptor() ([]byte, []int) {
|
||||
return file_transport_internet_finalmask_xdns_config_proto_rawDescGZIP(), []int{1}
|
||||
return file_transport_internet_finalmask_xdns_config_proto_rawDescGZIP(), []int{2}
|
||||
}
|
||||
|
||||
func (x *Config) GetDomains() []*DomainProto {
|
||||
@@ -144,7 +196,7 @@ func (x *Config) GetDomains() []*DomainProto {
|
||||
return nil
|
||||
}
|
||||
|
||||
func (x *Config) GetResolvers() []*serial.TypedMessage {
|
||||
func (x *Config) GetResolvers() []*ResolverProto {
|
||||
if x != nil {
|
||||
return x.Resolvers
|
||||
}
|
||||
@@ -158,94 +210,6 @@ func (x *Config) GetExtraPoll() int32 {
|
||||
return 0
|
||||
}
|
||||
|
||||
type TCPResolverProto struct {
|
||||
state protoimpl.MessageState `protogen:"open.v1"`
|
||||
Addr string `protobuf:"bytes,1,opt,name=addr,proto3" json:"addr,omitempty"`
|
||||
unknownFields protoimpl.UnknownFields
|
||||
sizeCache protoimpl.SizeCache
|
||||
}
|
||||
|
||||
func (x *TCPResolverProto) Reset() {
|
||||
*x = TCPResolverProto{}
|
||||
mi := &file_transport_internet_finalmask_xdns_config_proto_msgTypes[2]
|
||||
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
|
||||
ms.StoreMessageInfo(mi)
|
||||
}
|
||||
|
||||
func (x *TCPResolverProto) String() string {
|
||||
return protoimpl.X.MessageStringOf(x)
|
||||
}
|
||||
|
||||
func (*TCPResolverProto) ProtoMessage() {}
|
||||
|
||||
func (x *TCPResolverProto) ProtoReflect() protoreflect.Message {
|
||||
mi := &file_transport_internet_finalmask_xdns_config_proto_msgTypes[2]
|
||||
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 TCPResolverProto.ProtoReflect.Descriptor instead.
|
||||
func (*TCPResolverProto) Descriptor() ([]byte, []int) {
|
||||
return file_transport_internet_finalmask_xdns_config_proto_rawDescGZIP(), []int{2}
|
||||
}
|
||||
|
||||
func (x *TCPResolverProto) GetAddr() string {
|
||||
if x != nil {
|
||||
return x.Addr
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
type UDPResolverProto struct {
|
||||
state protoimpl.MessageState `protogen:"open.v1"`
|
||||
Addr string `protobuf:"bytes,1,opt,name=addr,proto3" json:"addr,omitempty"`
|
||||
unknownFields protoimpl.UnknownFields
|
||||
sizeCache protoimpl.SizeCache
|
||||
}
|
||||
|
||||
func (x *UDPResolverProto) Reset() {
|
||||
*x = UDPResolverProto{}
|
||||
mi := &file_transport_internet_finalmask_xdns_config_proto_msgTypes[3]
|
||||
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
|
||||
ms.StoreMessageInfo(mi)
|
||||
}
|
||||
|
||||
func (x *UDPResolverProto) String() string {
|
||||
return protoimpl.X.MessageStringOf(x)
|
||||
}
|
||||
|
||||
func (*UDPResolverProto) ProtoMessage() {}
|
||||
|
||||
func (x *UDPResolverProto) ProtoReflect() protoreflect.Message {
|
||||
mi := &file_transport_internet_finalmask_xdns_config_proto_msgTypes[3]
|
||||
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 UDPResolverProto.ProtoReflect.Descriptor instead.
|
||||
func (*UDPResolverProto) Descriptor() ([]byte, []int) {
|
||||
return file_transport_internet_finalmask_xdns_config_proto_rawDescGZIP(), []int{3}
|
||||
}
|
||||
|
||||
func (x *UDPResolverProto) GetAddr() string {
|
||||
if x != nil {
|
||||
return x.Addr
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
var File_transport_internet_finalmask_xdns_config_proto protoreflect.FileDescriptor
|
||||
|
||||
const file_transport_internet_finalmask_xdns_config_proto_rawDesc = "" +
|
||||
@@ -257,16 +221,15 @@ const file_transport_internet_finalmask_xdns_config_proto_rawDesc = "" +
|
||||
"\vlabel_limit\x18\x03 \x01(\x05R\n" +
|
||||
"labelLimit\x12\x14\n" +
|
||||
"\x05types\x18\x04 \x03(\x05R\x05types\x12\x14\n" +
|
||||
"\x05edns0\x18\x05 \x01(\x05R\x05edns0\"\xb6\x01\n" +
|
||||
"\x05edns0\x18\x05 \x01(\x05R\x05edns0\"7\n" +
|
||||
"\rResolverProto\x12\x12\n" +
|
||||
"\x04type\x18\x01 \x01(\tR\x04type\x12\x12\n" +
|
||||
"\x04addr\x18\x02 \x01(\tR\x04addr\"\xcb\x01\n" +
|
||||
"\x06Config\x12M\n" +
|
||||
"\adomains\x18\x01 \x03(\v23.xray.transport.internet.finalmask.xdns.DomainProtoR\adomains\x12>\n" +
|
||||
"\tresolvers\x18\x02 \x03(\v2 .xray.common.serial.TypedMessageR\tresolvers\x12\x1d\n" +
|
||||
"\adomains\x18\x01 \x03(\v23.xray.transport.internet.finalmask.xdns.DomainProtoR\adomains\x12S\n" +
|
||||
"\tresolvers\x18\x02 \x03(\v25.xray.transport.internet.finalmask.xdns.ResolverProtoR\tresolvers\x12\x1d\n" +
|
||||
"\n" +
|
||||
"extra_poll\x18\x03 \x01(\x05R\textraPoll\"&\n" +
|
||||
"\x10TCPResolverProto\x12\x12\n" +
|
||||
"\x04addr\x18\x01 \x01(\tR\x04addr\"&\n" +
|
||||
"\x10UDPResolverProto\x12\x12\n" +
|
||||
"\x04addr\x18\x01 \x01(\tR\x04addrB\x94\x01\n" +
|
||||
"extra_poll\x18\x03 \x01(\x05R\textraPollB\x94\x01\n" +
|
||||
"*com.xray.transport.internet.finalmask.xdnsP\x01Z;github.com/xtls/xray-core/transport/internet/finalmask/xdns\xaa\x02&Xray.Transport.Internet.Finalmask.Xdnsb\x06proto3"
|
||||
|
||||
var (
|
||||
@@ -281,17 +244,15 @@ func file_transport_internet_finalmask_xdns_config_proto_rawDescGZIP() []byte {
|
||||
return file_transport_internet_finalmask_xdns_config_proto_rawDescData
|
||||
}
|
||||
|
||||
var file_transport_internet_finalmask_xdns_config_proto_msgTypes = make([]protoimpl.MessageInfo, 4)
|
||||
var file_transport_internet_finalmask_xdns_config_proto_msgTypes = make([]protoimpl.MessageInfo, 3)
|
||||
var file_transport_internet_finalmask_xdns_config_proto_goTypes = []any{
|
||||
(*DomainProto)(nil), // 0: xray.transport.internet.finalmask.xdns.DomainProto
|
||||
(*Config)(nil), // 1: xray.transport.internet.finalmask.xdns.Config
|
||||
(*TCPResolverProto)(nil), // 2: xray.transport.internet.finalmask.xdns.TCPResolverProto
|
||||
(*UDPResolverProto)(nil), // 3: xray.transport.internet.finalmask.xdns.UDPResolverProto
|
||||
(*serial.TypedMessage)(nil), // 4: xray.common.serial.TypedMessage
|
||||
(*ResolverProto)(nil), // 1: xray.transport.internet.finalmask.xdns.ResolverProto
|
||||
(*Config)(nil), // 2: xray.transport.internet.finalmask.xdns.Config
|
||||
}
|
||||
var file_transport_internet_finalmask_xdns_config_proto_depIdxs = []int32{
|
||||
0, // 0: xray.transport.internet.finalmask.xdns.Config.domains:type_name -> xray.transport.internet.finalmask.xdns.DomainProto
|
||||
4, // 1: xray.transport.internet.finalmask.xdns.Config.resolvers:type_name -> xray.common.serial.TypedMessage
|
||||
1, // 1: xray.transport.internet.finalmask.xdns.Config.resolvers:type_name -> xray.transport.internet.finalmask.xdns.ResolverProto
|
||||
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
|
||||
@@ -310,7 +271,7 @@ func file_transport_internet_finalmask_xdns_config_proto_init() {
|
||||
GoPackagePath: reflect.TypeOf(x{}).PkgPath(),
|
||||
RawDescriptor: unsafe.Slice(unsafe.StringData(file_transport_internet_finalmask_xdns_config_proto_rawDesc), len(file_transport_internet_finalmask_xdns_config_proto_rawDesc)),
|
||||
NumEnums: 0,
|
||||
NumMessages: 4,
|
||||
NumMessages: 3,
|
||||
NumExtensions: 0,
|
||||
NumServices: 0,
|
||||
},
|
||||
|
||||
@@ -16,16 +16,13 @@ message DomainProto {
|
||||
int32 edns0 = 5;
|
||||
}
|
||||
|
||||
message ResolverProto {
|
||||
string type = 1;
|
||||
string addr = 2;
|
||||
}
|
||||
|
||||
message Config {
|
||||
repeated DomainProto domains = 1;
|
||||
repeated xray.common.serial.TypedMessage resolvers = 2;
|
||||
repeated ResolverProto resolvers = 2;
|
||||
int32 extra_poll = 3;
|
||||
}
|
||||
|
||||
message TCPResolverProto {
|
||||
string addr = 1;
|
||||
}
|
||||
|
||||
message UDPResolverProto {
|
||||
string addr = 1;
|
||||
}
|
||||
@@ -6,9 +6,8 @@ import (
|
||||
)
|
||||
|
||||
const (
|
||||
fragTTL = 8 * time.Second
|
||||
fragTTL = 4 * time.Second
|
||||
fragSize = 4096
|
||||
fragClientIDSize = 16384
|
||||
fragCount = 4096
|
||||
)
|
||||
|
||||
@@ -27,7 +26,6 @@ type FragEntry struct {
|
||||
|
||||
type FragManager struct {
|
||||
m map[FragKey]*FragEntry
|
||||
sizem map[ClientID]int
|
||||
ch chan struct{}
|
||||
mu sync.Mutex
|
||||
}
|
||||
@@ -35,7 +33,6 @@ type FragManager struct {
|
||||
func NewFragManager() *FragManager {
|
||||
m := &FragManager{
|
||||
m: make(map[FragKey]*FragEntry),
|
||||
sizem: make(map[ClientID]int),
|
||||
ch: make(chan struct{}),
|
||||
}
|
||||
go m.gc()
|
||||
@@ -51,9 +48,8 @@ func (m *FragManager) closed() bool {
|
||||
}
|
||||
}
|
||||
|
||||
func (m *FragManager) removeEntey(k FragKey, e *FragEntry) {
|
||||
m.sizem[k.clientID] -= e.size
|
||||
delete(m.m, k)
|
||||
func (m *FragManager) remove(key FragKey) {
|
||||
delete(m.m, key)
|
||||
}
|
||||
|
||||
func (m *FragManager) tryRemove() {
|
||||
@@ -70,7 +66,7 @@ func (m *FragManager) tryRemove() {
|
||||
first = false
|
||||
}
|
||||
}
|
||||
m.removeEntey(key, entry)
|
||||
m.remove(key)
|
||||
}
|
||||
|
||||
func (m *FragManager) gc() {
|
||||
@@ -84,7 +80,7 @@ func (m *FragManager) gc() {
|
||||
m.mu.Lock()
|
||||
for k, e := range m.m {
|
||||
if now.After(e.deadline) {
|
||||
m.removeEntey(k, e)
|
||||
m.remove(k)
|
||||
}
|
||||
}
|
||||
m.mu.Unlock()
|
||||
@@ -109,7 +105,7 @@ func (m *FragManager) Feed(out []byte, key FragKey, fragIdx, fragN byte, data []
|
||||
if entry == nil {
|
||||
m.tryRemove()
|
||||
} else {
|
||||
m.removeEntey(key, entry)
|
||||
m.remove(key)
|
||||
}
|
||||
entry = &FragEntry{
|
||||
data: make([][]byte, fragN),
|
||||
@@ -131,11 +127,6 @@ func (m *FragManager) Feed(out []byte, key FragKey, fragIdx, fragN byte, data []
|
||||
if entry.size+len(data) > fragSize {
|
||||
return 0
|
||||
}
|
||||
if entry.len < int(entry.total)-1 {
|
||||
if m.sizem[key.clientID]+len(data) > fragClientIDSize {
|
||||
return 0
|
||||
}
|
||||
}
|
||||
|
||||
cp := make([]byte, len(data))
|
||||
copy(cp, data)
|
||||
@@ -144,7 +135,6 @@ func (m *FragManager) Feed(out []byte, key FragKey, fragIdx, fragN byte, data []
|
||||
entry.size += len(data)
|
||||
entry.len++
|
||||
entry.deadline = now.Add(fragTTL)
|
||||
m.sizem[key.clientID] += len(data)
|
||||
|
||||
if entry.len < int(entry.total) {
|
||||
return 0
|
||||
@@ -154,7 +144,7 @@ func (m *FragManager) Feed(out []byte, key FragKey, fragIdx, fragN byte, data []
|
||||
for i := range entry.data {
|
||||
out = append(out, entry.data[i]...)
|
||||
}
|
||||
m.removeEntey(key, entry)
|
||||
m.remove(key)
|
||||
return len(out)
|
||||
}
|
||||
|
||||
|
||||
@@ -4,7 +4,6 @@ import (
|
||||
"errors"
|
||||
"net"
|
||||
|
||||
"github.com/xtls/xray-core/common/serial"
|
||||
"github.com/xtls/xray-core/transport/internet/finalmask"
|
||||
)
|
||||
|
||||
@@ -15,17 +14,13 @@ type Resolver interface {
|
||||
Close()
|
||||
}
|
||||
|
||||
func NewResolver(proto *serial.TypedMessage, dialer *finalmask.Dialer) (Resolver, error) {
|
||||
config, err := proto.GetInstance()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
switch v := config.(type) {
|
||||
case *TCPResolverProto:
|
||||
return NewTCPResolver(v, dialer)
|
||||
case *UDPResolverProto:
|
||||
return NewUDPResolver(v, dialer)
|
||||
func NewResolver(config *ResolverProto, dialer *finalmask.Dialer) (Resolver, error) {
|
||||
switch config.Type {
|
||||
case "tcp":
|
||||
return NewTCPResolver(config, dialer)
|
||||
case "udp":
|
||||
return NewUDPResolver(config, dialer)
|
||||
default:
|
||||
return nil, errors.New("unknown proto")
|
||||
return nil, errors.New("unknown type")
|
||||
}
|
||||
}
|
||||
|
||||
@@ -24,7 +24,7 @@ type TCPResolver struct {
|
||||
mu sync.Mutex
|
||||
}
|
||||
|
||||
func NewTCPResolver(config *TCPResolverProto, dialer *finalmask.Dialer) (Resolver, error) {
|
||||
func NewTCPResolver(config *ResolverProto, dialer *finalmask.Dialer) (Resolver, error) {
|
||||
dest, err := net.ParseDestination("tcp:" + config.Addr)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
@@ -130,13 +130,15 @@ func (r *TCPResolver) Send(p []byte) {
|
||||
|
||||
func (r *TCPResolver) Close() {
|
||||
r.mu.Lock()
|
||||
defer r.mu.Unlock()
|
||||
if r.closed() {
|
||||
r.mu.Unlock()
|
||||
return
|
||||
}
|
||||
close(r.closeCh)
|
||||
if r.conn != nil {
|
||||
_ = r.conn.Close()
|
||||
conn := r.conn
|
||||
r.mu.Unlock()
|
||||
if conn != nil {
|
||||
_ = conn.Close()
|
||||
}
|
||||
r.wg.Wait()
|
||||
close(r.readCh)
|
||||
|
||||
@@ -22,7 +22,7 @@ type UDPResolver struct {
|
||||
mu sync.Mutex
|
||||
}
|
||||
|
||||
func NewUDPResolver(config *UDPResolverProto, dialer *finalmask.Dialer) (Resolver, error) {
|
||||
func NewUDPResolver(config *ResolverProto, dialer *finalmask.Dialer) (Resolver, error) {
|
||||
dest, err := net.ParseDestination("udp:" + config.Addr)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
@@ -117,13 +117,15 @@ func (r *UDPResolver) Send(p []byte) {
|
||||
|
||||
func (r *UDPResolver) Close() {
|
||||
r.mu.Lock()
|
||||
defer r.mu.Unlock()
|
||||
if r.closed() {
|
||||
r.mu.Unlock()
|
||||
return
|
||||
}
|
||||
close(r.closeCh)
|
||||
if r.conn != nil {
|
||||
_ = r.conn.Close()
|
||||
conn := r.conn
|
||||
r.mu.Unlock()
|
||||
if conn != nil {
|
||||
_ = conn.Close()
|
||||
}
|
||||
r.wg.Wait()
|
||||
close(r.readCh)
|
||||
|
||||
@@ -52,6 +52,9 @@ func NewServer(c *Config, raw net.PacketConn) (net.PacketConn, error) {
|
||||
for j := range c.Domains[i].Types {
|
||||
types = append(types, uint16(c.Domains[i].Types[j]))
|
||||
}
|
||||
if len(types) == 0 {
|
||||
types = []uint16{1, 5, 16, 28}
|
||||
}
|
||||
domain, err := NewDomain(c.Domains[i].Name, int(c.Domains[i].LenLimit), int(c.Domains[i].LabelLimit), types, uint16(c.Domains[i].Edns0))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
|
||||
@@ -127,7 +127,7 @@ func getGrpcClient(ctx context.Context, dest net.Destination, streamSettings *in
|
||||
if streamSettings.FinalMask != nil {
|
||||
c, err = streamSettings.FinalMask.DialTCP(gctx, net.TCPDestination(address, port))
|
||||
} else {
|
||||
c, err = internet.DialSystem(ctx, dest, streamSettings.SocketSettings)
|
||||
c, err = internet.DialSystem(gctx, net.TCPDestination(address, port), streamSettings.SocketSettings)
|
||||
}
|
||||
if err == nil {
|
||||
if tlsConfig != nil {
|
||||
|
||||
@@ -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,6 +508,8 @@ func (c *Conn) ReadPacket(b []byte) (int, error) {
|
||||
return 0, err
|
||||
}
|
||||
}
|
||||
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")
|
||||
@@ -499,7 +518,8 @@ func (c *Conn) ReadPacket(b []byte) (int, error) {
|
||||
if contextID != 0 {
|
||||
continue
|
||||
}
|
||||
packet := data[n:]
|
||||
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)
|
||||
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,6 +242,24 @@ 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)
|
||||
}
|
||||
|
||||
var local []netip.Addr
|
||||
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("WARP needs an address")
|
||||
}
|
||||
ipConn.SetAssignedAddresses(prefixes)
|
||||
} else {
|
||||
if _, err := ipConn.RequestAddresses([]netip.Prefix{
|
||||
netip.PrefixFrom(netip.IPv4Unspecified(), 32),
|
||||
netip.PrefixFrom(netip.IPv6Unspecified(), 128),
|
||||
@@ -228,7 +267,6 @@ func establish(ctx context.Context, client tunnelClient, hconn httpConn, abort f
|
||||
ipConn.Close()
|
||||
return nil, err
|
||||
}
|
||||
var local []netip.Addr
|
||||
for len(local) == 0 {
|
||||
assigned, err := ipConn.ReceiveAddressAssignment(ctx)
|
||||
if err != nil {
|
||||
@@ -237,6 +275,7 @@ func establish(ctx context.Context, client tunnelClient, hconn httpConn, abort f
|
||||
}
|
||||
local = localAddrs(assigned)
|
||||
}
|
||||
}
|
||||
if !stop() {
|
||||
ipConn.Close()
|
||||
return nil, errors.New("no address assigned").Base(context.Cause(ctx))
|
||||
|
||||
@@ -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)
|
||||
if req.Method != http.MethodConnect || protocol != "" {
|
||||
field(":scheme", req.URL.Scheme)
|
||||
field(":path", req.URL.RequestURI())
|
||||
if protocol := req.Header.Get(":protocol"); protocol != "" {
|
||||
}
|
||||
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")
|
||||
}
|
||||
}
|
||||
@@ -19,6 +19,7 @@ import (
|
||||
"github.com/xtls/xray-core/common/platform/filesystem"
|
||||
"github.com/xtls/xray-core/common/protocol/tls/cert"
|
||||
"github.com/xtls/xray-core/transport/internet"
|
||||
"golang.org/x/crypto/x509roots/fallback/bundle"
|
||||
)
|
||||
|
||||
var globalSessionCache = tls.NewLRUClientSessionCache(128)
|
||||
@@ -578,3 +579,39 @@ func verifyChain(certs []*x509.Certificate, pinnedPeerCertSha256 [][]byte) (veri
|
||||
}
|
||||
return certNotFound, nil
|
||||
}
|
||||
|
||||
var bundleCertPool = sync.OnceValue(func() *x509.CertPool {
|
||||
pool := x509.NewCertPool()
|
||||
for r := range bundle.Roots() {
|
||||
cert, err := x509.ParseCertificate(r.Certificate)
|
||||
if err != nil {
|
||||
continue
|
||||
}
|
||||
if r.Constraint != nil {
|
||||
pool.AddCertWithConstraint(cert, r.Constraint)
|
||||
} else {
|
||||
pool.AddCert(cert)
|
||||
}
|
||||
}
|
||||
return pool
|
||||
})
|
||||
|
||||
var systemCertPool = sync.OnceValue(func() *x509.CertPool {
|
||||
pool, err := x509.SystemCertPool()
|
||||
if err != nil {
|
||||
// use bundle cert pool as fallback
|
||||
pool = bundleCertPool()
|
||||
}
|
||||
return pool
|
||||
})
|
||||
|
||||
// pool should not be modified directly, use CertPool.Clone() if needed.
|
||||
func loadCA(useSystem bool) *x509.CertPool {
|
||||
var pool *x509.CertPool
|
||||
if useSystem {
|
||||
pool = systemCertPool()
|
||||
} else {
|
||||
pool = bundleCertPool()
|
||||
}
|
||||
return pool
|
||||
}
|
||||
|
||||
@@ -206,6 +206,7 @@ type Config struct {
|
||||
EchConfigList string `protobuf:"bytes,19,opt,name=ech_config_list,json=echConfigList,proto3" json:"ech_config_list,omitempty"`
|
||||
EchSocketSettings *internet.SocketConfig `protobuf:"bytes,21,opt,name=ech_socket_settings,json=echSocketSettings,proto3" json:"ech_socket_settings,omitempty"`
|
||||
PinnedPeerCertSha256 [][]byte `protobuf:"bytes,22,rep,name=pinned_peer_cert_sha256,json=pinnedPeerCertSha256,proto3" json:"pinned_peer_cert_sha256,omitempty"`
|
||||
UseSystemCa bool `protobuf:"varint,23,opt,name=use_system_ca,json=useSystemCa,proto3" json:"use_system_ca,omitempty"`
|
||||
unknownFields protoimpl.UnknownFields
|
||||
sizeCache protoimpl.SizeCache
|
||||
}
|
||||
@@ -359,6 +360,13 @@ func (x *Config) GetPinnedPeerCertSha256() [][]byte {
|
||||
return nil
|
||||
}
|
||||
|
||||
func (x *Config) GetUseSystemCa() bool {
|
||||
if x != nil {
|
||||
return x.UseSystemCa
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
var File_transport_internet_tls_config_proto protoreflect.FileDescriptor
|
||||
|
||||
const file_transport_internet_tls_config_proto_rawDesc = "" +
|
||||
@@ -377,7 +385,7 @@ const file_transport_internet_tls_config_proto_rawDesc = "" +
|
||||
"\x05Usage\x12\x10\n" +
|
||||
"\fENCIPHERMENT\x10\x00\x12\x14\n" +
|
||||
"\x10AUTHORITY_VERIFY\x10\x01\x12\x13\n" +
|
||||
"\x0fAUTHORITY_ISSUE\x10\x02\"\xa6\x06\n" +
|
||||
"\x0fAUTHORITY_ISSUE\x10\x02\"\xca\x06\n" +
|
||||
"\x06Config\x12J\n" +
|
||||
"\vcertificate\x18\x02 \x03(\v2(.xray.transport.internet.tls.CertificateR\vcertificate\x12\x1f\n" +
|
||||
"\vserver_name\x18\x03 \x01(\tR\n" +
|
||||
@@ -398,7 +406,8 @@ const file_transport_internet_tls_config_proto_rawDesc = "" +
|
||||
"\x0fech_server_keys\x18\x12 \x01(\fR\rechServerKeys\x12&\n" +
|
||||
"\x0fech_config_list\x18\x13 \x01(\tR\rechConfigList\x12U\n" +
|
||||
"\x13ech_socket_settings\x18\x15 \x01(\v2%.xray.transport.internet.SocketConfigR\x11echSocketSettings\x125\n" +
|
||||
"\x17pinned_peer_cert_sha256\x18\x16 \x03(\fR\x14pinnedPeerCertSha256Bs\n" +
|
||||
"\x17pinned_peer_cert_sha256\x18\x16 \x03(\fR\x14pinnedPeerCertSha256\x12\"\n" +
|
||||
"\ruse_system_ca\x18\x17 \x01(\bR\vuseSystemCaBs\n" +
|
||||
"\x1fcom.xray.transport.internet.tlsP\x01Z0github.com/xtls/xray-core/transport/internet/tls\xaa\x02\x1bXray.Transport.Internet.Tlsb\x06proto3"
|
||||
|
||||
var (
|
||||
|
||||
@@ -86,4 +86,6 @@ message Config {
|
||||
SocketConfig ech_socket_settings = 21;
|
||||
|
||||
repeated bytes pinned_peer_cert_sha256 = 22;
|
||||
|
||||
bool use_system_ca = 23;
|
||||
}
|
||||
|
||||
@@ -5,50 +5,23 @@ package tls
|
||||
|
||||
import (
|
||||
"crypto/x509"
|
||||
"sync"
|
||||
|
||||
"github.com/xtls/xray-core/common/errors"
|
||||
)
|
||||
|
||||
type rootCertsCache struct {
|
||||
sync.Mutex
|
||||
pool *x509.CertPool
|
||||
}
|
||||
|
||||
func (c *rootCertsCache) load() (*x509.CertPool, error) {
|
||||
c.Lock()
|
||||
defer c.Unlock()
|
||||
|
||||
if c.pool != nil {
|
||||
return c.pool, nil
|
||||
}
|
||||
|
||||
pool, err := x509.SystemCertPool()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
c.pool = pool
|
||||
return pool, nil
|
||||
}
|
||||
|
||||
var rootCerts rootCertsCache
|
||||
|
||||
func (c *Config) getCertPool() (*x509.CertPool, error) {
|
||||
if c.DisableSystemRoot {
|
||||
return c.loadSelfCertPool()
|
||||
}
|
||||
|
||||
if len(c.Certificate) == 0 {
|
||||
return rootCerts.load()
|
||||
return loadCA(c.UseSystemCa), nil
|
||||
}
|
||||
|
||||
pool, err := x509.SystemCertPool()
|
||||
if err != nil {
|
||||
return nil, errors.New("system root").Base(err)
|
||||
}
|
||||
pool := loadCA(c.UseSystemCa).Clone()
|
||||
for _, cert := range c.Certificate {
|
||||
if !pool.AppendCertsFromPEM(cert.Certificate) {
|
||||
return nil, errors.New("append cert to root").Base(err)
|
||||
return nil, errors.New("append cert to root")
|
||||
}
|
||||
}
|
||||
return pool, nil
|
||||
|
||||
@@ -3,12 +3,30 @@
|
||||
|
||||
package tls
|
||||
|
||||
import "crypto/x509"
|
||||
import (
|
||||
"crypto/x509"
|
||||
|
||||
"github.com/xtls/xray-core/common/errors"
|
||||
)
|
||||
|
||||
func (c *Config) getCertPool() (*x509.CertPool, error) {
|
||||
if c.DisableSystemRoot {
|
||||
return c.loadSelfCertPool()
|
||||
}
|
||||
|
||||
// Windows should keep RootCAs nil for using the system CA.
|
||||
if c.UseSystemCa && len(c.Certificate) == 0 {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
if len(c.Certificate) == 0 {
|
||||
return loadCA(c.UseSystemCa), nil
|
||||
}
|
||||
pool := loadCA(c.UseSystemCa).Clone()
|
||||
for _, cert := range c.Certificate {
|
||||
if !pool.AppendCertsFromPEM(cert.Certificate) {
|
||||
return nil, errors.New("failed to append cert")
|
||||
}
|
||||
}
|
||||
return pool, nil
|
||||
}
|
||||
|
||||
@@ -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