Compare commits

..
2 Commits
Author SHA1 Message Date
Fangliding c84753dae6 fmt 2026-09-20 17:09:02 +08:00
Fangliding d17906c2f1 Optimize logger 2026-09-20 17:00:49 +08:00
210 changed files with 5398 additions and 27586 deletions
-22
View File
@@ -212,28 +212,6 @@ func (s *DNS) IsOwnLink(ctx context.Context) bool {
return false
}
// MayUseSystemResolver reports whether any name server configured here could
// still resolve through the system resolver. That is what happens when no name
// server is configured at all, and it is also what a name server pointed at
// "localhost" does. Callers that are about to redirect the system resolver need
// to know, because a resolution path that reaches it would then loop back to
// them.
//
// Any such server is enough: name servers can be selected per domain, so a
// single local one makes some query reach the system resolver even when
// independent upstreams are configured alongside it.
func (s *DNS) MayUseSystemResolver() bool {
if len(s.clients) == 0 {
return true
}
for _, client := range s.clients {
if _, isLocal := client.server.(*LocalNameServer); isLocal {
return true
}
}
return false
}
// LookupIP implements dns.Client.
func (s *DNS) LookupIP(domain string, option dns.IPOption) ([]net.IP, uint32, error) {
// Normalize the FQDN form query
-59
View File
@@ -1,59 +0,0 @@
package dns
import (
"context"
"testing"
"github.com/xtls/xray-core/common/net"
feature_dns "github.com/xtls/xray-core/features/dns"
)
// fakeServer stands in for any name server that is not the system resolver.
type fakeServer struct{}
func (fakeServer) Name() string { return "fake" }
func (fakeServer) IsDisableCache() bool { return false }
func (fakeServer) QueryIP(context.Context, string, feature_dns.IPOption) ([]net.IP, uint32, error) {
return nil, 0, nil
}
// Callers that are about to redirect the system resolver rely on this to tell
// whether any resolution path could still reach the system resolver, so the
// mixed shape has to be reported as reachable: a domain-specific rule can
// select the system resolver even when an independent upstream also exists.
func TestMayUseSystemResolver(t *testing.T) {
tests := []struct {
name string
clients []*Client
want bool
}{
{
name: "no clients at all",
want: true,
},
{
name: "only the system resolver",
clients: []*Client{{server: NewLocalNameServer()}},
want: true,
},
{
name: "the system resolver alongside an independent name server",
clients: []*Client{{server: fakeServer{}}, {server: NewLocalNameServer()}},
want: true,
},
{
name: "only independent name servers",
clients: []*Client{{server: fakeServer{}}},
want: false,
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
server := &DNS{clients: tt.clients}
if got := server.MayUseSystemResolver(); got != tt.want {
t.Errorf("MayUseSystemResolver() = %v, want %v", got, tt.want)
}
})
}
}
+1 -11
View File
@@ -5,8 +5,7 @@ import (
)
type windowsReader struct {
bufs []syscall.WSABuf
ready bool
bufs []syscall.WSABuf
}
func (r *windowsReader) Init(bs []*Buffer) {
@@ -16,7 +15,6 @@ func (r *windowsReader) Init(bs []*Buffer) {
for _, b := range bs {
r.bufs = append(r.bufs, syscall.WSABuf{Len: uint32(Size), Buf: &b.v[0]})
}
r.ready = false
}
func (r *windowsReader) Clear() {
@@ -27,14 +25,6 @@ func (r *windowsReader) Clear() {
}
func (r *windowsReader) Read(fd uintptr) int32 {
// On the first invocation, we return -1 to indicate "not ready"
// to make rawConn.Read wait for readability using the runtime's own mechanism
// because syscall.WSARecv() is a blocking call when used with nil OVERLAPPED
if !r.ready {
r.ready = true
return -1
}
var nBytes uint32
var flags uint32
err := syscall.WSARecv(syscall.Handle(fd), &r.bufs[0], uint32(len(r.bufs)), &nBytes, &flags, nil, nil)
+50 -33
View File
@@ -82,10 +82,19 @@ func (f *MphDomainMatcherFactory) BuildMatcher(rules []*DomainRule) (DomainMatch
}
g.Add(m, uint32(i))
case *DomainRule_Geosite:
err := loadSiteMatchers(v.Geosite, func(m strmatcher.Matcher) { g.Add(m, uint32(i)) })
domains, err := loadSiteWithAttrs(v.Geosite.File, v.Geosite.Code, v.Geosite.Attrs)
if err != nil {
return nil, err
}
for j, d := range domains {
domains[j] = nil // peak mem
m, err := parseDomain(d)
if err != nil {
errors.LogError(context.Background(), "ignore invalid geosite entry in ", v.Geosite.File, ":", v.Geosite.Code, " at index ", j, ", ", err)
continue
}
g.Add(m, uint32(i))
}
default:
panic("unknown domain rule type")
}
@@ -99,12 +108,12 @@ func (f *MphDomainMatcherFactory) BuildMatcher(rules []*DomainRule) (DomainMatch
return g, nil
}
type CompactMphDomainMatcherFactory struct {
type CompactDomainMatcherFactory struct {
sync.Mutex
shared *utils.WeakCacheMap[string, strmatcher.MphValueMatcher]
shared *utils.WeakCacheMap[string, strmatcher.LinearAnyMatcher]
}
func (f *CompactMphDomainMatcherFactory) getOrCreateFrom(rule *GeoSiteRule) (*strmatcher.MphValueMatcher, error) {
func (f *CompactDomainMatcherFactory) getOrCreateFrom(rule *GeoSiteRule) (strmatcher.MatcherSet, error) {
key := rule.File + ":" + rule.Code + "@" + rule.Attrs
f.Lock()
@@ -116,23 +125,33 @@ func (f *CompactMphDomainMatcherFactory) getOrCreateFrom(rule *GeoSiteRule) (*st
}
errors.LogDebug(context.Background(), "geodata geosite matcher cache MISS ", key)
s := strmatcher.NewMphValueMatcher()
if err := loadSiteMatchers(rule, func(m strmatcher.Matcher) { s.Add(m, 0) }); err != nil {
s := strmatcher.NewLinearAnyMatcher()
domains, err := loadSiteWithAttrs(rule.File, rule.Code, rule.Attrs)
if err != nil {
return nil, err
}
if err := s.Build(); err != nil {
return nil, err
for i, d := range domains {
domains[i] = nil // peak mem
m, err := parseDomain(d)
if err != nil {
errors.LogError(context.Background(), "ignore invalid geosite entry in ", rule.File, ":", rule.Code, " at index ", i, ", ", err)
continue
}
s.Add(m)
}
f.shared.Store(key, s)
return s, nil
return s, err
}
// BuildMatcher implements DomainMatcherFactory.
func (f *CompactMphDomainMatcherFactory) BuildMatcher(rules []*DomainRule) (DomainMatcher, error) {
func (f *CompactDomainMatcherFactory) BuildMatcher(rules []*DomainRule) (DomainMatcher, error) {
if len(rules) == 0 {
return nil, errors.New("empty domain rule list")
}
compact := new(CompactMphDomainMatcher)
compact := &CompactDomainMatcher{
matchers: make([]strmatcher.MatcherSet, 0, len(rules)),
values: make([]uint32, 0, len(rules)),
}
for i, r := range rules {
switch v := r.Value.(type) {
case *DomainRule_Custom:
@@ -149,7 +168,8 @@ func (f *CompactMphDomainMatcherFactory) BuildMatcher(rules []*DomainRule) (Doma
if err != nil {
return nil, err
}
compact.combiner.Add(m, uint32(i))
compact.matchers = append(compact.matchers, m)
compact.values = append(compact.values, uint32(i))
default:
panic("unknown domain rule type")
}
@@ -157,40 +177,37 @@ func (f *CompactMphDomainMatcherFactory) BuildMatcher(rules []*DomainRule) (Doma
return compact, nil
}
type CompactMphDomainMatcher struct {
type CompactDomainMatcher struct {
custom strmatcher.ValueMatcher
combiner strmatcher.MphValueMatcherCombiner
matchers []strmatcher.MatcherSet
values []uint32
}
// Match implements DomainMatcher.
func (c *CompactMphDomainMatcher) Match(input string) []uint32 {
result := c.combiner.Match(input)
func (c *CompactDomainMatcher) Match(input string) []uint32 {
var result []uint32
if c.custom != nil {
result = append(c.custom.Match(input), result...)
result = append(result, c.custom.Match(input)...)
}
for i, m := range c.matchers {
if m.MatchAny(input) {
result = append(result, c.values[i])
}
}
return result
}
// MatchAny implements DomainMatcher.
func (c *CompactMphDomainMatcher) MatchAny(input string) bool {
func (c *CompactDomainMatcher) MatchAny(input string) bool {
if c.custom != nil && c.custom.MatchAny(input) {
return true
}
return c.combiner.MatchAny(input)
}
// loadSiteMatchers calls add with a matcher for every domain of the geosite rule and logs the invalid ones.
func loadSiteMatchers(rule *GeoSiteRule, add func(strmatcher.Matcher)) error {
i := 0
return loadSite(rule.File, rule.Code, rule.Attrs, func(t Domain_Type, value []byte) {
m, err := parseDomain(&Domain{Type: t, Value: string(value)})
if err != nil {
errors.LogError(context.Background(), "ignore invalid geosite entry in ", rule.File, ":", rule.Code, " at index ", i, ", ", err)
} else {
add(m)
for _, m := range c.matchers {
if m.MatchAny(input) {
return true
}
i++
})
}
return false
}
func parseDomain(d *Domain) (strmatcher.Matcher, error) {
@@ -214,7 +231,7 @@ func parseDomain(d *Domain) (strmatcher.Matcher, error) {
func newDomainMatcherFactory() DomainMatcherFactory {
switch runtime.GOOS {
case "ios", "android":
return &CompactMphDomainMatcherFactory{shared: utils.NewWeakCacheMap[string, strmatcher.MphValueMatcher]()}
return &CompactDomainMatcherFactory{shared: utils.NewWeakCacheMap[string, strmatcher.LinearAnyMatcher]()}
default:
return &MphDomainMatcherFactory{shared: utils.NewWeakCacheMap[string, strmatcher.MphValueMatcher]()}
}
+2 -76
View File
@@ -4,7 +4,6 @@ import (
"path/filepath"
"reflect"
"slices"
"sync"
"testing"
"github.com/xtls/xray-core/common/geodata/strmatcher"
@@ -12,7 +11,7 @@ import (
)
func TestCompactDomainMatcher_PreservesCustomRuleIndices(t *testing.T) {
factory := &CompactMphDomainMatcherFactory{shared: utils.NewWeakCacheMap[string, strmatcher.MphValueMatcher]()}
factory := &CompactDomainMatcherFactory{shared: utils.NewWeakCacheMap[string, strmatcher.LinearAnyMatcher]()}
matcher, err := factory.BuildMatcher([]*DomainRule{
{Value: &DomainRule_Custom{Custom: &Domain{Type: Domain_Full, Value: "example.com"}}},
{Value: &DomainRule_Custom{Custom: &Domain{Type: Domain_Domain, Value: "example.com"}}},
@@ -33,7 +32,7 @@ func TestCompactDomainMatcher_PreservesCustomRuleIndices(t *testing.T) {
func TestCompactDomainMatcher_PreservesMixedRuleIndices(t *testing.T) {
t.Setenv("xray.location.asset", filepath.Join("..", "..", "resources"))
factory := &CompactMphDomainMatcherFactory{shared: utils.NewWeakCacheMap[string, strmatcher.MphValueMatcher]()}
factory := &CompactDomainMatcherFactory{shared: utils.NewWeakCacheMap[string, strmatcher.LinearAnyMatcher]()}
matcher, err := factory.BuildMatcher([]*DomainRule{
{Value: &DomainRule_Geosite{Geosite: &GeoSiteRule{File: DefaultGeoSiteDat, Code: "CN"}}},
{Value: &DomainRule_Custom{Custom: &Domain{Type: Domain_Full, Value: "163.com"}}},
@@ -73,76 +72,3 @@ func TestMphDomainMatcher_MatchReturnsDetachedSlice(t *testing.T) {
t.Fatalf("Match() after caller mutation = %v, want %v", gotAgain, []uint32{0, 1})
}
}
// DNS sorts every Match result in place, so a matcher must never hand out a
// slice it keeps, also when only its keyword or regex part matches.
func TestDomainMatcher_MatchResultsCanBeSortedConcurrently(t *testing.T) {
t.Setenv("xray.location.asset", filepath.Join("..", "..", "resources"))
rules := []*DomainRule{
{Value: &DomainRule_Custom{Custom: &Domain{Type: Domain_Full, Value: "example.com"}}},
{Value: &DomainRule_Custom{Custom: &Domain{Type: Domain_Domain, Value: "example.com"}}},
{Value: &DomainRule_Custom{Custom: &Domain{Type: Domain_Substr, Value: "exam"}}},
{Value: &DomainRule_Custom{Custom: &Domain{Type: Domain_Regex, Value: `^ex.*\.org$`}}},
{Value: &DomainRule_Custom{Custom: &Domain{Type: Domain_Substr, Value: "exam"}}},
{Value: &DomainRule_Geosite{Geosite: &GeoSiteRule{File: DefaultGeoSiteDat, Code: "CN"}}},
{Value: &DomainRule_Custom{Custom: &Domain{Type: Domain_Full, Value: "only.full.test"}}},
}
cases := []struct {
input string
want []uint32
}{
{"example.com", []uint32{0, 1, 2, 4}},
{"www.example.com", []uint32{1, 2, 4}},
{"exam.net", []uint32{2, 4}}, // keyword part only
{"example.org", []uint32{2, 3, 4}},
{"163.com", []uint32{5}},
{"www.163.com", []uint32{5}},
{"only.full.test", []uint32{6}}, // full part only
{"nomatch.test", nil},
}
factories := map[string]DomainMatcherFactory{
"mph": &MphDomainMatcherFactory{shared: utils.NewWeakCacheMap[string, strmatcher.MphValueMatcher]()},
"compact": &CompactMphDomainMatcherFactory{shared: utils.NewWeakCacheMap[string, strmatcher.MphValueMatcher]()},
}
for name, factory := range factories {
t.Run(name, func(t *testing.T) {
matcher, err := factory.BuildMatcher(rules)
if err != nil {
t.Fatalf("BuildMatcher() failed: %v", err)
}
for _, c := range cases {
got := matcher.Match(c.input)
if sorted := slices.Sorted(slices.Values(got)); !slices.Equal(sorted, c.want) {
t.Fatalf("Match(%q) = %v, want %v", c.input, sorted, c.want)
}
got = got[:cap(got)]
for j := range got {
got[j] = ^uint32(0)
}
if again := slices.Sorted(slices.Values(matcher.Match(c.input))); !slices.Equal(again, c.want) {
t.Fatalf("Match(%q) after caller mutation = %v, want %v", c.input, again, c.want)
}
}
var wg sync.WaitGroup
for range 8 {
wg.Add(1)
go func() {
defer wg.Done()
for range 500 {
for _, c := range cases {
got := matcher.Match(c.input)
slices.Sort(got)
if !slices.Equal(got, c.want) {
t.Errorf("Match(%q) = %v, want %v", c.input, got, c.want)
return
}
}
}
}()
}
wg.Wait()
})
}
}
+62 -213
View File
@@ -5,14 +5,11 @@ import (
"bytes"
"io"
"runtime"
"slices"
"strings"
"unicode/utf8"
"github.com/xtls/xray-core/common/errors"
"github.com/xtls/xray-core/common/platform/filesystem"
"google.golang.org/protobuf/encoding/protowire"
"google.golang.org/protobuf/proto"
)
@@ -55,56 +52,17 @@ func loadIP(file, code string) ([]*CIDR, error) {
return geoip.Cidr, nil
}
// loadSite calls fn, in file order, with the type and value of every domain of the geosite code
// that has all the "@"-separated attrs. It decodes the entry while reading the file instead of
// unmarshalling it into a []*Domain, so value is only valid during fn.
func loadSite(file, code, attrs string, fn func(Domain_Type, []byte)) error {
runtime.GC() // peak mem
r, err := filesystem.OpenAsset(file)
func loadSite(file, code string) ([]*Domain, error) {
bs, err := loadFile(file, code)
if err != nil {
return errors.New("failed to open ", file).Base(err)
return nil, err
}
defer r.Close()
br := bufio.NewReaderSize(r, 64*1024)
n, err := seek(br, []byte(code))
if err != nil {
return errors.New("failed to load code ", code, " from ", file).Base(err)
defer runtime.GC() // peak mem
var geosite GeoSite
if err := proto.Unmarshal(bs, &geosite); err != nil {
return nil, errors.New("error unmarshal Site in ", file, ":", code).Base(err)
}
loadErr := func(err error) error {
if err == io.EOF {
err = io.ErrUnexpectedEOF
}
return errors.New("failed to load code ", code, " from ", file).Base(err)
}
unmarshalErr := func(err error) error {
return errors.New("error unmarshal Site in ", file, ":", code).Base(err)
}
d := newSiteDecoder(attrs, fn)
for n > 0 {
w, err := br.Peek(min(n, br.Size()))
if err != nil {
return loadErr(err)
}
used, err := d.decode(w, len(w) < n)
if err != nil {
return unmarshalErr(err)
}
if used == 0 {
break // a field longer than the buffer
}
br.Discard(used)
n -= used
}
if n > 0 {
w := make([]byte, n)
if _, err := io.ReadFull(br, w); err != nil {
return loadErr(err)
}
if _, err := d.decode(w, false); err != nil {
return unmarshalErr(err)
}
}
return nil
return geosite.Domain, nil
}
func decodeVarint(br *bufio.Reader) (uint64, error) {
@@ -124,63 +82,68 @@ func decodeVarint(br *bufio.Reader) (uint64, error) {
}
func find(r io.Reader, code []byte, readBody bool) ([]byte, error) {
br := bufio.NewReaderSize(r, 64*1024)
bodyL, err := seek(br, code)
if err != nil || !readBody {
return nil, err
}
out := make([]byte, bodyL)
if _, err := io.ReadFull(br, out); err != nil {
return nil, err
}
return out, nil
}
// seek advances br to the body of the entry for code and returns the body length.
func seek(br *bufio.Reader, code []byte) (int, error) {
codeL := len(code)
if codeL == 0 {
return 0, errors.New("empty code")
return nil, errors.New("empty code")
}
br := bufio.NewReaderSize(r, 64*1024)
need := 2 + codeL // TODO: if code too long
prefixBuf := make([]byte, need)
for {
if _, err := br.ReadByte(); err != nil {
return 0, err
return nil, err
}
x, err := decodeVarint(br)
if err != nil {
return 0, err
return nil, err
}
bodyL := int(x)
if bodyL <= 0 {
return 0, errors.New("invalid body length: ", bodyL)
return nil, errors.New("invalid body length: ", bodyL)
}
// Peek no more than the buffer holds: a code longer than the buffer cannot match a single
// length byte anyway, so a short peek only skips it, as base find (io.ReadFull) does.
prefix, err := br.Peek(min(bodyL, need, br.Size()))
if err != nil {
if err == io.EOF && len(prefix) > 0 {
err = io.ErrUnexpectedEOF // as io.ReadFull
prefixL := bodyL
if prefixL > need {
prefixL = need
}
prefix := prefixBuf[:prefixL]
if _, err := io.ReadFull(br, prefix); err != nil {
return nil, err
}
match := false
if bodyL >= need {
if int(prefix[1]) == codeL && bytes.Equal(prefix[2:need], code) {
if !readBody {
return nil, nil
}
match = true
}
return 0, err
}
if bodyL >= need && len(prefix) >= need && int(prefix[1]) == codeL && bytes.Equal(prefix[2:], code) {
return bodyL, nil
remain := bodyL - prefixL
if match {
out := make([]byte, bodyL)
copy(out, prefix)
if remain > 0 {
if _, err := io.ReadFull(br, out[prefixL:]); err != nil {
return nil, err
}
}
return out, nil
}
if _, err := br.Discard(bodyL); err != nil {
return 0, err
if remain > 0 {
if _, err := br.Discard(remain); err != nil {
return nil, err
}
}
}
}
// AttributeMatcher, HasAttrMatcher, AllAttrsMatcher and NewAllAttrsMatcher are the exported
// attribute helpers that have been part of this package's API since #5814. The streaming loader
// above filters attributes itself without building a *Domain, so it does not use them, but they
// are kept for external callers. Their behaviour is unchanged.
type AttributeMatcher interface {
Match(*Domain) bool
}
@@ -222,137 +185,23 @@ func NewAllAttrsMatcher(attrs string) AttributeMatcher {
return m
}
var errInvalidUTF8 = errors.New("string field contains invalid UTF-8")
func loadSiteWithAttrs(file, code, attrs string) ([]*Domain, error) {
domains, err := loadSite(file, code)
if err != nil {
return nil, err
}
type siteDecoder struct {
want []string
has []bool
fn func(Domain_Type, []byte)
}
matcher := NewAllAttrsMatcher(attrs)
if matcher == nil {
return domains, nil
}
func newSiteDecoder(attrs string, fn func(Domain_Type, []byte)) *siteDecoder {
d := &siteDecoder{fn: fn}
if attrs != "" {
d.want = strings.Split(attrs, "@")
d.has = make([]bool, len(d.want))
filtered := make([]*Domain, 0, len(domains))
for _, d := range domains {
if matcher.Match(d) {
filtered = append(filtered, d)
}
}
return d
}
// decode walks the whole fields at the start of b, a part of an encoded GeoSite (see geodat.proto),
// calls fn for every domain that has all attrs and returns how many bytes it used. A field cut off
// by the end of b is an error unless more is set. It accepts and rejects what proto.Unmarshal does.
func (d *siteDecoder) decode(b []byte, more bool) (int, error) {
used := 0
for used < len(b) {
f, n, err := consumeField(b[used:])
if err == io.ErrUnexpectedEOF && more {
break
}
if err != nil {
return used, err
}
used += n
if f.typ != protowire.BytesType {
continue
}
switch f.num {
case 1: // code
if !utf8.Valid(f.v) {
return used, errInvalidUTF8
}
case 2: // domain
t, value, err := decodeDomain(f.v, d.want, d.has)
if err != nil {
return used, err
}
if !slices.Contains(d.has, false) {
d.fn(t, value)
}
}
}
return used, nil
}
// decodeDomain decodes an encoded Domain and sets has[i] if one of its attributes has the key want[i].
func decodeDomain(b []byte, want []string, has []bool) (t Domain_Type, value []byte, err error) {
clear(has)
for len(b) > 0 {
f, n, err := consumeField(b)
if err != nil {
return 0, nil, err
}
b = b[n:]
switch {
case f.num == 1 && f.typ == protowire.VarintType: // type
t = Domain_Type(f.x)
case f.num == 2 && f.typ == protowire.BytesType: // value
if !utf8.Valid(f.v) {
return 0, nil, errInvalidUTF8
}
value = f.v
case f.num == 3 && f.typ == protowire.BytesType: // attribute
key, err := decodeAttributeKey(f.v)
if err != nil {
return 0, nil, err
}
for i, w := range want {
if string(key) == w {
has[i] = true
}
}
}
}
return t, value, nil
}
// decodeAttributeKey returns the key of an encoded Domain.Attribute.
func decodeAttributeKey(b []byte) ([]byte, error) {
var key []byte
for len(b) > 0 {
f, n, err := consumeField(b)
if err != nil {
return nil, err
}
b = b[n:]
if f.num == 1 && f.typ == protowire.BytesType {
if !utf8.Valid(f.v) {
return nil, errInvalidUTF8
}
key = f.v
}
}
return key, nil
}
type protoField struct {
num protowire.Number
typ protowire.Type
v []byte // payload of a length-delimited field
x uint64 // value of a varint field
}
// consumeField parses the first field of an encoded message and returns it with its length.
func consumeField(b []byte) (protoField, int, error) {
num, typ, n := protowire.ConsumeTag(b)
if n < 0 {
return protoField{}, 0, protowire.ParseError(n)
}
if num > protowire.MaxValidNumber {
return protoField{}, 0, errors.New("invalid field number ", num)
}
f := protoField{num: num, typ: typ}
var m int
switch typ {
case protowire.BytesType:
f.v, m = protowire.ConsumeBytes(b[n:])
case protowire.VarintType:
f.x, m = protowire.ConsumeVarint(b[n:])
default:
m = protowire.ConsumeFieldValue(num, typ, b[n:])
}
if m < 0 {
return protoField{}, 0, protowire.ParseError(m)
}
return f, n + m, nil
return filtered, nil
}
-283
View File
@@ -1,283 +0,0 @@
package geodata
import (
"fmt"
"os"
"path/filepath"
"slices"
"strings"
"testing"
"google.golang.org/protobuf/encoding/protowire"
"google.golang.org/protobuf/proto"
)
type siteEntry struct {
Type Domain_Type
Value string
}
// unmarshalSite is what loadSite used to do: proto.Unmarshal, then keep the domains that have all attrs.
func unmarshalSite(b []byte, attrs string) ([]siteEntry, error) {
var site GeoSite
if err := proto.Unmarshal(b, &site); err != nil {
return nil, err
}
var entries []siteEntry
for _, d := range site.Domain {
ok := true
for _, key := range strings.Split(attrs, "@") {
ok = ok && (attrs == "" || slices.ContainsFunc(d.Attribute, func(a *Domain_Attribute) bool { return a.Key == key }))
}
if ok {
entries = append(entries, siteEntry{d.Type, d.Value})
}
}
return entries, nil
}
func checkDecodeSite(t *testing.T, name string, b []byte, attrs string) {
t.Helper()
want, wantErr := unmarshalSite(b, attrs)
var got []siteEntry
_, err := newSiteDecoder(attrs, func(typ Domain_Type, value []byte) {
got = append(got, siteEntry{typ, string(value)})
}).decode(b, false)
if (err == nil) != (wantErr == nil) {
t.Fatalf("%s@%s: error %v, proto.Unmarshal: %v", name, attrs, err, wantErr)
}
if err == nil && !slices.Equal(got, want) {
t.Fatalf("%s@%s: got %v, want %v", name, attrs, got, want)
}
}
func TestDecodeSiteMatchesUnmarshal(t *testing.T) {
bs, err := os.ReadFile(filepath.Join("..", "..", "resources", DefaultGeoSiteDat))
if err != nil {
t.Fatal(err)
}
for len(bs) > 0 {
num, typ, n := protowire.ConsumeTag(bs)
if n < 0 || num != 1 || typ != protowire.BytesType {
t.Fatal("unexpected GeoSiteList field")
}
entry, m := protowire.ConsumeBytes(bs[n:])
if m < 0 {
t.Fatal(protowire.ParseError(m))
}
bs = bs[n+m:]
var site GeoSite
if err := proto.Unmarshal(entry, &site); err != nil {
t.Fatal(err)
}
queries := []string{"", "none"}
for _, d := range site.Domain {
for _, a := range d.Attribute {
if !slices.Contains(queries, a.Key) {
queries = append(queries, a.Key, a.Key+"@none")
}
}
}
for _, attrs := range queries {
checkDecodeSite(t, site.Code, entry, attrs)
}
}
}
func TestDecodeSiteUnusualEncodings(t *testing.T) {
field := func(num protowire.Number, v []byte) []byte {
return protowire.AppendBytes(protowire.AppendTag(nil, num, protowire.BytesType), v)
}
typ := func(v Domain_Type) []byte {
return protowire.AppendVarint(protowire.AppendTag(nil, 1, protowire.VarintType), uint64(v))
}
value := func(s string) []byte { return field(2, []byte(s)) }
attr := func(keys ...string) []byte {
var b []byte
for _, k := range keys {
b = append(b, field(1, []byte(k))...)
}
return field(3, b)
}
domain := func(fields ...[]byte) []byte { return field(2, slices.Concat(fields...)) }
unknown := protowire.AppendFixed32(protowire.AppendTag(nil, 9, protowire.Fixed32Type), 1)
for name, b := range map[string][]byte{
"unknown field": domain(typ(Domain_Full), unknown, value("example.com")),
"repeated value": domain(value("a.com"), typ(Domain_Full), value("b.com")),
"repeated type": domain(typ(Domain_Full), value("a.com"), typ(Domain_Regex)),
"repeated key": domain(value("a.com"), attr("cn", "ads")),
"type as bytes": domain(field(1, []byte("x")), value("a.com")),
"no value": domain(typ(Domain_Domain), attr("cn")),
"truncated": domain(typ(Domain_Full), value("example.com"))[:10],
"invalid utf8": domain(value("example.\xff")),
"invalid key": domain(value("a.com"), attr("\xff")),
"bad field": protowire.AppendVarint(protowire.AppendTag(nil, protowire.MaxValidNumber+1, protowire.VarintType), 1),
"stray end group": protowire.AppendTag(nil, 5, protowire.EndGroupType),
} {
for _, attrs := range []string{"", "cn", "ads", "cn@ads"} {
checkDecodeSite(t, name, b, attrs)
}
}
}
// TestLoadSiteReadsInPieces covers what real lists never do: an entry far longer than the read
// buffer, with a field longer than the buffer in the middle, and a file cut short.
func TestLoadSiteReadsInPieces(t *testing.T) {
site := &GeoSite{Code: "BIG"}
for i := range 5000 {
d := &Domain{Type: Domain_Domain, Value: strings.Repeat("x", i%40) + ".example.com"}
if i%3 == 0 {
d.Attribute = []*Domain_Attribute{{Key: "cn"}}
}
if i == 2500 {
d = &Domain{Type: Domain_Regex, Value: strings.Repeat("a", 100_000)}
}
site.Domain = append(site.Domain, d)
}
list := &GeoSiteList{Entry: []*GeoSite{{Code: "SMALL", Domain: []*Domain{{Type: Domain_Full, Value: "a.com"}}}, site}}
bs, err := proto.Marshal(list)
if err != nil {
t.Fatal(err)
}
entry, err := proto.Marshal(site)
if err != nil {
t.Fatal(err)
}
dir := t.TempDir()
t.Setenv("xray.location.asset", dir)
write := func(b []byte) {
if err := os.WriteFile(filepath.Join(dir, "big.dat"), b, 0o644); err != nil {
t.Fatal(err)
}
}
for _, attrs := range []string{"", "cn"} {
want, _ := unmarshalSite(entry, attrs)
var got []siteEntry
write(bs)
err := loadSite("big.dat", "BIG", attrs, func(typ Domain_Type, value []byte) {
got = append(got, siteEntry{typ, string(value)})
})
if err != nil || !slices.Equal(got, want) {
t.Fatalf("attrs %q: %d entries, want %d, error %v", attrs, len(got), len(want), err)
}
for _, cut := range []int{30_000, len(bs) - 150_000, len(bs) - 1} {
write(bs[:cut])
if err := loadSite("big.dat", "BIG", attrs, func(Domain_Type, []byte) {}); err == nil {
t.Fatalf("file cut at %d of %d: no error", cut, len(bs))
}
}
}
}
// oneEntryGeoSiteFile wraps an encoded GeoSite as a one-entry GeoSiteList, the file loadSite reads.
func oneEntryGeoSiteFile(entry []byte) []byte {
return protowire.AppendBytes(protowire.AppendTag(nil, 1, protowire.BytesType), entry)
}
// TestLoadSiteWindowedMatchesSingleShot checks that the windowed reader in loadSite (its Peek/Discard
// loop, the more-break when a field is cut by a window edge, the used==0 fallback for a field longer
// than the buffer, and the tail path) reaches exactly the same result as decoding the whole entry at
// once, for a category several 64 KiB windows long, valid and then mutated near a window edge and
// early in the file: same error-or-not, and the same emitted (type, value) sequence when both accept.
func TestLoadSiteWindowedMatchesSingleShot(t *testing.T) {
const window = 64 * 1024
site := &GeoSite{Code: "BIG"}
for i := range 12000 { // ~250 KiB, four windows
d := &Domain{Type: Domain_Domain, Value: fmt.Sprintf("host%d.%s.example.com", i, strings.Repeat("y", i%30))}
if i%3 == 0 {
d.Attribute = []*Domain_Attribute{{Key: "cn"}}
}
site.Domain = append(site.Domain, d)
}
// a field longer than the buffer, straddling the third window, to force the used==0 fallback
site.Domain = slices.Insert(site.Domain, 8000, &Domain{Type: Domain_Regex, Value: strings.Repeat("a", 90_000)})
entry, err := proto.Marshal(site)
if err != nil {
t.Fatal(err)
}
dir := t.TempDir()
t.Setenv("xray.location.asset", dir)
// mutations of the encoded entry: unchanged, a byte flipped at several offsets (early windows and
// either side of a window edge), and truncations at the same places.
type mut struct {
name string
make func([]byte) []byte
}
muts := []mut{{"valid", func(b []byte) []byte { return b }}}
for _, off := range []int{3, 40, 4000, window - 1, window, window + 1, 2*window - 2, 2 * window} {
if off < len(entry) {
off := off
muts = append(muts, mut{fmt.Sprintf("flip@%d", off), func(b []byte) []byte {
c := slices.Clone(b)
c[off] ^= 0xff
return c
}})
muts = append(muts, mut{fmt.Sprintf("cut@%d", off), func(b []byte) []byte { return slices.Clone(b[:off]) }})
}
}
for _, attrs := range []string{"", "cn"} {
for _, m := range muts {
e := m.make(entry)
// single-shot reference: decode the whole entry in one call
var want []siteEntry
_, wantErr := newSiteDecoder(attrs, func(typ Domain_Type, value []byte) {
want = append(want, siteEntry{typ, string(value)})
}).decode(e, false)
// windowed: loadSite reads the file 64 KiB at a time
if err := os.WriteFile(filepath.Join(dir, "w.dat"), oneEntryGeoSiteFile(e), 0o644); err != nil {
t.Fatal(err)
}
var got []siteEntry
gotErr := loadSite("w.dat", "BIG", attrs, func(typ Domain_Type, value []byte) {
got = append(got, siteEntry{typ, string(value)})
})
if (gotErr == nil) != (wantErr == nil) {
t.Fatalf("%s attrs=%q: windowed err %v, single-shot err %v", m.name, attrs, gotErr, wantErr)
}
if gotErr == nil && !slices.Equal(got, want) {
t.Fatalf("%s attrs=%q: windowed got %d entries, single-shot %d", m.name, attrs, len(got), len(want))
}
}
}
}
// TestLoadSiteLongCode covers a geosite entry whose code is longer than the 64 KiB read buffer. seek
// must skip it (find compares a single length byte, so it never matches such a code) and still find a
// later entry, and looking the long code up must fail cleanly, like a missing code, not panic.
func TestLoadSiteLongCode(t *testing.T) {
longCode := strings.Repeat("Z", 70000)
list := &GeoSiteList{Entry: []*GeoSite{
{Code: "FIRST", Domain: []*Domain{{Type: Domain_Full, Value: "first.com"}}},
{Code: longCode, Domain: []*Domain{{Type: Domain_Full, Value: "huge.com"}}},
{Code: "AFTER", Domain: []*Domain{{Type: Domain_Domain, Value: "after.com"}}},
}}
bs, err := proto.Marshal(list)
if err != nil {
t.Fatal(err)
}
dir := t.TempDir()
t.Setenv("xray.location.asset", dir)
if err := os.WriteFile(filepath.Join(dir, "lc.dat"), bs, 0o644); err != nil {
t.Fatal(err)
}
collect := func(code string) ([]siteEntry, error) {
var got []siteEntry
err := loadSite("lc.dat", code, "", func(typ Domain_Type, value []byte) {
got = append(got, siteEntry{typ, string(value)})
})
return got, err
}
if got, err := collect("FIRST"); err != nil || !slices.Equal(got, []siteEntry{{Domain_Full, "first.com"}}) {
t.Fatalf("FIRST: %v %v", got, err)
}
if got, err := collect("AFTER"); err != nil || !slices.Equal(got, []siteEntry{{Domain_Domain, "after.com"}}) {
t.Fatalf("AFTER (past the oversized entry): %v %v", got, err)
}
if _, err := collect(longCode); err == nil {
t.Fatal("oversized code: expected a not-found error, got nil")
}
}
@@ -1,7 +1,6 @@
package strmatcher_test
import (
"regexp"
"strconv"
"testing"
@@ -73,64 +72,6 @@ func BenchmarkSubstrMatcher(b *testing.B) {
})
}
func BenchmarkRegexMatcher(b *testing.B) {
patterns := []string{ // taken from geosite
`(^|\.)91porn\.(best|com|cool|fun|group|party|plus|site|tw|work)$`,
`(^|\.)91porn[0-9]{3}\.me$`,
`(^|\.)apiproxy-device-prod-nlb-.+\.amazonaws\.com$`,
`(^|\.)dualstack\.apiproxy-.+\.amazonaws\.com$`,
`(^|\.)aqdk[0-9]{3}\.com$`,
`(^|\.)bilibili3(0[1-9]|1[0-2])\.xyz$`,
`(^|\.)byyum([3589]|2[235689]|3[34]|4[1-9]|5[1-79]|6[0134679])?\.com$`,
`(^|\.)fiftymvapi\..+$`,
`(^|\.)gossipfuli[0-9]{3,4}\.xyz$`,
`(^|\.)kpkuang\.(bond|fun|info|one|us)$`,
`(^|\.)rule34\.(asia|us|world|xxx|xyz)$`,
`(^|\.)[a-z][1-9][0-9][a-z]\.com$`,
`.+\.awsdns-[0-9][0-9]\.(co\.uk|com|net|org)$`,
`.+\.dkr\.ecr\.[^\.]+\.amazonaws\.com$`,
`^(.+\.)*zh\.okaapps\.com$`,
`^cdn\d-epicgames-\d+\.file\.myqcloud\.com$`,
`^chatgpt-async-webps-prod-\S+-\d+\.webpubsub\.azure\.com$`,
`^r+[0-9]+(---|\.)sn-(2x3|ni5|j5o)\w{5}\.googlevideo\.com$`,
`^speed\.(coe|open)\.ad\.[a-z]{2,6}\.prod\.hosts\.ooklaserver\.net$`,
`javdb\d+\.com$`,
}
domains := []string{
"www.google.com", "rr3---sn-4g5edndy.googlevideo.com", "r1---sn-2x3abcde.googlevideo.com", "i.ytimg.com",
"graph.facebook.com", "api.twitter.com", "www.baidu.com", "github.com", "objects.githubusercontent.com",
"login.microsoftonline.com", "e1234.dscb.akamaiedge.net", "d1a2b3c4d5e6f7.cloudfront.net",
"s3.us-east-1.amazonaws.com", "123456789012.dkr.ecr.us-east-1.amazonaws.com", "www.wikipedia.org",
"discord.com", "telegram.org", "store.steampowered.com", "www.91porn.com", "ns-1234.awsdns-12.org",
}
bench := func(b *testing.B, ctor func(pattern string) func(string) bool) {
var matchers []func(string) bool
for _, p := range patterns {
matchers = append(matchers, ctor(p))
}
b.ResetTimer()
for i := 0; i < b.N; i++ {
for _, d := range domains {
for _, match := range matchers {
_ = match(d)
}
}
}
}
b.Run("regexp", func(b *testing.B) {
bench(b, func(pattern string) func(string) bool {
return regexp.MustCompile(pattern).MatchString
})
})
b.Run("prefilter", func(b *testing.B) {
bench(b, func(pattern string) func(string) bool {
m, err := Regex.New(pattern)
common.Must(err)
return m.Match
})
})
}
// Utility functions for benchmark
func benchmarkMatcherType(b *testing.B, t Type, ctor func() MatcherGroup) {
+12 -8
View File
@@ -52,9 +52,7 @@ func (g *MphIndexMatcher) Add(matcher Matcher) uint32 {
func (g *MphIndexMatcher) Build() error {
if g.mph != nil {
runtime.GC() // peak mem
if err := g.mph.Build(); err != nil {
return err
}
g.mph.Build()
}
runtime.GC() // peak mem
if g.ac != nil {
@@ -66,17 +64,23 @@ func (g *MphIndexMatcher) Build() error {
// Match implements IndexMatcher.Match.
func (g *MphIndexMatcher) Match(input string) []uint32 {
var result []uint32
result := make([][]uint32, 0, 5)
if g.mph != nil {
result = g.mph.Match(input) // a new slice, returned without another copy
if matches := g.mph.Match(input); len(matches) > 0 {
result = append(result, matches)
}
}
if g.ac != nil {
result = append(result, g.ac.Match(input)...)
if matches := g.ac.Match(input); len(matches) > 0 {
result = append(result, matches)
}
}
if g.regex != nil {
result = append(result, g.regex.Match(input)...)
if matches := g.regex.Match(input); len(matches) > 0 {
result = append(result, matches)
}
}
return result
return CompositeMatches(result)
}
// MatchAny implements IndexMatcher.MatchAny.
@@ -78,10 +78,6 @@ func TestMphIndexMatcher(t *testing.T) {
Input: "example.com",
Output: []uint32{10, 4},
},
{
Input: "apis.org",
Output: []uint32{2, 6},
},
}
matcherGroup := NewMphIndexMatcher()
for _, rule := range rules {
@@ -91,13 +87,8 @@ func TestMphIndexMatcher(t *testing.T) {
}
matcherGroup.Build()
for _, test := range cases {
m := matcherGroup.Match(test.Input)
if !reflect.DeepEqual(m, test.Output) {
t.Error("unexpected output: ", m, " for test case ", test)
}
clear(m) // the caller owns the result, so this must not change the next one
if m := matcherGroup.Match(test.Input); !reflect.DeepEqual(m, test.Output) {
t.Error("unexpected output after clearing the previous one: ", m, " for test case ", test)
t.Error("unexpected output: ", m, " for test case ", test)
}
}
}
+136 -378
View File
@@ -1,440 +1,198 @@
package strmatcher
import (
"bytes"
"cmp"
"encoding/binary"
"errors"
"math"
"slices"
"math/bits"
"runtime"
"sort"
"strings"
"unsafe"
)
// Flags of a level1 slot, stored above the record offset.
const (
mphDomain = 1 << 31 // matches the pattern and its subdomains
mphFull = 1 << 30 // matches the pattern only
mphParent = 1 << 29 // matches subdomains only, from a pattern with a leading dot
mphOffMask = mphParent - 1
)
// PrimeRK is the prime base used in Rabin-Karp algorithm.
const PrimeRK = 16777619
// Kinds of an added pattern, indexes of mphKinds.
const (
mphKindFull = iota
mphKindParent
mphKindDomain
)
// mphKinds are the slot flags in the order Match reports their values.
var mphKinds = [...]uint32{mphFull, mphParent, mphDomain}
// mphMultipliers are odd multipliers for the suffix hash. Build moves to the next one if two patterns collide.
var mphMultipliers = [...]uint64{0x9e3779b97f4a7c15, 0xc2b2ae3d27d4eb4f, 0x165667b19e3779f9, 0x27d4eb2f165667c5}
var (
errMphCollision = errors.New("strmatcher: suffix hash collision in MphMatcherGroup")
errMphBuilt = errors.New("strmatcher: MphMatcherGroup is already built")
)
type mphEntry struct {
off uint32 // pattern start in buf
value uint32
n uint32 // pattern length
kind uint8
// RollingHash calculates the rolling murmurHash of given string based on a provided suffix hash.
func RollingHash(hash uint32, input string) uint32 {
for i := len(input) - 1; i >= 0; i-- {
hash = hash*PrimeRK + uint32(input[i])
}
return hash
}
// MphMatcherGroup is an implementation of MatcherGroup for Full and Domain matchers.
// Each distinct pattern is stored once as a record in arena: its length (255 means a uvarint length follows),
// its bytes and, if the group holds more than one distinct value, its values. A minimal perfect hash table
// built with hash, displace and compress (http://cmph.sourceforge.net/papers/esa09.pdf) maps a pattern to its
// record. Patterns are hashed from the right, so one pass over the input hashes all its parent domains.
type MphMatcherGroup struct {
arena string
level0 []uint16 // bucket -> seed
level1 []uint32 // slot -> flags | record offset
fp []uint8 // slot -> low byte of its pattern's hash, rejects most misses without reading arena
n0, n1 uint32
mul uint64 // multiplier of the suffix hash
single uint32 // the only value if !multi
multi bool
// MemHash is the hash function used by go map, it utilizes available hardware instructions(behaves
// as aeshash if aes instruction is available).
// With different seed, each MemHash<seed> performs as distinct hash functions.
func MemHash(seed uint32, input string) uint32 {
return uint32(strhash(unsafe.Pointer(&input), uintptr(seed))) // nosemgrep
}
buf []byte // build only, patterns in Add order
entries []mphEntry
const (
mphMatchTypeCount = 2 // Full and Domain
)
type mphRuleInfo struct {
rollingHash uint32
matchers [mphMatchTypeCount][]uint32
}
// MphMatcherGroup is an implementation of MatcherGroup.
// It implements Rabin-Karp algorithm and minimal perfect hash table for Full and Domain matcher.
type MphMatcherGroup struct {
rules []string // RuleIdx -> pattern string, index 0 reserved for failed lookup
values [][]uint32 // RuleIdx -> registered matcher values for the pattern (Full Matcher takes precedence)
level0 []uint32 // RollingHash & Mask -> seed for Memhash
level0Mask uint32 // Mask restricting RollingHash to 0 ~ len(level0)
level1 []uint32 // Memhash<seed> & Mask -> stored index for rules
level1Mask uint32 // Mask for restricting Memhash<seed> to 0 ~ len(level1)
ruleInfos *map[string]mphRuleInfo
}
func NewMphMatcherGroup() *MphMatcherGroup {
return new(MphMatcherGroup)
return &MphMatcherGroup{
rules: []string{""},
values: [][]uint32{nil},
level0: nil,
level0Mask: 0,
level1: nil,
level1Mask: 0,
ruleInfos: &map[string]mphRuleInfo{}, // Only used for building, destroyed after build complete
}
}
// AddFullMatcher implements MatcherGroupForFull.
func (g *MphMatcherGroup) AddFullMatcher(matcher FullMatcher, value uint32) {
g.add(matcher.Pattern(), mphKindFull, value)
pattern := strings.ToLower(matcher.Pattern())
g.addPattern(0, "", pattern, matcher.Type(), value)
}
// AddDomainMatcher implements MatcherGroupForDomain.
func (g *MphMatcherGroup) AddDomainMatcher(matcher DomainMatcher, value uint32) {
g.add(matcher.Pattern(), mphKindDomain, value)
pattern := strings.ToLower(matcher.Pattern())
hash := g.addPattern(0, "", pattern, matcher.Type(), value) // For full domain match
g.addPattern(hash, pattern, ".", matcher.Type(), value) // For partial domain match
}
func (g *MphMatcherGroup) add(pattern string, kind uint8, value uint32) {
if g.arena != "" {
panic(errMphBuilt)
}
pattern = strings.ToLower(pattern)
off := uint32(len(g.buf))
g.buf = append(g.buf, pattern...)
g.entries = append(g.entries, mphEntry{off: off, value: value, n: uint32(len(pattern)), kind: kind})
if len(pattern) > 0 && pattern[0] == '.' {
// ".x" has always matched "*.x" as well, so it also gets a parent-only record for "x"
g.entries = append(g.entries, mphEntry{off: off + 1, value: value, n: uint32(len(pattern) - 1), kind: mphKindParent})
func (g *MphMatcherGroup) addPattern(suffixHash uint32, suffixPattern string, pattern string, matcherType Type, value uint32) uint32 {
fullPattern := pattern + suffixPattern
info, found := (*g.ruleInfos)[fullPattern]
if !found {
info = mphRuleInfo{rollingHash: RollingHash(suffixHash, pattern)}
g.rules = append(g.rules, fullPattern)
g.values = append(g.values, nil)
}
info.matchers[matcherType] = append(info.matchers[matcherType], value)
(*g.ruleInfos)[fullPattern] = info
return info.rollingHash
}
func (g *MphMatcherGroup) key(i uint32) []byte {
e := &g.entries[i]
return g.buf[e.off : e.off+e.n]
}
// Build builds the hash table. It must be called once, after the last Add.
// Build builds a minimal perfect hash table for insert rules.
// Algorithm used: Hash, displace, and compress. See http://cmph.sourceforge.net/papers/esa09.pdf
func (g *MphMatcherGroup) Build() error {
if g.arena != "" {
return errMphBuilt
}
if uint64(len(g.buf)) > math.MaxUint32 {
return errors.New("too many rules for MphMatcherGroup")
}
recs := g.writeRecords()
if len(g.arena) > mphOffMask {
return errors.New("too many rules for MphMatcherGroup")
}
hashes := make([]uint64, len(recs))
for _, mul := range mphMultipliers {
for i, rec := range recs {
hashes[i] = mphMix(mphHash(mul, g.recKey(rec)))
}
g.mul = mul
if err := g.place(recs, hashes); err != errMphCollision {
return err
}
}
return errMphCollision
}
ruleCount := len(*g.ruleInfos)
g.level0 = make([]uint32, nextPow2(ruleCount/4))
g.level0Mask = uint32(len(g.level0) - 1)
g.level1 = make([]uint32, nextPow2(ruleCount))
g.level1Mask = uint32(len(g.level1) - 1)
// writeRecords writes one record per distinct pattern to arena and returns flags | offset of each.
func (g *MphMatcherGroup) writeRecords() []uint32 {
g.multi = false
if len(g.entries) > 0 {
g.single = g.entries[0].value
for _, e := range g.entries {
if e.value != g.single {
g.multi = true
break
}
}
// Create buckets based on all rule's rolling hash
buckets := make([][]uint32, len(g.level0))
for ruleIdx := 1; ruleIdx < len(g.rules); ruleIdx++ { // Traverse rules starting from index 1 (0 reserved for failed lookup)
ruleInfo := (*g.ruleInfos)[g.rules[ruleIdx]]
bucketIdx := ruleInfo.rollingHash & g.level0Mask
buckets[bucketIdx] = append(buckets[bucketIdx], uint32(ruleIdx))
g.values[ruleIdx] = append(ruleInfo.matchers[Full], ruleInfo.matchers[Domain]...) // nolint:gocritic
}
// Equal patterns become neighbours in Add order, so their values keep their priority
order := make([]uint32, len(g.entries))
for i := range order {
order[i] = uint32(i)
}
slices.SortFunc(order, func(a, b uint32) int {
return cmp.Or(bytes.Compare(g.key(a), g.key(b)), cmp.Compare(a, b))
})
g.ruleInfos = nil // Set ruleInfos nil to release memory
runtime.GC() // peak mem
size := len(g.buf) + len(g.entries) + 2
if g.multi {
size += 3 * len(g.entries)
// Sort buckets in descending order with respect to each bucket's size
bucketIdxs := make([]int, len(buckets))
for bucketIdx := range buckets {
bucketIdxs[bucketIdx] = bucketIdx
}
arena := make([]byte, 0, size)
recs := make([]uint32, 0, len(order))
var vals [len(mphKinds)][]uint32
for i := 0; i < len(order); {
k := g.key(order[i])
for t := range vals {
vals[t] = vals[t][:0]
}
for ; i < len(order) && bytes.Equal(g.key(order[i]), k); i++ {
e := &g.entries[order[i]]
if !slices.Contains(vals[e.kind], e.value) {
vals[e.kind] = append(vals[e.kind], e.value)
}
}
rec := uint32(len(arena))
if len(k) < 255 {
arena = append(arena, byte(len(k)))
} else {
arena = binary.AppendUvarint(append(arena, 255), uint64(len(k)))
}
arena = append(arena, k...)
for t, v := range vals {
if len(v) == 0 {
continue
}
rec |= mphKinds[t]
if g.multi {
arena = binary.AppendUvarint(arena, uint64(len(v)))
for _, x := range v {
arena = binary.AppendUvarint(arena, uint64(x))
sort.Slice(bucketIdxs, func(i, j int) bool { return len(buckets[bucketIdxs[i]]) > len(buckets[bucketIdxs[j]]) })
// Exercise Hash, Displace, and Compress algorithm to construct minimal perfect hash table
occupied := make([]bool, len(g.level1)) // Whether a second-level hash has been already used
hashedBucket := make([]uint32, 0, 4) // Second-level hashes for each rule in a specific bucket
for _, bucketIdx := range bucketIdxs {
bucket := buckets[bucketIdx]
hashedBucket = hashedBucket[:0]
seed := uint32(0)
for len(hashedBucket) != len(bucket) {
for _, ruleIdx := range bucket {
memHash := MemHash(seed, g.rules[ruleIdx]) & g.level1Mask
if occupied[memHash] { // Collision occurred with this seed
for _, hash := range hashedBucket { // Revert all values in this hashed bucket
occupied[hash] = false
g.level1[hash] = 0
}
hashedBucket = hashedBucket[:0]
seed++ // Try next seed
break
}
occupied[memHash] = true
g.level1[memHash] = ruleIdx // The final value in the hash table
hashedBucket = append(hashedBucket, memHash)
}
}
recs = append(recs, rec)
}
// Lookups may point one byte past a pattern, and an empty group needs a record at offset 0 for empty slots
arena = append(arena, 0)
if len(recs) == 0 {
arena = append(arena, 0)
}
g.buf, g.entries = nil, nil
if cap(arena)-len(arena) > len(arena)/32 {
arena = slices.Clone(arena)
}
g.arena = unsafe.String(unsafe.SliceData(arena), len(arena)) // arena is not written after this
return recs
}
// place fills level0, level1 and fp: records are bucketed by hash, and each bucket, largest first, gets
// the first seed that puts all its records in free slots.
func (g *MphMatcherGroup) place(recs []uint32, hashes []uint64) error {
r := len(recs)
n0, n1 := max(1, r/3), max(1, r+r/99)
g.n0, g.n1 = uint32(n0), uint32(n1)
g.level0 = make([]uint16, n0)
g.level1 = make([]uint32, n1)
g.fp = make([]uint8, n1)
start := make([]uint32, n0+1)
for _, h := range hashes {
start[g.bucket(h)+1]++
}
for b := range n0 {
start[b+1] += start[b]
}
members := make([]uint32, r)
fill := slices.Clone(start[:n0])
for i, h := range hashes {
b := g.bucket(h)
members[fill[b]] = uint32(i)
fill[b]++
}
fill = nil
buckets := make([]uint32, n0)
for b := range buckets {
buckets[b] = uint32(b)
}
slices.SortStableFunc(buckets, func(a, b uint32) int {
return cmp.Compare(start[b+1]-start[b], start[a+1]-start[a])
})
occupied := make([]uint64, (n1+63)/64)
var slots []uint32
next:
for _, b := range buckets {
m := members[start[b]:start[b+1]]
if len(m) == 0 {
break
}
for i := range m {
for j := range i {
if hashes[m[i]] == hashes[m[j]] {
return errMphCollision // no seed can separate them
}
}
}
search:
for seed := range math.MaxUint16 + 1 {
slots = slots[:0]
for _, ri := range m {
s := g.slot(hashes[ri], uint16(seed))
if occupied[s/64]&(1<<(s%64)) != 0 || slices.Contains(slots, s) {
continue search
}
slots = append(slots, s)
}
for k, ri := range m {
s := slots[k]
occupied[s/64] |= 1 << (s % 64)
g.level1[s] = recs[ri]
g.fp[s] = uint8(hashes[ri])
}
g.level0[b] = uint16(seed)
continue next
}
return errors.New("strmatcher: no seed found for a bucket in MphMatcherGroup")
g.level0[bucketIdx] = seed // Displacement value for this bucket
}
return nil
}
// mphHash is the suffix hash of s, taken from the right: the hash of s[i:] is the state after reading s[i].
func mphHash(mul uint64, s string) uint64 {
h := uint64(0)
for i := len(s) - 1; i >= 0; i-- {
h = h*mul + uint64(s[i])
}
return h
}
// mphMix spreads the weak low bits of a suffix hash.
func mphMix(h uint64) uint64 {
h ^= h >> 32
h *= 0xd6e8feb86659fd93
return h ^ h>>32
}
func (g *MphMatcherGroup) bucket(f uint64) uint32 {
return uint32(((f >> 32) * uint64(g.n0)) >> 32)
}
func (g *MphMatcherGroup) slot(f uint64, seed uint16) uint32 {
x := ((f ^ uint64(seed)*0x9e3779b97f4a7c15) * 0xc4ceb9fe1a85ec53) >> 32
return uint32((x * uint64(g.n1)) >> 32)
}
func (g *MphMatcherGroup) uvarint(p uint32) (x, next uint32) {
for shift := 0; ; shift += 7 {
c := g.arena[p]
p++
x |= uint32(c&0x7f) << shift
if c < 0x80 {
return x, p
}
}
}
// recSpan returns where the pattern of the record at off starts and how long it is.
func (g *MphMatcherGroup) recSpan(off uint32) (p, n uint32) {
n, p = uint32(g.arena[off]), off+1
if n == 255 {
n, p = g.uvarint(p)
}
return p, n
}
func (g *MphMatcherGroup) recKey(rec uint32) string {
p, n := g.recSpan(rec & mphOffMask)
return g.arena[p : p+n]
}
// lookup returns the level1 entry of s, or 0 if s is not a pattern. h is the suffix hash of s.
func (g *MphMatcherGroup) lookup(h uint64, s string) uint32 {
f := mphMix(h)
// bucket < n0 == len(level0) and slot < n1 == len(level1) == len(fp), skip the bounds checks
seed := *(*uint16)(unsafe.Add(unsafe.Pointer(unsafe.SliceData(g.level0)), uintptr(g.bucket(f))*2))
slot := uintptr(g.slot(f, seed))
if *(*uint8)(unsafe.Add(unsafe.Pointer(unsafe.SliceData(g.fp)), slot)) != uint8(f) {
return 0
}
e := *(*uint32)(unsafe.Add(unsafe.Pointer(unsafe.SliceData(g.level1)), slot*4))
if len(s) < 255 {
// A record whose length byte is len(s) has len(s) pattern bytes after it
p := unsafe.Add(unsafe.Pointer(unsafe.StringData(g.arena)), e&mphOffMask)
if int(*(*byte)(p)) == len(s) && unsafe.String((*byte)(unsafe.Add(p, 1)), len(s)) == s {
return e
}
return 0
}
if g.recKey(e) == s {
return e
// Lookup searches for input in minimal perfect hash table and returns its index. 0 indicates not found.
func (g *MphMatcherGroup) Lookup(rollingHash uint32, input string) uint32 {
i0 := rollingHash & g.level0Mask
seed := g.level0[i0]
i1 := MemHash(seed, input) & g.level1Mask
if n := g.level1[i1]; g.rules[n] == input {
return n
}
return 0
}
// appendValues appends the values of record e for the flags in want, in mphKinds order.
func (g *MphMatcherGroup) appendValues(dst []uint32, e, want uint32) []uint32 {
if !g.multi {
for _, flag := range mphKinds {
if e&want&flag != 0 {
dst = append(dst, g.single)
}
}
return dst
}
if e&want == 0 {
return dst
}
p, n := g.recSpan(e & mphOffMask)
p += n
for _, flag := range mphKinds {
if e&flag == 0 {
continue
}
var count, v uint32
for count, p = g.uvarint(p); count > 0; count-- {
v, p = g.uvarint(p)
if want&flag != 0 {
dst = append(dst, v)
}
}
}
return dst
}
// Match implements MatcherGroup.Match. Values of an exact match come first (Full, then Domain), then those of
// the parent domains, nearest first.
// Match implements MatcherGroup.Match.
func (g *MphMatcherGroup) Match(input string) []uint32 {
var stack [8]uint32
parents := stack[:0] // TLD side first
h, mul := uint64(0), g.mul
matches := make([][]uint32, 0, 5)
hash := uint32(0)
for i := len(input) - 1; i >= 0; i-- {
hash = hash*PrimeRK + uint32(input[i])
if input[i] == '.' {
if e := g.lookup(h, input[i+1:]); e&(mphDomain|mphParent) != 0 {
parents = append(parents, e)
if mphIdx := g.Lookup(hash, input[i:]); mphIdx != 0 {
matches = append(matches, g.values[mphIdx])
}
}
h = h*mul + uint64(input[i])
}
exact := g.lookup(h, input)
if exact&(mphFull|mphDomain) == 0 && len(parents) == 0 {
return nil
if mphIdx := g.Lookup(hash, input); mphIdx != 0 {
matches = append(matches, g.values[mphIdx])
}
result := g.appendValues(make([]uint32, 0, len(parents)+1), exact, mphFull|mphDomain)
for k := len(parents) - 1; k >= 0; k-- {
result = g.appendValues(result, parents[k], mphParent|mphDomain)
}
return result
return CompositeMatchesReverse(matches)
}
// MatchAny implements MatcherGroup.MatchAny.
func (g *MphMatcherGroup) MatchAny(input string) bool {
h, mul := uint64(0), g.mul
for i := len(input) - 1; i >= 0; i-- {
if input[i] == '.' && g.lookup(h, input[i+1:])&(mphDomain|mphParent) != 0 {
return true
}
h = h*mul + uint64(input[i])
}
return g.lookup(h, input)&(mphFull|mphDomain) != 0
}
// mphSuffix is the suffix hash of input[off:], a parent domain of the input.
type mphSuffix struct {
h uint64
off int
}
// mphSuffixes appends the suffix hashes of the parent domains of input to dst, TLD side first, and returns them
// with the hash of input itself: what MatchAny computes, computed once for several groups.
func mphSuffixes(dst []mphSuffix, mul uint64, input string) ([]mphSuffix, uint64) {
h := uint64(0)
hash := uint32(0)
for i := len(input) - 1; i >= 0; i-- {
hash = hash*PrimeRK + uint32(input[i])
if input[i] == '.' {
dst = append(dst, mphSuffix{h, i + 1})
if g.Lookup(hash, input[i:]) != 0 {
return true
}
}
h = h*mul + uint64(input[i])
}
return dst, h
return g.Lookup(hash, input) != 0
}
// matchAnyHashed is MatchAny with parents and h from mphSuffixes(_, mul, input).
func (g *MphMatcherGroup) matchAnyHashed(input string, parents []mphSuffix, h, mul uint64) bool {
if g.mul != mul {
return g.MatchAny(input) // built with a later multiplier after a collision
func nextPow2(v int) int {
if v <= 1 {
return 1
}
for _, p := range parents {
if g.lookup(p.h, input[p.off:])&(mphDomain|mphParent) != 0 {
return true
}
}
return g.lookup(h, input)&(mphFull|mphDomain) != 0
const MaxUInt = ^uint(0)
n := (MaxUInt >> bits.LeadingZeros(uint(v))) + 1
return int(n)
}
//go:noescape
//go:linkname strhash runtime.strhash
func strhash(p unsafe.Pointer, h uintptr) uintptr
@@ -1,108 +0,0 @@
package strmatcher
import (
"slices"
"testing"
)
func TestMphMatcherGroupHashCollision(t *testing.T) {
saved := mphMultipliers
defer func() { mphMultipliers = saved }()
mphMultipliers[0] = 1 // anagrams collide
g := NewMphMatcherGroup()
g.AddFullMatcher(FullMatcher("ab.com"), 1)
g.AddDomainMatcher(DomainMatcher("ba.com"), 2)
g.AddDomainMatcher(DomainMatcher("com"), 3)
if err := g.Build(); err != nil {
t.Fatal(err)
}
if g.mul != saved[1] {
t.Errorf("multiplier %#x, want the second one %#x", g.mul, saved[1])
}
for input, want := range map[string][]uint32{"ab.com": {1, 3}, "x.ba.com": {2, 3}, "x.ab.com": {3}, "ba.com": {2, 3}} {
if m := g.Match(input); !slices.Equal(m, want) {
t.Errorf("Match(%q) = %v, want %v", input, m, want)
}
}
// Thue-Morse strings of 2048 bytes and their complements collide for every odd multiplier
mphMultipliers = saved
a, b := make([]byte, 2048), make([]byte, 2048)
for i := range a {
a[i], b[i] = "ab"[bitsOnes(i)%2], "ba"[bitsOnes(i)%2]
}
g = NewMphMatcherGroup()
g.AddFullMatcher(FullMatcher(a), 1)
g.AddFullMatcher(FullMatcher(b), 1)
if err := g.Build(); err != errMphCollision {
t.Errorf("Build() = %v, want %v", err, errMphCollision)
}
}
func bitsOnes(i int) int {
n := 0
for ; i > 0; i &= i - 1 {
n++
}
return n
}
func TestMphValueMatcherCombiner(t *testing.T) {
build := func(matchers ...Matcher) *MphValueMatcher {
m := NewMphValueMatcher()
for _, x := range matchers {
m.Add(x, 0)
}
if err := m.Build(); err != nil {
t.Fatal(err)
}
return m
}
regex, err := Regex.New(`^a\d+\.net$`)
if err != nil {
t.Fatal(err)
}
saved := mphMultipliers
t.Cleanup(func() { mphMultipliers = saved })
mphMultipliers[0] = 1 // anagrams collide, so this one falls back to its own hash pass
collided := build(FullMatcher("ab.com"), DomainMatcher("ba.com"))
mphMultipliers = saved
if collided.mph.mul == mphMultipliers[0] {
t.Fatal("collided matcher uses the first multiplier")
}
matchers := []*MphValueMatcher{
build(DomainMatcher("example.com"), FullMatcher("full.org"), DomainMatcher(".dot.io")),
collided,
build(regex, SubstrMatcher("keyword")),
build(),
build(DomainMatcher("com"), DomainMatcher("a.b.c.d.e.f.g.h.i.j.k.l.m.n.o.p.q.r.s")),
}
var s MphValueMatcherCombiner
for i, m := range matchers {
s.Add(m, uint32(10+i))
}
inputs := []string{
"", ".", "..", "com", "example.com", "www.example.com", "xexample.com", "example.com.", "full.org", "x.full.org",
"dot.io", "x.dot.io", ".dot.io", "ab.com", "x.ab.com", "ba.com", "x.ba.com", "a12.net", "a12.net.x", "my-keyword.org",
"a.b.c.d.e.f.g.h.i.j.k.l.m.n.o.p.q.r.s", "0.a.b.c.d.e.f.g.h.i.j.k.l.m.n.o.p.q.r.s", "b.c.d.e.f.g.h.i.j.k.l.m.n.o.p.q.r.s",
"x.y.z.1.2.3.4.5.6.7.8.9.10.11.12.13.14.15.16.17.ab.com", "x.y.z.1.2.3.4.5.6.7.8.9.10.11.12.13.14.15.16.17.org",
}
for _, input := range inputs {
var want []uint32
for i, m := range matchers {
if m.MatchAny(input) {
want = append(want, uint32(10+i))
}
}
if got := s.Match(input); !slices.Equal(got, want) {
t.Errorf("Match(%q) = %v, want %v", input, got, want)
}
if got := s.MatchAny(input); got != (len(want) > 0) {
t.Errorf("MatchAny(%q) = %v", input, got)
}
}
if n := testing.AllocsPerRun(100, func() { s.MatchAny("www.a.b.c.example.org") }); n != 0 {
t.Errorf("MatchAny allocates %v times", n)
}
}
@@ -1,10 +1,7 @@
package strmatcher_test
import (
"math/rand"
"reflect"
"slices"
"strings"
"testing"
"github.com/xtls/xray-core/common"
@@ -279,142 +276,3 @@ func TestEmptyMphMatcherGroup(t *testing.T) {
t.Error("Expect [], but ", r)
}
}
func TestMphMatcherGroupRandom(t *testing.T) {
inputs := []string{""} // All strings over "ab." up to 7 bytes
for i := 0; len(inputs[i]) < 7; i++ {
for _, c := range []string{"a", "b", "."} {
inputs = append(inputs, inputs[i]+c)
}
}
for seed := int64(0); seed < 300; seed++ {
r := rand.New(rand.NewSource(seed))
g := NewMphMatcherGroup()
full, domain := map[string][]uint32{}, map[string][]uint32{} // Stored pattern -> values
for value := uint32(r.Intn(200)); value > 0; value-- {
pattern := make([]byte, r.Intn(8))
for i := range pattern {
pattern[i] = "ab."[r.Intn(3)]
}
if p := string(pattern); r.Intn(2) == 0 {
g.AddFullMatcher(FullMatcher(p), value)
full[p] = append(full[p], value)
} else {
g.AddDomainMatcher(DomainMatcher(p), value)
domain[p] = append(domain[p], value)
domain["."+p] = append(domain["."+p], value)
}
}
common.Must(g.Build())
for _, input := range inputs {
keys := []string{input} // Whole input first, then "." suffixes from longest to shortest
for i := range len(input) {
if input[i] == '.' {
keys = append(keys, input[i:])
}
}
var want []uint32
for _, k := range keys {
want = append(append(want, full[k]...), domain[k]...)
}
// Compared as sets: Match reports a value once per matching pattern, and orders them differently
// from want for patterns and inputs with a leading dot
m := g.Match(input)
if !slices.Equal(sortedSet(m), sortedSet(want)) {
t.Fatalf("seed %d: Match(%q) = %v, want %v", seed, input, m, want)
}
if m := g.MatchAny(input); m != (len(want) > 0) {
t.Fatalf("seed %d: MatchAny(%q) = %v", seed, input, m)
}
}
}
}
func TestMphMatcherGroupAppend(t *testing.T) {
g := NewMphMatcherGroup()
g.AddFullMatcher(FullMatcher("a.com"), 1)
g.AddFullMatcher(FullMatcher("b.com"), 2)
g.Build()
if m := append(g.Match("a.com"), 3); !slices.Equal(m, []uint32{1, 3}) {
t.Error("expect [1 3], but ", m)
}
if m := g.Match("b.com"); !slices.Equal(m, []uint32{2}) {
t.Error("expect [2], but ", m)
}
}
func sortedSet(v []uint32) []uint32 {
v = slices.Clone(v)
slices.Sort(v)
return slices.Compact(v)
}
func TestMphMatcherGroupLongPattern(t *testing.T) {
long := strings.Repeat("a", 300) + ".com"
for _, values := range [][4]uint32{{1, 2, 3, 4}, {7, 7, 7, 7}} {
g := NewMphMatcherGroup()
g.AddDomainMatcher(DomainMatcher(long), values[0])
g.AddFullMatcher(FullMatcher("x."+long), values[1])
g.AddFullMatcher(FullMatcher(long[:255]), values[2]) // the shortest pattern stored with a long length
g.AddFullMatcher(FullMatcher(long[:254]), values[3])
common.Must(g.Build())
cases := []struct {
input string
want []uint32
}{
{long, []uint32{values[0]}},
{"www." + long, []uint32{values[0]}},
{"x." + long, []uint32{values[1], values[0]}},
{long[1:], nil},
{"a" + long, nil},
{long[:255], []uint32{values[2]}},
{long[:254], []uint32{values[3]}},
{long[:256], nil},
{long[:253], nil},
}
for _, c := range cases {
if m := g.Match(c.input); !slices.Equal(m, c.want) {
t.Errorf("Match(%d bytes) = %v, want %v", len(c.input), m, c.want)
}
if m := g.MatchAny(c.input); m != (c.want != nil) {
t.Errorf("MatchAny(%d bytes) = %v", len(c.input), m)
}
}
}
// A pattern longer than 65535 bytes builds and matches: a record's length is a uvarint,
// so the only cap was the build-time length field, now widened to uint32.
huge := strings.Repeat("a", 70000)
g := NewMphMatcherGroup()
g.AddFullMatcher(FullMatcher(strings.Repeat("a", 65535)), 1)
g.AddDomainMatcher(DomainMatcher(huge+".com"), 2)
g.AddFullMatcher(FullMatcher("a.com"), 3)
common.Must(g.Build())
if !g.MatchAny(strings.Repeat("a", 65535)) || g.MatchAny(strings.Repeat("a", 65534)) {
t.Error("wrong answer for a 65535-byte pattern")
}
if m := g.Match(huge + ".com"); !slices.Equal(m, []uint32{2}) {
t.Errorf("Match(%d-byte input) = %v, want [2]", len(huge)+4, m)
}
if m := g.Match("x." + huge + ".com"); !slices.Equal(m, []uint32{2}) {
t.Errorf("Match(subdomain of a %d-byte pattern) = %v, want [2]", len(huge)+4, m)
}
if g.MatchAny(huge) { // the 70000-byte label on its own is not a rule
t.Error("unexpected match for the bare 70000-byte label")
}
}
func TestMphMatcherGroupBuildOnce(t *testing.T) {
g := NewMphMatcherGroup()
g.AddFullMatcher(FullMatcher("a.com"), 1)
common.Must(g.Build())
if err := g.Build(); err == nil || !g.MatchAny("a.com") {
t.Errorf("second Build() = %v, MatchAny(a.com) = %v", err, g.MatchAny("a.com"))
}
defer func() {
if recover() == nil {
t.Error("Add after Build did not panic")
}
}()
g.AddDomainMatcher(DomainMatcher("b.com"), 2)
}
+11 -281
View File
@@ -2,12 +2,9 @@ package strmatcher
import (
"errors"
"math/bits"
"regexp"
"regexp/syntax"
"slices"
"strings"
"unicode"
"unicode/utf8"
"golang.org/x/net/idna"
@@ -76,274 +73,7 @@ func (m SubstrMatcher) Match(s string) bool {
// RegexMatcher is an implementation of Matcher.
type RegexMatcher struct {
pattern *regexp.Regexp
literals []string // every match contains all of them, longest first
tail []byteSet // tail[i] holds the bytes a matching input can have i bytes before its end
rest *byteSet // the bytes it can have further before, nil if any
}
func newRegexMatcher(pattern string) (Matcher, error) {
regex, err := regexp.Compile(pattern)
if err != nil {
return nil, err
}
m := &RegexMatcher{pattern: regex}
if re, err := syntax.Parse(pattern, syntax.Perl); err == nil { // same flags as regexp.Compile
m.literals = requiredLiterals(re, nil)
slices.SortStableFunc(m.literals, func(a, b string) int { return len(b) - len(a) })
m.tail, m.rest = tailGuard(re)
}
return m, nil
}
// byteSet is a set of bytes. The bytes >= 0x80 share one bit with 0x7f.
type byteSet [4]uint32
func (s *byteSet) add(c byte) { c = min(c, 0x7f); s[c>>5] |= 1 << (c & 31) }
func (s *byteSet) has(c byte) bool { c = min(c, 0x7f); return s[c>>5]&(1<<(c&31)) != 0 }
func (s *byteSet) or(t *byteSet) {
for i := range s {
s[i] |= t[i]
}
}
var allBytes = byteSet{^uint32(0), ^uint32(0), ^uint32(0), ^uint32(0)}
// tailLen is how many positions before the end of the input tailGuard tells apart.
const tailLen = 8
// tailBudget caps how many repetition steps tailGuard walks. Only nested repeats can make the
// walk explode, so only they are charged: a flat pattern, however long, is walked once and keeps
// its guard.
const tailBudget = 100000
// tailWalk is a set of positions in the input, counted in bytes before its end.
type tailWalk struct {
at uint32 // bit i: exactly i bytes before the end, for i < tailLen
far bool // tailLen or more bytes before the end
free bool // not tied to the end of the input yet
}
func (w tailWalk) union(v tailWalk) tailWalk {
return tailWalk{w.at | v.at, w.far || v.far, w.free || v.free}
}
type tailBuilder struct {
tail [tailLen]byteSet
rest byteSet
void bool
work int
}
// tailGuard walks re backwards from the end of the input and collects the bytes an input
// matching re can have at each position before its end. It returns nil, nil when a branch
// of re does not end with $ or when nested repeats push the walk past tailBudget.
func tailGuard(re *syntax.Regexp) ([]byteSet, *byteSet) {
var b tailBuilder
w := b.walk(re, tailWalk{free: true})
b.stop(w)
if b.void {
return nil, nil
}
if w.at != 0 { // a match can start here, so any bytes can come before
for i := bits.TrailingZeros32(w.at); i < tailLen; i++ {
b.tail[i] = allBytes
}
}
if w.at != 0 || w.far {
b.rest = allBytes
}
n := tailLen
for n > 0 && b.tail[n-1] == b.rest {
n--
}
var tail []byteSet
if n > 0 {
tail = slices.Clone(b.tail[:n])
}
if b.rest != allBytes {
rest := b.rest
return tail, &rest
}
return tail, nil
}
// stop ends the paths of w. One that never met $ lets its match be followed by anything.
func (b *tailBuilder) stop(w tailWalk) {
if w.free {
b.void = true
}
}
func (b *tailBuilder) walk(re *syntax.Regexp, w tailWalk) tailWalk {
if w == (tailWalk{}) || b.void {
return w
}
switch re.Op {
case syntax.OpNoMatch:
return tailWalk{}
case syntax.OpLiteral:
for i := len(re.Rune) - 1; i >= 0; i-- {
var set byteSet
set.add(byte(min(re.Rune[i], utf8.RuneSelf)))
if re.Flags&syntax.FoldCase != 0 {
for f := unicode.SimpleFold(re.Rune[i]); f != re.Rune[i]; f = unicode.SimpleFold(f) {
set.add(byte(min(f, utf8.RuneSelf)))
}
}
w = b.step(w, &set)
}
return w
case syntax.OpCharClass:
var set byteSet
for i := 0; i+1 < len(re.Rune); i += 2 {
for r := min(re.Rune[i], utf8.RuneSelf); r <= min(re.Rune[i+1], utf8.RuneSelf); r++ {
set.add(byte(r))
}
}
return b.step(w, &set)
case syntax.OpAnyChar, syntax.OpAnyCharNotNL: // a domain has no \n to reject
return b.step(w, &allBytes)
case syntax.OpBeginText: // nothing comes before
b.stop(w)
return tailWalk{}
case syntax.OpEndText:
out := tailWalk{at: w.at & 1}
if w.free {
out.at = 1
}
return out
case syntax.OpCapture:
return b.walk(re.Sub[0], w)
case syntax.OpConcat:
for i := len(re.Sub) - 1; i >= 0; i-- {
w = b.walk(re.Sub[i], w)
}
return w
case syntax.OpAlternate:
var out tailWalk
for _, sub := range re.Sub {
out = out.union(b.walk(sub, w))
}
return out
case syntax.OpQuest:
return b.repeat(re.Sub[0], w, 1)
case syntax.OpStar:
return b.repeat(re.Sub[0], w, -1)
case syntax.OpPlus:
return b.repeat(re.Sub[0], b.walk(re.Sub[0], w), -1)
case syntax.OpRepeat:
for i := 0; i < re.Min; i++ {
if b.charge() {
return w
}
w = b.walk(re.Sub[0], w)
}
if re.Max < 0 {
return b.repeat(re.Sub[0], w, -1)
}
return b.repeat(re.Sub[0], w, re.Max-re.Min)
}
return w // empty match, line and word boundaries: no constraint
}
// charge counts one repetition step and reports whether the walk has run out of budget. Only
// repeats re-walk their body, so charging them alone bounds the blow-up of nested repeats while
// leaving a single linear pass, of any length, free.
func (b *tailBuilder) charge() bool {
b.work++
if b.work > tailBudget {
b.void = true
}
return b.void
}
// repeat walks back over up to n more repetitions of re, any number if n < 0.
func (b *tailBuilder) repeat(re *syntax.Regexp, w tailWalk, n int) tailWalk {
for ; n != 0; n-- {
if b.charge() {
return w
}
next := w.union(b.walk(re, w))
if next == w {
break
}
w = next
}
return w
}
// step walks back over one character whose last byte is in set. A character that can be
// non-ASCII can take up to 4 bytes, all >= 0x80; regexp matches an invalid byte as U+FFFD.
func (b *tailBuilder) step(w tailWalk, set *byteSet) tailWalk {
out := tailWalk{far: w.far, free: w.free}
if w.far {
b.rest.or(set)
}
width := 1
if set.has(0x80) {
width = utf8.UTFMax
}
for i := 0; i < tailLen; i++ {
if w.at&(1<<i) == 0 {
continue
}
b.tail[i].or(set)
for n := 1; n <= width; n++ {
if j := i + n; j < tailLen {
out.at |= 1 << j
if n < width {
b.tail[j].add(0x80)
}
} else {
out.far = true
if n < width {
b.rest.add(0x80)
}
}
}
}
return out
}
// mayMatch reports whether s passes the tail guard.
func (m *RegexMatcher) mayMatch(s string) bool {
n := len(s)
if m.rest == nil {
n = min(n, len(m.tail))
}
for i := 0; i < n; i++ {
set := m.rest
if i < len(m.tail) {
set = &m.tail[i]
}
if !set.has(s[len(s)-1-i]) {
return false
}
}
return true
}
// requiredLiterals appends to dst the case-sensitive strings that every match of re contains.
func requiredLiterals(re *syntax.Regexp, dst []string) []string {
switch re.Op {
case syntax.OpLiteral:
// regexp matches U+FFFD against invalid UTF-8 bytes, strings.Contains does not
if re.Flags&syntax.FoldCase == 0 && !slices.Contains(re.Rune, utf8.RuneError) {
dst = append(dst, string(re.Rune))
}
case syntax.OpCapture, syntax.OpPlus:
dst = requiredLiterals(re.Sub[0], dst)
case syntax.OpRepeat:
if re.Min > 0 {
dst = requiredLiterals(re.Sub[0], dst)
}
case syntax.OpConcat:
for _, sub := range re.Sub {
dst = requiredLiterals(sub, dst)
}
}
return dst
pattern *regexp.Regexp
}
func (*RegexMatcher) Type() Type {
@@ -359,14 +89,6 @@ func (m *RegexMatcher) String() string {
}
func (m *RegexMatcher) Match(s string) bool {
if !m.mayMatch(s) {
return false
}
for _, l := range m.literals {
if !strings.Contains(s, l) {
return false
}
}
return m.pattern.MatchString(s)
}
@@ -380,7 +102,11 @@ func (t Type) New(pattern string) (Matcher, error) {
case Domain:
return DomainMatcher(pattern), nil
case Regex: // 1. regex matching is case-sensitive
return newRegexMatcher(pattern)
regex, err := regexp.Compile(pattern)
if err != nil {
return nil, err
}
return &RegexMatcher{pattern: regex}, nil
default:
return nil, errors.New("unknown matcher type")
}
@@ -409,7 +135,11 @@ func (t Type) NewDomainPattern(pattern string) (Matcher, error) {
}
return DomainMatcher(pattern), nil
case Regex: // Regex's charset not in LDH subset
return newRegexMatcher(pattern)
regex, err := regexp.Compile(pattern)
if err != nil {
return nil, err
}
return &RegexMatcher{pattern: regex}, nil
default:
return nil, errors.New("unknown matcher type")
}
@@ -1,233 +0,0 @@
package strmatcher
import (
"hash/fnv"
"math/rand/v2"
"regexp"
"regexp/syntax"
"slices"
"strconv"
"strings"
"testing"
"unicode"
"unicode/utf8"
)
var regexLiteralCases = []struct {
pattern string
literals []string
}{
{`(^|\.)91porn\.(best|com)$`, []string{"91porn."}},
{`.+\.awsdns-cn-[0-9][0-9]\.(biz|com|net|top)$`, []string{".awsdns-cn-", "."}},
{`^r+[0-9]+(---|\.)sn-(2x3|ni5|j5o)\w{5}\.googlevideo\.com$`, []string{".googlevideo.com", "sn-", "r"}},
{`(?i)abc`, nil},
{`ab(?i:CD)ef`, []string{"ab", "ef"}},
{`(abc)?x`, []string{"x"}},
{`(abc)*x`, []string{"x"}},
{`x{0,3}yy`, []string{"yy"}},
{`(ab)+c{2}`, []string{"ab", "c"}},
{`abc|abd`, []string{"ab"}},
{`\Qa.b\E`, []string{"a.b"}},
{`a\x{FFFD}b`, nil},
{`^[^.]+$`, nil},
}
func TestRegexRequiredLiterals(t *testing.T) {
for _, test := range regexLiteralCases {
m, err := newRegexMatcher(test.pattern)
if err != nil {
t.Fatal(err)
}
if got := m.(*RegexMatcher).literals; !slices.Equal(got, test.literals) {
t.Errorf("%s: got %q, want %q", test.pattern, got, test.literals)
}
}
}
var regexTailCases = []struct {
pattern string
guard bool
match []string // inputs the pattern matches
reject []string // inputs the tail guard alone rejects
}{
{`^[a-z]([a-z0-9-]{0,61}[a-z0-9])?$`, true, []string{"a", "localhost", "x-1"}, []string{"www.example.com", "localhost.", "LOCALHOST", "a b"}},
{`(^|\.)[a-z][1-9][0-9][a-z]\.com$`, true, []string{"a12b.com", "x.q10z.com"}, []string{"google.com", "a12b.co", "a12b.com.", "ab12.com"}},
{`^hses[1-7]?\.akamaized\.net$`, true, []string{"hses.akamaized.net", "hses3.akamaized.net"}, []string{"xhses.akamaized.net", "www.hses.akamaized.net"}},
{`(?i)k\.net$`, true, []string{"k.net", "K.NET", "\u212a.net"}, []string{"x.net", "k.nex"}},
{`[^.]+\.cn$`, true, []string{"a.cn", "\xff.cn", "\u4e2d.cn"}, []string{"a.cnn", "a.c"}},
{`\x{FFFD}$`, true, []string{"\xff", "a\xc3", "\uFFFD"}, []string{"a", "\xff."}},
{`^.\.cn$`, true, []string{"a.cn", "\u4E2D.cn", "\xff.cn"}, []string{"ab.cn"}},
{`^$`, true, []string{""}, []string{"a"}},
{`(^|\.)youyuapi\..+$`, false, []string{"youyuapi.com"}, nil},
{`abc`, false, []string{"abc", "xabcx"}, nil},
{`^ab`, false, []string{"ab", "abc"}, nil},
{`a$|b`, false, []string{"a", "bx"}, nil},
{`(?m)a$`, false, []string{"a", "a\nb"}, nil},
{strings.Repeat(`(?:abcdefgh(?:a`, 20) + strings.Repeat(`)*)*`, 20) + `\.com$`, false, []string{".com", "abcdefgha.com"}, nil}, // over tailBudget
}
func TestRegexTailGuard(t *testing.T) {
for _, test := range regexTailCases {
m, err := newRegexMatcher(test.pattern)
if err != nil {
t.Fatal(err)
}
rm := m.(*RegexMatcher)
if guard := rm.tail != nil || rm.rest != nil; guard != test.guard {
t.Errorf("%s: guard %v, want %v", test.pattern, guard, test.guard)
}
for _, s := range test.match {
if !rm.pattern.MatchString(s) || !rm.Match(s) {
t.Errorf("%s: %q does not match", test.pattern, s)
}
}
for _, s := range test.reject {
if rm.pattern.MatchString(s) || rm.mayMatch(s) {
t.Errorf("%s: %q passes the guard", test.pattern, s)
}
}
}
}
// TestRegexTailGuardFlatAlternation checks that a long but non-recursive pattern keeps its
// guard. Only nested repeats are charged against tailBudget, so a flat alternation of many
// names, however large, is walked once and guarded; its guard is checked against regexp.
func TestRegexTailGuardFlatAlternation(t *testing.T) {
var sb strings.Builder
sb.WriteString("(?:")
for i := 0; i < 20000; i++ {
if i > 0 {
sb.WriteByte('|')
}
sb.WriteString("name")
sb.WriteString(strconv.Itoa(i))
}
sb.WriteString(`)\.example\.com$`)
m, err := newRegexMatcher(sb.String())
if err != nil {
t.Fatal(err)
}
rm := m.(*RegexMatcher)
if rm.tail == nil && rm.rest == nil {
t.Fatal("flat alternation of 20000 names lost its guard")
}
for _, s := range []string{"name0.example.com", "name19999.example.com", "x.name12345.example.com"} {
if !rm.pattern.MatchString(s) || !rm.Match(s) {
t.Errorf("%q should match", s)
}
}
for _, s := range []string{"name0.example.org", "name0.example.com.", "name0.example.con", "google.com"} {
if rm.pattern.MatchString(s) {
t.Fatalf("test bug: %q matches the pattern", s)
}
if rm.mayMatch(s) {
t.Errorf("%q should be rejected by the guard", s)
}
}
}
// sampleMatch appends a string that re matches, assertions aside, unless it runs out of
// budget, which it spends one per call so that nested repeats stay cheap.
func sampleMatch(sb *strings.Builder, re *syntax.Regexp, rnd *rand.Rand, budget *int) {
if *budget <= 0 {
return
}
*budget--
switch re.Op {
case syntax.OpLiteral:
for _, r := range re.Rune {
if re.Flags&syntax.FoldCase != 0 {
for n := rnd.IntN(4); n > 0; n-- {
r = unicode.SimpleFold(r)
}
}
sampleRune(sb, r, rnd)
}
case syntax.OpCharClass:
if len(re.Rune) > 0 {
i := rnd.IntN(len(re.Rune)/2) * 2
sampleRune(sb, re.Rune[i]+rnd.Int32N(min(re.Rune[i+1]-re.Rune[i]+1, 300)), rnd)
}
case syntax.OpAnyChar, syntax.OpAnyCharNotNL:
sampleRune(sb, []rune{'a', '.', '\n', 0xe9, 0x212a, utf8.RuneError}[rnd.IntN(6)], rnd)
case syntax.OpCapture:
sampleMatch(sb, re.Sub[0], rnd, budget)
case syntax.OpConcat:
for _, sub := range re.Sub {
sampleMatch(sb, sub, rnd, budget)
}
case syntax.OpAlternate:
sampleMatch(sb, re.Sub[rnd.IntN(len(re.Sub))], rnd, budget)
case syntax.OpQuest, syntax.OpStar, syntax.OpPlus, syntax.OpRepeat:
lo, hi := 0, 3
switch re.Op {
case syntax.OpQuest:
hi = 1
case syntax.OpPlus:
lo = 1
case syntax.OpRepeat:
lo, hi = re.Min, re.Min+3
if re.Max >= 0 {
hi = min(hi, re.Max)
}
}
for n := lo + rnd.IntN(hi-lo+1); n > 0; n-- {
sampleMatch(sb, re.Sub[0], rnd, budget)
}
}
}
func sampleRune(sb *strings.Builder, r rune, rnd *rand.Rand) {
if r == utf8.RuneError && rnd.IntN(2) == 0 {
sb.WriteByte(0x80 | byte(rnd.IntN(0x80))) // regexp matches an invalid byte as U+FFFD
return
}
sb.WriteRune(r)
}
func FuzzRegexMatcher(f *testing.F) {
inputs := []string{
"", "x", "yy", "abd", "ccc", "ABC", "abCDef", "abcdef", "abababcc", "a.b", "a\xffb", "a\uFFFDb",
"www.91porn.com", "ns1.awsdns-cn-01.top", "r1---sn-2x3abcde.googlevideo.com",
}
for _, test := range regexLiteralCases {
for _, s := range inputs {
f.Add(test.pattern, s)
}
}
for _, test := range regexTailCases {
for _, s := range append(test.match, test.reject...) {
f.Add(test.pattern, s)
}
}
f.Fuzz(func(t *testing.T, pattern, s string) {
re, err := regexp.Compile(pattern)
if err != nil {
return
}
m, _ := newRegexMatcher(pattern)
check := func(s string) {
if got, want := m.Match(s), re.MatchString(s); got != want {
t.Errorf("pattern %q, input %q: got %v, want %v", pattern, s, got, want)
}
}
check(s)
// random inputs seldom match, so also try strings built from the pattern
parsed, _ := syntax.Parse(pattern, syntax.Perl)
h := fnv.New64a()
h.Write([]byte(s))
rnd := rand.New(rand.NewPCG(h.Sum64(), 1))
for range 8 {
var sb strings.Builder
budget := 256
sampleMatch(&sb, parsed, rnd, &budget)
sample := sb.String()
check(sample)
check(s + sample)
if len(sample) > 0 && len(s) > 0 {
i := rnd.IntN(len(sample))
check(sample[:i] + s[:1] + sample[i+1:])
}
}
})
}
+12 -67
View File
@@ -46,9 +46,7 @@ func (g *MphValueMatcher) Add(matcher Matcher, value uint32) {
func (g *MphValueMatcher) Build() error {
if g.mph != nil {
runtime.GC() // peak mem
if err := g.mph.Build(); err != nil {
return err
}
g.mph.Build()
}
runtime.GC() // peak mem
if g.ac != nil {
@@ -60,17 +58,23 @@ func (g *MphValueMatcher) Build() error {
// Match implements ValueMatcher.Match.
func (g *MphValueMatcher) Match(input string) []uint32 {
var result []uint32
result := make([][]uint32, 0, 5)
if g.mph != nil {
result = g.mph.Match(input) // a new slice, returned without another copy
if matches := g.mph.Match(input); len(matches) > 0 {
result = append(result, matches)
}
}
if g.ac != nil {
result = append(result, g.ac.Match(input)...)
if matches := g.ac.Match(input); len(matches) > 0 {
result = append(result, matches)
}
}
if g.regex != nil {
result = append(result, g.regex.Match(input)...)
if matches := g.regex.Match(input); len(matches) > 0 {
result = append(result, matches)
}
}
return result
return CompositeMatches(result)
}
// MatchAny implements ValueMatcher.MatchAny.
@@ -83,62 +87,3 @@ func (g *MphValueMatcher) MatchAny(input string) bool {
}
return g.regex != nil && g.regex.MatchAny(input)
}
func (g *MphValueMatcher) matchAnyHashed(input string, parents []mphSuffix, h, mul uint64) bool {
if g.mph != nil && g.mph.matchAnyHashed(input, parents, h, mul) {
return true
}
if g.ac != nil && g.ac.MatchAny(input) {
return true
}
return g.regex != nil && g.regex.MatchAny(input)
}
// MphValueMatcherCombiner combines several built MphValueMatchers, each bound to one value, and matches an input
// against them as their MatchAny would, hashing the input once for all of them.
type MphValueMatcherCombiner struct {
matchers []*MphValueMatcher
values []uint32
}
// Add adds a built matcher that stands for value.
func (s *MphValueMatcherCombiner) Add(m *MphValueMatcher, value uint32) {
s.matchers = append(s.matchers, m)
s.values = append(s.values, value)
}
// Match returns the values of the matchers that match input, in Add order.
func (s *MphValueMatcherCombiner) Match(input string) []uint32 {
if len(s.matchers) == 0 {
return nil
}
var stack [16]mphSuffix
mul := mphMultipliers[0]
parents, h := mphSuffixes(stack[:0], mul, input)
var result []uint32
for i, m := range s.matchers {
if m.matchAnyHashed(input, parents, h, mul) {
result = append(result, s.values[i])
}
}
return result
}
// MatchAny returns true as soon as one matcher matches input.
func (s *MphValueMatcherCombiner) MatchAny(input string) bool {
switch len(s.matchers) {
case 0:
return false
case 1:
return s.matchers[0].MatchAny(input) // nothing to share, and it stops at the first matching suffix
}
var stack [16]mphSuffix
mul := mphMultipliers[0]
parents, h := mphSuffixes(stack[:0], mul, input)
for _, m := range s.matchers {
if m.matchAnyHashed(input, parents, h, mul) {
return true
}
}
return false
}
-20
View File
@@ -1,20 +0,0 @@
package net
// PacketConnWrapper wraps a PacketConn into a Conn with a fixed destination address.
type PacketConnWrapper struct {
PacketConn
Dest Addr
}
func (c *PacketConnWrapper) Read(p []byte) (int, error) {
n, _, err := c.PacketConn.ReadFrom(p)
return n, err
}
func (c *PacketConnWrapper) Write(p []byte) (int, error) {
return c.PacketConn.WriteTo(p, c.Dest)
}
func (c *PacketConnWrapper) RemoteAddr() Addr {
return c.Dest
}
+53
View File
@@ -0,0 +1,53 @@
package singbridge
import (
M "github.com/sagernet/sing/common/metadata"
N "github.com/sagernet/sing/common/network"
"github.com/xtls/xray-core/common/errors"
"github.com/xtls/xray-core/common/net"
)
func ToNetwork(network string) net.Network {
switch N.NetworkName(network) {
case N.NetworkTCP:
return net.Network_TCP
case N.NetworkUDP:
return net.Network_UDP
default:
return net.Network_Unknown
}
}
func ToDestination(socksaddr M.Socksaddr, network net.Network) (net.Destination, error) {
// IsFqdn() implicitly checks if the domain name is valid
if socksaddr.IsFqdn() {
return net.Destination{
Network: network,
Address: net.DomainAddress(socksaddr.Fqdn),
Port: net.Port(socksaddr.Port),
}, nil
}
// IsIP() implicitly checks if the IP address is valid
if socksaddr.IsIP() {
return net.Destination{
Network: network,
Address: net.IPAddress(socksaddr.Addr.AsSlice()),
Port: net.Port(socksaddr.Port),
}, nil
}
return net.Destination{}, errors.New("invalid socks address: ", socksaddr)
}
func ToSocksaddr(destination net.Destination) M.Socksaddr {
var addr M.Socksaddr
switch destination.Address.Family() {
case net.AddressFamilyDomain:
addr.Fqdn = destination.Address.Domain()
default:
addr.Addr = M.AddrFromIP(destination.Address.IP())
}
addr.Port = uint16(destination.Port)
return addr
}
+72
View File
@@ -0,0 +1,72 @@
package singbridge
import (
"context"
"os"
M "github.com/sagernet/sing/common/metadata"
N "github.com/sagernet/sing/common/network"
"github.com/xtls/xray-core/common/net"
"github.com/xtls/xray-core/common/net/cnc"
"github.com/xtls/xray-core/common/session"
"github.com/xtls/xray-core/proxy"
"github.com/xtls/xray-core/transport"
"github.com/xtls/xray-core/transport/internet"
"github.com/xtls/xray-core/transport/pipe"
)
var _ N.Dialer = (*XrayDialer)(nil)
type XrayDialer struct {
internet.Dialer
}
func NewDialer(dialer internet.Dialer) *XrayDialer {
return &XrayDialer{dialer}
}
func (d *XrayDialer) DialContext(ctx context.Context, network string, destination M.Socksaddr) (net.Conn, error) {
dest, err := ToDestination(destination, ToNetwork(network))
if err != nil {
return nil, err
}
return d.Dialer.Dial(ctx, dest)
}
func (d *XrayDialer) ListenPacket(ctx context.Context, destination M.Socksaddr) (net.PacketConn, error) {
return nil, os.ErrInvalid
}
type XrayOutboundDialer struct {
outbound proxy.Outbound
dialer internet.Dialer
}
func NewOutboundDialer(outbound proxy.Outbound, dialer internet.Dialer) *XrayOutboundDialer {
return &XrayOutboundDialer{outbound, dialer}
}
func (d *XrayOutboundDialer) DialContext(ctx context.Context, network string, destination M.Socksaddr) (net.Conn, error) {
dest, err := ToDestination(destination, ToNetwork(network))
if err != nil {
return nil, err
}
outbounds := session.OutboundsFromContext(ctx)
if len(outbounds) == 0 {
outbounds = []*session.Outbound{{}}
ctx = session.ContextWithOutbounds(ctx, outbounds)
}
ob := outbounds[len(outbounds)-1]
ob.Target = dest
opts := []pipe.Option{pipe.WithSizeLimit(64 * 1024)}
uplinkReader, uplinkWriter := pipe.New(opts...)
downlinkReader, downlinkWriter := pipe.New(opts...)
conn := cnc.NewConnection(cnc.ConnectionInputMulti(downlinkWriter), cnc.ConnectionOutputMulti(uplinkReader))
go d.outbound.Process(ctx, &transport.Link{Reader: downlinkReader, Writer: uplinkWriter}, d.dialer)
return conn, nil
}
func (d *XrayOutboundDialer) ListenPacket(ctx context.Context, destination M.Socksaddr) (net.PacketConn, error) {
return nil, os.ErrInvalid
}
+10
View File
@@ -0,0 +1,10 @@
package singbridge
import E "github.com/sagernet/sing/common/exceptions"
func ReturnError(err error) error {
if E.IsClosedOrCanceled(err) {
return nil
}
return err
}
+58
View File
@@ -0,0 +1,58 @@
package singbridge
import (
"context"
"io"
M "github.com/sagernet/sing/common/metadata"
N "github.com/sagernet/sing/common/network"
"github.com/xtls/xray-core/common/buf"
"github.com/xtls/xray-core/common/errors"
"github.com/xtls/xray-core/common/net"
"github.com/xtls/xray-core/features/routing"
"github.com/xtls/xray-core/transport"
)
var (
_ N.TCPConnectionHandler = (*Dispatcher)(nil)
_ N.UDPConnectionHandler = (*Dispatcher)(nil)
)
type Dispatcher struct {
upstream routing.Dispatcher
newErrorFunc func(values ...any) *errors.Error
}
func NewDispatcher(dispatcher routing.Dispatcher, newErrorFunc func(values ...any) *errors.Error) *Dispatcher {
return &Dispatcher{
upstream: dispatcher,
newErrorFunc: newErrorFunc,
}
}
func (d *Dispatcher) NewConnection(ctx context.Context, conn net.Conn, metadata M.Metadata) error {
dest, err := ToDestination(metadata.Destination, net.Network_TCP)
if err != nil {
return err
}
xConn := NewConn(conn)
return d.upstream.DispatchLink(ctx, dest, &transport.Link{
Reader: xConn,
Writer: xConn,
})
}
func (d *Dispatcher) NewPacketConnection(ctx context.Context, conn N.PacketConn, metadata M.Metadata) error {
dest, err := ToDestination(metadata.Destination, net.Network_UDP)
if err != nil {
return err
}
return d.upstream.DispatchLink(ctx, dest, &transport.Link{
Reader: buf.NewPacketReader(conn.(io.Reader)),
Writer: buf.NewWriter(conn.(io.Writer)),
})
}
func (d *Dispatcher) NewError(ctx context.Context, err error) {
errors.LogInfo(ctx, err.Error())
}
+70
View File
@@ -0,0 +1,70 @@
package singbridge
import (
"context"
"github.com/sagernet/sing/common/logger"
"github.com/xtls/xray-core/common/errors"
)
var _ logger.ContextLogger = (*XrayLogger)(nil)
type XrayLogger struct {
newError func(values ...any) *errors.Error
}
func NewLogger(newErrorFunc func(values ...any) *errors.Error) *XrayLogger {
return &XrayLogger{
newErrorFunc,
}
}
func (l *XrayLogger) Trace(args ...any) {
}
func (l *XrayLogger) Debug(args ...any) {
errors.LogDebug(context.Background(), args...)
}
func (l *XrayLogger) Info(args ...any) {
errors.LogInfo(context.Background(), args...)
}
func (l *XrayLogger) Warn(args ...any) {
errors.LogWarning(context.Background(), args...)
}
func (l *XrayLogger) Error(args ...any) {
errors.LogError(context.Background(), args...)
}
func (l *XrayLogger) Fatal(args ...any) {
}
func (l *XrayLogger) Panic(args ...any) {
}
func (l *XrayLogger) TraceContext(ctx context.Context, args ...any) {
}
func (l *XrayLogger) DebugContext(ctx context.Context, args ...any) {
errors.LogDebug(ctx, args...)
}
func (l *XrayLogger) InfoContext(ctx context.Context, args ...any) {
errors.LogInfo(ctx, args...)
}
func (l *XrayLogger) WarnContext(ctx context.Context, args ...any) {
errors.LogWarning(ctx, args...)
}
func (l *XrayLogger) ErrorContext(ctx context.Context, args ...any) {
errors.LogError(ctx, args...)
}
func (l *XrayLogger) FatalContext(ctx context.Context, args ...any) {
}
func (l *XrayLogger) PanicContext(ctx context.Context, args ...any) {
}
+107
View File
@@ -0,0 +1,107 @@
package singbridge
import (
"context"
"time"
B "github.com/sagernet/sing/common/buf"
"github.com/sagernet/sing/common/bufio"
M "github.com/sagernet/sing/common/metadata"
"github.com/xtls/xray-core/common"
"github.com/xtls/xray-core/common/buf"
"github.com/xtls/xray-core/common/net"
"github.com/xtls/xray-core/common/signal"
"github.com/xtls/xray-core/transport"
)
func CopyPacketConn(ctx context.Context, inboundConn net.Conn, link *transport.Link, destination net.Destination, serverConn net.PacketConn) error {
cancel := func() {
common.Interrupt(link.Reader)
common.Interrupt(serverConn)
}
conn := &PacketConnWrapper{
Reader: link.Reader,
Writer: link.Writer,
Dest: destination,
Conn: inboundConn,
T: signal.CancelAfterInactivity(ctx, cancel, 300*time.Second),
}
return ReturnError(bufio.CopyPacketConn(ctx, conn, bufio.NewPacketConn(serverConn)))
}
type PacketConnWrapper struct {
buf.Reader
buf.Writer
net.Conn
Dest net.Destination
cached buf.MultiBuffer
// A simple patch to avoid goroutine leak since sing infra cannot awake read block by write err
T *signal.ActivityTimer
}
func (w *PacketConnWrapper) ReadPacket(buffer *B.Buffer) (addr M.Socksaddr, err error) {
w.T.Update()
defer func() {
if err != nil {
// uplinkonly
w.T.SetTimeout(2 * time.Second)
}
}()
if w.cached != nil {
mb, bb := buf.SplitFirst(w.cached)
if bb == nil {
w.cached = nil
} else {
buffer.Write(bb.Bytes())
w.cached = mb
var destination net.Destination
if bb.UDP != nil {
destination = *bb.UDP
} else {
destination = w.Dest
}
bb.Release()
return ToSocksaddr(destination), nil
}
}
mb, err := w.ReadMultiBuffer()
nb, bb := buf.SplitFirst(mb)
if bb == nil {
return M.Socksaddr{}, nil
} else {
buffer.Write(bb.Bytes())
w.cached = nb
var destination net.Destination
if bb.UDP != nil {
destination = *bb.UDP
} else {
destination = w.Dest
}
bb.Release()
return ToSocksaddr(destination), nil
}
}
func (w *PacketConnWrapper) WritePacket(buffer *B.Buffer, destination M.Socksaddr) (err error) {
w.T.Update()
defer func() {
if err != nil {
// downlinkonly
w.T.SetTimeout(5 * time.Second)
}
}()
endpoint, err := ToDestination(destination, net.Network_UDP)
if err != nil {
return err
}
vBuf := buf.New()
vBuf.Write(buffer.Bytes())
vBuf.UDP = &endpoint
return w.WriteMultiBuffer(buf.MultiBuffer{vBuf})
}
func (w *PacketConnWrapper) Close() error {
buf.ReleaseMulti(w.cached)
return nil
}
+81
View File
@@ -0,0 +1,81 @@
package singbridge
import (
"context"
"io"
"net"
"time"
"github.com/sagernet/sing/common/bufio"
"github.com/xtls/xray-core/common"
"github.com/xtls/xray-core/common/buf"
"github.com/xtls/xray-core/common/signal"
"github.com/xtls/xray-core/transport"
)
func CopyConn(ctx context.Context, inboundConn net.Conn, link *transport.Link, serverConn net.Conn) error {
conn := &PipeConnWrapper{
W: link.Writer,
Conn: inboundConn,
}
if ir, ok := link.Reader.(io.Reader); ok {
conn.R = ir
} else {
conn.R = &buf.BufferedReader{Reader: link.Reader}
}
cancel := func() {
common.Interrupt(link.Reader)
common.Interrupt(serverConn)
}
conn.T = signal.CancelAfterInactivity(ctx, cancel, 300*time.Second)
return ReturnError(bufio.CopyConn(ctx, conn, serverConn))
}
type PipeConnWrapper struct {
R io.Reader
W buf.Writer
net.Conn
// A simple patch to avoid goroutine leak since sing infra cannot awake read block by write err
T *signal.ActivityTimer
}
func (w *PipeConnWrapper) Close() error {
return nil
}
func (w *PipeConnWrapper) Read(b []byte) (n int, err error) {
w.T.Update()
n, err = w.R.Read(b)
if err != nil {
// uplinkonly
w.T.SetTimeout(2 * time.Second)
}
return
}
func (w *PipeConnWrapper) Write(p []byte) (n int, err error) {
w.T.Update()
n = len(p)
var mb buf.MultiBuffer
pLen := len(p)
for pLen > 0 {
buffer := buf.New()
if pLen > buf.Size {
_, err = buffer.Write(p[:buf.Size])
p = p[buf.Size:]
} else {
buffer.Write(p)
}
pLen -= int(buffer.Len())
mb = append(mb, buffer)
}
err = w.W.WriteMultiBuffer(mb)
if err != nil {
n = 0
buf.ReleaseMulti(mb)
// downlinkonly
w.T.SetTimeout(5 * time.Second)
}
return
}
+66
View File
@@ -0,0 +1,66 @@
package singbridge
import (
"time"
"github.com/sagernet/sing/common"
"github.com/sagernet/sing/common/bufio"
N "github.com/sagernet/sing/common/network"
"github.com/xtls/xray-core/common/buf"
"github.com/xtls/xray-core/common/net"
)
var (
_ buf.Reader = (*Conn)(nil)
_ buf.TimeoutReader = (*Conn)(nil)
_ buf.Writer = (*Conn)(nil)
)
type Conn struct {
net.Conn
writer N.VectorisedWriter
}
func NewConn(conn net.Conn) *Conn {
writer, _ := bufio.CreateVectorisedWriter(conn)
return &Conn{
Conn: conn,
writer: writer,
}
}
func (c *Conn) ReadMultiBuffer() (buf.MultiBuffer, error) {
buffer, err := buf.ReadBuffer(c.Conn)
if err != nil {
return nil, err
}
return buf.MultiBuffer{buffer}, nil
}
func (c *Conn) ReadMultiBufferTimeout(duration time.Duration) (buf.MultiBuffer, error) {
err := c.SetReadDeadline(time.Now().Add(duration))
if err != nil {
return nil, err
}
defer c.SetReadDeadline(time.Time{})
return c.ReadMultiBuffer()
}
func (c *Conn) WriteMultiBuffer(bufferList buf.MultiBuffer) error {
defer buf.ReleaseMulti(bufferList)
if c.writer != nil {
bytesList := make([][]byte, len(bufferList))
for i, buffer := range bufferList {
bytesList[i] = buffer.Bytes()
}
return common.Error(bufio.WriteVectorised(c.writer, bytesList))
}
// Since this conn is only used by tun, we don't force buffer writes to merge.
for _, buffer := range bufferList {
_, err := c.Conn.Write(buffer.Bytes())
if err != nil {
return err
}
}
return nil
}
+1 -1
View File
@@ -20,7 +20,7 @@ import (
var (
Version_x byte = 26
Version_y byte = 9
Version_z byte = 30
Version_z byte = 9
)
var (
+1 -1
View File
@@ -13,7 +13,7 @@ type FakeDNSEngine interface {
var (
FakeIPv4Pool = "198.18.0.0/15"
FakeIPv6Pool = "2001:2::/48"
FakeIPv6Pool = "fc00::/18"
)
type FakeDNSEngineRev0 interface {
-3
View File
@@ -97,9 +97,6 @@ func New() *Client {
r := &net.Resolver{
PreferGo: true,
Dial: func(ctx context.Context, network, address string) (net.Conn, error) {
if internet.IsSkippedDNSServer(address) {
return nil, errors.New("skipped DNS server ", address)
}
return d.DialContext(ctx, network, address)
},
}
-23
View File
@@ -1,23 +0,0 @@
package localdns
import (
"context"
"net/netip"
"testing"
"github.com/xtls/xray-core/transport/internet"
)
func TestSkippedDNSServers(t *testing.T) {
internet.SkipDNSServers([]netip.Addr{netip.MustParseAddr("203.0.113.53")})
t.Cleanup(func() { internet.SkipDNSServers(nil) })
c := New()
if _, err := c.r.Dial(context.Background(), "udp", "203.0.113.53:53"); err == nil {
t.Error("a skipped DNS server was dialed")
}
conn, err := c.r.Dial(context.Background(), "udp", "127.0.0.1:53")
if err != nil {
t.Fatal(err)
}
conn.Close()
}
+10 -8
View File
@@ -18,19 +18,21 @@ require (
github.com/pires/go-proxyproto v0.15.0
github.com/refraction-networking/utls v1.8.3-0.20260301010127-aa6edf4b11af
github.com/robfig/cron/v3 v3.0.1
github.com/sagernet/sing v0.5.1
github.com/sagernet/sing-shadowsocks v0.2.7
github.com/stretchr/testify v1.12.1
github.com/vishvananda/netlink v1.3.1
github.com/xtls/reality v0.0.0-20260908062103-8cdf7bf9c7f0
go4.org/netipx v0.0.0-20231129151722-fdeea329fbba
golang.org/x/crypto v0.57.0
golang.org/x/crypto v0.55.0
golang.org/x/exp v0.0.0-20240506185415-9bf2ced13842
golang.org/x/net v0.59.0
golang.org/x/sync v0.23.0
golang.org/x/sys v0.48.0
golang.org/x/net v0.58.0
golang.org/x/sync v0.22.0
golang.org/x/sys v0.47.0
golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2
golang.zx2c4.com/wireguard v0.0.0-20250521234502-f333402bd9cb
golang.zx2c4.com/wireguard/windows v1.1.1
google.golang.org/grpc v1.84.0
golang.zx2c4.com/wireguard/windows v1.0.1
google.golang.org/grpc v1.83.2
google.golang.org/protobuf v1.36.12
gvisor.dev/gvisor v0.0.0-20260122175437-89a5d21be8f0
h12.io/socks v1.0.3
@@ -55,9 +57,9 @@ 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/text v0.42.0 // indirect
golang.org/x/text v0.41.0 // indirect
golang.org/x/time v0.14.0 // indirect
golang.org/x/tools v0.49.0 // indirect
google.golang.org/genproto/googleapis/rpc v0.0.0-20260706201446-f0a921348800 // indirect
google.golang.org/genproto/googleapis/rpc v0.0.0-20260526163538-3dc84a4a5aaa // indirect
gopkg.in/yaml.v2 v2.4.0 // indirect
)
+38 -16
View File
@@ -2,10 +2,16 @@ github.com/andybalholm/brotli v1.0.6 h1:Yf9fFpf49Zrxb9NlQaluyE92/+X7UVHlhMNJN2sx
github.com/andybalholm/brotli v1.0.6/go.mod h1:fO7iG3H7G2nSZ7m0zPUDn85XEX2GTukHGRSepvi9Eig=
github.com/apernet/quic-go v0.61.1-0.20260806010916-184d081eef3e h1:5mgtR5gwIgBKMiGI1QdXldZZ+SNor06Nbu1wCBulQBg=
github.com/apernet/quic-go v0.61.1-0.20260806010916-184d081eef3e/go.mod h1:x7qxEvX6MCVtDuBKHj3E+88+BtrbEMuAL5qGUKItjW8=
github.com/cespare/xxhash/v2 v2.3.0 h1:UL815xU9SqsFlibzuggzjXhog7bL6oX9BbNZnL2UFvs=
github.com/cespare/xxhash/v2 v2.3.0/go.mod h1:VGX0DQ3Q6kWi7AoAeZDth3/j3BFtOZR5XLFGgcrjCOs=
github.com/cloudflare/circl v1.6.5 h1:O64F26HEqNhznd/hrC5KZXVKYuKM2rx4deZDTc4ihQA=
github.com/cloudflare/circl v1.6.5/go.mod h1:h5LNyxAc5nTue9DS5jT+48en2PSDYt3zdGnz5OstK6c=
github.com/ghodss/yaml v1.0.1-0.20220118164431-d8423dcdf344 h1:Arcl6UOIS/kgO2nW3A65HN+7CMjSDP/gofXL4CZt1V4=
github.com/ghodss/yaml v1.0.1-0.20220118164431-d8423dcdf344/go.mod h1:GIjDIg/heH5DOkXY3YJ/wNhfHsQHoXGjl8G8amsYQ1I=
github.com/go-logr/logr v1.4.3 h1:CjnDlHq8ikf6E492q6eKboGOC0T8CDaOvkHCIg8idEI=
github.com/go-logr/logr v1.4.3/go.mod h1:9T104GzyrTigFIr8wt5mBrctHMim0Nb2HLGrmQ40KvY=
github.com/go-logr/stdr v1.2.2 h1:hSWxHoqTgW2S2qGc0LTAI563KZ5YKYRhT3MFKZMbjag=
github.com/go-logr/stdr v1.2.2/go.mod h1:mMo/vtBO5dYbehREoey6XUKy/eSumjCCveDpRre4VKE=
github.com/go-quicktest/qt v1.102.0 h1:HSQxCeh5YZH3EL3W39ixjtyaEhcWSXQHtHnMBzSs474=
github.com/go-quicktest/qt v1.102.0/go.mod h1:p4lGIVX+8Wa6ZPNDvqcxq36XpUDLh42FLetFU7odllI=
github.com/golang/mock v1.7.0-rc.1 h1:YojYx61/OLFsiv6Rw1Z96LpldJIy31o+UHmwAUMJ6/U=
@@ -70,6 +76,10 @@ github.com/robfig/cron/v3 v3.0.1 h1:WdRxkvbJztn8LMz/QEvLN5sBU+xKpSqwwUO1Pjr4qDs=
github.com/robfig/cron/v3 v3.0.1/go.mod h1:eQICP3HwyT7UooqI/z+Ov+PtYAWygg1TEWWzGIFLtro=
github.com/rogpeppe/go-internal v1.16.0 h1:O9DK+vNMDVGLr2BeZqmpLeMjiMNkuXfcqntWbZV6S5g=
github.com/rogpeppe/go-internal v1.16.0/go.mod h1:DrUVZyrJU+txYW5/1kwtXQSMFio52ZOxX7yM1VHvnxs=
github.com/sagernet/sing v0.5.1 h1:mhL/MZVq0TjuvHcpYcFtmSD1BFOxZ/+8ofbNZcg1k1Y=
github.com/sagernet/sing v0.5.1/go.mod h1:ARkL0gM13/Iv5VCZmci/NuoOlePoIsW0m7BWfln/Hak=
github.com/sagernet/sing-shadowsocks v0.2.7 h1:zaopR1tbHEw5Nk6FAkM05wCslV6ahVegEZaKMv9ipx8=
github.com/sagernet/sing-shadowsocks v0.2.7/go.mod h1:0rIKJZBR65Qi0zwdKezt4s57y/Tl1ofkaq6NlkzVuyE=
github.com/stretchr/testify v1.12.1 h1:EuwCh5fleGS7H32xRwO3wRGT7DxrDhLAT6FF8MpWDWE=
github.com/stretchr/testify v1.12.1/go.mod h1:MDEgiDPPsNp5cuIrHPPCyornHKgEVbtFUmoNlxoYthg=
github.com/vishvananda/netlink v1.3.1 h1:3AEMt62VKqz90r0tmNhog0r/PpWKmrEShJU0wJW6bV0=
@@ -81,6 +91,18 @@ github.com/wlynxg/anet v0.0.5/go.mod h1:eay5PRQr7fIVAMbTbchTnO9gG65Hg/uYGdc7mguH
github.com/xtls/reality v0.0.0-20260908062103-8cdf7bf9c7f0 h1:rb+fKQFhz+5I2PPuQsNYxI5mUU840XWYtRF0ZBjvkws=
github.com/xtls/reality v0.0.0-20260908062103-8cdf7bf9c7f0/go.mod h1:DsJblcWDGt76+FVqBVwbwRhxyyNJsGV48gJLch0OOWI=
github.com/yuin/goldmark v1.4.1/go.mod h1:mwnBkeHKe2W/ZEtQ+71ViKU8L12m81fl3OWwC1Zlc8k=
go.opentelemetry.io/auto/sdk v1.2.1 h1:jXsnJ4Lmnqd11kwkBV2LgLoFMZKizbCi5fNZ/ipaZ64=
go.opentelemetry.io/auto/sdk v1.2.1/go.mod h1:KRTj+aOaElaLi+wW1kO/DZRXwkF4C5xPbEe3ZiIhN7Y=
go.opentelemetry.io/otel v1.44.0 h1:JjwHmHpA4iZ3wBxluu2fbbE7j4kqlE8jXyAyPXH7HqU=
go.opentelemetry.io/otel v1.44.0/go.mod h1:BMgjTHL9WPRlRjL2oZCBTL4whCGtXch2H4BhOPIAyYc=
go.opentelemetry.io/otel/metric v1.44.0 h1:1w0gILTcHdr3YI+ixLyjemwrVnsMURbTZFrSYCdDdmc=
go.opentelemetry.io/otel/metric v1.44.0/go.mod h1:8O7hanEPBNgEMmybD3s2VBKcgWOCsA6tzHBPODAiquo=
go.opentelemetry.io/otel/sdk v1.44.0 h1:nHYwb9lK+fJPU/dnT6s7W7Z8itMWyqrnVfbheVYrZ58=
go.opentelemetry.io/otel/sdk v1.44.0/go.mod h1:Osuydd3Se74nqjAKxid74N5eC+jfEqfTegHRnq58oK0=
go.opentelemetry.io/otel/sdk/metric v1.44.0 h1:3LlKgI+VjbVsjNRFZJZAJ30WjXC5VkNRks6si09iEfI=
go.opentelemetry.io/otel/sdk/metric v1.44.0/go.mod h1:5B5pMARnXxKhltooO4xUuCBorl65a4EpnTalObqOigA=
go.opentelemetry.io/otel/trace v1.44.0 h1:jxF5CsGYCe74MCRx2X4g7WsY/VBKRqqpNvXlX/6gtIk=
go.opentelemetry.io/otel/trace v1.44.0/go.mod h1:oLl1jrMQAVo6v3GAggN+1VH9VIz9iUSvW53sW1Q8PIE=
go.uber.org/mock v0.5.2 h1:LbtPTcP8A5k9WPXj54PPPbjcI4Y6lhyOZXn+VS7wNko=
go.uber.org/mock v0.5.2/go.mod h1:wLlUxC2vVTPTaE3UD51E0BGOAElKrILxhVSDYQLld5o=
go.yaml.in/yaml/v3 v3.0.5 h1:N6y/pJk8buWs9NY5ERU2HSMfm+IuD/OtfdAnq6kESPw=
@@ -89,8 +111,8 @@ go4.org/netipx v0.0.0-20231129151722-fdeea329fbba h1:0b9z3AuHCjxk0x/opv64kcgZLBs
go4.org/netipx v0.0.0-20231129151722-fdeea329fbba/go.mod h1:PLyyIXexvUFg3Owu6p/WfdlivPbZJsZdgWZlrGope/Y=
golang.org/x/crypto v0.0.0-20190308221718-c2843e01d9a2/go.mod h1:djNgcEr1/C05ACkg1iLfiJU5Ep61QUkGW8qpdssI0+w=
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 v0.55.0 h1:+KWHjbgOaAQ66dh/YlkZKHlz9ZUlq61AFirAR9ntP8M=
golang.org/x/crypto v0.55.0/go.mod h1:uq0V9dE/fzQuJtbnL+2EhWOE63vo164FY8xqEnV9xis=
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=
@@ -99,12 +121,12 @@ golang.org/x/mod v0.5.1/go.mod h1:5OXOZSfqPIIbmVBIIKWRFfZjPR0E5r58TLhUjH0a2Ro=
golang.org/x/net v0.0.0-20190404232315-eb5bcb51f2a3/go.mod h1:t9HGtf8HONx5eT2rtn7q6eTqICYqUVnKs3thJo3Qplg=
golang.org/x/net v0.0.0-20190620200207-3b0461eec859/go.mod h1:z5CRVTTTmAJ677TzLLGU+0bjPO0LkuOLi4/5GtJWs/s=
golang.org/x/net v0.0.0-20211015210444-4f30a5c0130f/go.mod h1:9nx3DQGgdP8bBQD5qxJ1jj9UTztislL4KSBs9R2vV5Y=
golang.org/x/net v0.59.0 h1:5zfYln+w5XCxwrnMMJPufRgNoXEaGxl0wo5GqPXyues=
golang.org/x/net v0.59.0/go.mod h1:2DA/G1UfVbCpQPeWTmMPGY7Cs2PkBkwu743bVX5PIVg=
golang.org/x/net v0.58.0 h1:ynWG7rqYi4ccpTEuPZ2QGWHktVEM9DMCj9yzDE0Q7To=
golang.org/x/net v0.58.0/go.mod h1:YwCddHnFlT7eLQqVprV19OnhLGtc5xOKgE0RyqgfWAU=
golang.org/x/sync v0.0.0-20190423024810-112230192c58/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
golang.org/x/sync v0.0.0-20210220032951-036812b2e83c/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
golang.org/x/sync v0.23.0 h1:KameEIfc1IkluZyXWLn39Wd4tURc6GbCiISGiZm2bQk=
golang.org/x/sync v0.23.0/go.mod h1:sUUOizhqBxiL6pEWpqNLUiaJn1ShEbZ6BBqskPbjZm0=
golang.org/x/sync v0.22.0 h1:SZjpbeLmrCk4xhRSZFNZW5gFUeCeFgjekvI/+gfScek=
golang.org/x/sync v0.22.0/go.mod h1:9xrNwdLfx4jkKbNva9FpL6vEN7evnE43NNNJQ2LF3+0=
golang.org/x/sys v0.0.0-20190215142949-d0b11bdaac8a/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY=
golang.org/x/sys v0.0.0-20190412213103-97732733099d/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs=
golang.org/x/sys v0.0.0-20201119102817-f84b799fce68/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs=
@@ -112,14 +134,14 @@ golang.org/x/sys v0.0.0-20210423082822-04245dca01da/go.mod h1:h1NjWce9XRLGQEsW7w
golang.org/x/sys v0.0.0-20211019181941-9d821ace8654/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
golang.org/x/sys v0.2.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
golang.org/x/sys v0.10.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
golang.org/x/sys v0.48.0 h1:bbX/i/6MgT9BVLM9RT1thmxL04yeTAhbEz4SyadbXoo=
golang.org/x/sys v0.48.0/go.mod h1:hNLxWAXmnKAxqDtdwIYC4bM9oQPEecfsnNMuSxOs3og=
golang.org/x/sys v0.47.0 h1:o7XGOvZQCADBQQ4Y7VNq2dRWQR7JmOUW8Kxx4ZsNgWs=
golang.org/x/sys v0.47.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw=
golang.org/x/term v0.0.0-20201126162022-7de9c90e9dd1/go.mod h1:bj7SfCRtBDWHUb9snDiAeCFNEtKQo2Wmx5Cou7ajbmo=
golang.org/x/text v0.3.0/go.mod h1:NqM8EUOU14njkJ3fqMW+pc6Ldnwhi/IjpwHt7yyuwOQ=
golang.org/x/text v0.3.6/go.mod h1:5Zoc/QRtKVWzQhOtBMvqHzDpF6irO9z98xDceosuGiQ=
golang.org/x/text v0.3.7/go.mod h1:u+2+/6zg+i71rQMx5EYifcz6MCKuco9NR6JIITiCfzQ=
golang.org/x/text v0.42.0 h1:JbOZXgfeCPU9gacVtYliJqOhD+zhrEqK4LfdpmlUZqI=
golang.org/x/text v0.42.0/go.mod h1:ojzP1Z+2QtioaF8DTtO8K5q7JWVVYwZKenzujK0Zd0E=
golang.org/x/text v0.41.0 h1:vz/seA0lnX87Othu2f/0L24RcgrXD9/YFTSuGjj3rH8=
golang.org/x/text v0.41.0/go.mod h1:jvf1O8ajNzZqhSrQBPbutR/EB83Cc0CFrezNQIwbb5M=
golang.org/x/time v0.14.0 h1:MRx4UaLrDotUKUdCIqzPC48t1Y9hANFKIRpNx+Te8PI=
golang.org/x/time v0.14.0/go.mod h1:eL/Oa2bBBK0TkX57Fyni+NgnyQQN4LitPmob2Hjnqw4=
golang.org/x/tools v0.0.0-20180917221912-90fa682c2a6e/go.mod h1:n7NCudcB/nEzxVGmLbDWY5pfWTLqBcC2KZ6jyYvM4mQ=
@@ -135,14 +157,14 @@ golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2 h1:B82qJJgjvYKsXS9jeu
golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2/go.mod h1:deeaetjYA+DHMHg+sMSMI58GrEteJUUzzw7en6TJQcI=
golang.zx2c4.com/wireguard v0.0.0-20250521234502-f333402bd9cb h1:whnFRlWMcXI9d+ZbWg+4sHnLp52d5yiIPUxMBSt4X9A=
golang.zx2c4.com/wireguard v0.0.0-20250521234502-f333402bd9cb/go.mod h1:rpwXGsirqLqN2L0JDJQlwOboGHmptD5ZD6T2VmcqhTw=
golang.zx2c4.com/wireguard/windows v1.1.1 h1:8/H97U1v1PNDNcBsMZgU3KFuND9MQdTsU2NOwmCXArE=
golang.zx2c4.com/wireguard/windows v1.1.1/go.mod h1:+fbT3FFdX4zzYDLwJh5+HPEcNN/3HyNdzhNSVsQM+zs=
golang.zx2c4.com/wireguard/windows v1.0.1 h1:eOxiDVbywPC+ZQqvdCK7x+ZwWXKbYv50TtH8ysFIbw8=
golang.zx2c4.com/wireguard/windows v1.0.1/go.mod h1:+fbT3FFdX4zzYDLwJh5+HPEcNN/3HyNdzhNSVsQM+zs=
gonum.org/v1/gonum v0.17.0 h1:VbpOemQlsSMrYmn7T2OUvQ4dqxQXU+ouZFQsZOx50z4=
gonum.org/v1/gonum v0.17.0/go.mod h1:El3tOrEuMpv2UdMrbNlKEh9vd86bmQ6vqIcDwxEOc1E=
google.golang.org/genproto/googleapis/rpc v0.0.0-20260706201446-f0a921348800 h1:qEHAMpSaUhtD0p3NbEEI83HwNGFxEwaSJ1G9PLnCBZE=
google.golang.org/genproto/googleapis/rpc v0.0.0-20260706201446-f0a921348800/go.mod h1:4Hqkh8ycfw05ld/3BWL7rJOSfebL2Q+DVDeRgYgxUU8=
google.golang.org/grpc v1.84.0 h1:soMyaPJ8pAak5PIQ0DGBUir0XRo2fRoMqhNWMLlLxO0=
google.golang.org/grpc v1.84.0/go.mod h1:ljCht0DrxQrXBDRTZp52Qxh3Ffk8CdYm2sj4O2QN2C0=
google.golang.org/genproto/googleapis/rpc v0.0.0-20260526163538-3dc84a4a5aaa h1:mZHHdPZl0dbGHCflZgAq/Q468DWVFcU2whhB2KAo8fk=
google.golang.org/genproto/googleapis/rpc v0.0.0-20260526163538-3dc84a4a5aaa/go.mod h1:4Hqkh8ycfw05ld/3BWL7rJOSfebL2Q+DVDeRgYgxUU8=
google.golang.org/grpc v1.83.2 h1:EManeRomTObA0BU7I8vXgg/78uE5MJ9M8B39EX2WscU=
google.golang.org/grpc v1.83.2/go.mod h1:YPI1hK3kDked6iHvgX3tR0y+nX/qpMFKhPgFsokw1S8=
google.golang.org/protobuf v1.36.12 h1:pJOKDDOyeXErUroCihFAd5LQuwXBSpVnKGrj5o/fwxc=
google.golang.org/protobuf v1.36.12/go.mod h1:HTf+CrKn2C3g5S8VImy6tdcUvCska2kB7j23XfzDpco=
gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0=
-103
View File
@@ -1,103 +0,0 @@
package conf
import (
"net/netip"
"strings"
"github.com/xtls/xray-core/common/errors"
"github.com/xtls/xray-core/common/protocol"
"github.com/xtls/xray-core/common/serial"
"github.com/xtls/xray-core/proxy/masque"
"google.golang.org/protobuf/proto"
)
type MasqueClientConfig struct {
Address *Address `json:"address"`
Port uint16 `json:"port"`
RemoteDNS []string `json:"remoteDNS"`
}
func (c *MasqueClientConfig) Build() (proto.Message, error) {
if c.Address == nil {
return nil, errors.New(`MASQUE: "address" is not set`)
}
if c.Port == 0 {
return nil, errors.New(`MASQUE: "port" is not set`)
}
for _, s := range c.RemoteDNS {
if _, err := netip.ParseAddr(s); err != nil {
return nil, errors.New(`MASQUE: invalid "remoteDNS" `, s).Base(err)
}
}
return &masque.ClientConfig{
Server: &protocol.ServerEndpoint{
Address: c.Address.Build(),
Port: uint32(c.Port),
},
RemoteDns: c.RemoteDNS,
}, nil
}
type MasqueUserConfig struct {
Pass string `json:"pass"`
Level uint32 `json:"level"`
Email string `json:"email"`
}
type MasqueServerConfig struct {
Users []*MasqueUserConfig `json:"users"`
Clients []*MasqueUserConfig `json:"clients"`
Address []string `json:"address"`
MTU uint32 `json:"mtu"`
}
func (c *MasqueServerConfig) Build() (proto.Message, error) {
if c.Clients != nil {
c.Users = c.Clients
}
config := &masque.ServerConfig{
Address: c.Address,
Mtu: c.MTU,
}
emails := make(map[string]bool)
for _, user := range c.Users {
if user.Email == "" {
return nil, errors.New(`MASQUE: "email" is empty`)
}
if strings.Contains(user.Email, ":") {
return nil, errors.New(`MASQUE: invalid "email" `, user.Email)
}
if user.Pass == "" {
return nil, errors.New(`MASQUE: "pass" of `, user.Email, ` is empty`)
}
email := strings.ToLower(user.Email)
if emails[email] {
return nil, errors.New(`MASQUE: duplicate "email" `, user.Email)
}
emails[email] = true
config.Users = append(config.Users, &protocol.User{
Email: user.Email,
Level: user.Level,
Account: serial.ToTypedMessage(&masque.Account{Password: user.Pass}),
})
}
if len(c.Address) == 0 {
return nil, errors.New(`MASQUE: "address" is not set`)
}
var v4, v6 bool
for _, s := range c.Address {
prefix, err := netip.ParsePrefix(s)
if err != nil {
return nil, errors.New(`MASQUE: invalid "address" `, s).Base(err)
}
if prefix.Addr().Is4() && v4 || prefix.Addr().Is6() && v6 {
return nil, errors.New(`MASQUE: "address" takes at most one IPv4 and one IPv6 prefix`)
}
v4 = v4 || prefix.Addr().Is4()
v6 = v6 || prefix.Addr().Is6()
}
if c.MTU != 0 && (c.MTU < 1280 || c.MTU > 65535) {
return nil, errors.New(`MASQUE: "mtu" must be between 1280 and 65535`)
}
return config, nil
}
-188
View File
@@ -1,188 +0,0 @@
package conf_test
import (
"encoding/json"
"testing"
"github.com/xtls/xray-core/common/protocol"
"github.com/xtls/xray-core/common/serial"
. "github.com/xtls/xray-core/infra/conf"
masqueproxy "github.com/xtls/xray-core/proxy/masque"
"github.com/xtls/xray-core/transport/internet/masque"
)
func TestMasqueConfig(t *testing.T) {
creator := func() Buildable {
return new(MasqueConfig)
}
runMultiTestCase(t, []TestCase{
{
Input: `{}`,
Parser: loadJSON(creator),
Output: &masque.Config{Path: "/.well-known/masque/ip/*/*/"},
},
{
Input: `{
"host": "example.com:8443",
"path": "/.well-known/masque/ip/{target}/{ipproto}/",
"headers": {"Authorization": "Basic dTpw"}
}`,
Parser: loadJSON(creator),
Output: &masque.Config{
Host: "example.com:8443",
Path: "/.well-known/masque/ip/*/*/",
Headers: map[string]string{"Authorization": "Basic dTpw"},
},
},
{
Input: `{"path": "/masque/ip{?target,ipproto}"}`,
Parser: loadJSON(creator),
Output: &masque.Config{Path: "/masque/ip?target=*&ipproto=*"},
},
{
Input: `{"user": "u", "pass": "p:q", "headers": {"X-Token": "a"}}`,
Parser: loadJSON(creator),
Output: &masque.Config{
Path: "/.well-known/masque/ip/*/*/",
Headers: map[string]string{"Authorization": "Basic dTpwOnE=", "X-Token": "a"},
},
},
})
for _, input := range []string{
`{"path": "/masque/{target}/{ipproto}/{dns}"}`,
`{"path": "masque"}`,
`{"host": "example.com/path"}`,
`{"headers": {"host": "example.com"}}`,
`{"headers": {"Capsule-Protocol": "?0"}}`,
`{"headers": {"X Token": "a"}}`,
`{"headers": {"X-Token": "a\r\nb"}}`,
`{"user": "u:v", "pass": "p"}`,
`{"user": "u", "pass": "p", "headers": {"authorization": "Basic dTpw"}}`,
} {
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)
if err := json.Unmarshal([]byte(s), c); err != nil {
return err
}
_, err := c.Build()
return err
}
if err := build(`{
"protocol": "masque",
"settings": {"address": "example.com", "port": 443},
"streamSettings": {"network": "masque", "security": "tls"},
"mux": {"enabled": false, "concurrency": -1}
}`); err != nil {
t.Error(err)
}
for _, input := range []string{
`{"protocol": "masque", "settings": {"address": "example.com"}, "streamSettings": {"network": "masque", "security": "tls"}}`,
`{"protocol": "masque", "settings": {"address": "example.com", "port": 443}, "streamSettings": {"network": "masque", "security": "tls"}, "mux": {"enabled": true}}`,
`{"protocol": "masque", "settings": {"address": "example.com", "port": 443}, "streamSettings": {"network": "masque", "security": "tls"}, "mux": {"enabled": true, "concurrency": -1}}`,
`{"protocol": "freedom", "streamSettings": {"network": "masque", "security": "tls"}}`,
} {
if err := build(input); err == nil {
t.Errorf("expected an error for %s", input)
}
}
}
func TestMasqueServerConfig(t *testing.T) {
creator := func() Buildable {
return new(MasqueServerConfig)
}
runMultiTestCase(t, []TestCase{
{
Input: `{
"users": [{"email": "u@example.com", "pass": "p", "level": 1}],
"address": ["10.13.0.1/24", "fd13::1/64"],
"mtu": 1400
}`,
Parser: loadJSON(creator),
Output: &masqueproxy.ServerConfig{
Users: []*protocol.User{{
Email: "u@example.com",
Level: 1,
Account: serial.ToTypedMessage(&masqueproxy.Account{Password: "p"}),
}},
Address: []string{"10.13.0.1/24", "fd13::1/64"},
Mtu: 1400,
},
},
{
Input: `{"clients": [{"email": "u", "pass": "p:q"}], "address": ["10.13.0.1/24"]}`,
Parser: loadJSON(creator),
Output: &masqueproxy.ServerConfig{
Users: []*protocol.User{{
Email: "u",
Account: serial.ToTypedMessage(&masqueproxy.Account{Password: "p:q"}),
}},
Address: []string{"10.13.0.1/24"},
},
},
{
Input: `{"address": ["10.13.0.1/24"]}`,
Parser: loadJSON(creator),
Output: &masqueproxy.ServerConfig{
Address: []string{"10.13.0.1/24"},
},
},
})
for _, input := range []string{
`{"users": [{"email": "u:v", "pass": "p"}], "address": ["10.13.0.1/24"]}`,
`{"users": [{"email": "", "pass": "p"}], "address": ["10.13.0.1/24"]}`,
`{"users": [{"pass": "p"}], "address": ["10.13.0.1/24"]}`,
`{"users": [{"email": "u", "pass": ""}], "address": ["10.13.0.1/24"]}`,
`{"users": [{"email": "u", "pass": "p"}, {"email": "U", "pass": "q"}], "address": ["10.13.0.1/24"]}`,
`{"users": [{"email": "u", "pass": "p"}]}`,
`{"users": [{"email": "u", "pass": "p"}], "address": ["10.13.0.1"]}`,
`{"users": [{"email": "u", "pass": "p"}], "address": ["10.13.0.1/24", "10.14.0.1/24"]}`,
`{"users": [{"email": "u", "pass": "p"}], "address": ["fd13::1/64", "fd14::1/64"]}`,
`{"users": [{"email": "u", "pass": "p"}], "address": ["10.13.0.1/24"], "mtu": 1000}`,
`{"users": [{"email": "u", "pass": "p"}], "address": ["10.13.0.1/24"], "mtu": 70000}`,
} {
if _, err := loadJSON(creator)(input); err == nil {
t.Errorf("expected an error for %s", input)
}
}
}
func TestMasqueInboundConfig(t *testing.T) {
build := func(s string) error {
c := new(InboundDetourConfig)
if err := json.Unmarshal([]byte(s), c); err != nil {
return err
}
_, err := c.Build()
return err
}
if err := build(`{
"protocol": "masque",
"port": 443,
"settings": {"users": [{"email": "u@example.com", "pass": "p"}], "address": ["10.13.0.1/24"]},
"streamSettings": {"network": "masque", "security": "tls"}
}`); err != nil {
t.Error(err)
}
if err := build(`{
"protocol": "vless",
"port": 443,
"settings": {"users": [{"id": "27848739-7e62-4138-9fd3-098a63964b6b"}], "decryption": "none"},
"streamSettings": {"network": "masque", "security": "tls"}
}`); err == nil {
t.Error("expected an error for the masque transport on a vless inbound")
}
}
+56 -37
View File
@@ -3,6 +3,8 @@ package conf
import (
"strings"
"github.com/sagernet/sing-shadowsocks/shadowaead_2022"
C "github.com/sagernet/sing/common"
"github.com/xtls/xray-core/common/errors"
"github.com/xtls/xray-core/common/protocol"
"github.com/xtls/xray-core/common/serial"
@@ -53,7 +55,7 @@ func (v *ShadowsocksServerConfig) Build() (proto.Message, error) {
v.Users = v.Clients
}
if _, err := shadowsocks_2022.GetCipherMethod(v.Cipher); err == nil {
if C.Contains(shadowaead_2022.List, v.Cipher) {
return buildShadowsocks2022(v)
}
@@ -109,14 +111,12 @@ func (v *ShadowsocksServerConfig) Build() (proto.Message, error) {
}
func buildShadowsocks2022(v *ShadowsocksServerConfig) (proto.Message, error) {
v.Cipher = strings.ToLower(v.Cipher)
if len(v.Users) == 0 {
config := new(shadowsocks_2022.ServerConfig)
config.Method = v.Cipher
config.Key = v.Password
config.Network = v.NetworkList.Build()
config.Email = v.Email
config.Level = int32(v.Level)
return config, nil
}
@@ -171,7 +171,6 @@ func buildShadowsocks2022(v *ShadowsocksServerConfig) (proto.Message, error) {
Email: user.Email,
Address: user.Address.Build(),
Port: uint32(user.Port),
Level: int32(user.Level),
})
}
return config, nil
@@ -215,43 +214,63 @@ func (v *ShadowsocksClientConfig) Build() (proto.Message, error) {
return nil, errors.New(`Shadowsocks settings: "servers" should have one and only one member. Multiple endpoints in "servers" should use multiple Shadowsocks outbounds and routing balancer instead`)
}
server := v.Servers[0]
if server.Address == nil {
return nil, errors.New("Shadowsocks server address is not set.")
}
if server.Port == 0 {
return nil, errors.New("Invalid Shadowsocks port.")
}
if server.Password == "" {
return nil, errors.New("Shadowsocks password is not specified.")
if len(v.Servers) == 1 {
server := v.Servers[0]
if C.Contains(shadowaead_2022.List, server.Cipher) {
if server.Address == nil {
return nil, errors.New("Shadowsocks server address is not set.")
}
if server.Port == 0 {
return nil, errors.New("Invalid Shadowsocks port.")
}
if server.Password == "" {
return nil, errors.New("Shadowsocks password is not specified.")
}
config := new(shadowsocks_2022.ClientConfig)
config.Address = server.Address.Build()
config.Port = uint32(server.Port)
config.Method = server.Cipher
config.Key = server.Password
return config, nil
}
}
if _, err := shadowsocks_2022.GetCipherMethod(server.Cipher); err == nil {
config := new(shadowsocks_2022.ClientConfig)
config.Address = server.Address.Build()
config.Port = uint32(server.Port)
config.Method = server.Cipher
config.Key = server.Password
return config, nil
}
config := new(shadowsocks.ClientConfig)
account := &shadowsocks.Account{
Password: server.Password,
for _, server := range v.Servers {
if C.Contains(shadowaead_2022.List, server.Cipher) {
return nil, errors.New("Shadowsocks 2022 accept no multi servers")
}
if server.Address == nil {
return nil, errors.New("Shadowsocks server address is not set.")
}
if server.Port == 0 {
return nil, errors.New("Invalid Shadowsocks port.")
}
if server.Password == "" {
return nil, errors.New("Shadowsocks password is not specified.")
}
account := &shadowsocks.Account{
Password: server.Password,
}
account.CipherType = cipherFromString(server.Cipher)
if account.CipherType == shadowsocks.CipherType_UNKNOWN {
return nil, errors.New("unknown cipher method: ", server.Cipher)
}
ss := &protocol.ServerEndpoint{
Address: server.Address.Build(),
Port: uint32(server.Port),
User: &protocol.User{
Level: uint32(server.Level),
Email: server.Email,
Account: serial.ToTypedMessage(account),
},
}
config.Server = ss
break
}
account.CipherType = cipherFromString(server.Cipher)
if account.CipherType == shadowsocks.CipherType_UNKNOWN {
return nil, errors.New("unknown cipher method: ", server.Cipher)
}
ss := &protocol.ServerEndpoint{
Address: server.Address.Build(),
Port: uint32(server.Port),
User: &protocol.User{
Level: uint32(server.Level),
Email: server.Email,
Account: serial.ToTypedMessage(account),
},
}
config.Server = ss
return config, nil
}
+44 -193
View File
@@ -1,7 +1,6 @@
package conf
import (
"context"
"crypto/x509"
"encoding/base64"
"encoding/hex"
@@ -15,7 +14,7 @@ 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"
"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"
@@ -83,7 +82,7 @@ var (
"noise": func() interface{} { return new(NoiseMask) },
"salamander": func() interface{} { return new(Salamander) },
"sudoku": func() interface{} { return new(Sudoku) },
"xdns": func() interface{} { return new(XDNS) },
"xdns": func() interface{} { return new(Xdns) },
"xicmp": func() interface{} { return new(Xicmp) },
"realm": func() interface{} { return new(Realm) },
"udphop": func() interface{} { return new(UDPHop) },
@@ -310,27 +309,14 @@ type NoiseMask struct {
}
func (c *NoiseMask) Build() (proto.Message, error) {
noiseSlice := make([]*noise.Item, 0, len(c.Noise))
for _, item := range c.Noise {
if len(item.Packet) > 0 && item.Rand.To > 0 {
return nil, errors.New("len(item.Packet) > 0 && item.Rand.To > 0")
}
if strings.ToLower(item.Type) == "exp" {
var exp string
if err := json.Unmarshal(item.Packet, &exp); err != nil {
return nil, errors.New(`"packet" of noise "type": "exp" must be a string`).Base(err)
}
segments, err := parseNoiseExp(exp)
if err != nil {
return nil, err
}
noiseSlice = append(noiseSlice, &noise.Item{
Segments: segments,
DelayMin: int64(item.Delay.From),
DelayMax: int64(item.Delay.To),
})
continue
}
}
noiseSlice := make([]*noise.Item, 0, len(c.Noise))
for _, item := range c.Noise {
if item.RandRange == nil {
item.RandRange = &Int32Range{From: 0, To: 255}
}
@@ -359,88 +345,6 @@ func (c *NoiseMask) Build() (proto.Message, error) {
}, nil
}
var noiseExpPattern = regexp.MustCompile(`<\s*([a-z]+)(?:\s+([^>]*?))?\s*>`)
func parseNoiseExp(exp string) ([]*noise.Segment, error) {
var segments []*noise.Segment
matches := noiseExpPattern.FindAllStringSubmatchIndex(exp, -1)
last := 0
for _, m := range matches {
if strings.TrimSpace(exp[last:m[0]]) != "" {
return nil, errors.New("invalid noise exp near ", exp[last:m[0]])
}
last = m[1]
key := exp[m[2]:m[3]]
arg := ""
if m[4] >= 0 {
arg = exp[m[4]:m[5]]
}
segment, err := buildNoiseSegment(key, arg)
if err != nil {
return nil, err
}
segments = append(segments, segment)
}
if strings.TrimSpace(exp[last:]) != "" {
return nil, errors.New("invalid noise exp near ", exp[last:])
}
if len(segments) == 0 {
return nil, errors.New("empty noise exp: ", exp)
}
return segments, nil
}
func buildNoiseSegment(key, arg string) (*noise.Segment, error) {
sizeSegment := func(kind noise.Segment_Kind) (*noise.Segment, error) {
if arg == "" {
return nil, errors.New("<", key, "> in noise exp needs a size")
}
lo, hi, err := ParseRangeString(arg)
if err != nil {
return nil, err
}
if lo < 0 || hi < lo || hi > 65535 {
return nil, errors.New("invalid size in noise exp: ", arg)
}
return &noise.Segment{Kind: kind, MinSize: int64(lo), MaxSize: int64(hi)}, nil
}
switch key {
case "b":
hexStr := strings.TrimPrefix(strings.TrimPrefix(strings.Join(strings.Fields(arg), ""), "0x"), "0X")
if len(hexStr) == 0 {
return nil, errors.New("empty bytes in noise exp")
}
raw, err := hex.DecodeString(hexStr)
if err != nil {
return nil, errors.New("invalid hex in noise exp: ", arg).Base(err)
}
return &noise.Segment{Kind: noise.Segment_BYTES, Bytes: raw}, nil
case "r":
return sizeSegment(noise.Segment_RANDOM)
case "rc":
return sizeSegment(noise.Segment_RANDOM_ASCII)
case "rd":
return sizeSegment(noise.Segment_RANDOM_DIGIT)
case "t":
if arg != "" {
return nil, errors.New("<t> in noise exp takes no argument")
}
return &noise.Segment{Kind: noise.Segment_TIMESTAMP}, nil
case "c":
if arg != "" {
return nil, errors.New("<c> in noise exp takes no argument")
}
return &noise.Segment{Kind: noise.Segment_COUNTER}, nil
case "n":
if arg != "" {
return nil, errors.New("<n> in noise exp takes no argument")
}
return &noise.Segment{Kind: noise.Segment_NONCE}, nil
default:
return nil, errors.New("unknown <", key, "> in noise exp")
}
}
type UDPItem struct {
Rand int32 `json:"rand"`
RandRange *Int32Range `json:"randRange"`
@@ -791,88 +695,32 @@ func (c *Sudoku) Build() (proto.Message, error) {
}, nil
}
type XDNSDomain struct {
Name string `json:"name"`
LenLimit int32 `json:"lenLimit"`
LabelLimit int32 `json:"labelLimit"`
Types []int32 `json:"types"`
Edns0 int32 `json:"edns0"`
type Xdns struct {
Domain json.RawMessage `json:"domain"`
Domains []string `json:"domains"`
Resolvers []string `json:"resolvers"`
}
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"`
}
type XDNS struct {
Domains []XDNSDomain `json:"domains"`
Resolvers []XDNSResolver `json:"resolvers"`
ExtraPoll int32 `json:"extraPoll"`
}
func (c *XDNS) Build() (proto.Message, error) {
var domains []*xdns.DomainProto
var resolvers []*serial.TypedMessage
for i := range c.Domains {
if c.Domains[i].LenLimit == 0 {
c.Domains[i].LenLimit = 255
}
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))
if err != nil {
return nil, err
}
errors.LogInfo(context.Background(), domain.Show())
domains = append(domains, &xdns.DomainProto{
Name: c.Domains[i].Name,
LenLimit: c.Domains[i].LenLimit,
LabelLimit: c.Domains[i].LabelLimit,
Types: c.Domains[i].Types,
Edns0: c.Domains[i].Edns0,
})
func (c *Xdns) Build() (proto.Message, error) {
if c.Domain != nil {
return nil, errors.PrintRemovedFeatureError("domain", "domains(server) & resolvers(client)")
}
for i := range c.Resolvers {
config, err := xdnsLoader.LoadWithID(c.Resolvers[i].Settings, c.Resolvers[i].Type)
if err != nil {
return nil, err
}
pm, err := config.(interface{ Build() (proto.Message, error) }).Build()
if err != nil {
return nil, err
}
resolvers = append(resolvers, serial.ToTypedMessage(pm))
if len(c.Domains) == 0 && len(c.Resolvers) == 0 {
return nil, errors.New("empty domains & empty resolvers")
}
if c.ExtraPoll < 0 || c.ExtraPoll > 3 {
return nil, errors.New("c.ExtraPoll < 0 || c.ExtraPoll > 3")
for _, r := range c.Resolvers {
if !strings.Contains(r, "+udp://") {
return nil, errors.New("invalid resolver ", r)
}
}
return &xdns.Config{Domains: domains, Resolvers: resolvers, ExtraPoll: c.ExtraPoll}, nil
return &xdns.Config{
Domains: c.Domains,
Resolvers: c.Resolvers,
}, nil
}
type XMC struct {
@@ -1061,13 +909,22 @@ func (c *Realm) Build() (proto.Message, error) {
}
type UDPHop struct {
Mode string `json:"mode"`
Interval Int32Range `json:"interval"`
RemoteIPs []string `json:"remoteIPs"`
RemotePorts PortList `json:"remotePorts"`
Sockopt *SocketConfig `json:"sockopt"`
Mode string `json:"mode"`
Interval Int32Range `json:"interval"`
RemotePorts PortList `json:"remotePorts"`
RemoteIPs []string `json:"remoteIPs"`
}
func (c *UDPHop) Build() (proto.Message, error) {
var sockopt *internet.SocketConfig
if c.Sockopt != nil {
var err error
sockopt, err = c.Sockopt.Build()
if err != nil {
return nil, err
}
}
var local, remote, remoteOnce bool
for _, mode := range strings.Split(c.Mode, ",") {
switch strings.ToLower(mode) {
@@ -1095,21 +952,15 @@ func (c *UDPHop) Build() (proto.Message, error) {
}
return nil, errors.New("invalid ip ", ip)
}
interval := c.Interval
if interval.From == 0 && interval.To == 0 {
interval.From, interval.To = 30, 30
}
if interval.From < 5 {
return nil, errors.New("interval must be at least 5")
}
return &udphop.Config{
Sockopt: sockopt,
Local: local,
Remote: remote,
RemoteOnce: remoteOnce,
IntervalMin: int64(interval.From),
IntervalMax: int64(interval.To),
RemoteIPs: remoteIPs,
IntervalMin: int64(c.Interval.From),
IntervalMax: int64(c.Interval.To),
RemotePorts: c.RemotePorts.Build().Ports(),
RemoteIPs: remoteIPs,
}, nil
}
@@ -1,136 +0,0 @@
package conf
import (
"encoding/json"
"testing"
"github.com/xtls/xray-core/transport/internet/finalmask/noise"
)
func expPacket(exp string) json.RawMessage {
b, _ := json.Marshal(exp)
return b
}
func buildNoiseExp(exp string) (*noise.Config, error) {
msg, err := (&NoiseMask{Noise: []NoiseItem{{Type: "exp", Packet: expPacket(exp)}}}).Build()
if err != nil {
return nil, err
}
return msg.(*noise.Config), nil
}
func TestNoiseExp(t *testing.T) {
cfg, err := buildNoiseExp("<b 0d0a0d0a><t><r 24><rc 20-40><rd 8><c><n>")
if err != nil {
t.Fatal(err)
}
segments := cfg.Items[0].Segments
if len(segments) != 7 {
t.Fatalf("got %d segments, want 7", len(segments))
}
want := []struct {
kind noise.Segment_Kind
bytes []byte
min, max int64
}{
{noise.Segment_BYTES, []byte{0x0d, 0x0a, 0x0d, 0x0a}, 0, 0},
{noise.Segment_TIMESTAMP, nil, 0, 0},
{noise.Segment_RANDOM, nil, 24, 24},
{noise.Segment_RANDOM_ASCII, nil, 20, 40},
{noise.Segment_RANDOM_DIGIT, nil, 8, 8},
{noise.Segment_COUNTER, nil, 0, 0},
{noise.Segment_NONCE, nil, 0, 0},
}
for i, w := range want {
s := segments[i]
if s.Kind != w.kind || s.MinSize != w.min || s.MaxSize != w.max || string(s.Bytes) != string(w.bytes) {
t.Errorf("segment %d = %+v, want %+v", i, s, w)
}
}
}
func TestNoiseExpStripsHexPrefix(t *testing.T) {
cfg, err := buildNoiseExp("<b 0x16030100>")
if err != nil {
t.Fatal(err)
}
if got := cfg.Items[0].Segments[0].Bytes; string(got) != string([]byte{0x16, 0x03, 0x01, 0x00}) {
t.Errorf("got %x", got)
}
}
func TestNoiseExpWhitespace(t *testing.T) {
if _, err := buildNoiseExp(" <b 00> <t> "); err != nil {
t.Errorf("surrounding whitespace should be allowed: %v", err)
}
cfg, err := buildNoiseExp("<b 0d 0a 0d 0a>")
if err != nil {
t.Fatal(err)
}
if got := cfg.Items[0].Segments[0].Bytes; string(got) != "\r\n\r\n" {
t.Errorf("got %x", got)
}
}
func TestNoiseExpRejects(t *testing.T) {
for _, exp := range []string{
"<x 1>",
"<b>",
"<b zz>",
"<b 0d0>",
"<r>",
"<r -1>",
"<r 40-20>",
"<r 70000>",
"<t 5>",
"<n 5>",
"garbage<t>",
"<t> tail",
"<t><b>",
} {
if _, err := buildNoiseExp(exp); err == nil {
t.Errorf("expected an error for %q", exp)
}
}
}
func TestNoiseExpConflicts(t *testing.T) {
if _, err := (&NoiseMask{Noise: []NoiseItem{{Type: "exp", Packet: expPacket("<t>"), Rand: Int32Range{From: 10, To: 20}}}}).Build(); err == nil {
t.Error("exp with rand should be rejected")
}
for _, packet := range []string{``, `[1, 2]`, `5`} {
if _, err := (&NoiseMask{Noise: []NoiseItem{{Type: "exp", Packet: json.RawMessage(packet)}}}).Build(); err == nil {
t.Errorf("expected an error for packet %q", packet)
}
}
}
func TestNoiseExpFromJSON(t *testing.T) {
var mask NoiseMask
if err := json.Unmarshal([]byte(`{"noise": [
{"type": "exp", "packet": "<b 504f5354><rd 10-20>", "delay": "1-3"},
{"type": "EXP", "packet": "<t>"},
{"type": "str", "packet": "<t>"},
{"rand": "10-20"}
]}`), &mask); err != nil {
t.Fatal(err)
}
msg, err := mask.Build()
if err != nil {
t.Fatal(err)
}
items := msg.(*noise.Config).Items
if len(items[0].Segments) != 2 || items[0].DelayMin != 1 || items[0].DelayMax != 3 {
t.Errorf("item 0 = %+v", items[0])
}
if len(items[1].Segments) != 1 || items[1].Segments[0].Kind != noise.Segment_TIMESTAMP {
t.Errorf("item 1 = %+v", items[1])
}
if len(items[2].Segments) != 0 || string(items[2].Packet) != "<t>" {
t.Errorf("item 2 = %+v", items[2])
}
if len(items[3].Segments) != 0 || items[3].RandMin != 10 || items[3].RandMax != 20 {
t.Errorf("item 3 = %+v", items[3])
}
}
-26
View File
@@ -36,10 +36,6 @@ func (p TransportProtocol) Build() (string, error) {
return "", errors.PrintRemovedFeatureError("QUIC transport (without web service, etc.)", "XHTTP stream-one H3")
case "hysteria":
return "hysteria", nil
case "masque":
return "masque", nil
case "xdrive":
return "xdrive", nil
default:
return "", errors.New("Config: unknown transport protocol: ", p)
}
@@ -63,8 +59,6 @@ type StreamConfig struct {
WSSettings *WebSocketConfig `json:"wsSettings"`
HTTPUPGRADESettings *HttpUpgradeConfig `json:"httpupgradeSettings"`
HysteriaSettings *HysteriaConfig `json:"hysteriaSettings"`
MASQUESettings *MasqueConfig `json:"masqueSettings"`
XDRIVESettings *XDriveConfig `json:"xdriveSettings"`
SocketSettings *SocketConfig `json:"sockopt"`
}
@@ -198,26 +192,6 @@ func (c *StreamConfig) Build() (*internet.StreamConfig, error) {
Settings: serial.ToTypedMessage(hs),
})
}
if c.MASQUESettings != nil {
ms, err := c.MASQUESettings.Build()
if err != nil {
return nil, errors.New("Failed to build MASQUE config.").Base(err)
}
config.TransportSettings = append(config.TransportSettings, &internet.TransportConfig{
ProtocolName: "masque",
Settings: serial.ToTypedMessage(ms),
})
}
if c.XDRIVESettings != nil {
xs, err := c.XDRIVESettings.Build()
if err != nil {
return nil, errors.New("Failed to build XDRIVE config.").Base(err)
}
config.TransportSettings = append(config.TransportSettings, &internet.TransportConfig{
ProtocolName: "xdrive",
Settings: serial.ToTypedMessage(xs),
})
}
if c.SocketSettings != nil {
ss, err := c.SocketSettings.Build()
if err != nil {
-109
View File
@@ -1,9 +1,7 @@
package conf
import (
"encoding/base64"
"encoding/json"
"maps"
"math/big"
"net/url"
"sort"
@@ -22,12 +20,9 @@ import (
"github.com/xtls/xray-core/transport/internet/httpupgrade"
"github.com/xtls/xray-core/transport/internet/hysteria"
"github.com/xtls/xray-core/transport/internet/kcp"
"github.com/xtls/xray-core/transport/internet/masque"
"github.com/xtls/xray-core/transport/internet/splithttp"
"github.com/xtls/xray-core/transport/internet/tcp"
"github.com/xtls/xray-core/transport/internet/websocket"
"github.com/xtls/xray-core/transport/internet/xdrive"
"golang.org/x/net/http/httpguts"
"google.golang.org/protobuf/proto"
)
@@ -790,63 +785,6 @@ func (c *HysteriaConfig) Build() (proto.Message, error) {
return config, nil
}
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"`
}
func (c *MasqueConfig) Build() (proto.Message, error) {
path := c.Path
if path == "" {
path = masque.DefaultPath
}
path = strings.NewReplacer(
"{target}", "*", "{ipproto}", "*",
"{?target,ipproto}", "?target=*&ipproto=*", "{?ipproto,target}", "?ipproto=*&target=*",
"{&target,ipproto}", "&target=*&ipproto=*", "{&ipproto,target}", "&ipproto=*&target=*",
).Replace(path)
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)
}
}
for k, v := range c.Headers {
if !httpguts.ValidHeaderFieldName(k) || !httpguts.ValidHeaderFieldValue(v) {
return nil, errors.New(`invalid header in "headers": `, strconv.Quote(k))
}
switch strings.ToLower(k) {
case "host", "capsule-protocol":
return nil, errors.New(`"headers" can't contain "`, k, `"`)
case "authorization":
if c.User != "" || c.Pass != "" {
return nil, errors.New(`"headers" can't contain "`, k, `" when "user" or "pass" is set`)
}
}
}
headers := c.Headers
if c.User != "" || c.Pass != "" {
if strings.Contains(c.User, ":") {
return nil, errors.New(`invalid "user": `, c.User)
}
headers = maps.Clone(c.Headers)
if headers == nil {
headers = make(map[string]string)
}
headers["Authorization"] = "Basic " + base64.StdEncoding.EncodeToString([]byte(c.User+":"+c.Pass))
}
return &masque.Config{
Host: c.Host,
Path: path,
Headers: headers,
}, nil
}
func readFileOrString(f string, s []string) ([]byte, error) {
if len(f) > 0 {
return filesystem.ReadCert(f)
@@ -856,50 +794,3 @@ func readFileOrString(f string, s []string) ([]byte, error) {
}
return nil, errors.New("both file and bytes are empty.")
}
type XDriveConfig struct {
RemoteFolder string `json:"remoteFolder"`
Service string `json:"service"`
Secrets []string `json:"secrets"`
SegmentBytes uint32 `json:"segmentBytes"`
FlushIntervalMs uint32 `json:"flushIntervalMs"`
PollIntervalMs uint32 `json:"pollIntervalMs"`
MaxPollIntervalMs uint32 `json:"maxPollIntervalMs"`
SessionTTLSeconds uint32 `json:"sessionTtlSeconds"`
Concurrency uint32 `json:"concurrency"`
EagerWindowMs uint32 `json:"eagerWindowMs"`
HoleTimeoutMs uint32 `json:"holeTimeoutMs"`
Template json.RawMessage `json:"template"`
}
// Build implements Buildable.
func (c *XDriveConfig) Build() (proto.Message, error) {
switch c.Service {
case "local":
case "Google Drive":
if len(c.Secrets) != 3 {
return nil, errors.New("Google Drive needs 3 secrets in order of ClientID, ClientSecret, RefreshToken")
}
case "template":
if len(c.Template) == 0 {
return nil, errors.New(`service "template" needs a "template" object`)
}
default:
return nil, errors.New("unsupported service")
}
config := &xdrive.Config{
RemoteFolder: c.RemoteFolder,
Service: c.Service,
Secrets: c.Secrets,
SegmentBytes: c.SegmentBytes,
FlushIntervalMs: c.FlushIntervalMs,
PollIntervalMs: c.PollIntervalMs,
MaxPollIntervalMs: c.MaxPollIntervalMs,
SessionTtlSeconds: c.SessionTTLSeconds,
Concurrency: c.Concurrency,
EagerWindowMs: c.EagerWindowMs,
HoleTimeoutMs: c.HoleTimeoutMs,
Template: string(c.Template),
}
return config, nil
}
-73
View File
@@ -291,76 +291,3 @@ func TestHeaderCustomUDPBuildRejectsExprWithoutArgs(t *testing.T) {
t.Fatalf("expected transform arg rejection, got %v", err)
}
}
func TestXDriveStreamConfig(t *testing.T) {
config := new(StreamConfig)
if err := json.Unmarshal([]byte(`{
"method": "xdrive",
"xdriveSettings": {
"remoteFolder": "/tmp/xdrive",
"service": "local"
}
}`), config); err != nil {
t.Fatalf("Unmarshal: %v", err)
}
built, err := config.Build()
if err != nil {
t.Fatalf("Build: %v", err)
}
if built.ProtocolName != "xdrive" {
t.Fatalf("ProtocolName is %q, want %q", built.ProtocolName, "xdrive")
}
if len(built.TransportSettings) != 1 || built.TransportSettings[0].ProtocolName != "xdrive" {
t.Fatalf("TransportSettings is %v, want a single xdrive entry", built.TransportSettings)
}
}
func TestXDriveRejectsUnknownService(t *testing.T) {
config := new(XDriveConfig)
if err := json.Unmarshal([]byte(`{"remoteFolder": "/tmp/xdrive", "service": "Dropbox"}`), config); err != nil {
t.Fatalf("Unmarshal: %v", err)
}
if _, err := config.Build(); err == nil {
t.Fatal("Build accepted an unsupported service")
}
}
func TestXDriveTemplateStreamConfig(t *testing.T) {
config := new(StreamConfig)
if err := json.Unmarshal([]byte(`{
"method": "xdrive",
"xdriveSettings": {
"remoteFolder": "folder",
"service": "template",
"secrets": ["user", "pass"],
"template": {
"flatten": true,
"auth": {"type": "basic", "username": "{secret0}", "password": "{secret1}"},
"put": {"method": "PUT", "url": "https://dav.example/{folder}/{name}"},
"get": {"method": "GET", "url": "https://dav.example/{folder}/{name}"},
"delete": {"method": "DELETE", "url": "https://dav.example/{folder}/{name}"},
"list": {"method": "PROPFIND", "url": "https://dav.example/{folder}/", "namesRegex": "<d:href>/folder/([^<]+)</d:href>"}
}
}
}`), config); err != nil {
t.Fatalf("Unmarshal: %v", err)
}
built, err := config.Build()
if err != nil {
t.Fatalf("Build: %v", err)
}
if built.ProtocolName != "xdrive" {
t.Fatalf("ProtocolName is %q, want xdrive", built.ProtocolName)
}
}
func TestXDriveTemplateNeedsTemplate(t *testing.T) {
config := new(XDriveConfig)
if err := json.Unmarshal([]byte(`{"remoteFolder": "f", "service": "template"}`), config); err != nil {
t.Fatalf("Unmarshal: %v", err)
}
if _, err := config.Build(); err == nil {
t.Fatal("Build accepted a template service without a template")
}
}
-32
View File
@@ -5,12 +5,8 @@ import (
"fmt"
"math/big"
"net"
"runtime"
"slices"
"strconv"
"strings"
"github.com/xtls/xray-core/common/errors"
"github.com/xtls/xray-core/proxy/tun"
"google.golang.org/protobuf/proto"
)
@@ -24,8 +20,6 @@ type TunConfig struct {
UserLevel uint32 `json:"userLevel"`
AutoSystemRoutingTable []string `json:"autoSystemRoutingTable"`
AutoOutboundsInterface *string `json:"autoOutboundsInterface"`
AutoSystemDnsToGateway bool `json:"autoSystemDnsToGateway"`
AutoSystemWfpBlockLeak []string `json:"autoSystemWfpBlockLeak"`
}
func (v *TunConfig) Build() (proto.Message, error) {
@@ -37,32 +31,6 @@ func (v *TunConfig) Build() (proto.Message, error) {
DNS: v.DNS,
UserLevel: v.UserLevel,
AutoSystemRoutingTable: v.AutoSystemRoutingTable,
AutoSystemDnsToGateway: v.AutoSystemDnsToGateway,
}
for _, leak := range v.AutoSystemWfpBlockLeak {
switch leak := strings.ToLower(leak); leak {
case "dns", "misconfigtun":
config.AutoSystemWfpBlockLeak = append(config.AutoSystemWfpBlockLeak, leak)
default:
return nil, errors.New("unknown autoSystemWfpBlockLeak value: ", leak)
}
}
// Each option needs other settings on the system it takes effect on: the
// filters go along with the routes of autoSystemRoutingTable, "dns" lets
// DNS through the TUN only, and autoSystemDnsToGateway points the system
// DNS at the gateway.
switch runtime.GOOS {
case "windows":
if len(config.AutoSystemWfpBlockLeak) > 0 && len(v.AutoSystemRoutingTable) == 0 {
return nil, errors.New("autoSystemWfpBlockLeak needs autoSystemRoutingTable to be set")
}
if slices.Contains(config.AutoSystemWfpBlockLeak, "dns") && len(v.DNS) == 0 {
return nil, errors.New(`autoSystemWfpBlockLeak "dns" needs dns to be set`)
}
case "linux":
if v.AutoSystemDnsToGateway && len(v.Gateway) == 0 {
return nil, errors.New("autoSystemDnsToGateway needs gateway to be set")
}
}
if v.AutoOutboundsInterface != nil {
config.AutoOutboundsInterface = *v.AutoOutboundsInterface
-71
View File
@@ -1,71 +0,0 @@
package conf_test
import (
"encoding/json"
"runtime"
"testing"
. "github.com/xtls/xray-core/infra/conf"
"github.com/xtls/xray-core/proxy/tun"
)
func TestTunConfigAutoSystem(t *testing.T) {
creator := func() Buildable {
return new(TunConfig)
}
runMultiTestCase(t, []TestCase{
{
Input: `{"name": "xray0"}`,
Parser: loadJSON(creator),
Output: &tun.Config{Name: "xray0", Desc: "Wintun", MTU: 1500},
},
{
Input: `{"name": "xray0", "gateway": ["10.0.0.1/24"], "autoSystemDnsToGateway": true}`,
Parser: loadJSON(creator),
Output: &tun.Config{Name: "xray0", Desc: "Wintun", MTU: 1500, Gateway: []string{"10.0.0.1/24"}, AutoSystemDnsToGateway: true},
},
{
Input: `{"name": "xray0", "dns": ["1.1.1.1"], "autoSystemRoutingTable": ["0.0.0.0/0"], "autoSystemWfpBlockLeak": ["dns", "misconfigtun"]}`,
Parser: loadJSON(creator),
Output: &tun.Config{Name: "xray0", Desc: "Wintun", MTU: 1500, DNS: []string{"1.1.1.1"}, AutoSystemRoutingTable: []string{"0.0.0.0/0"}, AutoOutboundsInterface: "auto", AutoSystemWfpBlockLeak: []string{"dns", "misconfigtun"}},
},
{
Input: `{"name": "xray0", "dns": ["1.1.1.1"], "autoSystemRoutingTable": ["0.0.0.0/0"], "autoSystemWfpBlockLeak": ["DNS"]}`,
Parser: loadJSON(creator),
Output: &tun.Config{Name: "xray0", Desc: "Wintun", MTU: 1500, DNS: []string{"1.1.1.1"}, AutoSystemRoutingTable: []string{"0.0.0.0/0"}, AutoOutboundsInterface: "auto", AutoSystemWfpBlockLeak: []string{"dns"}},
},
})
}
// TestTunConfigAutoSystemNeeds checks that an option is rejected without the
// setting it needs, only on the system it takes effect on.
func TestTunConfigAutoSystemNeeds(t *testing.T) {
for _, c := range []struct {
input string
goos string // where it is rejected
}{
{`{"name": "xray0", "autoSystemWfpBlockLeak": ["misconfigtun"]}`, "windows"},
{`{"name": "xray0", "autoSystemRoutingTable": ["0.0.0.0/0"], "autoSystemWfpBlockLeak": ["misconfigtun"]}`, ""},
{`{"name": "xray0", "autoSystemRoutingTable": ["0.0.0.0/0"], "autoSystemWfpBlockLeak": ["dns"]}`, "windows"},
{`{"name": "xray0", "autoSystemDnsToGateway": true}`, "linux"},
} {
config := new(TunConfig)
if err := json.Unmarshal([]byte(c.input), config); err != nil {
t.Fatal(err)
}
if _, err := config.Build(); (err != nil) != (runtime.GOOS == c.goos) {
t.Errorf("%s on %s: error = %v", c.input, runtime.GOOS, err)
}
}
}
func TestTunConfigAutoSystemWfpBlockLeakUnknown(t *testing.T) {
config := new(TunConfig)
if err := json.Unmarshal([]byte(`{"name": "xray0", "autoSystemWfpBlockLeak": ["dns", "ip"]}`), config); err != nil {
t.Fatal(err)
}
if _, err := config.Build(); err == nil {
t.Error("an unknown autoSystemWfpBlockLeak value was accepted")
}
}
+23 -7
View File
@@ -59,13 +59,14 @@ func (c *WireGuardPeerConfig) Build() (*wireguard.PeerConfig, error) {
type WireGuardConfig struct {
IsClient bool `json:""`
NoKernelTun bool `json:"noKernelTun"`
SecretKey string `json:"secretKey"`
Address []string `json:"address"`
Peers []*WireGuardPeerConfig `json:"peers"`
MTU int32 `json:"mtu"`
Reserved []byte `json:"reserved"`
DNS []string `json:"remoteDNS"`
NoKernelTun bool `json:"noKernelTun"`
SecretKey string `json:"secretKey"`
Address []string `json:"address"`
Peers []*WireGuardPeerConfig `json:"peers"`
MTU int32 `json:"mtu"`
Reserved []byte `json:"reserved"`
DomainStrategy string `json:"domainStrategy"`
DNS []string `json:"remoteDNS"`
}
func (c *WireGuardConfig) Build() (proto.Message, error) {
@@ -124,6 +125,21 @@ func (c *WireGuardConfig) Build() (proto.Message, error) {
}
config.Reserved = c.Reserved
switch strings.ToLower(c.DomainStrategy) {
case "forceip", "":
config.DomainStrategy = wireguard.DeviceConfig_FORCE_IP
case "forceipv4":
config.DomainStrategy = wireguard.DeviceConfig_FORCE_IP4
case "forceipv6":
config.DomainStrategy = wireguard.DeviceConfig_FORCE_IP6
case "forceipv4v6":
config.DomainStrategy = wireguard.DeviceConfig_FORCE_IP46
case "forceipv6v4":
config.DomainStrategy = wireguard.DeviceConfig_FORCE_IP64
default:
return nil, errors.New("unsupported domain strategy: ", c.DomainStrategy)
}
config.IsClient = c.IsClient
config.NoKernelTun = c.NoKernelTun
config.DNS = c.DNS
-14
View File
@@ -16,7 +16,6 @@ import (
"github.com/xtls/xray-core/common/serial"
core "github.com/xtls/xray-core/core"
"github.com/xtls/xray-core/proxy/freedom"
"github.com/xtls/xray-core/proxy/masque"
"github.com/xtls/xray-core/transport/internet"
)
@@ -33,7 +32,6 @@ var (
"trojan": func() interface{} { return new(TrojanServerConfig) },
"wireguard": func() interface{} { return &WireGuardConfig{IsClient: false} },
"hysteria": func() interface{} { return new(HysteriaServerConfig) },
"masque": func() interface{} { return new(MasqueServerConfig) },
"tun": func() interface{} { return new(TunConfig) },
}, "protocol", "settings")
@@ -50,7 +48,6 @@ var (
"vmess": func() interface{} { return new(VMessOutboundConfig) },
"trojan": func() interface{} { return new(TrojanClientConfig) },
"hysteria": func() interface{} { return new(HysteriaClientConfig) },
"masque": func() interface{} { return new(MasqueClientConfig) },
"dns": func() interface{} { return new(DNSOutboundConfig) },
"wireguard": func() interface{} { return &WireGuardConfig{IsClient: true} },
}, "protocol", "settings")
@@ -206,9 +203,6 @@ func (c *InboundDetourConfig) Build() (*core.InboundHandlerConfig, error) {
if err != nil {
return nil, errors.New("failed to build inbound handler for protocol ", c.Protocol).Base(err)
}
if _, ok := ts.(*masque.ServerConfig); !ok && receiverSettings.StreamSettings != nil && receiverSettings.StreamSettings.ProtocolName == "masque" {
return nil, errors.New("the masque transport can only be used by the masque inbound")
}
return &core.InboundHandlerConfig{
Tag: c.Tag,
@@ -344,14 +338,6 @@ func (c *OutboundDetourConfig) Build() (*core.OutboundHandlerConfig, error) {
return nil, err
}
if _, ok := ts.(*masque.ClientConfig); ok {
if ms := senderSettings.MultiplexSettings; ms != nil && ms.Enabled {
return nil, errors.New(`masque outbound does not support "mux"`)
}
} else if senderSettings.StreamSettings != nil && senderSettings.StreamSettings.ProtocolName == "masque" {
return nil, errors.New("the masque transport can only be used by the masque outbound")
}
if fc, ok := ts.(*freedom.Config); ok {
if senderSettings.StreamSettings != nil &&
senderSettings.StreamSettings.SocketSettings != nil &&
@@ -12,8 +12,6 @@ import (
"github.com/xtls/xray-core/core"
"github.com/xtls/xray-core/infra/conf"
"github.com/xtls/xray-core/infra/conf/serial"
"github.com/xtls/xray-core/proxy/hysteria"
"github.com/xtls/xray-core/proxy/masque"
"github.com/xtls/xray-core/proxy/shadowsocks"
"github.com/xtls/xray-core/proxy/shadowsocks_2022"
"github.com/xtls/xray-core/proxy/trojan"
@@ -90,10 +88,6 @@ func extractInboundUsers(inb *core.InboundHandlerConfig) []*protocol.User {
return ty.Users
case *shadowsocks_2022.MultiUserServerConfig:
return ty.Users
case *masque.ServerConfig:
return ty.Users
case *hysteria.ServerConfig:
return ty.Users
default:
fmt.Println("unsupported inbound type")
}
-3
View File
@@ -41,7 +41,6 @@ import (
_ "github.com/xtls/xray-core/proxy/freedom"
_ "github.com/xtls/xray-core/proxy/http"
_ "github.com/xtls/xray-core/proxy/loopback"
_ "github.com/xtls/xray-core/proxy/masque"
_ "github.com/xtls/xray-core/proxy/shadowsocks"
_ "github.com/xtls/xray-core/proxy/socks"
_ "github.com/xtls/xray-core/proxy/trojan"
@@ -55,14 +54,12 @@ import (
_ "github.com/xtls/xray-core/transport/internet/grpc"
_ "github.com/xtls/xray-core/transport/internet/httpupgrade"
_ "github.com/xtls/xray-core/transport/internet/kcp"
_ "github.com/xtls/xray-core/transport/internet/masque"
_ "github.com/xtls/xray-core/transport/internet/reality"
_ "github.com/xtls/xray-core/transport/internet/splithttp"
_ "github.com/xtls/xray-core/transport/internet/tcp"
_ "github.com/xtls/xray-core/transport/internet/tls"
_ "github.com/xtls/xray-core/transport/internet/udp"
_ "github.com/xtls/xray-core/transport/internet/websocket"
_ "github.com/xtls/xray-core/transport/internet/xdrive"
// Transport headers
_ "github.com/xtls/xray-core/transport/internet/headers/http"
+4 -4
View File
@@ -467,7 +467,7 @@ func NewPacketReader(conn net.Conn, h *Handler, defaultRule *FinalRule, UDPOverr
if statConn != nil {
counter = statConn.ReadCounter
}
if c, ok := iConn.(*net.PacketConnWrapper); ok {
if c, ok := iConn.(*internet.PacketConnWrapper); ok {
isOverridden := false
if UDPOverride.Address != nil || UDPOverride.Port != 0 {
isOverridden = true
@@ -487,7 +487,7 @@ func NewPacketReader(conn net.Conn, h *Handler, defaultRule *FinalRule, UDPOverr
}
type PacketReader struct {
*net.PacketConnWrapper
*internet.PacketConnWrapper
stats.Counter
Handler *Handler
DefaultRule *FinalRule
@@ -542,7 +542,7 @@ func NewPacketWriter(conn net.Conn, h *Handler, defaultRule *FinalRule, UDPOverr
if statConn != nil {
counter = statConn.WriteCounter
}
if c, ok := iConn.(*net.PacketConnWrapper); ok {
if c, ok := iConn.(*internet.PacketConnWrapper); ok {
// If DialDest is a domain, it will be resolved in dialer
// check this behavior and add it to map
resolvedUDPAddr := utils.NewTypedSyncMap[string, net.Address]()
@@ -563,7 +563,7 @@ func NewPacketWriter(conn net.Conn, h *Handler, defaultRule *FinalRule, UDPOverr
}
type PacketWriter struct {
*net.PacketConnWrapper
*internet.PacketConnWrapper
stats.Counter
*Handler
DefaultRule *FinalRule
+3 -3
View File
@@ -236,14 +236,14 @@ type UDPReader struct {
func (r *UDPReader) ReadFrom(p []byte) (n int, addr *net.Destination, err error) {
for {
var packet [1500]byte
var buf [hysteria.MaxDatagramFrameSize]byte
n, err := r.reader.Read(packet[:])
n, err := r.reader.Read(buf[:])
if err != nil {
return 0, nil, err
}
msg, err := ParseUDPMessage(packet[:n])
msg, err := ParseUDPMessage(buf[:n])
if err != nil {
continue
}
-108
View File
@@ -1,108 +0,0 @@
package masque
import (
"crypto/subtle"
"strings"
"sync"
"github.com/xtls/xray-core/common/errors"
"github.com/xtls/xray-core/common/protocol"
"google.golang.org/protobuf/proto"
)
func (a *Account) AsAccount() (protocol.Account, error) {
return &MemoryAccount{Password: a.Password}, nil
}
type MemoryAccount struct {
Password string
}
func (a *MemoryAccount) Equals(other protocol.Account) bool {
b, ok := other.(*MemoryAccount)
return ok && a.Password == b.Password
}
func (a *MemoryAccount) ToProto() proto.Message {
return &Account{Password: a.Password}
}
type validator struct {
mu sync.RWMutex
users map[string]*protocol.MemoryUser
}
func newValidator() *validator {
return &validator{users: make(map[string]*protocol.MemoryUser)}
}
func (v *validator) add(user *protocol.MemoryUser) error {
account, ok := user.Account.(*MemoryAccount)
if !ok {
return errors.New("not a MASQUE account")
}
if user.Email == "" || strings.Contains(user.Email, ":") {
return errors.New("invalid email ", user.Email)
}
if account.Password == "" {
return errors.New("empty password for ", user.Email)
}
email := strings.ToLower(user.Email)
v.mu.Lock()
defer v.mu.Unlock()
if _, found := v.users[email]; found {
return errors.New("user ", user.Email, " already exists")
}
v.users[email] = user
return nil
}
func (v *validator) delByEmail(email string) (*protocol.MemoryUser, error) {
key := strings.ToLower(email)
v.mu.Lock()
defer v.mu.Unlock()
user, found := v.users[key]
if !found {
return nil, errors.New("user ", email, " not found")
}
delete(v.users, key)
return user, nil
}
func (v *validator) contains(user *protocol.MemoryUser) bool {
v.mu.RLock()
defer v.mu.RUnlock()
return v.users[strings.ToLower(user.Email)] == user
}
func (v *validator) get(email, password string) *protocol.MemoryUser {
v.mu.RLock()
user := v.users[strings.ToLower(email)]
v.mu.RUnlock()
if user == nil || subtle.ConstantTimeCompare([]byte(user.Account.(*MemoryAccount).Password), []byte(password)) != 1 {
return nil
}
return user
}
func (v *validator) getByEmail(email string) *protocol.MemoryUser {
v.mu.RLock()
defer v.mu.RUnlock()
return v.users[strings.ToLower(email)]
}
func (v *validator) getAll() []*protocol.MemoryUser {
v.mu.RLock()
defer v.mu.RUnlock()
users := make([]*protocol.MemoryUser, 0, len(v.users))
for _, user := range v.users {
users = append(users, user)
}
return users
}
func (v *validator) count() int64 {
v.mu.RLock()
defer v.mu.RUnlock()
return int64(len(v.users))
}
-49
View File
@@ -1,49 +0,0 @@
package masque
import (
"testing"
"github.com/stretchr/testify/require"
"github.com/xtls/xray-core/common/protocol"
)
func TestValidator(t *testing.T) {
v := newValidator()
user := &protocol.MemoryUser{Email: "U@example.com", Account: &MemoryAccount{Password: "p"}}
require.NoError(t, v.add(user))
for _, u := range []*protocol.MemoryUser{
{Email: "u@example.com", Account: &MemoryAccount{Password: "other"}},
{Account: &MemoryAccount{Password: "p"}},
{Email: "a:b", Account: &MemoryAccount{Password: "p"}},
{Email: "b@example.com", Account: &MemoryAccount{}},
} {
require.Error(t, v.add(u), u.Email)
}
require.Equal(t, user, v.get("u@example.com", "p"))
require.Equal(t, user, v.get("U@EXAMPLE.COM", "p"))
require.Nil(t, v.get("u@example.com", "x"))
require.Nil(t, v.get("x@example.com", "p"))
require.Nil(t, v.get("", ""))
require.Equal(t, user, v.getByEmail("u@example.com"))
require.Equal(t, []*protocol.MemoryUser{user}, v.getAll())
require.Equal(t, int64(1), v.count())
require.True(t, v.contains(user))
removed, err := v.delByEmail("u@EXAMPLE.com")
require.NoError(t, err)
require.Equal(t, user, removed)
_, err = v.delByEmail("u@example.com")
require.Error(t, err)
require.False(t, v.contains(user))
require.Nil(t, v.get("u@example.com", "p"))
require.Zero(t, v.count())
}
func TestAccount(t *testing.T) {
account, err := (&Account{Password: "p"}).AsAccount()
require.NoError(t, err)
require.True(t, account.Equals(&MemoryAccount{Password: "p"}))
require.False(t, account.Equals(&MemoryAccount{Password: "x"}))
require.Equal(t, &Account{Password: "p"}, account.ToProto())
}
-328
View File
@@ -1,328 +0,0 @@
package masque
import (
"context"
go_errors "errors"
"io"
"net/netip"
"slices"
"sync"
"sync/atomic"
"time"
"golang.zx2c4.com/wireguard/tun"
"github.com/xtls/xray-core/common"
"github.com/xtls/xray-core/common/buf"
"github.com/xtls/xray-core/common/errors"
"github.com/xtls/xray-core/common/net"
"github.com/xtls/xray-core/common/protocol"
"github.com/xtls/xray-core/common/session"
"github.com/xtls/xray-core/common/signal"
"github.com/xtls/xray-core/common/task"
"github.com/xtls/xray-core/core"
"github.com/xtls/xray-core/features/policy"
"github.com/xtls/xray-core/proxy/wireguard"
"github.com/xtls/xray-core/transport"
"github.com/xtls/xray-core/transport/internet"
"github.com/xtls/xray-core/transport/internet/masque"
"github.com/xtls/xray-core/transport/internet/stat"
"github.com/xtls/xray-core/transport/internet/tls"
)
const (
establishTimeout = 10 * time.Second
retryInterval = time.Second
)
type Client struct {
server *protocol.ServerSpec
policyManager policy.Manager
remoteDNS []netip.Addr
ctx context.Context
cancel context.CancelFunc
tunnel atomic.Pointer[tunnel]
mu sync.Mutex
lastErr error
lastErrAt time.Time
}
func NewClient(ctx context.Context, config *ClientConfig) (*Client, error) {
v := core.MustFromContext(ctx)
p := v.GetFeature(policy.ManagerType()).(policy.Manager)
streamSettings := session.StreamSettingsFromContext(ctx).(*internet.MemoryStreamConfig)
if _, ok := streamSettings.ProtocolSettings.(*masque.Config); !ok {
return nil, errors.New("not masque transport")
}
if tls.ConfigFromStreamSettings(streamSettings) == nil {
return nil, errors.New(`MASQUE requires "security": "tls"`)
}
if config.Server == nil {
return nil, errors.New(`no target server found`)
}
server, err := protocol.NewServerSpecFromPB(config.Server)
if err != nil {
return nil, errors.New("failed to get server spec").Base(err)
}
dns := config.RemoteDns
if len(dns) == 0 {
dns = []string{"1.1.1.1", "1.0.0.1", "2606:4700:4700::1111", "2606:4700:4700::1001"}
}
remoteDNS := make([]netip.Addr, 0, len(dns))
for _, s := range dns {
addr, err := netip.ParseAddr(s)
if err != nil {
return nil, errors.New("invalid remote DNS server ", s).Base(err)
}
remoteDNS = append(remoteDNS, addr)
}
c := &Client{
server: server,
policyManager: p,
remoteDNS: remoteDNS,
}
c.ctx, c.cancel = context.WithCancel(context.Background())
return c, nil
}
func (c *Client) Process(ctx context.Context, link *transport.Link, dialer internet.Dialer) error {
outbounds := session.OutboundsFromContext(ctx)
ob := outbounds[len(outbounds)-1]
if !ob.Target.IsValid() {
return errors.New("target not specified")
}
ob.Name = "masque"
ob.CanSpliceCopy = 3
t, err := c.getTunnel(ctx, dialer)
if err != nil {
return errors.New("failed to establish CONNECT-IP tunnel").Base(err)
}
var newCtx context.Context
var newCancel context.CancelFunc
if session.TimeoutOnlyFromContext(ctx) {
newCtx, newCancel = context.WithCancel(context.Background())
}
sessionPolicy := c.policyManager.ForLevel(0)
ctx, cancel := context.WithCancel(ctx)
timer := signal.CancelAfterInactivity(ctx, func() {
cancel()
if newCancel != nil {
newCancel()
}
}, sessionPolicy.Timeouts.ConnectionIdle)
if newCtx != nil {
ctx = newCtx
}
var reader buf.Reader
var writer buf.Writer
switch ob.Target.Network {
case net.Network_TCP:
var conn net.Conn
var err error
if sessionPolicy.Timeouts.Handshake != 0 {
timeoutCtx, timeoutCancel := context.WithTimeout(ctx, sessionPolicy.Timeouts.Handshake)
conn, err = t.tnet.DialContext(timeoutCtx, "tcp", ob.Target.NetAddr())
timeoutCancel()
} else {
conn, err = t.tnet.Dial("tcp", ob.Target.NetAddr())
}
if err != nil {
return errors.New("failed to create TCP connection").Base(err)
}
defer conn.Close()
reader = buf.NewReader(conn)
writer = buf.NewWriter(conn)
case net.Network_UDP:
conn, err := t.tnet.Dial("udp", ob.Target.NetAddr())
if err != nil {
return errors.New("failed to create UDP connection").Base(err)
}
defer conn.Close()
uc := &wireguard.UDPConnClient{
PacketConn: conn.(*net.PacketConnWrapper).PacketConn,
Dest: conn.RemoteAddr().(*net.UDPAddr),
}
reader = uc
writer = uc
default:
panic(ob.Target.Network)
}
requestFunc := func() error {
defer timer.SetTimeout(sessionPolicy.Timeouts.DownlinkOnly)
return buf.Copy(link.Reader, writer, buf.UpdateActivity(timer))
}
responseFunc := func() error {
defer timer.SetTimeout(sessionPolicy.Timeouts.UplinkOnly)
return buf.Copy(reader, link.Writer, buf.UpdateActivity(timer))
}
responseDonePost := task.OnSuccess(responseFunc, task.Close(link.Writer))
if err := task.Run(ctx, requestFunc, responseDonePost); err != nil {
common.Interrupt(link.Reader)
common.Interrupt(link.Writer)
return errors.New("connection ends").Base(err)
}
return nil
}
func (c *Client) getTunnel(ctx context.Context, dialer internet.Dialer) (*tunnel, error) {
c.mu.Lock()
defer c.mu.Unlock()
if c.ctx.Err() != nil {
return nil, errors.New("closed")
}
if t := c.tunnel.Load(); t != nil {
select {
case <-t.done:
default:
return t, nil
}
}
if err := ctx.Err(); err != nil {
return nil, err
}
if c.lastErr != nil && time.Since(c.lastErrAt) < retryInterval {
return nil, c.lastErr
}
t, err := c.establish(ctx, dialer)
if err != nil {
c.lastErr, c.lastErrAt = err, time.Now()
return nil, err
}
c.lastErr = nil
c.tunnel.Store(t)
if c.ctx.Err() != nil {
if c.tunnel.CompareAndSwap(t, nil) {
t.close()
}
return nil, errors.New("closed")
}
return t, nil
}
func (c *Client) establish(ctx context.Context, dialer internet.Dialer) (*tunnel, error) {
ctx, cancel := context.WithTimeout(context.WithoutCancel(ctx), establishTimeout)
defer cancel()
defer context.AfterFunc(c.ctx, cancel)()
conn, err := dialer.Dial(ctx, c.server.Destination)
if err != nil {
return nil, err
}
mconn, ok := stat.TryUnwrapStatsConn(conn).(*masque.Conn)
if !ok {
conn.Close()
return nil, errors.New("not a CONNECT-IP connection")
}
t, err := newTunnel(conn, mconn.LocalAddrs(), c.remoteDNS)
if err != nil {
conn.Close()
return nil, err
}
errors.LogInfo(ctx, "MASQUE: tunnel established from ", mconn.LocalAddrs())
return t, nil
}
func (c *Client) Close() error {
c.cancel()
if t := c.tunnel.Swap(nil); t != nil {
t.close()
}
return nil
}
type tunnel struct {
conn stat.Connection
dev tun.Device
tnet *wireguard.Net
done chan struct{}
closeOnce sync.Once
}
func newTunnel(conn stat.Connection, local []netip.Addr, remoteDNS []netip.Addr) (*tunnel, error) {
var dns []netip.Addr
for _, addr := range remoteDNS {
if slices.ContainsFunc(local, func(l netip.Addr) bool { return l.Is4() == addr.Is4() }) {
dns = append(dns, addr)
}
}
if len(dns) == 0 {
errors.LogWarning(context.Background(), "MASQUE: no remote DNS server is reachable from the assigned addresses ", local, ", domain names will fail to resolve")
dns = remoteDNS
}
dev, tnet, _, err := wireguard.CreateNetTUN(local, dns, masque.MinPacketSize, true)
if err != nil {
return nil, err
}
t := &tunnel{
conn: conn,
dev: dev,
tnet: tnet,
done: make(chan struct{}),
}
go t.readFromTunnel()
go t.writeToTunnel()
return t, nil
}
func (t *tunnel) readFromTunnel() {
defer t.close()
b := make([]byte, buf.Size)
for {
n, err := t.conn.Read(b)
if err != nil {
if go_errors.Is(err, io.ErrShortBuffer) {
continue
}
errors.LogInfoInner(context.Background(), err, "MASQUE: tunnel closed")
return
}
t.dev.Write([][]byte{b[:n]}, 0)
}
}
func (t *tunnel) writeToTunnel() {
bufs := [][]byte{make([]byte, masque.MinPacketSize)}
sizes := []int{0}
for {
if _, err := t.dev.Read(bufs, sizes, 0); err != nil {
return
}
if _, err := t.conn.Write(bufs[0][:sizes[0]]); err != nil {
var ptb *masque.PacketTooBigError
if go_errors.As(err, &ptb) {
go t.dev.Write([][]byte{ptb.ICMP}, 0)
}
}
}
}
func (t *tunnel) close() {
t.closeOnce.Do(func() {
close(t.done)
t.conn.Close()
t.dev.Close()
})
}
func init() {
common.Must(common.RegisterConfig((*ClientConfig)(nil), func(ctx context.Context, config interface{}) (interface{}, error) {
return NewClient(ctx, config.(*ClientConfig))
}))
}
-250
View File
@@ -1,250 +0,0 @@
// Code generated by protoc-gen-go. DO NOT EDIT.
// versions:
// protoc-gen-go v1.36.11
// protoc v6.33.5
// source: proxy/masque/config.proto
package masque
import (
protocol "github.com/xtls/xray-core/common/protocol"
protoreflect "google.golang.org/protobuf/reflect/protoreflect"
protoimpl "google.golang.org/protobuf/runtime/protoimpl"
reflect "reflect"
sync "sync"
unsafe "unsafe"
)
const (
// Verify that this generated code is sufficiently up-to-date.
_ = protoimpl.EnforceVersion(20 - protoimpl.MinVersion)
// Verify that runtime/protoimpl is sufficiently up-to-date.
_ = protoimpl.EnforceVersion(protoimpl.MaxVersion - 20)
)
type ClientConfig struct {
state protoimpl.MessageState `protogen:"open.v1"`
Server *protocol.ServerEndpoint `protobuf:"bytes,1,opt,name=server,proto3" json:"server,omitempty"`
RemoteDns []string `protobuf:"bytes,2,rep,name=remote_dns,json=remoteDns,proto3" json:"remote_dns,omitempty"`
unknownFields protoimpl.UnknownFields
sizeCache protoimpl.SizeCache
}
func (x *ClientConfig) Reset() {
*x = ClientConfig{}
mi := &file_proxy_masque_config_proto_msgTypes[0]
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
ms.StoreMessageInfo(mi)
}
func (x *ClientConfig) String() string {
return protoimpl.X.MessageStringOf(x)
}
func (*ClientConfig) ProtoMessage() {}
func (x *ClientConfig) ProtoReflect() protoreflect.Message {
mi := &file_proxy_masque_config_proto_msgTypes[0]
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 ClientConfig.ProtoReflect.Descriptor instead.
func (*ClientConfig) Descriptor() ([]byte, []int) {
return file_proxy_masque_config_proto_rawDescGZIP(), []int{0}
}
func (x *ClientConfig) GetServer() *protocol.ServerEndpoint {
if x != nil {
return x.Server
}
return nil
}
func (x *ClientConfig) GetRemoteDns() []string {
if x != nil {
return x.RemoteDns
}
return nil
}
type Account struct {
state protoimpl.MessageState `protogen:"open.v1"`
Password string `protobuf:"bytes,1,opt,name=password,proto3" json:"password,omitempty"`
unknownFields protoimpl.UnknownFields
sizeCache protoimpl.SizeCache
}
func (x *Account) Reset() {
*x = Account{}
mi := &file_proxy_masque_config_proto_msgTypes[1]
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
ms.StoreMessageInfo(mi)
}
func (x *Account) String() string {
return protoimpl.X.MessageStringOf(x)
}
func (*Account) ProtoMessage() {}
func (x *Account) ProtoReflect() protoreflect.Message {
mi := &file_proxy_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 Account.ProtoReflect.Descriptor instead.
func (*Account) Descriptor() ([]byte, []int) {
return file_proxy_masque_config_proto_rawDescGZIP(), []int{1}
}
func (x *Account) GetPassword() string {
if x != nil {
return x.Password
}
return ""
}
type ServerConfig struct {
state protoimpl.MessageState `protogen:"open.v1"`
Users []*protocol.User `protobuf:"bytes,1,rep,name=users,proto3" json:"users,omitempty"`
Address []string `protobuf:"bytes,2,rep,name=address,proto3" json:"address,omitempty"`
Mtu uint32 `protobuf:"varint,3,opt,name=mtu,proto3" json:"mtu,omitempty"`
unknownFields protoimpl.UnknownFields
sizeCache protoimpl.SizeCache
}
func (x *ServerConfig) Reset() {
*x = ServerConfig{}
mi := &file_proxy_masque_config_proto_msgTypes[2]
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
ms.StoreMessageInfo(mi)
}
func (x *ServerConfig) String() string {
return protoimpl.X.MessageStringOf(x)
}
func (*ServerConfig) ProtoMessage() {}
func (x *ServerConfig) ProtoReflect() protoreflect.Message {
mi := &file_proxy_masque_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 ServerConfig.ProtoReflect.Descriptor instead.
func (*ServerConfig) Descriptor() ([]byte, []int) {
return file_proxy_masque_config_proto_rawDescGZIP(), []int{2}
}
func (x *ServerConfig) GetUsers() []*protocol.User {
if x != nil {
return x.Users
}
return nil
}
func (x *ServerConfig) GetAddress() []string {
if x != nil {
return x.Address
}
return nil
}
func (x *ServerConfig) GetMtu() uint32 {
if x != nil {
return x.Mtu
}
return 0
}
var File_proxy_masque_config_proto protoreflect.FileDescriptor
const file_proxy_masque_config_proto_rawDesc = "" +
"\n" +
"\x19proxy/masque/config.proto\x12\x11xray.proxy.masque\x1a!common/protocol/server_spec.proto\x1a\x1acommon/protocol/user.proto\"k\n" +
"\fClientConfig\x12<\n" +
"\x06server\x18\x01 \x01(\v2$.xray.common.protocol.ServerEndpointR\x06server\x12\x1d\n" +
"\n" +
"remote_dns\x18\x02 \x03(\tR\tremoteDns\"%\n" +
"\aAccount\x12\x1a\n" +
"\bpassword\x18\x01 \x01(\tR\bpassword\"l\n" +
"\fServerConfig\x120\n" +
"\x05users\x18\x01 \x03(\v2\x1a.xray.common.protocol.UserR\x05users\x12\x18\n" +
"\aaddress\x18\x02 \x03(\tR\aaddress\x12\x10\n" +
"\x03mtu\x18\x03 \x01(\rR\x03mtuBU\n" +
"\x15com.xray.proxy.masqueP\x01Z&github.com/xtls/xray-core/proxy/masque\xaa\x02\x11Xray.Proxy.Masqueb\x06proto3"
var (
file_proxy_masque_config_proto_rawDescOnce sync.Once
file_proxy_masque_config_proto_rawDescData []byte
)
func file_proxy_masque_config_proto_rawDescGZIP() []byte {
file_proxy_masque_config_proto_rawDescOnce.Do(func() {
file_proxy_masque_config_proto_rawDescData = protoimpl.X.CompressGZIP(unsafe.Slice(unsafe.StringData(file_proxy_masque_config_proto_rawDesc), len(file_proxy_masque_config_proto_rawDesc)))
})
return file_proxy_masque_config_proto_rawDescData
}
var file_proxy_masque_config_proto_msgTypes = make([]protoimpl.MessageInfo, 3)
var file_proxy_masque_config_proto_goTypes = []any{
(*ClientConfig)(nil), // 0: xray.proxy.masque.ClientConfig
(*Account)(nil), // 1: xray.proxy.masque.Account
(*ServerConfig)(nil), // 2: xray.proxy.masque.ServerConfig
(*protocol.ServerEndpoint)(nil), // 3: xray.common.protocol.ServerEndpoint
(*protocol.User)(nil), // 4: xray.common.protocol.User
}
var file_proxy_masque_config_proto_depIdxs = []int32{
3, // 0: xray.proxy.masque.ClientConfig.server:type_name -> xray.common.protocol.ServerEndpoint
4, // 1: xray.proxy.masque.ServerConfig.users:type_name -> xray.common.protocol.User
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_proxy_masque_config_proto_init() }
func file_proxy_masque_config_proto_init() {
if File_proxy_masque_config_proto != nil {
return
}
type x struct{}
out := protoimpl.TypeBuilder{
File: protoimpl.DescBuilder{
GoPackagePath: reflect.TypeOf(x{}).PkgPath(),
RawDescriptor: unsafe.Slice(unsafe.StringData(file_proxy_masque_config_proto_rawDesc), len(file_proxy_masque_config_proto_rawDesc)),
NumEnums: 0,
NumMessages: 3,
NumExtensions: 0,
NumServices: 0,
},
GoTypes: file_proxy_masque_config_proto_goTypes,
DependencyIndexes: file_proxy_masque_config_proto_depIdxs,
MessageInfos: file_proxy_masque_config_proto_msgTypes,
}.Build()
File_proxy_masque_config_proto = out.File
file_proxy_masque_config_proto_goTypes = nil
file_proxy_masque_config_proto_depIdxs = nil
}
-25
View File
@@ -1,25 +0,0 @@
syntax = "proto3";
package xray.proxy.masque;
option csharp_namespace = "Xray.Proxy.Masque";
option go_package = "github.com/xtls/xray-core/proxy/masque";
option java_package = "com.xray.proxy.masque";
option java_multiple_files = true;
import "common/protocol/server_spec.proto";
import "common/protocol/user.proto";
message ClientConfig {
xray.common.protocol.ServerEndpoint server = 1;
repeated string remote_dns = 2;
}
message Account {
string password = 1;
}
message ServerConfig {
repeated xray.common.protocol.User users = 1;
repeated string address = 2;
uint32 mtu = 3;
}
-80
View File
@@ -1,80 +0,0 @@
package masque
import (
"net/netip"
"sync"
"github.com/xtls/xray-core/common/errors"
)
type addressPool struct {
mu sync.Mutex
prefix netip.Prefix
server netip.Addr
first netip.Addr
last netip.Addr
next netip.Addr
used map[netip.Addr]struct{}
}
func newAddressPool(address netip.Prefix) (*addressPool, error) {
server := address.Addr()
if server.Is4In6() || server.Zone() != "" {
return nil, errors.New("invalid address ", address)
}
prefix := address.Masked()
last := lastAddr(prefix)
if server == prefix.Addr() || server.Is4() && server == last {
return nil, errors.New("address ", address, " is not a host address")
}
if server.Is4() {
last = last.Prev()
}
first := prefix.Addr().Next()
if first == last {
return nil, errors.New("address ", address, " leaves no addresses to assign")
}
return &addressPool{
prefix: prefix,
server: server,
first: first,
last: last,
next: first,
used: make(map[netip.Addr]struct{}),
}, nil
}
func lastAddr(prefix netip.Prefix) netip.Addr {
b := prefix.Addr().AsSlice()
for i := prefix.Bits(); i < len(b)*8; i++ {
b[i/8] |= 1 << (7 - i%8)
}
addr, _ := netip.AddrFromSlice(b)
return addr
}
func (p *addressPool) allocate() (netip.Addr, bool) {
p.mu.Lock()
defer p.mu.Unlock()
for addr := p.next; ; {
next := addr.Next()
if addr == p.last {
next = p.first
}
if _, found := p.used[addr]; !found && addr != p.server {
p.used[addr] = struct{}{}
p.next = next
return addr, true
}
if next == p.next {
return netip.Addr{}, false
}
addr = next
}
}
func (p *addressPool) release(addr netip.Addr) {
p.mu.Lock()
defer p.mu.Unlock()
delete(p.used, addr)
}
-59
View File
@@ -1,59 +0,0 @@
package masque
import (
"net/netip"
"testing"
"github.com/stretchr/testify/require"
)
func allocateAll(p *addressPool) []netip.Addr {
var addrs []netip.Addr
for {
addr, ok := p.allocate()
if !ok {
return addrs
}
addrs = append(addrs, addr)
}
}
func TestAddressPool(t *testing.T) {
p, err := newAddressPool(netip.MustParsePrefix("10.0.0.1/29"))
require.NoError(t, err)
var want []netip.Addr
for _, s := range []string{"10.0.0.2", "10.0.0.3", "10.0.0.4", "10.0.0.5", "10.0.0.6"} {
want = append(want, netip.MustParseAddr(s))
}
require.Equal(t, want, allocateAll(p))
p.release(netip.MustParseAddr("10.0.0.4"))
addr, ok := p.allocate()
require.True(t, ok)
require.Equal(t, netip.MustParseAddr("10.0.0.4"), addr)
_, ok = p.allocate()
require.False(t, ok)
p, err = newAddressPool(netip.MustParsePrefix("fd00::1/126"))
require.NoError(t, err)
require.Equal(t, []netip.Addr{netip.MustParseAddr("fd00::2"), netip.MustParseAddr("fd00::3")}, allocateAll(p))
p, err = newAddressPool(netip.MustParsePrefix("10.0.0.2/30"))
require.NoError(t, err)
require.Equal(t, []netip.Addr{netip.MustParseAddr("10.0.0.1")}, allocateAll(p))
}
func TestAddressPoolRejects(t *testing.T) {
for _, s := range []string{
"10.0.0.0/24",
"10.0.0.255/24",
"10.0.0.1/31",
"10.0.0.1/32",
"fd00::1/127",
"fd00::1/128",
"::ffff:10.0.0.1/120",
} {
_, err := newAddressPool(netip.MustParsePrefix(s))
require.Error(t, err, s)
}
}
-550
View File
@@ -1,550 +0,0 @@
package masque
import (
"context"
go_errors "errors"
"io"
stdnet "net"
"net/http"
"net/netip"
"slices"
"sync"
"golang.zx2c4.com/wireguard/tun"
"github.com/xtls/xray-core/common"
"github.com/xtls/xray-core/common/buf"
c "github.com/xtls/xray-core/common/ctx"
"github.com/xtls/xray-core/common/errors"
"github.com/xtls/xray-core/common/log"
"github.com/xtls/xray-core/common/net"
"github.com/xtls/xray-core/common/protocol"
"github.com/xtls/xray-core/common/session"
"github.com/xtls/xray-core/core"
"github.com/xtls/xray-core/features/routing"
"github.com/xtls/xray-core/proxy/wireguard"
"github.com/xtls/xray-core/transport"
"github.com/xtls/xray-core/transport/internet"
"github.com/xtls/xray-core/transport/internet/masque"
"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"
)
const (
authenticateHeader = `Basic realm="masque", charset="UTF-8"`
tunnelQueueSize = 512
)
type Server struct {
validator *validator
dispatcher routing.Dispatcher
ctx context.Context
tag string
sniffing session.SniffingRequest
mtu int
dev tun.Device
pools []*addressPool
local []netip.Addr
mu sync.RWMutex
tunnels map[netip.Addr]*serverTunnel
closed bool
started bool
}
type serverTunnel struct {
conn stat.Connection
ipConn *connectip.Conn
user *protocol.MemoryUser
addrs []netip.Addr
queue chan *buf.Buffer
done chan struct{}
mu sync.Mutex
conns map[net.Conn]struct{}
}
func newServerTunnel(conn stat.Connection, user *protocol.MemoryUser) *serverTunnel {
return &serverTunnel{
conn: conn,
user: user,
queue: make(chan *buf.Buffer, tunnelQueueSize),
done: make(chan struct{}),
conns: make(map[net.Conn]struct{}),
}
}
func (t *serverTunnel) send(b *buf.Buffer) bool {
select {
case <-t.done:
return false
default:
}
select {
case t.queue <- b:
return true
default:
return false
}
}
func (t *serverTunnel) track(conn net.Conn) bool {
t.mu.Lock()
defer t.mu.Unlock()
if t.conns == nil {
return false
}
t.conns[conn] = struct{}{}
return true
}
func (t *serverTunnel) untrack(conn net.Conn) {
t.mu.Lock()
delete(t.conns, conn)
t.mu.Unlock()
}
func (t *serverTunnel) close() {
t.mu.Lock()
conns := t.conns
if conns != nil {
t.conns = nil
close(t.done)
}
t.mu.Unlock()
for conn := range conns {
conn.Close()
}
}
func NewServer(ctx context.Context, config *ServerConfig) (*Server, error) {
v := core.MustFromContext(ctx)
streamSettings := session.StreamSettingsFromContext(ctx).(*internet.MemoryStreamConfig)
if _, ok := streamSettings.ProtocolSettings.(*masque.Config); !ok {
return nil, errors.New("not masque transport")
}
if tls.ConfigFromStreamSettings(streamSettings) == nil {
return nil, errors.New(`MASQUE requires "security": "tls"`)
}
users := newValidator()
for _, user := range config.Users {
u, err := user.ToMemoryUser()
if err != nil {
return nil, errors.New("failed to get MASQUE user").Base(err)
}
if err := users.add(u); err != nil {
return nil, errors.New("failed to add user").Base(err)
}
}
var pools []*addressPool
var local []netip.Addr
for _, s := range config.Address {
prefix, err := netip.ParsePrefix(s)
if err != nil {
return nil, errors.New("invalid address ", s).Base(err)
}
if slices.ContainsFunc(local, func(addr netip.Addr) bool { return addr.Is4() == prefix.Addr().Is4() }) {
return nil, errors.New("only one address per IP family is supported")
}
pool, err := newAddressPool(prefix)
if err != nil {
return nil, err
}
pools = append(pools, pool)
local = append(local, prefix.Addr())
}
if len(pools) == 0 {
return nil, errors.New("no address to assign")
}
mtu := int(config.Mtu)
if mtu == 0 {
mtu = masque.MinPacketSize
}
dev, _, gstack, err := wireguard.CreateNetTUN(local, nil, mtu, false)
if err != nil {
return nil, err
}
s := &Server{
validator: users,
dispatcher: v.GetFeature(routing.DispatcherType()).(routing.Dispatcher),
ctx: core.ToBackgroundDetachedContext(ctx),
mtu: mtu,
dev: dev,
pools: pools,
local: local,
tunnels: make(map[netip.Addr]*serverTunnel),
}
if inbound := session.InboundFromContext(ctx); inbound != nil {
s.tag = inbound.Tag
}
if content := session.ContentFromContext(ctx); content != nil {
s.sniffing = content.SniffingRequest
}
wireguard.CreateForwarder(gstack, s.handleConnection)
return s, nil
}
func (s *Server) Start() error {
s.mu.Lock()
defer s.mu.Unlock()
if s.started || s.closed {
return nil
}
s.started = true
go s.readFromStack()
return nil
}
func (s *Server) Close() error {
s.mu.Lock()
if s.closed {
s.mu.Unlock()
return nil
}
s.closed = true
var tunnels []*serverTunnel
for _, t := range s.tunnels {
if !slices.Contains(tunnels, t) {
tunnels = append(tunnels, t)
}
}
s.mu.Unlock()
for _, t := range tunnels {
t.conn.Close()
}
return s.dev.Close()
}
func (s *Server) AddUser(ctx context.Context, user *protocol.MemoryUser) error {
return s.validator.add(user)
}
func (s *Server) RemoveUser(ctx context.Context, email string) error {
user, err := s.validator.delByEmail(email)
if err != nil {
return err
}
s.mu.RLock()
var conns []stat.Connection
for _, t := range s.tunnels {
if t.user == user && !slices.Contains(conns, t.conn) {
conns = append(conns, t.conn)
}
}
s.mu.RUnlock()
for _, conn := range conns {
conn.Close()
}
return nil
}
func (s *Server) GetUser(ctx context.Context, email string) *protocol.MemoryUser {
return s.validator.getByEmail(email)
}
func (s *Server) GetUsers(ctx context.Context) []*protocol.MemoryUser {
return s.validator.getAll()
}
func (s *Server) GetUsersCount(context.Context) int64 {
return s.validator.count()
}
func (s *Server) Network() []net.Network {
return []net.Network{net.Network_TCP}
}
func (s *Server) Process(ctx context.Context, network net.Network, conn stat.Connection, dispatcher routing.Dispatcher) error {
sconn, ok := stat.TryUnwrapStatsConn(conn).(*masque.ServerConn)
if !ok {
return errors.New("not a MASQUE connection")
}
inbound := session.InboundFromContext(ctx)
inbound.Name = "masque"
inbound.CanSpliceCopy = 3
name, pass, _ := sconn.Request().BasicAuth()
user := s.validator.get(name, pass)
if user == nil {
sconn.Reject(http.StatusUnauthorized, http.Header{"WWW-Authenticate": {authenticateHeader}})
log.Record(&log.AccessMessage{
From: conn.RemoteAddr(),
To: "",
Status: log.AccessRejected,
Reason: errors.New("invalid credentials"),
})
return errors.New("MASQUE: authentication failed for ", name)
}
inbound.User = user
t := newServerTunnel(conn, user)
for _, pool := range s.pools {
if addr, ok := pool.allocate(); ok {
t.addrs = append(t.addrs, addr)
}
}
defer s.release(t)
if len(t.addrs) == 0 {
sconn.Reject(http.StatusServiceUnavailable, nil)
return errors.New("MASQUE: no address left to assign")
}
ipConn, err := sconn.Accept()
if err != nil {
return errors.New("MASQUE: failed to accept the tunnel").Base(err)
}
t.ipConn = ipConn
if !s.register(t) {
return errors.New("MASQUE: server closed")
}
if !s.validator.contains(user) {
return errors.New("MASQUE: user ", name, " was removed")
}
go s.writeToTunnel(t)
prefixes := make([]netip.Prefix, len(t.addrs))
for i, addr := range t.addrs {
prefixes[i] = netip.PrefixFrom(addr, addr.BitLen())
}
if err := ipConn.AssignAddresses(prefixes); err != nil {
return err
}
if err := ipConn.AdvertiseRoute(fullRoutes(t.addrs)); err != nil {
return err
}
go serveAddressRequests(t)
ctx = log.ContextWithAccessMessage(ctx, &log.AccessMessage{
From: conn.RemoteAddr(),
To: "",
Status: log.AccessAccepted,
Email: user.Email,
})
errors.LogInfo(ctx, "MASQUE: tunnel from ", inbound.Source, " assigned ", t.addrs)
return s.readFromTunnel(t)
}
func fullRoutes(addrs []netip.Addr) []connectip.IPRoute {
var routes []connectip.IPRoute
if slices.ContainsFunc(addrs, netip.Addr.Is4) {
routes = append(routes, connectip.IPRoute{StartIP: netip.IPv4Unspecified(), EndIP: netip.AddrFrom4([4]byte{255, 255, 255, 255})})
}
if slices.ContainsFunc(addrs, netip.Addr.Is6) {
routes = append(routes, connectip.IPRoute{StartIP: netip.IPv6Unspecified(), EndIP: netip.AddrFrom16([16]byte{0: 0xff, 1: 0xff, 2: 0xff, 3: 0xff, 4: 0xff, 5: 0xff, 6: 0xff, 7: 0xff, 8: 0xff, 9: 0xff, 10: 0xff, 11: 0xff, 12: 0xff, 13: 0xff, 14: 0xff, 15: 0xff})})
}
return routes
}
func serveAddressRequests(t *serverTunnel) {
for {
req, err := t.ipConn.ReceiveAddressRequest(context.Background())
if err != nil {
return
}
assigned := make([]netip.Prefix, len(req.Prefixes))
used := make(map[netip.Addr]bool)
for i, requested := range req.Prefixes {
for _, addr := range t.addrs {
if addr.Is4() == requested.Addr().Is4() && !used[addr] {
used[addr] = true
assigned[i] = netip.PrefixFrom(addr, addr.BitLen())
break
}
}
}
var additional []netip.Prefix
for _, addr := range t.addrs {
if !used[addr] {
additional = append(additional, netip.PrefixFrom(addr, addr.BitLen()))
}
}
if err := req.Respond(assigned, additional); err != nil {
return
}
}
}
func (s *Server) register(t *serverTunnel) bool {
s.mu.Lock()
defer s.mu.Unlock()
if s.closed {
return false
}
for _, addr := range t.addrs {
s.tunnels[addr] = t
}
return true
}
func (s *Server) release(t *serverTunnel) {
s.mu.Lock()
for _, addr := range t.addrs {
if s.tunnels[addr] == t {
delete(s.tunnels, addr)
}
}
s.mu.Unlock()
t.close()
for _, addr := range t.addrs {
for _, pool := range s.pools {
if pool.prefix.Contains(addr) {
pool.release(addr)
}
}
}
}
func (s *Server) lookup(addr netip.Addr) *serverTunnel {
s.mu.RLock()
defer s.mu.RUnlock()
return s.tunnels[addr]
}
func (s *Server) inPool(addr netip.Addr) bool {
return slices.ContainsFunc(s.pools, func(pool *addressPool) bool { return pool.prefix.Contains(addr) })
}
func (s *Server) readFromTunnel(t *serverTunnel) error {
b := make([]byte, 1<<16)
for {
n, err := t.conn.Read(b)
if err != nil {
if go_errors.Is(err, io.ErrShortBuffer) {
continue
}
if go_errors.Is(err, stdnet.ErrClosed) || go_errors.Is(err, io.EOF) {
return nil
}
return err
}
dst, ok := packetDestination(b[:n])
if !ok || dst.IsLinkLocalUnicast() || dst.IsMulticast() {
continue
}
if other := s.lookup(dst); other != nil {
if other != t {
packet := buf.NewWithSize(int32(n))
packet.Write(b[:n])
if !other.send(packet) {
packet.Release()
}
}
continue
}
if s.inPool(dst) && !slices.Contains(s.local, dst) {
continue
}
s.dev.Write([][]byte{b[:n]}, 0)
}
}
func (s *Server) readFromStack() {
sizes := []int{0}
var b *buf.Buffer
for {
if b == nil {
b = buf.NewWithSize(int32(s.mtu))
}
b.Clear()
if _, err := s.dev.Read([][]byte{b.Extend(int32(s.mtu))}, sizes, 0); err != nil {
b.Release()
return
}
b.Resize(0, int32(sizes[0]))
dst, ok := packetDestination(b.Bytes())
if !ok {
continue
}
if t := s.lookup(dst); t != nil && t.send(b) {
b = nil
}
}
}
func (s *Server) writeToTunnel(t *serverTunnel) {
for {
select {
case b := <-t.queue:
_, err := t.conn.Write(b.Bytes())
b.Release()
if ptb, ok := go_errors.AsType[*masque.PacketTooBigError](err); ok {
s.dev.Write([][]byte{ptb.ICMP}, 0)
}
case <-t.done:
return
}
}
}
func packetDestination(packet []byte) (netip.Addr, bool) {
if len(packet) == 0 {
return netip.Addr{}, false
}
switch packet[0] >> 4 {
case 4:
if len(packet) >= 20 {
return netip.AddrFrom4([4]byte(packet[16:20])), true
}
case 6:
if len(packet) >= 40 {
return netip.AddrFrom16([16]byte(packet[24:40])), true
}
}
return netip.Addr{}, false
}
func (s *Server) handleConnection(conn net.Conn, dest net.Destination) {
defer conn.Close()
source := net.DestinationFromAddr(conn.RemoteAddr())
addr, _ := netip.AddrFromSlice(source.Address.IP())
t := s.lookup(addr.Unmap())
if t == nil || !t.track(conn) {
errors.LogInfo(s.ctx, "MASQUE: no tunnel for ", source, " to ", dest)
return
}
defer t.untrack(conn)
ctx, cancel := context.WithCancel(s.ctx)
defer cancel()
ctx = c.ContextWithID(ctx, session.NewID())
inbound := session.Inbound{
Name: "masque",
Tag: s.tag,
CanSpliceCopy: 3,
Source: source,
User: t.user,
}
ctx = session.ContextWithInbound(ctx, &inbound)
ctx = session.ContextWithContent(ctx, &session.Content{
SniffingRequest: s.sniffing,
})
ctx = session.SubContextFromMuxInbound(ctx)
ctx = log.ContextWithAccessMessage(ctx, &log.AccessMessage{
From: source,
To: dest,
Status: log.AccessAccepted,
Email: t.user.Email,
})
errors.LogInfo(ctx, "processing from ", source, " to ", dest)
link := &transport.Link{
Reader: &buf.TimeoutWrapperReader{Reader: buf.NewReader(conn)},
Writer: buf.NewWriter(conn),
}
if err := s.dispatcher.DispatchLink(ctx, dest, link); err != nil {
errors.LogError(ctx, errors.New("connection closed").Base(err))
}
}
func init() {
common.Must(common.RegisterConfig((*ServerConfig)(nil), func(ctx context.Context, config interface{}) (interface{}, error) {
return NewServer(ctx, config.(*ServerConfig))
}))
}
-328
View File
@@ -1,328 +0,0 @@
package masque
import (
"bytes"
"context"
"io"
"net/netip"
"os"
"sync"
"testing"
"time"
"github.com/stretchr/testify/require"
"github.com/xtls/xray-core/common/buf"
"github.com/xtls/xray-core/common/net"
"github.com/xtls/xray-core/common/protocol"
"golang.zx2c4.com/wireguard/tun"
)
type fakeTunnelConn struct {
mu sync.Mutex
reads chan []byte
written [][]byte
closed bool
stall chan struct{}
}
func newFakeTunnelConn() *fakeTunnelConn {
return &fakeTunnelConn{reads: make(chan []byte, 16)}
}
func (c *fakeTunnelConn) Read(b []byte) (int, error) {
p, ok := <-c.reads
if !ok {
return 0, io.EOF
}
return copy(b, p), nil
}
func (c *fakeTunnelConn) Write(b []byte) (int, error) {
if c.stall != nil {
<-c.stall
}
c.mu.Lock()
defer c.mu.Unlock()
c.written = append(c.written, bytes.Clone(b))
return len(b), nil
}
func (c *fakeTunnelConn) Close() error {
c.mu.Lock()
defer c.mu.Unlock()
if !c.closed {
c.closed = true
close(c.reads)
}
return nil
}
func (c *fakeTunnelConn) packets() [][]byte {
c.mu.Lock()
defer c.mu.Unlock()
return c.written
}
func (c *fakeTunnelConn) isClosed() bool {
c.mu.Lock()
defer c.mu.Unlock()
return c.closed
}
func (c *fakeTunnelConn) LocalAddr() net.Addr { return &net.TCPAddr{} }
func (c *fakeTunnelConn) RemoteAddr() net.Addr { return &net.TCPAddr{} }
func (c *fakeTunnelConn) SetDeadline(t time.Time) error { return nil }
func (c *fakeTunnelConn) SetReadDeadline(t time.Time) error { return nil }
func (c *fakeTunnelConn) SetWriteDeadline(t time.Time) error { return nil }
type fakeDevice struct {
mu sync.Mutex
reads chan []byte
written [][]byte
closed bool
}
func (d *fakeDevice) File() *os.File { return nil }
func (d *fakeDevice) MTU() (int, error) { return 1280, nil }
func (d *fakeDevice) Name() (string, error) { return "fake", nil }
func (d *fakeDevice) Events() <-chan tun.Event { return nil }
func (d *fakeDevice) BatchSize() int { return 1 }
func (d *fakeDevice) Read(bufs [][]byte, sizes []int, offset int) (int, error) {
p, ok := <-d.reads
if !ok {
return 0, os.ErrClosed
}
sizes[0] = copy(bufs[0][offset:], p)
return 1, nil
}
func (d *fakeDevice) Write(bufs [][]byte, offset int) (int, error) {
d.mu.Lock()
defer d.mu.Unlock()
for _, b := range bufs {
d.written = append(d.written, bytes.Clone(b[offset:]))
}
return len(bufs), nil
}
func (d *fakeDevice) Close() error {
d.mu.Lock()
defer d.mu.Unlock()
if !d.closed {
d.closed = true
close(d.reads)
}
return nil
}
func (d *fakeDevice) packets() [][]byte {
d.mu.Lock()
defer d.mu.Unlock()
return d.written
}
func ipPacket(src, dst string) []byte {
s, d := netip.MustParseAddr(src), netip.MustParseAddr(dst)
if s.Is4() {
b := make([]byte, 20)
b[0] = 0x45
b[8] = 64
copy(b[12:16], s.AsSlice())
copy(b[16:20], d.AsSlice())
return b
}
b := make([]byte, 40)
b[0] = 0x60
b[7] = 64
copy(b[8:24], s.AsSlice())
copy(b[24:40], d.AsSlice())
return b
}
func newTestServer(t *testing.T) (*Server, *fakeDevice) {
t.Helper()
pool4, err := newAddressPool(netip.MustParsePrefix("10.14.0.1/24"))
require.NoError(t, err)
pool6, err := newAddressPool(netip.MustParsePrefix("fd14::1/64"))
require.NoError(t, err)
dev := &fakeDevice{reads: make(chan []byte, tunnelQueueSize*2)}
s := &Server{
mtu: 1280,
dev: dev,
pools: []*addressPool{pool4, pool6},
local: []netip.Addr{netip.MustParseAddr("10.14.0.1"), netip.MustParseAddr("fd14::1")},
tunnels: make(map[netip.Addr]*serverTunnel),
}
return s, dev
}
func addTunnel(t *testing.T, s *Server) (*serverTunnel, *fakeTunnelConn) {
t.Helper()
return addUserTunnel(t, s, &protocol.MemoryUser{})
}
func addUserTunnel(t *testing.T, s *Server, user *protocol.MemoryUser) (*serverTunnel, *fakeTunnelConn) {
t.Helper()
conn := newFakeTunnelConn()
tunnel := newServerTunnel(conn, user)
for _, pool := range s.pools {
addr, ok := pool.allocate()
require.True(t, ok)
tunnel.addrs = append(tunnel.addrs, addr)
}
require.True(t, s.register(tunnel))
go s.writeToTunnel(tunnel)
t.Cleanup(tunnel.close)
return tunnel, conn
}
func TestServerRoutesTunnelPackets(t *testing.T) {
s, dev := newTestServer(t)
a, aConn := addTunnel(t, s)
b, bConn := addTunnel(t, s)
require.Equal(t, []netip.Addr{netip.MustParseAddr("10.14.0.2"), netip.MustParseAddr("fd14::2")}, a.addrs)
require.Equal(t, []netip.Addr{netip.MustParseAddr("10.14.0.3"), netip.MustParseAddr("fd14::3")}, b.addrs)
toB := ipPacket("10.14.0.2", "10.14.0.3")
toB6 := ipPacket("fd14::2", "fd14::3")
toServer := ipPacket("10.14.0.2", "10.14.0.1")
toInternet := ipPacket("fd14::2", "2001:db8::1")
for _, p := range [][]byte{
toB,
toB6,
ipPacket("10.14.0.2", "10.14.0.9"),
ipPacket("fd14::2", "fd14::99"),
ipPacket("fd14::2", "fe80::1"),
ipPacket("fd14::2", "ff02::1"),
ipPacket("10.14.0.2", "224.0.0.251"),
ipPacket("10.14.0.2", "10.14.0.2"),
toServer,
toInternet,
} {
aConn.reads <- p
}
aConn.Close()
require.NoError(t, s.readFromTunnel(a))
require.Eventually(t, func() bool { return len(bConn.packets()) == 2 }, time.Second, time.Millisecond)
require.Equal(t, [][]byte{toB, toB6}, bConn.packets())
require.Equal(t, [][]byte{toServer, toInternet}, dev.packets())
require.Empty(t, aConn.packets())
}
func TestServerRoutesStackPackets(t *testing.T) {
s, dev := newTestServer(t)
_, aConn := addTunnel(t, s)
_, bConn := addTunnel(t, s)
require.NoError(t, s.Start())
toA := ipPacket("192.0.2.1", "10.14.0.2")
toB := ipPacket("2001:db8::1", "fd14::3")
dev.reads <- toA
dev.reads <- ipPacket("192.0.2.1", "10.14.0.9")
dev.reads <- toB
require.Eventually(t, func() bool {
return len(aConn.packets()) == 1 && len(bConn.packets()) == 1
}, time.Second, time.Millisecond)
require.Equal(t, [][]byte{toA}, aConn.packets())
require.Equal(t, [][]byte{toB}, bConn.packets())
require.NoError(t, s.Close())
require.True(t, aConn.isClosed())
require.True(t, bConn.isClosed())
require.False(t, s.register(&serverTunnel{}))
}
func TestServerSlowTunnelDoesNotBlockOthers(t *testing.T) {
s, dev := newTestServer(t)
_, aConn := addTunnel(t, s)
_, bConn := addTunnel(t, s)
aConn.stall = make(chan struct{})
defer close(aConn.stall)
require.NoError(t, s.Start())
defer s.Close()
for range tunnelQueueSize + 10 {
dev.reads <- ipPacket("192.0.2.1", "10.14.0.2")
}
toB := ipPacket("192.0.2.1", "10.14.0.3")
dev.reads <- toB
require.Eventually(t, func() bool { return len(bConn.packets()) == 1 }, time.Second, time.Millisecond)
require.Equal(t, [][]byte{toB}, bConn.packets())
}
func TestServerClosesTunnelConnections(t *testing.T) {
s, _ := newTestServer(t)
a, _ := addTunnel(t, s)
conn := newFakeTunnelConn()
require.True(t, a.track(conn))
other := newFakeTunnelConn()
require.True(t, a.track(other))
a.untrack(other)
s.release(a)
require.True(t, conn.isClosed())
require.False(t, other.isClosed())
require.False(t, a.track(newFakeTunnelConn()))
require.False(t, a.send(buf.New()))
}
func TestServerReleasesAddresses(t *testing.T) {
s, _ := newTestServer(t)
a, _ := addTunnel(t, s)
s.release(a)
require.Nil(t, s.lookup(netip.MustParseAddr("10.14.0.2")))
b, _ := addTunnel(t, s)
require.Equal(t, []netip.Addr{netip.MustParseAddr("10.14.0.3"), netip.MustParseAddr("fd14::3")}, b.addrs)
for range 250 {
addTunnel(t, s)
}
c, _ := addTunnel(t, s)
require.Equal(t, netip.MustParseAddr("10.14.0.254"), c.addrs[0])
addr, ok := s.pools[0].allocate()
require.True(t, ok)
require.Equal(t, netip.MustParseAddr("10.14.0.2"), addr)
_, ok = s.pools[0].allocate()
require.False(t, ok)
}
func TestServerRemoveUserClosesTunnels(t *testing.T) {
s, _ := newTestServer(t)
s.validator = newValidator()
alice := &protocol.MemoryUser{Email: "a@example.com", Account: &MemoryAccount{Password: "p"}}
bob := &protocol.MemoryUser{Email: "b@example.com", Account: &MemoryAccount{Password: "p"}}
require.NoError(t, s.AddUser(context.Background(), alice))
require.NoError(t, s.AddUser(context.Background(), bob))
_, aConn := addUserTunnel(t, s, alice)
_, bConn := addUserTunnel(t, s, bob)
require.NoError(t, s.RemoveUser(context.Background(), "a@example.com"))
require.True(t, aConn.isClosed())
require.False(t, bConn.isClosed())
require.Error(t, s.RemoveUser(context.Background(), "a@example.com"))
require.Nil(t, s.validator.get("a@example.com", "p"))
require.Equal(t, bob, s.validator.get("b@example.com", "p"))
}
func TestPacketDestination(t *testing.T) {
v4 := make([]byte, 20)
v4[0] = 0x45
copy(v4[16:20], []byte{192, 0, 2, 1})
addr, ok := packetDestination(v4)
require.True(t, ok)
require.Equal(t, netip.MustParseAddr("192.0.2.1"), addr)
v6 := make([]byte, 40)
v6[0] = 0x60
dst := netip.MustParseAddr("2001:db8::1").As16()
copy(v6[24:40], dst[:])
addr, ok = packetDestination(v6)
require.True(t, ok)
require.Equal(t, netip.MustParseAddr("2001:db8::1"), addr)
for _, b := range [][]byte{nil, v4[:19], v6[:39], {0x50}} {
_, ok = packetDestination(b)
require.False(t, ok)
}
}
-15
View File
@@ -277,7 +277,6 @@ func (w *VisionReader) ReadMultiBuffer() (buf.MultiBuffer, error) {
w.ob.CanSpliceCopy = 1
}
}
SuppressOuterCloseNotify(w.conn)
readerConn, readCounter, _ := UnwrapRawConn(w.conn)
w.directReadCounter = readCounter
w.Reader = buf.NewReader(readerConn)
@@ -341,7 +340,6 @@ func (w *VisionWriter) WriteMultiBuffer(mb buf.MultiBuffer) error {
// w.ob.CanSpliceCopy = 1
// }
}
SuppressOuterCloseNotify(w.conn)
rawConn, _, writerCounter := UnwrapRawConn(w.conn)
w.Writer = buf.NewWriter(rawConn)
w.directWriteCounter = writerCounter
@@ -671,19 +669,6 @@ func XtlsFilterTls(buffer buf.MultiBuffer, trafficState *TrafficState, ctx conte
}
}
type CloseNotifySuppressor interface {
SuppressCloseNotify()
}
// Close our local TLS conn instance might send a incorrect close_notify alert
// if the XTLS direct copy mode is enabled and cause TLS BAD_RECORD_MAC on users' browser
// Close the underlying connection directly to avoid this issue.
func SuppressOuterCloseNotify(conn net.Conn) {
if suppressor, ok := stat.TryUnwrapStatsConn(conn).(CloseNotifySuppressor); ok {
suppressor.SuppressCloseNotify()
}
}
// UnwrapRawConn support unwrap encryption, stats, mask wrappers, tls, utls, reality, proxyproto, uds-wrapper conn and get raw tcp/uds conn from it
func UnwrapRawConn(conn net.Conn) (net.Conn, stats.Counter, stats.Counter) {
var readCounter, writerCounter stats.Counter
-55
View File
@@ -1,55 +0,0 @@
package shadowsocks_2022
import (
"crypto/aes"
"crypto/cipher"
"errors"
"strings"
"golang.org/x/crypto/chacha20poly1305"
)
type CipherMethod struct {
Name string
KeySaltLength int
IsChaCha bool
}
var methods = map[string]*CipherMethod{
MethodAES128GCM: {Name: MethodAES128GCM, KeySaltLength: 16, IsChaCha: false},
MethodAES256GCM: {Name: MethodAES256GCM, KeySaltLength: 32, IsChaCha: false},
MethodChaCha20Poly1305: {Name: MethodChaCha20Poly1305, KeySaltLength: 32, IsChaCha: true},
}
func GetCipherMethod(name string) (*CipherMethod, error) {
name = strings.ToLower(name)
if m, ok := methods[name]; ok {
return m, nil
}
return nil, errors.New("unknown shadowsocks 2022 method")
}
// NewAEAD creates standard stream AEAD cipher instance (AES-GCM or ChaCha20-Poly1305)
func (m *CipherMethod) NewAEAD(key []byte) (cipher.AEAD, error) {
if m.IsChaCha {
return chacha20poly1305.New(key)
}
block, err := aes.NewCipher(key)
if err != nil {
return nil, err
}
return cipher.NewGCM(block)
}
// NewBlock creates standard 16-byte block cipher for AES header encryption/decryption
func (m *CipherMethod) NewBlock(key []byte) (cipher.Block, error) {
return aes.NewCipher(key)
}
// NewUDPCipher creates AEAD cipher for UDP packets (XChaCha20-Poly1305 with 24-byte nonce)
func (m *CipherMethod) NewUDPCipher(key []byte) (cipher.AEAD, error) {
if m.IsChaCha {
return chacha20poly1305.NewX(key)
}
return nil, errors.New("shadowsocks-2022: udp separate AEAD cipher only available for chacha20 method")
}
+4 -12
View File
@@ -1,9 +1,6 @@
package shadowsocks_2022
import (
"bytes"
"encoding/base64"
"google.golang.org/protobuf/proto"
"github.com/xtls/xray-core/common/protocol"
@@ -11,31 +8,26 @@ import (
// MemoryAccount is an account type converted from Account.
type MemoryAccount struct {
Key []byte
Key string
}
// AsAccount implements protocol.AsAccount.
func (u *Account) AsAccount() (protocol.Account, error) {
keyStr := u.GetKey()
raw, err := base64.StdEncoding.DecodeString(keyStr)
if err != nil {
raw = []byte(keyStr)
}
return &MemoryAccount{
Key: raw,
Key: u.GetKey(),
}, nil
}
// Equals implements protocol.Account.Equals().
func (a *MemoryAccount) Equals(another protocol.Account) bool {
if account, ok := another.(*MemoryAccount); ok {
return bytes.Equal(a.Key, account.Key)
return a.Key == account.Key
}
return false
}
func (a *MemoryAccount) ToProto() proto.Message {
return &Account{
Key: base64.StdEncoding.EncodeToString(a.Key),
Key: a.Key,
}
}
+128 -130
View File
@@ -4,16 +4,23 @@ import (
"context"
"time"
shadowsocks "github.com/sagernet/sing-shadowsocks"
"github.com/sagernet/sing-shadowsocks/shadowaead_2022"
C "github.com/sagernet/sing/common"
B "github.com/sagernet/sing/common/buf"
"github.com/sagernet/sing/common/bufio"
E "github.com/sagernet/sing/common/exceptions"
M "github.com/sagernet/sing/common/metadata"
N "github.com/sagernet/sing/common/network"
"github.com/xtls/xray-core/common"
"github.com/xtls/xray-core/common/antireplay"
"github.com/xtls/xray-core/common/buf"
"github.com/xtls/xray-core/common/errors"
"github.com/xtls/xray-core/common/log"
"github.com/xtls/xray-core/common/net"
"github.com/xtls/xray-core/common/protocol"
"github.com/xtls/xray-core/common/session"
"github.com/xtls/xray-core/core"
"github.com/xtls/xray-core/features/policy"
"github.com/xtls/xray-core/common/signal"
"github.com/xtls/xray-core/common/singbridge"
"github.com/xtls/xray-core/features/routing"
"github.com/xtls/xray-core/transport/internet/stat"
)
@@ -25,13 +32,10 @@ func init() {
}
type Inbound struct {
networks []net.Network
method *CipherMethod
psk []byte
user *protocol.MemoryUser
saltFilter *antireplay.ReplayFilter[[32]byte]
udpCodec *UDPServerCodec
policyManager policy.Manager
networks []net.Network
service shadowsocks.Service
email string
level int
}
func NewServer(ctx context.Context, config *ServerConfig) (*Inbound, error) {
@@ -42,35 +46,20 @@ func NewServer(ctx context.Context, config *ServerConfig) (*Inbound, error) {
net.Network_UDP,
}
}
method, err := GetCipherMethod(config.Method)
if err != nil {
return nil, errors.New("unsupported method: ", config.Method).Base(err)
inbound := &Inbound{
networks: networks,
email: config.Email,
level: int(config.Level),
}
psk, err := ParseKey(config.Key, method.KeySaltLength)
if err != nil {
return nil, err
if !C.Contains(shadowaead_2022.List, config.Method) {
return nil, errors.New("unsupported method ", config.Method)
}
udpCodec, err := NewUDPServerCodec(method, psk, 500*time.Second)
service, err := shadowaead_2022.NewServiceWithPassword(config.Method, config.Key, 500, inbound, nil)
if err != nil {
return nil, err
return nil, errors.New("create service").Base(err)
}
v := core.MustFromContext(ctx)
return &Inbound{
networks: networks,
method: method,
psk: psk,
saltFilter: antireplay.NewMapFilter[[32]byte](60),
user: &protocol.MemoryUser{
Email: config.Email,
Level: uint32(config.Level),
},
udpCodec: udpCodec,
policyManager: v.GetFeature(policy.ManagerType()).(policy.Manager),
}, nil
inbound.service = service
return inbound, nil
}
func (i *Inbound) Network() []net.Network {
@@ -81,105 +70,114 @@ func (i *Inbound) Process(ctx context.Context, network net.Network, connection s
inbound := session.InboundFromContext(ctx)
inbound.Name = "shadowsocks-2022"
inbound.CanSpliceCopy = 3
inbound.User = i.user
var metadata M.Metadata
if inbound.Source.IsValid() {
metadata.Source = M.ParseSocksaddr(inbound.Source.NetAddr())
}
ctx = session.ContextWithDispatcher(ctx, dispatcher)
if network == net.Network_TCP {
return i.processTCP(ctx, connection, dispatcher)
}
return i.processUDP(ctx, connection, dispatcher)
}
func (i *Inbound) processTCP(ctx context.Context, conn net.Conn, dispatcher routing.Dispatcher) error {
defer conn.Close()
sessionPolicy := i.policyManager.ForLevel(0)
if err := conn.SetReadDeadline(time.Now().Add(sessionPolicy.Timeouts.Handshake)); err != nil {
return errors.New("unable to set read deadline").Base(err)
}
// 1. Single read call for Salt + Fixed-length header chunk per SIP022 §3.1.4
headerLen := i.method.KeySaltLength + RequestHeaderFixedChunkLength + AEADTagSize
headerBuf := make([]byte, headerLen)
n, err := conn.Read(headerBuf)
if err != nil || n < headerLen {
ResetTCPConn(conn)
return errors.New("failed to read complete handshake header")
}
var salt [32]byte
copy(salt[:i.method.KeySaltLength], headerBuf[:i.method.KeySaltLength])
saltSlice := salt[:i.method.KeySaltLength]
fixedChunk := headerBuf[i.method.KeySaltLength:]
reader, reqHeader, err := InitServerStream(conn, i.method, i.psk, saltSlice, salt, fixedChunk, i.saltFilter)
if err != nil {
ResetTCPConn(conn)
return err
}
dest := reqHeader.Destination
writer := NewServerStreamWriter(conn, i.method, i.psk, saltSlice)
ctx = log.ContextWithAccessMessage(ctx, &log.AccessMessage{
From: conn.RemoteAddr(),
To: dest,
Status: log.AccessAccepted,
Email: i.user.Email,
})
errors.LogInfo(ctx, "tunneling request to ", dest)
link, err := dispatcher.Dispatch(ctx, dest)
if err != nil {
return err
}
if len(reqHeader.EarlyData) > 0 {
mb := buf.MergeBytes(nil, reqHeader.EarlyData)
if err := link.Writer.WriteMultiBuffer(mb); err != nil {
return err
}
}
return TransportTCP(ctx, i.policyManager.ForLevel(uint32(i.user.Level)), reader, writer, link)
}
func (i *Inbound) processUDP(ctx context.Context, conn stat.Connection, dispatcher routing.Dispatcher) error {
reader := buf.NewPacketReader(conn)
for {
mb, err := reader.ReadMultiBuffer()
if err != nil {
buf.ReleaseMulti(mb)
return err
}
for _, b := range mb {
decoded, err := i.udpCodec.DecodePacket(b.Bytes())
b.Release()
if err != nil || decoded.HeaderType != HeaderTypeClient {
continue
}
sessionItem := i.udpCodec.GetSession(decoded.SessionID)
if sessionItem.User == nil {
sessionItem.Lock()
if sessionItem.User == nil {
sessionItem.User = i.user
}
sessionItem.Unlock()
}
link, err := sessionItem.EnsureLink(ctx, conn, decoded.Destination, dispatcher, i.policyManager, func(dest net.Destination, payload []byte) ([]byte, error) {
return i.udpCodec.EncodeServerPacket(decoded.SessionID, dest, payload)
})
return singbridge.ReturnError(i.service.NewConnection(ctx, connection, metadata))
} else {
reader := buf.NewReader(connection)
pc := &natPacketConn{connection}
for {
mb, err := reader.ReadMultiBuffer()
if err != nil {
continue
buf.ReleaseMulti(mb)
return singbridge.ReturnError(err)
}
for _, buffer := range mb {
packet := B.As(buffer.Bytes()).ToOwned()
buffer.Release()
err = i.service.NewPacket(ctx, pc, packet, metadata)
if err != nil {
packet.Release()
buf.ReleaseMulti(mb)
return err
}
}
payloadBuf := buf.New()
payloadBuf.Write(decoded.Payload)
payloadBuf.UDP = &decoded.Destination
_ = link.Writer.WriteMultiBuffer(buf.MultiBuffer{payloadBuf})
}
}
}
func (i *Inbound) NewConnection(ctx context.Context, conn net.Conn, metadata M.Metadata) error {
inbound := session.InboundFromContext(ctx)
inbound.User = &protocol.MemoryUser{
Email: i.email,
Level: uint32(i.level),
}
ctx = log.ContextWithAccessMessage(ctx, &log.AccessMessage{
From: metadata.Source,
To: metadata.Destination,
Status: log.AccessAccepted,
Email: i.email,
})
errors.LogInfo(ctx, "tunnelling request to tcp:", metadata.Destination)
dispatcher := session.DispatcherFromContext(ctx)
destination, err := singbridge.ToDestination(metadata.Destination, net.Network_TCP)
if err != nil {
return err
}
link, err := dispatcher.Dispatch(ctx, destination)
if err != nil {
return err
}
return singbridge.CopyConn(ctx, nil, link, conn)
}
func (i *Inbound) NewPacketConnection(ctx context.Context, conn N.PacketConn, metadata M.Metadata) error {
inbound := session.InboundFromContext(ctx)
inbound.User = &protocol.MemoryUser{
Email: i.email,
Level: uint32(i.level),
}
ctx = log.ContextWithAccessMessage(ctx, &log.AccessMessage{
From: metadata.Source,
To: metadata.Destination,
Status: log.AccessAccepted,
Email: i.email,
})
errors.LogInfo(ctx, "tunnelling request to udp:", metadata.Destination)
dispatcher := session.DispatcherFromContext(ctx)
destination, err := singbridge.ToDestination(metadata.Destination, net.Network_UDP)
if err != nil {
return err
}
link, err := dispatcher.Dispatch(ctx, destination)
if err != nil {
return err
}
outConn := &singbridge.PacketConnWrapper{
Reader: link.Reader,
Writer: link.Writer,
Dest: destination,
T: signal.CancelAfterInactivity(ctx, func() {
common.Interrupt(link.Reader)
}, 300*time.Second),
}
return bufio.CopyPacketConn(ctx, conn, outConn)
}
func (i *Inbound) NewError(ctx context.Context, err error) {
if E.IsClosed(err) {
return
}
errors.LogWarning(ctx, err.Error())
}
type natPacketConn struct {
net.Conn
}
func (c *natPacketConn) ReadPacket(buffer *B.Buffer) (addr M.Socksaddr, err error) {
_, err = buffer.ReadFrom(c)
return
}
func (c *natPacketConn) WritePacket(buffer *B.Buffer, addr M.Socksaddr) error {
_, err := buffer.WriteTo(c)
return err
}
+177 -243
View File
@@ -2,26 +2,30 @@ package shadowsocks_2022
import (
"context"
"crypto/cipher"
"encoding/binary"
"encoding/base64"
"strconv"
"strings"
"sync"
"sync/atomic"
"time"
"github.com/sagernet/sing-shadowsocks/shadowaead_2022"
C "github.com/sagernet/sing/common"
A "github.com/sagernet/sing/common/auth"
B "github.com/sagernet/sing/common/buf"
"github.com/sagernet/sing/common/bufio"
E "github.com/sagernet/sing/common/exceptions"
M "github.com/sagernet/sing/common/metadata"
N "github.com/sagernet/sing/common/network"
"github.com/xtls/xray-core/common"
"github.com/xtls/xray-core/common/antireplay"
"github.com/xtls/xray-core/common/buf"
"github.com/xtls/xray-core/common/errors"
"github.com/xtls/xray-core/common/log"
"github.com/xtls/xray-core/common/net"
"github.com/xtls/xray-core/common/protocol"
"github.com/xtls/xray-core/common/session"
"github.com/xtls/xray-core/common/utils"
"github.com/xtls/xray-core/common/signal"
"github.com/xtls/xray-core/common/singbridge"
"github.com/xtls/xray-core/common/uuid"
"github.com/xtls/xray-core/core"
"github.com/xtls/xray-core/features/policy"
"github.com/xtls/xray-core/features/routing"
"github.com/xtls/xray-core/transport/internet/stat"
)
@@ -34,16 +38,9 @@ func init() {
type MultiUserInbound struct {
sync.Mutex
networks []net.Network
method *CipherMethod
masterPSK []byte
usersByHash *utils.TypedSyncMap[[AESBlockSize]byte, *protocol.MemoryUser]
usersByEmail *utils.TypedSyncMap[string, *protocol.MemoryUser]
userCount atomic.Int64
saltFilter *antireplay.ReplayFilter[[32]byte]
udpSessions *UDPSessionManager
udpMasterCipher cipher.Block
policyManager policy.Manager
networks []net.Network
users []*protocol.MemoryUser
service *shadowaead_2022.MultiService[int]
}
func NewMultiServer(ctx context.Context, config *MultiUserServerConfig) (*MultiUserInbound, error) {
@@ -54,131 +51,138 @@ func NewMultiServer(ctx context.Context, config *MultiUserServerConfig) (*MultiU
net.Network_UDP,
}
}
method, err := GetCipherMethod(config.Method)
if err != nil {
return nil, err
}
if method.IsChaCha {
return nil, errors.New("shadowsocks 2022 multi-user: only aes methods are supported")
}
masterPSK, err := ParseKey(config.Key, method.KeySaltLength)
if err != nil {
return nil, err
}
masterBlock, err := method.NewBlock(masterPSK)
if err != nil {
return nil, err
}
v := core.MustFromContext(ctx)
i := &MultiUserInbound{
networks: networks,
method: method,
masterPSK: masterPSK,
usersByHash: utils.NewTypedSyncMap[[AESBlockSize]byte, *protocol.MemoryUser](),
usersByEmail: utils.NewTypedSyncMap[string, *protocol.MemoryUser](),
saltFilter: antireplay.NewMapFilter[[32]byte](60),
udpSessions: NewUDPSessionManager(500 * time.Second),
udpMasterCipher: masterBlock,
policyManager: v.GetFeature(policy.ManagerType()).(policy.Manager),
}
for idx, user := range config.Users {
memUsers := []*protocol.MemoryUser{}
for i, user := range config.Users {
if user.Email == "" {
u := uuid.New()
user.Email = "unnamed-user-" + strconv.Itoa(idx) + "-" + u.String()
user.Email = "unnamed-user-" + strconv.Itoa(i) + "-" + u.String()
}
memUser, err := user.ToMemoryUser()
u, err := user.ToMemoryUser()
if err != nil {
return nil, errors.New("failed to parse shadowsocks user").Base(err)
}
if err := i.AddUser(ctx, memUser); err != nil {
return nil, err
return nil, errors.New("failed to get shadowsocks user").Base(err)
}
memUsers = append(memUsers, u)
}
return i, nil
inbound := &MultiUserInbound{
networks: networks,
users: memUsers,
}
if config.Key == "" {
return nil, errors.New("missing key")
}
psk, err := base64.StdEncoding.DecodeString(config.Key)
if err != nil {
return nil, errors.New("parse config").Base(err)
}
service, err := shadowaead_2022.NewMultiService[int](config.Method, psk, 500, inbound, nil)
if err != nil {
return nil, errors.New("create service").Base(err)
}
err = service.UpdateUsersWithPasswords(
C.MapIndexed(memUsers, func(index int, it *protocol.MemoryUser) int { return index }),
C.Map(memUsers, func(it *protocol.MemoryUser) string { return it.Account.(*MemoryAccount).Key }),
)
if err != nil {
return nil, errors.New("create service").Base(err)
}
inbound.service = service
return inbound, nil
}
// AddUser implements proxy.UserManager.AddUser()
// AddUser implements proxy.UserManager.AddUser().
func (i *MultiUserInbound) AddUser(ctx context.Context, u *protocol.MemoryUser) error {
i.Lock()
defer i.Unlock()
var emailKey string
if u.Email != "" {
emailKey = strings.ToLower(u.Email)
if _, exists := i.usersByEmail.Load(emailKey); exists {
return errors.New("user ", u.Email, " already exists")
for idx := range i.users {
if i.users[idx].Email == u.Email {
return errors.New("User ", u.Email, " already exists.")
}
}
}
i.users = append(i.users, u)
memAcc, ok := u.Account.(*MemoryAccount)
if !ok {
return errors.New("missing or invalid user account")
}
if len(memAcc.Key) != i.method.KeySaltLength {
return ErrBadKey
}
pskHash := DeriveUserPSKHash(memAcc.Key)
i.usersByHash.Store(pskHash, u)
if emailKey != "" {
i.usersByEmail.Store(emailKey, u)
}
i.userCount.Add(1)
// sync to multi service
// Considering implements shadowsocks2022 in xray-core may have better performance.
i.service.UpdateUsersWithPasswords(
C.MapIndexed(i.users, func(index int, it *protocol.MemoryUser) int { return index }),
C.Map(i.users, func(it *protocol.MemoryUser) string { return it.Account.(*MemoryAccount).Key }),
)
return nil
}
// RemoveUser implements proxy.UserManager.RemoveUser()
// RemoveUser implements proxy.UserManager.RemoveUser().
func (i *MultiUserInbound) RemoveUser(ctx context.Context, email string) error {
if email == "" {
return errors.New("email must not be empty")
return errors.New("Email must not be empty.")
}
i.Lock()
defer i.Unlock()
emailKey := strings.ToLower(email)
u, loaded := i.usersByEmail.LoadAndDelete(emailKey)
if !loaded {
return errors.New("user ", email, " not found")
idx := -1
for ii, u := range i.users {
if strings.EqualFold(u.Email, email) {
idx = ii
break
}
}
pskHash := DeriveUserPSKHash(u.Account.(*MemoryAccount).Key)
i.usersByHash.Delete(pskHash)
i.userCount.Add(-1)
if idx == -1 {
return errors.New("User ", email, " not found.")
}
ulen := len(i.users)
i.users[idx] = i.users[ulen-1]
i.users[ulen-1] = nil
i.users = i.users[:ulen-1]
// sync to multi service
// Considering implements shadowsocks2022 in xray-core may have better performance.
i.service.UpdateUsersWithPasswords(
C.MapIndexed(i.users, func(index int, it *protocol.MemoryUser) int { return index }),
C.Map(i.users, func(it *protocol.MemoryUser) string { return it.Account.(*MemoryAccount).Key }),
)
return nil
}
// GetUser implements proxy.UserManager.GetUser()
// GetUser implements proxy.UserManager.GetUser().
func (i *MultiUserInbound) GetUser(ctx context.Context, email string) *protocol.MemoryUser {
if email == "" {
return nil
}
u, _ := i.usersByEmail.Load(strings.ToLower(email))
return u
i.Lock()
defer i.Unlock()
for _, u := range i.users {
if strings.EqualFold(u.Email, email) {
return u
}
}
return nil
}
// GetUsers implements proxy.UserManager.GetUsers()
// GetUsers implements proxy.UserManager.GetUsers().
func (i *MultiUserInbound) GetUsers(ctx context.Context) []*protocol.MemoryUser {
var users []*protocol.MemoryUser
i.usersByEmail.Range(func(_ string, user *protocol.MemoryUser) bool {
users = append(users, user)
return true
})
return users
i.Lock()
defer i.Unlock()
dst := make([]*protocol.MemoryUser, len(i.users))
copy(dst, i.users)
return dst
}
// GetUsersCount implements proxy.UserManager.GetUsersCount()
// GetUsersCount implements proxy.UserManager.GetUsersCount().
func (i *MultiUserInbound) GetUsersCount(context.Context) int64 {
return i.userCount.Load()
i.Lock()
defer i.Unlock()
return int64(len(i.users))
}
func (i *MultiUserInbound) Network() []net.Network {
@@ -190,167 +194,97 @@ func (i *MultiUserInbound) Process(ctx context.Context, network net.Network, con
inbound.Name = "shadowsocks-2022-multi"
inbound.CanSpliceCopy = 3
if network == net.Network_TCP {
return i.processTCP(ctx, connection, dispatcher)
var metadata M.Metadata
if inbound.Source.IsValid() {
metadata.Source = M.ParseSocksaddr(inbound.Source.NetAddr())
}
ctx = session.ContextWithDispatcher(ctx, dispatcher)
if network == net.Network_TCP {
return singbridge.ReturnError(i.service.NewConnection(ctx, connection, metadata))
} else {
reader := buf.NewReader(connection)
pc := &natPacketConn{connection}
for {
mb, err := reader.ReadMultiBuffer()
if err != nil {
buf.ReleaseMulti(mb)
return singbridge.ReturnError(err)
}
for _, buffer := range mb {
packet := B.As(buffer.Bytes()).ToOwned()
buffer.Release()
err = i.service.NewPacket(ctx, pc, packet, metadata)
if err != nil {
packet.Release()
buf.ReleaseMulti(mb)
return err
}
}
}
}
return i.processUDP(ctx, connection, dispatcher)
}
func (i *MultiUserInbound) processTCP(ctx context.Context, conn net.Conn, dispatcher routing.Dispatcher) error {
defer conn.Close()
sessionPolicy := i.policyManager.ForLevel(0)
if err := conn.SetReadDeadline(time.Now().Add(sessionPolicy.Timeouts.Handshake)); err != nil {
return errors.New("unable to set read deadline").Base(err)
}
// 1. Single read call for Salt + EIH + Fixed-length header chunk per SIP022 §3.1.4
headerLen := i.method.KeySaltLength + AESBlockSize + RequestHeaderFixedChunkLength + AEADTagSize
headerBuf := make([]byte, headerLen)
n, err := conn.Read(headerBuf)
if err != nil || n < headerLen {
ResetTCPConn(conn)
return errors.New("failed to read complete handshake header")
}
var salt [32]byte
copy(salt[:i.method.KeySaltLength], headerBuf[:i.method.KeySaltLength])
saltSlice := salt[:i.method.KeySaltLength]
eih := headerBuf[i.method.KeySaltLength : i.method.KeySaltLength+AESBlockSize]
fixedChunk := headerBuf[i.method.KeySaltLength+AESBlockSize:]
decryptedHash, err := DecryptEIH(i.method, i.masterPSK, saltSlice, eih)
if err != nil {
ResetTCPConn(conn)
return err
}
// Lookup user
user, ok := i.usersByHash.Load(decryptedHash)
if !ok {
ResetTCPConn(conn)
return ErrInvalidRequest
}
userPSK := user.Account.(*MemoryAccount).Key
reader, reqHeader, err := InitServerStream(conn, i.method, userPSK, saltSlice, salt, fixedChunk, i.saltFilter)
if err != nil {
ResetTCPConn(conn)
return err
}
dest := reqHeader.Destination
writer := NewServerStreamWriter(conn, i.method, userPSK, saltSlice)
// Dispatch Connection to Xray routing with matched User
func (i *MultiUserInbound) NewConnection(ctx context.Context, conn net.Conn, metadata M.Metadata) error {
inbound := session.InboundFromContext(ctx)
userInt, _ := A.UserFromContext[int](ctx)
user := i.users[userInt]
inbound.User = user
ctx = log.ContextWithAccessMessage(ctx, &log.AccessMessage{
From: conn.RemoteAddr(),
To: dest,
From: metadata.Source,
To: metadata.Destination,
Status: log.AccessAccepted,
Email: user.Email,
})
errors.LogInfo(ctx, "tunneling request to ", dest, " for user ", user.Email)
link, err := dispatcher.Dispatch(ctx, dest)
errors.LogInfo(ctx, "tunnelling request to tcp:", metadata.Destination)
dispatcher := session.DispatcherFromContext(ctx)
destination, err := singbridge.ToDestination(metadata.Destination, net.Network_TCP)
if err != nil {
return err
}
if len(reqHeader.EarlyData) > 0 {
mb := buf.MergeBytes(nil, reqHeader.EarlyData)
if err := link.Writer.WriteMultiBuffer(mb); err != nil {
return err
}
link, err := dispatcher.Dispatch(ctx, destination)
if err != nil {
return err
}
return TransportTCP(ctx, i.policyManager.ForLevel(user.Level), reader, writer, link)
return singbridge.CopyConn(ctx, conn, link, conn)
}
func (i *MultiUserInbound) processUDP(ctx context.Context, conn stat.Connection, dispatcher routing.Dispatcher) error {
reader := buf.NewPacketReader(conn)
for {
mb, err := reader.ReadMultiBuffer()
if err != nil {
buf.ReleaseMulti(mb)
return err
}
for _, b := range mb {
// In multi-user UDP:
// Packet header is 16 bytes: Encrypted(SessionID + PacketID)
// Followed by 16 bytes EIH
packetBytes := b.Bytes()
if len(packetBytes) < 32+1+8+2 {
b.Release()
continue
}
var rawHeader [16]byte
i.udpMasterCipher.Decrypt(rawHeader[:], packetBytes[:16])
sessionID := binary.BigEndian.Uint64(rawHeader[:8])
packetID := binary.BigEndian.Uint64(rawHeader[8:16])
sessionItem := i.udpSessions.GetOrCreate(sessionID)
if !sessionItem.CheckPacketID(packetID) {
b.Release()
continue
}
var userPSK []byte
var currentUser *protocol.MemoryUser
sessionItem.Lock()
currentUser = sessionItem.User
userPSK = sessionItem.UserPSK
sessionItem.Unlock()
if currentUser == nil {
// Decrypt EIH
decryptedHash := DecryptUDPEIH(i.udpMasterCipher, rawHeader[:], packetBytes[16:32])
user, ok := i.usersByHash.Load(decryptedHash)
if !ok {
b.Release()
continue
}
currentUser = user
userPSK = user.Account.(*MemoryAccount).Key
}
decoded, err := sessionItem.DecryptAESPayload(i.method, userPSK, sessionID, packetID, rawHeader[:], packetBytes[32:])
b.Release()
if err != nil {
continue
}
sessionItem.Lock()
if sessionItem.User == nil {
sessionItem.User = currentUser
sessionItem.UserPSK = userPSK
}
sessionItem.Unlock()
link, err := sessionItem.EnsureLink(ctx, conn, decoded.Destination, dispatcher, i.policyManager, func(replyDest net.Destination, payload []byte) ([]byte, error) {
return i.encodeServerUDPPacket(sessionID, userPSK, replyDest, payload)
})
if err != nil {
continue
}
pBuf := buf.New()
pBuf.Write(decoded.Payload)
pBuf.UDP = &decoded.Destination
_ = link.Writer.WriteMultiBuffer(buf.MultiBuffer{pBuf})
}
func (i *MultiUserInbound) NewPacketConnection(ctx context.Context, conn N.PacketConn, metadata M.Metadata) error {
inbound := session.InboundFromContext(ctx)
userInt, _ := A.UserFromContext[int](ctx)
user := i.users[userInt]
inbound.User = user
ctx = log.ContextWithAccessMessage(ctx, &log.AccessMessage{
From: metadata.Source,
To: metadata.Destination,
Status: log.AccessAccepted,
Email: user.Email,
})
errors.LogInfo(ctx, "tunnelling request to udp:", metadata.Destination)
dispatcher := session.DispatcherFromContext(ctx)
destination, err := singbridge.ToDestination(metadata.Destination, net.Network_UDP)
if err != nil {
return err
}
link, err := dispatcher.Dispatch(ctx, destination)
if err != nil {
return err
}
outConn := &singbridge.PacketConnWrapper{
Reader: link.Reader,
Writer: link.Writer,
Dest: destination,
T: signal.CancelAfterInactivity(ctx, func() {
common.Interrupt(link.Reader)
}, 300*time.Second),
}
return bufio.CopyPacketConn(ctx, conn, outConn)
}
func (i *MultiUserInbound) encodeServerUDPPacket(clientSessionID uint64, userPSK []byte, dest net.Destination, payload []byte) ([]byte, error) {
return i.udpSessions.EncodeServerPacket(i.method, userPSK, clientSessionID, dest, payload)
func (i *MultiUserInbound) NewError(ctx context.Context, err error) {
if E.IsClosed(err) {
return
}
errors.LogWarning(ctx, err.Error())
}
+135 -206
View File
@@ -2,11 +2,18 @@ package shadowsocks_2022
import (
"context"
"crypto/cipher"
"encoding/binary"
"strconv"
"strings"
"time"
"github.com/sagernet/sing-shadowsocks/shadowaead_2022"
C "github.com/sagernet/sing/common"
A "github.com/sagernet/sing/common/auth"
B "github.com/sagernet/sing/common/buf"
"github.com/sagernet/sing/common/bufio"
E "github.com/sagernet/sing/common/exceptions"
M "github.com/sagernet/sing/common/metadata"
N "github.com/sagernet/sing/common/network"
"github.com/xtls/xray-core/common"
"github.com/xtls/xray-core/common/buf"
"github.com/xtls/xray-core/common/errors"
@@ -14,9 +21,9 @@ import (
"github.com/xtls/xray-core/common/net"
"github.com/xtls/xray-core/common/protocol"
"github.com/xtls/xray-core/common/session"
"github.com/xtls/xray-core/common/signal"
"github.com/xtls/xray-core/common/singbridge"
"github.com/xtls/xray-core/common/uuid"
"github.com/xtls/xray-core/core"
"github.com/xtls/xray-core/features/policy"
"github.com/xtls/xray-core/features/routing"
"github.com/xtls/xray-core/transport/internet/stat"
)
@@ -27,21 +34,10 @@ func init() {
}))
}
type relayDest struct {
destination net.Destination
email string
level uint32
blockCipher cipher.Block
}
type RelayInbound struct {
networks []net.Network
method *CipherMethod
relayPSK []byte
relayBlock cipher.Block
destinations map[[AESBlockSize]byte]*relayDest
udpSessions *UDPSessionManager
policyManager policy.Manager
networks []net.Network
destinations []*RelayDestination
service *shadowaead_2022.RelayService[int]
}
func NewRelayServer(ctx context.Context, config *RelayServerConfig) (*RelayInbound, error) {
@@ -52,62 +48,39 @@ func NewRelayServer(ctx context.Context, config *RelayServerConfig) (*RelayInbou
net.Network_UDP,
}
}
method, err := GetCipherMethod(config.Method)
inbound := &RelayInbound{
networks: networks,
destinations: config.Destinations,
}
if !C.Contains(shadowaead_2022.List, config.Method) || !strings.Contains(config.Method, "aes") {
return nil, errors.New("unsupported method ", config.Method)
}
service, err := shadowaead_2022.NewRelayServiceWithPassword[int](config.Method, config.Key, 500, inbound)
if err != nil {
return nil, err
}
if method.IsChaCha {
return nil, errors.New("shadowsocks 2022 relay: only aes methods are supported")
return nil, errors.New("create service").Base(err)
}
relayPSK, err := ParseKey(config.Key, method.KeySaltLength)
if err != nil {
return nil, err
}
relayBlock, err := method.NewBlock(relayPSK)
if err != nil {
return nil, err
}
v := core.MustFromContext(ctx)
i := &RelayInbound{
networks: networks,
method: method,
relayPSK: relayPSK,
relayBlock: relayBlock,
destinations: make(map[[AESBlockSize]byte]*relayDest),
udpSessions: NewUDPSessionManager(500 * time.Second),
policyManager: v.GetFeature(policy.ManagerType()).(policy.Manager),
}
for idx, d := range config.Destinations {
if d.Email == "" {
for i, destination := range config.Destinations {
if destination.Email == "" {
u := uuid.New()
d.Email = "unnamed-destination-" + strconv.Itoa(idx) + "-" + u.String()
}
destKey, err := ParseKey(d.Key, method.KeySaltLength)
if err != nil {
return nil, err
}
destBlock, err := method.NewBlock(destKey)
if err != nil {
return nil, err
}
hash := DeriveUserPSKHash(destKey)
i.destinations[hash] = &relayDest{
destination: net.TCPDestination(d.Address.AsAddress(), net.Port(d.Port)),
email: d.Email,
level: uint32(d.Level),
blockCipher: destBlock,
destination.Email = "unnamed-destination-" + strconv.Itoa(i) + "-" + u.String()
}
}
return i, nil
err = service.UpdateUsersWithPasswords(
C.MapIndexed(config.Destinations, func(index int, it *RelayDestination) int { return index }),
C.Map(config.Destinations, func(it *RelayDestination) string { return it.Key }),
C.Map(config.Destinations, func(it *RelayDestination) M.Socksaddr {
return singbridge.ToSocksaddr(net.Destination{
Address: it.Address.AsAddress(),
Port: net.Port(it.Port),
})
}),
)
if err != nil {
return nil, errors.New("create service").Base(err)
}
inbound.service = service
return inbound, nil
}
func (i *RelayInbound) Network() []net.Network {
@@ -119,147 +92,103 @@ func (i *RelayInbound) Process(ctx context.Context, network net.Network, connect
inbound.Name = "shadowsocks-2022-relay"
inbound.CanSpliceCopy = 3
var metadata M.Metadata
if inbound.Source.IsValid() {
metadata.Source = M.ParseSocksaddr(inbound.Source.NetAddr())
}
ctx = session.ContextWithDispatcher(ctx, dispatcher)
if network == net.Network_TCP {
return i.processTCP(ctx, connection, dispatcher)
}
return i.processUDP(ctx, connection, dispatcher)
}
func (i *RelayInbound) processTCP(ctx context.Context, conn net.Conn, dispatcher routing.Dispatcher) error {
defer conn.Close()
sessionPolicy := i.policyManager.ForLevel(0)
if err := conn.SetReadDeadline(time.Now().Add(sessionPolicy.Timeouts.Handshake)); err != nil {
return errors.New("unable to set read deadline").Base(err)
}
// Read initial handshake in a single read call per SIP022 §3.1.3 & §3.1.4
needed := i.method.KeySaltLength + AESBlockSize
requestHeader := buf.New()
n, err := requestHeader.ReadFrom(conn)
if err != nil {
requestHeader.Release()
ResetTCPConn(conn)
return err
}
if int(n) < needed {
requestHeader.Release()
ResetTCPConn(conn)
return ErrInvalidRequest
}
headerSlice := requestHeader.Bytes()
salt := headerSlice[:i.method.KeySaltLength]
eih := headerSlice[i.method.KeySaltLength:needed]
decryptedHash, err := DecryptEIH(i.method, i.relayPSK, salt, eih)
if err != nil {
requestHeader.Release()
ResetTCPConn(conn)
return err
}
targetDest, ok := i.destinations[decryptedHash]
if !ok {
requestHeader.Release()
ResetTCPConn(conn)
return ErrInvalidRequest
}
conn.SetReadDeadline(time.Time{})
inbound := session.InboundFromContext(ctx)
inbound.User = &protocol.MemoryUser{
Email: targetDest.email,
Level: targetDest.level,
}
ctx = log.ContextWithAccessMessage(ctx, &log.AccessMessage{
From: conn.RemoteAddr(),
To: targetDest.destination,
Status: log.AccessAccepted,
Email: targetDest.email,
})
errors.LogInfo(ctx, "relaying connection to ", targetDest.destination)
link, err := dispatcher.Dispatch(ctx, targetDest.destination)
if err != nil {
requestHeader.Release()
return err
}
// Unwrap outer EIH: send client salt and remaining handshake bytes to next hop
// in a single write call, satisfying downstream server's single-read handshake expectation (SIP022 §3.1.3).
var saltCopy [32]byte
copy(saltCopy[:i.method.KeySaltLength], salt)
copy(requestHeader.Bytes()[AESBlockSize:AESBlockSize+i.method.KeySaltLength], saltCopy[:i.method.KeySaltLength])
requestHeader.Advance(AESBlockSize)
if err := link.Writer.WriteMultiBuffer(buf.MultiBuffer{requestHeader}); err != nil {
return err
}
return TransportTCP(ctx, i.policyManager.ForLevel(targetDest.level), buf.NewReader(conn), buf.NewWriter(conn), link)
}
func (i *RelayInbound) processUDP(ctx context.Context, conn stat.Connection, dispatcher routing.Dispatcher) error {
reader := buf.NewPacketReader(conn)
for {
mb, err := reader.ReadMultiBuffer()
if err != nil {
buf.ReleaseMulti(mb)
return err
}
for _, b := range mb {
data := b.Bytes()
if len(data) < 2*AESBlockSize {
b.Release()
continue
}
var packetHeader [AESBlockSize]byte
i.relayBlock.Decrypt(packetHeader[:], data[:AESBlockSize])
eiHeader := DecryptUDPEIH(i.relayBlock, packetHeader[:], data[AESBlockSize:2*AESBlockSize])
targetDest, ok := i.destinations[eiHeader]
if !ok {
b.Release()
continue
}
// Extract sessionID from raw packetHeader for session-level link caching before re-encrypting
sessionID := binary.BigEndian.Uint64(packetHeader[:8])
// Re-encrypt packetHeader with next hop block cipher
targetDest.blockCipher.Encrypt(packetHeader[:], packetHeader[:])
// Strip outer EIH: replace second block with re-encrypted packetHeader and advance
copy(data[AESBlockSize:2*AESBlockSize], packetHeader[:])
b.Advance(int32(AESBlockSize))
dest := targetDest.destination
dest.Network = net.Network_UDP
sessionItem := i.udpSessions.GetOrCreate(sessionID)
if sessionItem.User == nil {
sessionItem.Lock()
if sessionItem.User == nil {
sessionItem.User = &protocol.MemoryUser{
Email: targetDest.email,
Level: targetDest.level,
}
}
sessionItem.Unlock()
}
link, err := sessionItem.EnsureLink(ctx, conn, dest, dispatcher, i.policyManager, nil)
return singbridge.ReturnError(i.service.NewConnection(ctx, connection, metadata))
} else {
reader := buf.NewReader(connection)
pc := &natPacketConn{connection}
for {
mb, err := reader.ReadMultiBuffer()
if err != nil {
b.Release()
continue
buf.ReleaseMulti(mb)
return singbridge.ReturnError(err)
}
for _, buffer := range mb {
packet := B.As(buffer.Bytes()).ToOwned()
buffer.Release()
err = i.service.NewPacket(ctx, pc, packet, metadata)
if err != nil {
packet.Release()
buf.ReleaseMulti(mb)
return err
}
}
_ = link.Writer.WriteMultiBuffer(buf.MultiBuffer{b})
}
}
}
func (i *RelayInbound) NewConnection(ctx context.Context, conn net.Conn, metadata M.Metadata) error {
inbound := session.InboundFromContext(ctx)
userInt, _ := A.UserFromContext[int](ctx)
user := i.destinations[userInt]
inbound.User = &protocol.MemoryUser{
Email: user.Email,
Level: uint32(user.Level),
}
ctx = log.ContextWithAccessMessage(ctx, &log.AccessMessage{
From: metadata.Source,
To: metadata.Destination,
Status: log.AccessAccepted,
Email: user.Email,
})
errors.LogInfo(ctx, "tunnelling request to tcp:", metadata.Destination)
dispatcher := session.DispatcherFromContext(ctx)
destination, err := singbridge.ToDestination(metadata.Destination, net.Network_TCP)
if err != nil {
return err
}
link, err := dispatcher.Dispatch(ctx, destination)
if err != nil {
return err
}
return singbridge.CopyConn(ctx, nil, link, conn)
}
func (i *RelayInbound) NewPacketConnection(ctx context.Context, conn N.PacketConn, metadata M.Metadata) error {
inbound := session.InboundFromContext(ctx)
userInt, _ := A.UserFromContext[int](ctx)
user := i.destinations[userInt]
inbound.User = &protocol.MemoryUser{
Email: user.Email,
Level: uint32(user.Level),
}
ctx = log.ContextWithAccessMessage(ctx, &log.AccessMessage{
From: metadata.Source,
To: metadata.Destination,
Status: log.AccessAccepted,
Email: user.Email,
})
errors.LogInfo(ctx, "tunnelling request to udp:", metadata.Destination)
dispatcher := session.DispatcherFromContext(ctx)
destination, err := singbridge.ToDestination(metadata.Destination, net.Network_UDP)
if err != nil {
return err
}
link, err := dispatcher.Dispatch(ctx, destination)
if err != nil {
return err
}
outConn := &singbridge.PacketConnWrapper{
Reader: link.Reader,
Writer: link.Writer,
Dest: destination,
T: signal.CancelAfterInactivity(ctx, func() {
common.Interrupt(link.Reader)
}, 300*time.Second),
}
return bufio.CopyPacketConn(ctx, conn, outConn)
}
func (i *RelayInbound) NewError(ctx context.Context, err error) {
if E.IsClosed(err) {
return
}
errors.LogWarning(ctx, err.Error())
}
-74
View File
@@ -1,74 +0,0 @@
package shadowsocks_2022
import (
"encoding/base64"
"strings"
"lukechampine.com/blake3"
)
const (
ContextSessionSubKey = "shadowsocks 2022 session subkey"
ContextIdentitySubKey = "shadowsocks 2022 identity subkey"
)
// ParseKey decodes a base64 or raw PSK key string and validates its length
func ParseKey(key string, keyLength int) ([]byte, error) {
raw, err := base64.StdEncoding.DecodeString(key)
if err != nil {
raw = []byte(key)
}
if len(raw) != keyLength {
return nil, ErrBadKey
}
return raw, nil
}
func ParsePSKList(password string, keyLength int) ([][]byte, error) {
parts := strings.Split(password, ":")
pskList := make([][]byte, len(parts))
for i, part := range parts {
norm, err := ParseKey(part, keyLength)
if err != nil {
return nil, err
}
pskList[i] = norm
}
return pskList, nil
}
func deriveSubKey(ctx string, psk, salt []byte, keyLength int) []byte {
var keyMaterial [64]byte
kmLen := len(psk) + len(salt)
copy(keyMaterial[:], psk)
copy(keyMaterial[len(psk):], salt)
out := make([]byte, keyLength)
blake3.DeriveKey(out, ctx, keyMaterial[:kmLen])
return out
}
func DeriveSessionSubKey(psk, salt []byte, keyLength int) []byte {
return deriveSubKey(ContextSessionSubKey, psk, salt, keyLength)
}
func DeriveIdentitySubKey(psk, salt []byte, keyLength int) []byte {
return deriveSubKey(ContextIdentitySubKey, psk, salt, keyLength)
}
func DeriveUserPSKHash(userPSK []byte) [AESBlockSize]byte {
h := blake3.Sum512(userPSK)
var out [AESBlockSize]byte
copy(out[:], h[:AESBlockSize])
return out
}
func DecryptEIH(method *CipherMethod, key, salt, eih []byte) ([AESBlockSize]byte, error) {
identitySubkey := DeriveIdentitySubKey(key, salt, method.KeySaltLength)
block, err := method.NewBlock(identitySubkey)
if err != nil {
return [AESBlockSize]byte{}, err
}
var decryptedHash [AESBlockSize]byte
block.Decrypt(decryptedHash[:], eih)
return decryptedHash, nil
}
+84 -161
View File
@@ -2,19 +2,21 @@ package shadowsocks_2022
import (
"context"
"crypto/rand"
"io"
"time"
shadowsocks "github.com/sagernet/sing-shadowsocks"
"github.com/sagernet/sing-shadowsocks/shadowaead_2022"
C "github.com/sagernet/sing/common"
B "github.com/sagernet/sing/common/buf"
"github.com/sagernet/sing/common/bufio"
N "github.com/sagernet/sing/common/network"
"github.com/xtls/xray-core/common"
"github.com/xtls/xray-core/common/buf"
"github.com/xtls/xray-core/common/errors"
"github.com/xtls/xray-core/common/net"
"github.com/xtls/xray-core/common/retry"
"github.com/xtls/xray-core/common/session"
"github.com/xtls/xray-core/common/signal"
"github.com/xtls/xray-core/common/task"
"github.com/xtls/xray-core/core"
"github.com/xtls/xray-core/features/policy"
"github.com/xtls/xray-core/common/singbridge"
"github.com/xtls/xray-core/transport"
"github.com/xtls/xray-core/transport/internet"
)
@@ -26,51 +28,42 @@ func init() {
}
type Outbound struct {
server net.Destination
method *CipherMethod
pskList [][]byte
finalPSK []byte
udpCodec *UDPPacketCodec
policyManager policy.Manager
ctx context.Context
server net.Destination
method shadowsocks.Method
}
func NewClient(ctx context.Context, config *ClientConfig) (*Outbound, error) {
method, err := GetCipherMethod(config.Method)
if err != nil {
return nil, errors.New("unsupported method: ", config.Method).Base(err)
}
pskList, err := ParsePSKList(config.Key, method.KeySaltLength)
if err != nil {
return nil, errors.New("invalid key: ", config.Key).Base(err)
}
if method.IsChaCha && len(pskList) > 1 {
return nil, errors.New("multi-key is not supported for chacha20-poly1305")
}
finalPSK := pskList[len(pskList)-1]
udpCodec, err := NewUDPPacketCodec(method, pskList)
if err != nil {
return nil, errors.New("failed to create udp packet codec").Base(err)
}
v := core.MustFromContext(ctx)
return &Outbound{
o := &Outbound{
ctx: ctx,
server: net.Destination{
Address: config.Address.AsAddress(),
Port: net.Port(config.Port),
Network: net.Network_TCP,
},
method: method,
pskList: pskList,
finalPSK: finalPSK,
udpCodec: udpCodec,
policyManager: v.GetFeature(policy.ManagerType()).(policy.Manager),
}, nil
}
if C.Contains(shadowaead_2022.List, config.Method) {
if config.Key == "" {
return nil, errors.New("missing psk")
}
method, err := shadowaead_2022.NewWithPassword(config.Method, config.Key, nil)
if err != nil {
return nil, errors.New("create method").Base(err)
}
o.method = method
} else {
return nil, errors.New("unknown method ", config.Method)
}
return o, nil
}
func (o *Outbound) Process(ctx context.Context, link *transport.Link, dialer internet.Dialer) error {
var inboundConn net.Conn
inbound := session.InboundFromContext(ctx)
if inbound != nil {
inboundConn = inbound.Conn
}
outbounds := session.OutboundsFromContext(ctx)
ob := outbounds[len(outbounds)-1]
if !ob.Target.IsValid() {
@@ -85,140 +78,70 @@ func (o *Outbound) Process(ctx context.Context, link *transport.Link, dialer int
serverDestination := o.server
serverDestination.Network = network
var conn net.Conn
if err := retry.ExponentialBackoff(5, 100).On(func() error {
rawConn, err := dialer.Dial(ctx, serverDestination)
if err != nil {
return err
}
conn = rawConn
return nil
}); err != nil {
return errors.New("failed to find an available destination").Base(err)
connection, err := dialer.Dial(ctx, serverDestination)
if err != nil {
return errors.New("failed to connect to server").Base(err)
}
defer conn.Close()
defer connection.Close()
var newCtx context.Context
var newCancel context.CancelFunc
if session.TimeoutOnlyFromContext(ctx) {
newCtx, newCancel = context.WithCancel(context.Background())
}
sessionPolicy := o.policyManager.ForLevel(0)
ctx, cancel := context.WithCancel(ctx)
timer := signal.CancelAfterInactivity(ctx, func() {
cancel()
if newCancel != nil {
newCancel()
}
}, sessionPolicy.Timeouts.ConnectionIdle)
ctx = policy.ContextWithBufferPolicy(ctx, sessionPolicy.Buffer)
if newCtx != nil {
ctx = newCtx
ctx, _ = context.WithCancel(context.Background())
}
if network == net.Network_TCP {
var clientSalt [32]byte
clientSaltSlice := clientSalt[:o.method.KeySaltLength]
if _, err := io.ReadFull(rand.Reader, clientSaltSlice); err != nil {
return errors.New("failed to generate client salt").Base(err)
}
requestDone := func() error {
defer timer.SetTimeout(sessionPolicy.Timeouts.DownlinkOnly)
var initialPayload []byte
var firstBuf *buf.Buffer
var remainingMB buf.MultiBuffer
if timeoutReader, ok := link.Reader.(buf.TimeoutReader); ok {
if mb, err := timeoutReader.ReadMultiBufferTimeout(0); err == nil && !mb.IsEmpty() {
remainingMB, firstBuf = buf.SplitFirst(mb)
initialPayload = firstBuf.Bytes()
serverConn := o.method.DialEarlyConn(connection, singbridge.ToSocksaddr(destination))
var handshake bool
if timeoutReader, isTimeoutReader := link.Reader.(buf.TimeoutReader); isTimeoutReader {
mb, err := timeoutReader.ReadMultiBufferTimeout(time.Millisecond * 100)
if err != nil && err != buf.ErrNotTimeoutReader && err != buf.ErrReadTimeout {
return errors.New("read payload").Base(err)
}
payload := B.New()
for {
payload.Reset()
nb, n := buf.SplitBytes(mb, payload.FreeBytes())
if n > 0 {
payload.Truncate(n)
_, err = serverConn.Write(payload.Bytes())
if err != nil {
payload.Release()
return errors.New("write payload").Base(err)
}
handshake = true
}
}
bodyWriter, err := WriteTCPRequest(conn, o.method, o.pskList, destination, clientSaltSlice, initialPayload)
if firstBuf != nil {
firstBuf.Release()
}
if err != nil {
buf.ReleaseMulti(remainingMB)
return errors.New("failed to write request").Base(err)
}
if !remainingMB.IsEmpty() {
if err := bodyWriter.WriteMultiBuffer(remainingMB); err != nil {
return err
if nb.IsEmpty() {
break
}
mb = nb
}
return buf.Copy(link.Reader, bodyWriter, buf.UpdateActivity(timer))
payload.Release()
}
responseDone := func() error {
defer timer.SetTimeout(sessionPolicy.Timeouts.UplinkOnly)
responseReader, err := ReadTCPResponse(conn, o.method, o.finalPSK, clientSaltSlice)
if !handshake {
_, err = serverConn.Write(nil)
if err != nil {
return err
return errors.New("client handshake").Base(err)
}
}
return singbridge.CopyConn(ctx, inboundConn, link, serverConn)
} else {
var packetConn N.PacketConn
if pc, isPacketConn := inboundConn.(N.PacketConn); isPacketConn {
packetConn = pc
} else if nc, isNetPacket := inboundConn.(net.PacketConn); isNetPacket {
packetConn = bufio.NewPacketConn(nc)
} else {
packetConn = &singbridge.PacketConnWrapper{
Reader: link.Reader,
Writer: link.Writer,
Conn: inboundConn,
Dest: destination,
T: signal.CancelAfterInactivity(ctx, func() {
common.Interrupt(link.Reader)
}, 300*time.Second),
}
return buf.Copy(responseReader, link.Writer, buf.UpdateActivity(timer))
}
responseDoneAndCloseWriter := task.OnSuccess(responseDone, task.Close(link.Writer))
if err := task.Run(ctx, requestDone, responseDoneAndCloseWriter); err != nil {
return errors.New("connection ends").Base(err)
}
return nil
serverConn := o.method.DialPacketConn(connection)
return singbridge.ReturnError(bufio.CopyPacketConn(ctx, packetConn, serverConn))
}
if network == net.Network_UDP {
session, err := o.udpCodec.NewClientSession()
if err != nil {
return errors.New("failed to create client udp session").Base(err)
}
requestDone := func() error {
defer timer.SetTimeout(sessionPolicy.Timeouts.DownlinkOnly)
writer := &UDPWriter{
Writer: conn,
Destination: destination,
Session: session,
}
if err := buf.Copy(link.Reader, writer, buf.UpdateActivity(timer)); err != nil {
return errors.New("failed to transport all UDP request").Base(err)
}
return nil
}
responseDone := func() error {
defer timer.SetTimeout(sessionPolicy.Timeouts.UplinkOnly)
reader := &UDPReader{
Reader: conn,
Session: session,
}
if err := buf.Copy(reader, link.Writer, buf.UpdateActivity(timer)); err != nil {
return errors.New("failed to transport all UDP response").Base(err)
}
return nil
}
responseDoneAndCloseWriter := task.OnSuccess(responseDone, task.Close(link.Writer))
if err := task.Run(ctx, requestDone, responseDoneAndCloseWriter); err != nil {
return errors.New("connection ends").Base(err)
}
return nil
}
return errors.New("unsupported network: ", network)
}
-766
View File
@@ -1,766 +0,0 @@
package shadowsocks_2022
import (
"crypto/cipher"
"crypto/rand"
"encoding/binary"
"io"
"math"
mrand "math/rand/v2"
"sync/atomic"
"time"
"github.com/xtls/xray-core/common/buf"
"github.com/xtls/xray-core/common/errors"
"github.com/xtls/xray-core/common/net"
)
type UDPCodec struct {
method *CipherMethod
pskList [][]byte
psk []byte
blockCipher cipher.Block
blockCiphers []cipher.Block
chachaCipher cipher.AEAD
sessions *UDPSessionManager
}
type (
UDPPacketCodec = UDPCodec
UDPServerCodec = UDPCodec
)
func newUDPCodec(method *CipherMethod, psk []byte) (*UDPCodec, error) {
c := &UDPCodec{
method: method,
psk: psk,
}
var err error
if method.IsChaCha {
c.chachaCipher, err = method.NewUDPCipher(psk)
} else {
c.blockCipher, err = method.NewBlock(psk)
}
if err != nil {
return nil, err
}
return c, nil
}
func NewUDPPacketCodec(method *CipherMethod, pskList [][]byte) (*UDPCodec, error) {
if method.IsChaCha && len(pskList) > 1 {
return nil, errors.New("multi-key is not supported for chacha20-poly1305")
}
finalPSK := pskList[len(pskList)-1]
c, err := newUDPCodec(method, finalPSK)
if err != nil {
return nil, err
}
c.pskList = pskList
if len(pskList) > 1 {
c.blockCiphers = make([]cipher.Block, len(pskList))
for i, psk := range pskList {
c.blockCiphers[i], err = method.NewBlock(psk)
if err != nil {
return nil, err
}
}
}
return c, nil
}
func NewUDPServerCodec(method *CipherMethod, psk []byte, sessionTimeout time.Duration) (*UDPCodec, error) {
c, err := newUDPCodec(method, psk)
if err != nil {
return nil, err
}
c.sessions = NewUDPSessionManager(sessionTimeout)
return c, nil
}
func (c *UDPCodec) Sessions() *UDPSessionManager {
return c.sessions
}
func (c *UDPCodec) GetSession(sessionID uint64) *ServerUDPSession {
if c.sessions == nil {
return nil
}
return c.sessions.GetOrCreate(sessionID)
}
type DecodedUDPPacket struct {
SessionID uint64
PacketID uint64
HeaderType byte
Timestamp uint64
ClientSessionID uint64
Destination net.Destination
Payload []byte
}
func DecryptUDPEIH(block cipher.Block, rawHeader, eih []byte) [AESBlockSize]byte {
var decryptedHash [AESBlockSize]byte
block.Decrypt(decryptedHash[:], eih)
for k := 0; k < AESBlockSize; k++ {
decryptedHash[k] ^= rawHeader[k]
}
return decryptedHash
}
func ParseAddressPort(data []byte) (net.Destination, int, error) {
if len(data) < 1 {
return net.Destination{}, 0, ErrPacketTooShort
}
switch data[0] {
case 1: // IPv4
if len(data) < 1+4+2 {
return net.Destination{}, 0, ErrPacketTooShort
}
ip := net.IPAddress(data[1:5])
port := binary.BigEndian.Uint16(data[5:7])
return net.UDPDestination(ip, net.Port(port)), 7, nil
case 4: // IPv6
if len(data) < 1+16+2 {
return net.Destination{}, 0, ErrPacketTooShort
}
ip := net.IPAddress(data[1:17])
port := binary.BigEndian.Uint16(data[17:19])
return net.UDPDestination(ip, net.Port(port)), 19, nil
case 3: // Domain
if len(data) < 2 {
return net.Destination{}, 0, ErrPacketTooShort
}
domainLen := int(data[1])
if len(data) < 2+domainLen+2 {
return net.Destination{}, 0, ErrPacketTooShort
}
domain := string(data[2 : 2+domainLen])
port := binary.BigEndian.Uint16(data[2+domainLen : 2+domainLen+2])
return net.UDPDestination(net.DomainAddress(domain), net.Port(port)), 2 + domainLen + 2, nil
default:
return net.Destination{}, 0, errors.New("unknown address type")
}
}
func parsePlainUDPPacket(sessionID, packetID uint64, bodyPlain []byte) (DecodedUDPPacket, error) {
if len(bodyPlain) < 1+8+2 {
return DecodedUDPPacket{}, ErrPacketTooShort
}
headerType := bodyPlain[0]
if headerType != HeaderTypeClient && headerType != HeaderTypeServer {
return DecodedUDPPacket{}, ErrBadHeaderType
}
epoch := binary.BigEndian.Uint64(bodyPlain[1:9])
diff := int(math.Abs(float64(time.Now().Unix() - int64(epoch))))
if diff > 30 {
return DecodedUDPPacket{}, ErrBadTimestamp
}
offset := 9
var clientSessionID uint64
if headerType == HeaderTypeServer {
if len(bodyPlain) < offset+8+2 {
return DecodedUDPPacket{}, ErrPacketTooShort
}
clientSessionID = binary.BigEndian.Uint64(bodyPlain[offset : offset+8])
offset += 8
}
paddingLen := int(binary.BigEndian.Uint16(bodyPlain[offset : offset+2]))
offset += 2
if len(bodyPlain) < offset+paddingLen {
return DecodedUDPPacket{}, ErrNoPadding
}
offset += paddingLen
dest, addrLen, err := ParseAddressPort(bodyPlain[offset:])
if err != nil {
return DecodedUDPPacket{}, err
}
payload := bodyPlain[offset+addrLen:]
return DecodedUDPPacket{
SessionID: sessionID,
PacketID: packetID,
HeaderType: headerType,
Timestamp: epoch,
ClientSessionID: clientSessionID,
Destination: dest,
Payload: payload,
}, nil
}
func (c *UDPCodec) DecodePacket(data []byte) (DecodedUDPPacket, error) {
if len(data) < PacketMinimalHeaderSize {
return DecodedUDPPacket{}, ErrPacketTooShort
}
if c.method.IsChaCha {
if len(data) < PacketNonceSize+AEADTagSize {
return DecodedUDPPacket{}, ErrPacketTooShort
}
nonce := data[:PacketNonceSize]
ciphertext := data[PacketNonceSize:]
plain, err := c.chachaCipher.Open(nil, nonce, ciphertext, nil)
if err != nil {
return DecodedUDPPacket{}, errors.New("failed to decrypt chacha udp packet").Base(err)
}
if len(plain) < 16+1+8+2 {
return DecodedUDPPacket{}, ErrPacketTooShort
}
sessionID := binary.BigEndian.Uint64(plain[:8])
packetID := binary.BigEndian.Uint64(plain[8:16])
sessionItem := c.sessions.GetOrCreate(sessionID)
if !sessionItem.CheckPacketID(packetID) {
return DecodedUDPPacket{}, ErrPacketIdNotUnique
}
decoded, err := parsePlainUDPPacket(sessionID, packetID, plain[16:])
if err != nil {
return DecodedUDPPacket{}, err
}
if decoded.HeaderType != HeaderTypeClient {
return DecodedUDPPacket{}, ErrBadHeaderType
}
sessionItem.AddPacketID(packetID)
return decoded, nil
}
// AES mode
var rawHeader [16]byte
c.blockCipher.Decrypt(rawHeader[:], data[:16])
sessionID := binary.BigEndian.Uint64(rawHeader[:8])
packetID := binary.BigEndian.Uint64(rawHeader[8:16])
sessionItem := c.sessions.GetOrCreate(sessionID)
if !sessionItem.CheckPacketID(packetID) {
return DecodedUDPPacket{}, ErrPacketIdNotUnique
}
return sessionItem.DecryptAESPayload(c.method, c.psk, sessionID, packetID, rawHeader[:], data[16:])
}
func (s *ServerUDPSession) DecryptAESPayload(method *CipherMethod, psk []byte, sessionID, packetID uint64, rawHeader, bodyCipher []byte) (DecodedUDPPacket, error) {
bodyAead := s.clientBodyCipher
isNewCipher := false
if bodyAead == nil {
bodyKey := DeriveSessionSubKey(psk, rawHeader[:8], method.KeySaltLength)
var err error
bodyAead, err = method.NewAEAD(bodyKey)
if err != nil {
return DecodedUDPPacket{}, err
}
isNewCipher = true
}
bodyNonce := rawHeader[4:16]
bodyPlain, err := bodyAead.Open(nil, bodyNonce, bodyCipher, nil)
if err != nil {
return DecodedUDPPacket{}, errors.New("failed to decrypt aes udp body").Base(err)
}
decoded, err := parsePlainUDPPacket(sessionID, packetID, bodyPlain)
if err != nil {
return DecodedUDPPacket{}, err
}
if decoded.HeaderType != HeaderTypeClient {
return DecodedUDPPacket{}, ErrBadHeaderType
}
s.AddPacketID(packetID)
if isNewCipher {
s.clientBodyCipher = bodyAead
}
return decoded, nil
}
func (s *ServerUDPSession) EnsureServerState(method *CipherMethod, psk []byte) error {
s.Lock()
defer s.Unlock()
if s.ServerSessionID != 0 {
return nil
}
var sidBuf [8]byte
for {
if _, err := io.ReadFull(rand.Reader, sidBuf[:]); err != nil {
return err
}
s.ServerSessionID = binary.BigEndian.Uint64(sidBuf[:])
if s.ServerSessionID != 0 {
break
}
}
if method.IsChaCha {
var err error
s.serverChaCha, err = method.NewUDPCipher(psk)
return err
}
var err error
s.serverHeaderBlock, err = method.NewBlock(psk)
if err != nil {
s.ServerSessionID = 0
return err
}
bodyKey := DeriveSessionSubKey(psk, sidBuf[:], method.KeySaltLength)
s.serverBodyCipher, err = method.NewAEAD(bodyKey)
if err != nil {
s.ServerSessionID = 0
return err
}
return nil
}
func (s *ServerUDPSession) EncodeServerPacket(method *CipherMethod, clientSessionID uint64, dest net.Destination, payload []byte) ([]byte, error) {
serverSessionID := s.ServerSessionID
serverPacketID := s.ServerPacketID.Add(1) - 1
if method.IsChaCha {
var nonce [PacketNonceSize]byte
if _, err := io.ReadFull(rand.Reader, nonce[:]); err != nil {
return nil, err
}
plainBuf := buf.New()
defer plainBuf.Release()
var hdr [16 + 1 + 8 + 8 + 2]byte
binary.BigEndian.PutUint64(hdr[0:8], serverSessionID)
binary.BigEndian.PutUint64(hdr[8:16], serverPacketID)
hdr[16] = HeaderTypeServer
binary.BigEndian.PutUint64(hdr[17:25], uint64(time.Now().Unix()))
binary.BigEndian.PutUint64(hdr[25:33], clientSessionID)
binary.BigEndian.PutUint16(hdr[33:35], 0)
plainBuf.Write(hdr[:])
if err := WriteAddressPort(plainBuf, dest); err != nil {
return nil, err
}
plainBuf.Write(payload)
sealed := s.serverChaCha.Seal(nil, nonce[:], plainBuf.Bytes(), nil)
res := make([]byte, PacketNonceSize+len(sealed))
copy(res[:PacketNonceSize], nonce[:])
copy(res[PacketNonceSize:], sealed)
return res, nil
}
// AES mode
var rawHeader [16]byte
binary.BigEndian.PutUint64(rawHeader[:8], serverSessionID)
binary.BigEndian.PutUint64(rawHeader[8:16], serverPacketID)
var encryptedHeader [16]byte
s.serverHeaderBlock.Encrypt(encryptedHeader[:], rawHeader[:])
bodyBuf := buf.New()
defer bodyBuf.Release()
var hdr [1 + 8 + 8 + 2]byte
hdr[0] = HeaderTypeServer
binary.BigEndian.PutUint64(hdr[1:9], uint64(time.Now().Unix()))
binary.BigEndian.PutUint64(hdr[9:17], clientSessionID)
binary.BigEndian.PutUint16(hdr[17:19], 0)
bodyBuf.Write(hdr[:])
if err := WriteAddressPort(bodyBuf, dest); err != nil {
return nil, err
}
bodyBuf.Write(payload)
bodyNonce := rawHeader[4:16]
sealedBody := s.serverBodyCipher.Seal(nil, bodyNonce, bodyBuf.Bytes(), nil)
res := make([]byte, 16+len(sealedBody))
copy(res[:16], encryptedHeader[:])
copy(res[16:], sealedBody)
return res, nil
}
func (c *UDPCodec) EncodeServerPacket(clientSessionID uint64, dest net.Destination, payload []byte) ([]byte, error) {
return c.sessions.EncodeServerPacket(c.method, c.psk, clientSessionID, dest, payload)
}
type serverSessionState struct {
sessionID uint64
window *SlidingWindow
cipher cipher.AEAD
lastSeen atomic.Int64
}
func (st *serverSessionState) check(packetID uint64) bool {
if st.window == nil {
st.window = new(SlidingWindow)
}
return st.window.Check(packetID)
}
func (st *serverSessionState) add(packetID uint64) {
if st.window == nil {
st.window = new(SlidingWindow)
}
st.window.Add(packetID)
}
type ClientUDPSession struct {
codec *UDPCodec
clientSessionID uint64
nextPacketID atomic.Uint64
clientBodyCipher cipher.AEAD
current atomic.Pointer[serverSessionState]
old atomic.Pointer[serverSessionState]
}
func (c *UDPCodec) NewClientSession() (*ClientUDPSession, error) {
var sessID [8]byte
if _, err := io.ReadFull(rand.Reader, sessID[:]); err != nil {
return nil, err
}
clientSessionID := binary.BigEndian.Uint64(sessID[:])
var clientBodyCipher cipher.AEAD
var err error
if !c.method.IsChaCha {
finalPSK := c.psk
clientBodyKey := DeriveSessionSubKey(finalPSK, sessID[:], c.method.KeySaltLength)
clientBodyCipher, err = c.method.NewAEAD(clientBodyKey)
if err != nil {
return nil, err
}
}
return &ClientUDPSession{
codec: c,
clientSessionID: clientSessionID,
clientBodyCipher: clientBodyCipher,
}, nil
}
func (s *ClientUDPSession) getServerSession(sessionID uint64, now int64) (*serverSessionState, error) {
cur := s.current.Load()
if cur != nil && cur.sessionID == sessionID {
return cur, nil
}
old := s.old.Load()
if old != nil && old.sessionID == sessionID {
if now-old.lastSeen.Load() > 60 {
s.old.CompareAndSwap(old, nil)
return nil, errors.New("old server session expired")
}
return old, nil
}
// New server session:
// Spec §3.2.4: reject newer server sessions when the last packet received from the old session is less than 1 minute old.
if old != nil && now-old.lastSeen.Load() < 60 {
return nil, errors.New("newer server session rejected: old session is less than 1 minute old")
}
var bodyAead cipher.AEAD
if !s.codec.method.IsChaCha {
var sessBytes [8]byte
binary.BigEndian.PutUint64(sessBytes[:], sessionID)
bodyKey := DeriveSessionSubKey(s.codec.psk, sessBytes[:], s.codec.method.KeySaltLength)
var err error
bodyAead, err = s.codec.method.NewAEAD(bodyKey)
if err != nil {
return nil, err
}
}
newState := &serverSessionState{
sessionID: sessionID,
cipher: bodyAead,
}
newState.lastSeen.Store(now)
if cur == nil {
s.current.CompareAndSwap(nil, newState)
return s.current.Load(), nil
}
s.old.Store(cur)
s.current.Store(newState)
return newState, nil
}
func (s *ClientUDPSession) ClientSessionID() uint64 {
return s.clientSessionID
}
func (s *ClientUDPSession) EncodePacket(dest net.Destination, payload []byte) (*buf.Buffer, error) {
packetID := s.nextPacketID.Add(1) - 1
sessID := s.clientSessionID
var paddingLen int
if dest.Port == 53 && len(payload) < MaxPaddingLength {
paddingLen = mrand.IntN(MaxPaddingLength) + 1
}
addrPortLen := AddrPortLength(dest)
if s.codec.method.IsChaCha {
totalLen := PacketNonceSize + 27 + paddingLen + addrPortLen + len(payload) + AEADTagSize
if totalLen > buf.Size {
return nil, ErrPacketTooLarge
}
outBuf := buf.New()
var nonce [PacketNonceSize]byte
if _, err := io.ReadFull(rand.Reader, nonce[:]); err != nil {
outBuf.Release()
return nil, err
}
outBuf.Write(nonce[:])
var hdr [16 + 1 + 8 + 2]byte
binary.BigEndian.PutUint64(hdr[0:8], sessID)
binary.BigEndian.PutUint64(hdr[8:16], packetID)
hdr[16] = HeaderTypeClient
binary.BigEndian.PutUint64(hdr[17:25], uint64(time.Now().Unix()))
binary.BigEndian.PutUint16(hdr[25:27], uint16(paddingLen))
outBuf.Write(hdr[:])
if paddingLen > 0 {
outBuf.Write(zeroPadding[:paddingLen])
}
if err := WriteAddressPort(outBuf, dest); err != nil {
outBuf.Release()
return nil, err
}
outBuf.Write(payload)
plainBytes := outBuf.Bytes()[PacketNonceSize:]
outBuf.Extend(int32(s.codec.chachaCipher.Overhead()))
s.codec.chachaCipher.Seal(plainBytes[:0], nonce[:], plainBytes, nil)
return outBuf, nil
}
// AES mode
var sessBytes [8]byte
binary.BigEndian.PutUint64(sessBytes[:], sessID)
var rawHeader [16]byte
copy(rawHeader[:8], sessBytes[:])
binary.BigEndian.PutUint64(rawHeader[8:16], packetID)
eihCount := 0
if len(s.codec.pskList) > 1 {
eihCount = len(s.codec.pskList) - 1
}
totalLen := 16 + eihCount*16 + 11 + paddingLen + addrPortLen + len(payload) + AEADTagSize
if totalLen > buf.Size {
return nil, ErrPacketTooLarge
}
outBuf := buf.New()
if len(s.codec.pskList) > 1 {
var encryptedHeader [16]byte
s.codec.blockCiphers[0].Encrypt(encryptedHeader[:], rawHeader[:])
outBuf.Write(encryptedHeader[:])
for i := 0; i < len(s.codec.pskList)-1; i++ {
nextPSK := s.codec.pskList[i+1]
pskHash := DeriveUserPSKHash(nextPSK)
var eihPlain [16]byte
for k := 0; k < 16; k++ {
eihPlain[k] = pskHash[k] ^ rawHeader[k]
}
var encryptedEIH [16]byte
s.codec.blockCiphers[i].Encrypt(encryptedEIH[:], eihPlain[:])
outBuf.Write(encryptedEIH[:])
}
} else {
var encryptedHeader [16]byte
s.codec.blockCipher.Encrypt(encryptedHeader[:], rawHeader[:])
outBuf.Write(encryptedHeader[:])
}
bodyAead := s.clientBodyCipher
var hdr [1 + 8 + 2]byte
hdr[0] = HeaderTypeClient
binary.BigEndian.PutUint64(hdr[1:9], uint64(time.Now().Unix()))
binary.BigEndian.PutUint16(hdr[9:11], uint16(paddingLen))
outBuf.Write(hdr[:])
if paddingLen > 0 {
outBuf.Write(zeroPadding[:paddingLen])
}
if err := WriteAddressPort(outBuf, dest); err != nil {
outBuf.Release()
return nil, err
}
outBuf.Write(payload)
headerOffset := 16 + eihCount*16
plainBytes := outBuf.Bytes()[headerOffset:]
bodyNonce := rawHeader[4:16]
outBuf.Extend(int32(bodyAead.Overhead()))
bodyAead.Seal(plainBytes[:0], bodyNonce, plainBytes, nil)
return outBuf, nil
}
func (s *ClientUDPSession) DecodePacket(data []byte) (DecodedUDPPacket, error) {
if len(data) < PacketMinimalHeaderSize {
return DecodedUDPPacket{}, ErrPacketTooShort
}
if s.codec.method.IsChaCha {
if len(data) < PacketNonceSize+AEADTagSize {
return DecodedUDPPacket{}, ErrPacketTooShort
}
nonce := data[:PacketNonceSize]
ciphertext := data[PacketNonceSize:]
plain, err := s.codec.chachaCipher.Open(nil, nonce, ciphertext, nil)
if err != nil {
return DecodedUDPPacket{}, errors.New("failed to decrypt chacha udp packet").Base(err)
}
if len(plain) < 16+1+8+2 {
return DecodedUDPPacket{}, ErrPacketTooShort
}
sessionID := binary.BigEndian.Uint64(plain[:8])
packetID := binary.BigEndian.Uint64(plain[8:16])
now := time.Now().Unix()
st, err := s.getServerSession(sessionID, now)
if err != nil {
return DecodedUDPPacket{}, err
}
if !st.check(packetID) {
return DecodedUDPPacket{}, ErrPacketIdNotUnique
}
decoded, err := parsePlainUDPPacket(sessionID, packetID, plain[16:])
if err != nil {
return DecodedUDPPacket{}, err
}
if decoded.HeaderType != HeaderTypeServer {
return DecodedUDPPacket{}, ErrBadHeaderType
}
if decoded.ClientSessionID != s.clientSessionID {
return DecodedUDPPacket{}, errors.New("client session ID mismatch")
}
st.add(packetID)
st.lastSeen.Store(now)
return decoded, nil
}
// AES mode
var rawHeader [16]byte
s.codec.blockCipher.Decrypt(rawHeader[:], data[:16])
sessionID := binary.BigEndian.Uint64(rawHeader[:8])
packetID := binary.BigEndian.Uint64(rawHeader[8:16])
now := time.Now().Unix()
st, err := s.getServerSession(sessionID, now)
if err != nil {
return DecodedUDPPacket{}, err
}
if !st.check(packetID) {
return DecodedUDPPacket{}, ErrPacketIdNotUnique
}
bodyAead := st.cipher
bodyNonce := rawHeader[4:16]
bodyCipher := data[16:]
bodyPlain, err := bodyAead.Open(nil, bodyNonce, bodyCipher, nil)
if err != nil {
return DecodedUDPPacket{}, errors.New("failed to decrypt aes udp body").Base(err)
}
decoded, err := parsePlainUDPPacket(sessionID, packetID, bodyPlain)
if err != nil {
return DecodedUDPPacket{}, err
}
if decoded.HeaderType != HeaderTypeServer {
return DecodedUDPPacket{}, ErrBadHeaderType
}
if decoded.ClientSessionID != s.clientSessionID {
return DecodedUDPPacket{}, errors.New("client session ID mismatch")
}
st.add(packetID)
st.lastSeen.Store(now)
return decoded, nil
}
type UDPWriter struct {
Writer io.Writer
Destination net.Destination
Session *ClientUDPSession
}
func (w *UDPWriter) WriteMultiBuffer(mb buf.MultiBuffer) error {
for {
mb2, b := buf.SplitFirst(mb)
mb = mb2
if b == nil {
break
}
dest := w.Destination
if b.UDP != nil {
dest = *b.UDP
}
pktBuf, err := w.Session.EncodePacket(dest, b.Bytes())
b.Release()
if err != nil {
buf.ReleaseMulti(mb)
return err
}
_, writeErr := w.Writer.Write(pktBuf.Bytes())
pktBuf.Release()
if writeErr != nil {
buf.ReleaseMulti(mb)
return writeErr
}
}
return nil
}
type UDPReader struct {
Reader io.Reader
Session *ClientUDPSession
}
func (r *UDPReader) ReadMultiBuffer() (buf.MultiBuffer, error) {
for {
buffer := buf.New()
_, err := buffer.ReadFrom(r.Reader)
if err != nil {
buffer.Release()
return nil, err
}
decoded, err := r.Session.DecodePacket(buffer.Bytes())
if err != nil {
buffer.Release()
continue
}
buffer.Clear()
buffer.Write(decoded.Payload)
dest := decoded.Destination
buffer.UDP = &dest
return buf.MultiBuffer{buffer}, nil
}
}
-376
View File
@@ -1,376 +0,0 @@
package shadowsocks_2022_test
import (
"context"
"crypto/rand"
"encoding/base64"
"encoding/binary"
"errors"
"io"
gonet "net"
"sync"
"sync/atomic"
"testing"
"time"
"github.com/xtls/xray-core/common"
"github.com/xtls/xray-core/common/buf"
"github.com/xtls/xray-core/common/net"
"github.com/xtls/xray-core/features/routing"
. "github.com/xtls/xray-core/proxy/shadowsocks_2022"
"github.com/xtls/xray-core/transport"
"lukechampine.com/blake3"
)
// encodeRelayClientUDPPacket encodes a Shadowsocks-2022 UDP packet with 1 layer of EIH (Relay)
func encodeRelayClientUDPPacket(relayKey, destKey []byte, sessionID, packetID uint64, dest net.Destination, payload []byte) ([]byte, error) {
method, err := GetCipherMethod(MethodAES128GCM)
if err != nil {
return nil, err
}
relayBlock, err := method.NewBlock(relayKey)
if err != nil {
return nil, err
}
// 1. Plain packet header: sessionID (8B) + packetID (8B)
var rawHeader [16]byte
binary.BigEndian.PutUint64(rawHeader[:8], sessionID)
binary.BigEndian.PutUint64(rawHeader[8:16], packetID)
// Encrypt packetHeader under relayKey
var encPacketHeader [16]byte
relayBlock.Encrypt(encPacketHeader[:], rawHeader[:])
// 2. EI Header: blake3(destKey)[:16] ^ rawHeader
var destHash [16]byte
hash512 := blake3.Sum512(destKey)
copy(destHash[:], hash512[:16])
var eiHeader [16]byte
for i := 0; i < 16; i++ {
eiHeader[i] = destHash[i] ^ rawHeader[i]
}
var encEIHeader [16]byte
relayBlock.Encrypt(encEIHeader[:], eiHeader[:])
// 3. Payload under destination server's AEAD
bodyKey := DeriveSessionSubKey(destKey, rawHeader[:8], 16)
bodyAead, err := method.NewAEAD(bodyKey)
if err != nil {
return nil, err
}
bodyNonce := rawHeader[4:16]
outBuf := buf.New()
defer outBuf.Release()
// VarHeader: client type (1) + timestamp (8) + paddingLen (2) + padding + dest + payload
var hdr [1 + 8 + 2]byte
hdr[0] = HeaderTypeClient
binary.BigEndian.PutUint64(hdr[1:9], uint64(time.Now().Unix()))
binary.BigEndian.PutUint16(hdr[9:11], 0)
outBuf.Write(hdr[:])
if err := WriteAddressPort(outBuf, dest); err != nil {
return nil, err
}
outBuf.Write(payload)
plainBytes := outBuf.Bytes()
outBuf.Extend(int32(bodyAead.Overhead()))
bodyAead.Seal(plainBytes[:0], bodyNonce, plainBytes, nil)
// Full packet: encPacketHeader (16B) + encEIHeader (16B) + sealedBody
packet := make([]byte, 0, 32+outBuf.Len())
packet = append(packet, encPacketHeader[:]...)
packet = append(packet, encEIHeader[:]...)
packet = append(packet, outBuf.Bytes()...)
return packet, nil
}
func TestRelayUDPSessionStabilityAndDispatch(t *testing.T) {
relayKey := []byte("0123456789abcdef")
destKey := []byte("fedcba9876543210")
relayKeyB64 := base64.StdEncoding.EncodeToString(relayKey)
destKeyB64 := base64.StdEncoding.EncodeToString(destKey)
config := &RelayServerConfig{
Method: MethodAES128GCM,
Key: relayKeyB64,
Destinations: []*RelayDestination{
{
Key: destKeyB64,
Address: &net.IPOrDomain{Address: &net.IPOrDomain_Ip{Ip: []byte{127, 0, 0, 1}}},
Port: 8388,
Email: "dest@example.com",
},
},
}
inbound, err := NewRelayServer(newTestContext(), config)
if err != nil {
t.Fatalf("failed to create RelayServer: %v", err)
}
sessionID := uint64(0x1122334455667788)
dest := net.UDPDestination(net.LocalHostIP, 8388)
pkt1, err := encodeRelayClientUDPPacket(relayKey, destKey, sessionID, 1, dest, []byte("xray packet 1"))
if err != nil {
t.Fatalf("failed to encode pkt1: %v", err)
}
pkt2, err := encodeRelayClientUDPPacket(relayKey, destKey, sessionID, 2, dest, []byte("xray packet 2"))
if err != nil {
t.Fatalf("failed to encode pkt2: %v", err)
}
var dispatchCount atomic.Int32
var receivedPackets [][]byte
var mu sync.Mutex
disp := &dummyDispatcher{
onDispatch: func(ctx context.Context, d net.Destination) (*transport.Link, error) {
dispatchCount.Add(1)
linkR, linkW := gonet.Pipe()
t.Cleanup(func() {
linkW.Close()
linkR.Close()
})
link := &transport.Link{
Reader: buf.NewReader(linkR),
Writer: &customWriter{
write: func(mb buf.MultiBuffer) error {
mu.Lock()
defer mu.Unlock()
for _, b := range mb {
cpy := make([]byte, b.Len())
copy(cpy, b.Bytes())
receivedPackets = append(receivedPackets, cpy)
b.Release()
}
return nil
},
},
}
return link, nil
},
}
clientConn, serverConn := gonet.Pipe()
defer clientConn.Close()
defer serverConn.Close()
inboundConn := &dummyStatConn{Conn: serverConn}
ctx, cancel := context.WithCancel(newTestContext())
defer cancel()
go func() {
_ = inbound.Process(ctx, net.Network_UDP, inboundConn, disp)
}()
// Send Packet 1
_, err = clientConn.Write(pkt1)
if err != nil {
t.Fatalf("write pkt1 failed: %v", err)
}
time.Sleep(50 * time.Millisecond)
// Send Packet 2 (same sessionID, packetID=2)
_, err = clientConn.Write(pkt2)
if err != nil {
t.Fatalf("write pkt2 failed: %v", err)
}
time.Sleep(50 * time.Millisecond)
// Check dispatch count: For the SAME UDP session, Dispatch MUST be called exactly ONCE!
if count := dispatchCount.Load(); count != 1 {
t.Fatalf("CRITICAL BUG CONFIRMED: expected dispatchCount = 1 for same session, got %d (sessionID was corrupted by Encrypt!)", count)
}
// Verify downstream destination can decode both packets
method, err := GetCipherMethod(MethodAES128GCM)
common.Must(err)
destCodec, err := NewUDPServerCodec(method, destKey, 300*time.Second)
common.Must(err)
mu.Lock()
pkts := receivedPackets
mu.Unlock()
if len(pkts) != 2 {
t.Fatalf("expected 2 received packets at destination, got %d", len(pkts))
}
dec1, err := destCodec.DecodePacket(pkts[0])
if err != nil {
t.Fatalf("dest failed to decode packet 1: %v", err)
}
if dec1.SessionID != sessionID || dec1.PacketID != 1 || string(dec1.Payload) != "xray packet 1" {
t.Fatalf("dec1 mismatch: sess=%x, pktID=%d, payload=%s", dec1.SessionID, dec1.PacketID, string(dec1.Payload))
}
dec2, err := destCodec.DecodePacket(pkts[1])
if err != nil {
t.Fatalf("dest failed to decode packet 2: %v", err)
}
if dec2.SessionID != sessionID || dec2.PacketID != 2 || string(dec2.Payload) != "xray packet 2" {
t.Fatalf("dec2 mismatch: sess=%x, pktID=%d, payload=%s", dec2.SessionID, dec2.PacketID, string(dec2.Payload))
}
}
type customWriter struct {
write func(mb buf.MultiBuffer) error
}
func (w *customWriter) WriteMultiBuffer(mb buf.MultiBuffer) error {
return w.write(mb)
}
func (w *customWriter) Close() error {
return nil
}
func (w *customWriter) Interrupt() {}
type dummyDispatcher struct {
onDispatch func(ctx context.Context, dest net.Destination) (*transport.Link, error)
}
func (d *dummyDispatcher) Dispatch(ctx context.Context, dest net.Destination) (*transport.Link, error) {
if d.onDispatch != nil {
return d.onDispatch(ctx, dest)
}
return nil, errors.New("not handled")
}
func (d *dummyDispatcher) DispatchLink(ctx context.Context, dest net.Destination, link *transport.Link) error {
return nil
}
func (d *dummyDispatcher) Start() error { return nil }
func (d *dummyDispatcher) Close() error { return nil }
func (d *dummyDispatcher) Type() interface{} { return routing.DispatcherType() }
type dummyStatConn struct {
gonet.Conn
}
func (c *dummyStatConn) ReadMultiBuffer() (buf.MultiBuffer, error) {
b := buf.New()
_, err := b.ReadFrom(c.Conn)
return buf.MultiBuffer{b}, err
}
func (c *dummyStatConn) WriteMultiBuffer(mb buf.MultiBuffer) error {
defer buf.ReleaseMulti(mb)
for _, b := range mb {
if _, err := c.Conn.Write(b.Bytes()); err != nil {
return err
}
}
return nil
}
func TestRelayTCPHandshakeForwarding(t *testing.T) {
methods := []string{MethodAES128GCM, MethodAES256GCM}
for _, methodName := range methods {
t.Run(methodName, func(t *testing.T) {
method, err := GetCipherMethod(methodName)
common.Must(err)
relayKey := make([]byte, method.KeySaltLength)
destKey := make([]byte, method.KeySaltLength)
_, _ = io.ReadFull(rand.Reader, relayKey)
_, _ = io.ReadFull(rand.Reader, destKey)
targetPort := uint32(54321)
relayConfig := &RelayServerConfig{
Method: methodName,
Key: base64.StdEncoding.EncodeToString(relayKey),
Destinations: []*RelayDestination{
{
Key: base64.StdEncoding.EncodeToString(destKey),
Address: net.NewIPOrDomain(net.LocalHostIP),
Port: targetPort,
Email: "test@xray.com",
},
},
}
testCtx := newTestContext()
inbound, err := NewRelayServer(testCtx, relayConfig)
common.Must(err)
targetDest := net.TCPDestination(net.LocalHostIP, net.Port(targetPort))
downstreamR, downstreamW := gonet.Pipe()
defer downstreamR.Close()
defer downstreamW.Close()
disp := &dummyDispatcher{
onDispatch: func(ctx context.Context, dest net.Destination) (*transport.Link, error) {
inLink := &transport.Link{
Reader: buf.NewReader(downstreamR),
Writer: &customWriter{
write: func(mb buf.MultiBuffer) error {
defer buf.ReleaseMulti(mb)
for _, b := range mb {
if _, err := downstreamW.Write(b.Bytes()); err != nil {
return err
}
}
return nil
},
},
}
return inLink, nil
},
}
clientConn, relayConn := gonet.Pipe()
defer clientConn.Close()
defer relayConn.Close()
go func() {
_ = inbound.Process(testCtx, net.Network_TCP, &dummyStatConn{Conn: relayConn}, disp)
}()
clientSalt := make([]byte, method.KeySaltLength)
_, _ = io.ReadFull(rand.Reader, clientSalt)
pskList := [][]byte{relayKey, destKey}
go func() {
_, err := WriteTCPRequest(clientConn, method, pskList, targetDest, clientSalt, []byte("relay payload"))
if err != nil {
t.Errorf("WriteTCPRequest failed: %v", err)
}
}()
// Downstream server must be able to read Salt + Fixed chunk in a single Read call!
headerLen := method.KeySaltLength + RequestHeaderFixedChunkLength + AEADTagSize
headerBuf := make([]byte, headerLen)
n, err := downstreamR.Read(headerBuf)
if err != nil {
t.Fatalf("downstream failed to read handshake: %v", err)
}
if n < headerLen {
t.Fatalf("downstream expected single read >= %d bytes, got %d", headerLen, n)
}
// Verify downstream can decode the fixed chunk and subsequent payload
sessionKey := DeriveSessionSubKey(destKey, headerBuf[:method.KeySaltLength], method.KeySaltLength)
aead, err := method.NewAEAD(sessionKey)
common.Must(err)
reader := NewStreamReader(downstreamR, aead)
reqHeader, err := ReadClientRequestHeaderWithFixed(reader, headerBuf[method.KeySaltLength:])
if err != nil {
t.Fatalf("downstream failed to parse client request header: %v", err)
}
if string(reqHeader.EarlyData) != "relay payload" {
t.Fatalf("payload mismatch: expected 'relay payload', got '%s'", string(reqHeader.EarlyData))
}
})
}
}
-183
View File
@@ -1,183 +0,0 @@
package shadowsocks_2022
import (
"crypto/cipher"
"sync"
"sync/atomic"
"time"
"github.com/xtls/xray-core/common/net"
"github.com/xtls/xray-core/common/protocol"
"github.com/xtls/xray-core/common/signal"
"github.com/xtls/xray-core/common/utils"
"github.com/xtls/xray-core/transport"
)
const (
swBlockBitLog = 6 // 1<<6 == 64 bits
swBlockBits = 1 << swBlockBitLog // 64
swRingBlocks = 1 << 7 // 128
swBlockMask = swRingBlocks - 1 // 127
swBitMask = swBlockBits - 1 // 63
swSize = (swRingBlocks - 1) * swBlockBits // 8128
)
type SlidingWindow struct {
last uint64
ring [swRingBlocks]uint64
}
func (f *SlidingWindow) Reset() {
*f = SlidingWindow{}
}
func (f *SlidingWindow) Check(counter uint64) bool {
switch {
case counter > f.last:
return true
case f.last-counter > swSize:
return false
}
blockIndex := (counter >> swBlockBitLog) & swBlockMask
bitIndex := counter & swBitMask
return (f.ring[blockIndex]>>bitIndex)&1 == 0
}
func (f *SlidingWindow) Add(counter uint64) {
blockIndex := counter >> swBlockBitLog
if counter > f.last {
lastBlockIndex := f.last >> swBlockBitLog
diff := int(blockIndex - lastBlockIndex)
if diff > swRingBlocks {
diff = swRingBlocks
}
for i := 0; i < diff; i++ {
lastBlockIndex = (lastBlockIndex + 1) & swBlockMask
f.ring[lastBlockIndex] = 0
}
f.last = counter
}
blockIndex &= swBlockMask
bitIndex := counter & swBitMask
f.ring[blockIndex] |= 1 << bitIndex
}
func (f *SlidingWindow) CheckAndAdd(counter uint64) bool {
if !f.Check(counter) {
return false
}
f.Add(counter)
return true
}
type ServerUDPSession struct {
sync.Mutex
SessionID uint64
Window *SlidingWindow
User *protocol.MemoryUser
UserPSK []byte
LastActive atomic.Int64 // Unix timestamp in seconds
clientBodyCipher cipher.AEAD
ServerSessionID uint64
ServerPacketID atomic.Uint64
serverBodyCipher cipher.AEAD
serverHeaderBlock cipher.Block
serverChaCha cipher.AEAD
manager *UDPSessionManager
link atomic.Pointer[transport.Link]
timer *signal.ActivityTimer
currentConn atomic.Value // stores stat.Connection
}
func (s *ServerUDPSession) CheckPacketID(packetID uint64) bool {
s.Lock()
defer s.Unlock()
if s.Window == nil {
s.Window = new(SlidingWindow)
}
return s.Window.Check(packetID)
}
func (s *ServerUDPSession) AddPacketID(packetID uint64) {
s.Lock()
defer s.Unlock()
if s.Window == nil {
s.Window = new(SlidingWindow)
}
s.Window.Add(packetID)
}
type UDPSessionManager struct {
sessions *utils.TypedSyncMap[uint64, *ServerUDPSession]
timeout time.Duration
lastClean atomic.Int64 // Unix timestamp in seconds
}
func NewUDPSessionManager(timeout time.Duration) *UDPSessionManager {
return &UDPSessionManager{
sessions: utils.NewTypedSyncMap[uint64, *ServerUDPSession](),
timeout: timeout,
}
}
func (m *UDPSessionManager) GetOrCreate(sessionID uint64) *ServerUDPSession {
now := time.Now().Unix()
if s, ok := m.sessions.Load(sessionID); ok {
s.LastActive.Store(now)
return s
}
s := &ServerUDPSession{
SessionID: sessionID,
manager: m,
}
s.LastActive.Store(now)
actual, loaded := m.sessions.LoadOrStore(sessionID, s)
if loaded {
actual.LastActive.Store(now)
return actual
}
// Trigger cleanup if at least 30 seconds have passed since last cleanup
last := m.lastClean.Load()
if now-last > 30 && m.lastClean.CompareAndSwap(last, now) {
go m.cleanup(now)
}
return s
}
func (m *UDPSessionManager) cleanup(now int64) {
timeoutSec := int64(m.timeout.Seconds())
if timeoutSec <= 0 {
timeoutSec = 60
}
m.sessions.Range(func(k uint64, v *ServerUDPSession) bool {
if now-v.LastActive.Load() > timeoutSec {
m.sessions.Delete(k)
v.Close()
}
return true
})
}
func (m *UDPSessionManager) Delete(sessionID uint64) {
m.sessions.Delete(sessionID)
}
func (m *UDPSessionManager) EncodeServerPacket(method *CipherMethod, psk []byte, clientSessionID uint64, dest net.Destination, payload []byte) ([]byte, error) {
sessionItem := m.GetOrCreate(clientSessionID)
if err := sessionItem.EnsureServerState(method, psk); err != nil {
return nil, err
}
return sessionItem.EncodeServerPacket(method, clientSessionID, dest, payload)
}
-192
View File
@@ -1,193 +1 @@
package shadowsocks_2022
import (
"context"
"github.com/xtls/xray-core/common"
"github.com/xtls/xray-core/common/buf"
"github.com/xtls/xray-core/common/errors"
"github.com/xtls/xray-core/common/log"
"github.com/xtls/xray-core/common/net"
"github.com/xtls/xray-core/common/session"
"github.com/xtls/xray-core/common/signal"
"github.com/xtls/xray-core/features/policy"
"github.com/xtls/xray-core/features/routing"
"github.com/xtls/xray-core/proxy"
"github.com/xtls/xray-core/transport"
"github.com/xtls/xray-core/transport/internet/stat"
)
func (s *ServerUDPSession) UpdateConn(conn stat.Connection) {
if s.currentConn.Load() == nil {
s.currentConn.Store(conn)
}
if s.timer != nil {
s.timer.Update()
}
}
func (s *ServerUDPSession) WriteToClient(b []byte) error {
connVal := s.currentConn.Load()
if connVal == nil {
return errors.New("client connection closed")
}
conn, ok := connVal.(stat.Connection)
if !ok || conn == nil {
return errors.New("client connection closed")
}
_, err := conn.Write(b)
return err
}
func (s *ServerUDPSession) Close() {
if s.timer != nil {
s.timer.SetTimeout(0)
}
if link := s.link.Load(); link != nil {
common.Interrupt(link.Reader)
common.Interrupt(link.Writer)
}
}
func (s *ServerUDPSession) EnsureLink(
ctx context.Context,
conn stat.Connection,
dest net.Destination,
dispatcher routing.Dispatcher,
policyManager policy.Manager,
responseEncoder func(dest net.Destination, payload []byte) ([]byte, error),
) (*transport.Link, error) {
s.UpdateConn(conn)
if link := s.link.Load(); link != nil {
return link, nil
}
s.Lock()
defer s.Unlock()
if link := s.link.Load(); link != nil {
return link, nil
}
sessCtx, cancel := context.WithCancel(ctx)
inbound := session.InboundFromContext(sessCtx)
if inbound != nil && s.User != nil {
inbound.User = s.User
}
var email string
var level uint32
if s.User != nil {
email = s.User.Email
level = s.User.Level
}
sessCtx = log.ContextWithAccessMessage(sessCtx, &log.AccessMessage{
From: conn.RemoteAddr(),
To: dest,
Status: log.AccessAccepted,
Email: email,
})
link, err := dispatcher.Dispatch(sessCtx, dest)
if err != nil {
cancel()
return nil, err
}
s.link.Store(link)
sessionPolicy := policyManager.ForLevel(level)
s.timer = signal.CancelAfterInactivity(sessCtx, func() {
if s.manager != nil {
s.manager.Delete(s.SessionID)
}
s.Close()
cancel()
}, sessionPolicy.Timeouts.ConnectionIdle)
go handleUDPResponse(s, link, dest, responseEncoder)
return link, nil
}
// ResetTCPConn sets SO_LINGER to 0 per SIP022 §3.1.4 to consistently send RST on close
// when handshake or header validation fails.
func ResetTCPConn(conn net.Conn) {
rawConn, _, _ := proxy.UnwrapRawConn(conn)
if tcpConn, ok := rawConn.(*net.TCPConn); ok {
_ = tcpConn.SetLinger(0)
}
}
func handleUDPResponse(s *ServerUDPSession, link *transport.Link, fallbackDest net.Destination, encode func(dest net.Destination, payload []byte) ([]byte, error)) {
defer func() {
if s.timer != nil {
s.timer.SetTimeout(0)
}
}()
for {
resMb, err := link.Reader.ReadMultiBuffer()
if err != nil {
return
}
if s.timer != nil {
s.timer.Update()
}
for i, rb := range resMb {
b := rb.Bytes()
if encode != nil {
replyDest := fallbackDest
if rb.UDP != nil {
replyDest = *rb.UDP
}
encPacket, err := encode(replyDest, b)
rb.Release()
if err != nil {
continue
}
if err := s.WriteToClient(encPacket); err != nil {
buf.ReleaseMulti(resMb[i+1:])
return
}
} else {
err := s.WriteToClient(b)
rb.Release()
if err != nil {
buf.ReleaseMulti(resMb[i+1:])
return
}
}
}
}
}
const (
HeaderTypeClient = 0
HeaderTypeServer = 1
MaxPaddingLength = 900
PacketNonceSize = 24
MaxPacketSize = 65535
RequestHeaderFixedChunkLength = 1 + 8 + 2 // Type (1B) + Timestamp (8B) + VarHeaderLen (2B)
PacketMinimalHeaderSize = 30
StreamNonceSize = 12
AESBlockSize = 16
AEADTagSize = 16
)
var zeroPadding [MaxPaddingLength]byte
const (
MethodAES128GCM = "2022-blake3-aes-128-gcm"
MethodAES256GCM = "2022-blake3-aes-256-gcm"
MethodChaCha20Poly1305 = "2022-blake3-chacha20-poly1305"
)
var (
ErrBadKey = errors.New("bad key")
ErrBadHeaderType = errors.New("bad header type")
ErrBadTimestamp = errors.New("bad timestamp")
ErrSaltNotUnique = errors.New("salt not unique")
ErrPacketIdNotUnique = errors.New("packet id not unique")
ErrPacketTooShort = errors.New("packet too short")
ErrPacketTooLarge = errors.New("packet too large")
ErrNoPadding = errors.New("bad request: missing payload or padding")
ErrInvalidRequest = errors.New("invalid request")
)
@@ -1,498 +0,0 @@
package shadowsocks_2022_test
import (
"bytes"
"context"
"crypto/rand"
"encoding/base64"
"encoding/binary"
"io"
gonet "net"
"sync"
"testing"
"time"
"github.com/google/go-cmp/cmp"
"github.com/xtls/xray-core/common"
"github.com/xtls/xray-core/common/buf"
"github.com/xtls/xray-core/common/net"
"github.com/xtls/xray-core/common/protocol"
"github.com/xtls/xray-core/common/serial"
"github.com/xtls/xray-core/common/session"
"github.com/xtls/xray-core/core"
. "github.com/xtls/xray-core/proxy/shadowsocks_2022"
)
func newTestContext() context.Context {
v, err := core.New(&core.Config{})
common.Must(err)
ctx := context.WithValue(context.Background(), core.XrayKey(1), v)
ctx = session.ContextWithInbound(ctx, &session.Inbound{})
return ctx
}
func generateRandomKey(size int) string {
b := make([]byte, size)
_, _ = rand.Read(b)
return base64.StdEncoding.EncodeToString(b)
}
func TestKDF(t *testing.T) {
// Test ParseKey
if _, err := ParseKey("", 16); err != ErrBadKey {
t.Fatalf("expected ErrBadKey for empty key, got %v", err)
}
shortKey := base64.StdEncoding.EncodeToString([]byte("short"))
if _, err := ParseKey(shortKey, 16); err != ErrBadKey {
t.Fatalf("expected ErrBadKey for short key, got %v", err)
}
exactKey := []byte("0123456789abcdef")
exactKeyB64 := base64.StdEncoding.EncodeToString(exactKey)
normExact, err := ParseKey(exactKeyB64, 16)
if err != nil || !bytes.Equal(normExact, exactKey) {
t.Fatalf("unexpected parsed exact key: %v, err: %v", normExact, err)
}
longKey := base64.StdEncoding.EncodeToString([]byte("0123456789abcdef_longer_key_for_testing"))
if _, err := ParseKey(longKey, 16); err != ErrBadKey {
t.Fatalf("expected ErrBadKey for long key, got %v", err)
}
// Test Session Subkey determinism
salt := []byte("random_salt_1234")
k1 := DeriveSessionSubKey(normExact, salt, 16)
k2 := DeriveSessionSubKey(normExact, salt, 16)
if !bytes.Equal(k1, k2) {
t.Fatal("DeriveSessionSubKey should be deterministic")
}
// Identity subkey must differ from session subkey with same inputs
idKey := DeriveIdentitySubKey(normExact, salt, 16)
if bytes.Equal(k1, idKey) {
t.Fatal("DeriveIdentitySubKey must differ from DeriveSessionSubKey")
}
// User PSK hash
h1 := DeriveUserPSKHash(normExact)
h2 := DeriveUserPSKHash(normExact)
if h1 != h2 {
t.Fatal("DeriveUserPSKHash should be deterministic")
}
}
func TestSlidingWindow(t *testing.T) {
var window SlidingWindow
if !window.Check(1) {
t.Fatal("packet 1 should be accepted")
}
window.Add(1)
if window.Check(1) {
t.Fatal("duplicate packet 1 should be rejected")
}
if !window.Check(100) {
t.Fatal("packet 100 should be accepted")
}
window.Add(100)
if window.Check(100) {
t.Fatal("duplicate packet 100 should be rejected")
}
if !window.Check(50) {
t.Fatal("out-of-order packet 50 within window should be accepted")
}
window.Add(50)
if window.Check(50) {
t.Fatal("duplicate packet 50 should be rejected")
}
// Check packet far behind window (> 8128)
window.Add(10000)
if window.Check(1) {
t.Fatal("packet 1 should be rejected as behind window")
}
}
func TestTCPStream(t *testing.T) {
methods := []struct {
name string
keySize int
}{
{MethodAES128GCM, 16},
{MethodAES256GCM, 32},
{MethodChaCha20Poly1305, 32},
}
dest := net.TCPDestination(net.LocalHostIP, net.Port(8080))
testPayload := []byte("Hello, Shadowsocks 2022 Native Implementation!")
for _, m := range methods {
t.Run(m.name, func(t *testing.T) {
rawKey := make([]byte, m.keySize)
_, _ = rand.Read(rawKey)
method, err := GetCipherMethod(m.name)
common.Must(err)
clientConn, serverConn := gonet.Pipe()
defer clientConn.Close()
defer serverConn.Close()
var wg sync.WaitGroup
wg.Add(2)
var receivedDest net.Destination
var receivedPayload []byte
// Server goroutine
go func() {
defer wg.Done()
salt := make([]byte, method.KeySaltLength)
_, err := io.ReadFull(serverConn, salt)
common.Must(err)
sessionKey := DeriveSessionSubKey(rawKey, salt, method.KeySaltLength)
aead, err := method.NewAEAD(sessionKey)
common.Must(err)
reader := NewStreamReader(serverConn, aead)
// Read fixed chunk (11 + 16 bytes)
var fixedBuf [RequestHeaderFixedChunkLength + AEADTagSize]byte
_, err = io.ReadFull(serverConn, fixedBuf[:])
common.Must(err)
plainFixed, err := aead.Open(fixedBuf[:0], reader.Nonce(), fixedBuf[:], nil)
common.Must(err)
IncreaseNonce(reader.Nonce())
if plainFixed[0] != HeaderTypeClient {
t.Errorf("expected client header type, got %d", plainFixed[0])
}
// Read variable chunk
varLen := int(plainFixed[9])<<8 | int(plainFixed[10])
varBuf := make([]byte, varLen+AEADTagSize)
_, err = io.ReadFull(serverConn, varBuf)
common.Must(err)
plainVar, err := aead.Open(varBuf[:0], reader.Nonce(), varBuf, nil)
common.Must(err)
IncreaseNonce(reader.Nonce())
dest, addrLen, err := ParseAddressPort(plainVar)
common.Must(err)
receivedDest = net.TCPDestination(dest.Address, dest.Port)
plainVar = plainVar[addrLen:]
padLen := int(binary.BigEndian.Uint16(plainVar[:2]))
receivedPayload = plainVar[2+padLen:]
// Server sends response stream with receivedPayload as first payload
writer := NewServerStreamWriter(serverConn, method, rawKey, salt)
pBuf := buf.New()
pBuf.Write(receivedPayload)
_ = writer.WriteMultiBuffer(buf.MultiBuffer{pBuf})
// Read and echo additional stream data
mb, err := reader.ReadMultiBuffer()
common.Must(err)
_ = writer.WriteMultiBuffer(mb)
_ = writer.Close()
}()
// Client goroutine
go func() {
defer wg.Done()
clientSalt := make([]byte, method.KeySaltLength)
common.Must2(io.ReadFull(rand.Reader, clientSalt))
writer, err := WriteTCPRequest(clientConn, method, [][]byte{rawKey}, dest, clientSalt, testPayload)
common.Must(err)
reader, err := ReadTCPResponse(clientConn, method, rawKey, clientSalt)
common.Must(err)
// The first ReadMultiBuffer drains initialPayload from reader cache
mbInit, err := reader.ReadMultiBuffer()
common.Must(err)
if !bytes.Equal(mbInit[0].Bytes(), testPayload) {
t.Errorf("drained initial payload mismatch: got %s, want %s", mbInit[0].Bytes(), testPayload)
}
buf.ReleaseMulti(mbInit)
// Send additional stream data
streamData := []byte("stream chunk test")
_ = writer.WriteMultiBuffer(buf.MultiBuffer{buf.FromBytes(streamData)})
mb, err := reader.ReadMultiBuffer()
common.Must(err)
if !bytes.Equal(mb[0].Bytes(), streamData) {
t.Errorf("echoed stream data mismatch: got %s, want %s", mb[0].Bytes(), streamData)
}
buf.ReleaseMulti(mb)
}()
wg.Wait()
if receivedDest.NetAddr() != dest.NetAddr() {
t.Errorf("destination mismatch: got %s, want %s", receivedDest.NetAddr(), dest.NetAddr())
}
if diff := cmp.Diff(receivedPayload, testPayload); diff != "" {
t.Errorf("payload mismatch: %s", diff)
}
})
}
}
func TestUDPCodec(t *testing.T) {
methods := []string{
MethodAES128GCM,
MethodAES256GCM,
MethodChaCha20Poly1305,
}
dest := net.UDPDestination(net.LocalHostIP, net.Port(53))
payload := []byte("DNS query payload")
for _, methodName := range methods {
t.Run(methodName, func(t *testing.T) {
method, err := GetCipherMethod(methodName)
common.Must(err)
psk := make([]byte, method.KeySaltLength)
_, _ = rand.Read(psk)
clientCodec, err := NewUDPPacketCodec(method, [][]byte{psk})
common.Must(err)
serverCodec, err := NewUDPServerCodec(method, psk, time.Minute)
common.Must(err)
session, err := clientCodec.NewClientSession()
common.Must(err)
pktBuf, err := session.EncodePacket(dest, payload)
common.Must(err)
defer pktBuf.Release()
rawCopy := make([]byte, pktBuf.Len())
copy(rawCopy, pktBuf.Bytes())
decoded, err := serverCodec.DecodePacket(pktBuf.Bytes())
common.Must(err)
if decoded.HeaderType != HeaderTypeClient {
t.Errorf("expected header type %d, got %d", HeaderTypeClient, decoded.HeaderType)
}
if decoded.Destination.Port != dest.Port {
t.Errorf("port mismatch: got %d, want %d", decoded.Destination.Port, dest.Port)
}
if !bytes.Equal(decoded.Payload, payload) {
t.Errorf("payload mismatch: got %s, want %s", decoded.Payload, payload)
}
// Replay same packet wire bytes should fail with ErrPacketIdNotUnique
_, err = serverCodec.DecodePacket(rawCopy)
if err != ErrPacketIdNotUnique {
t.Fatalf("expected ErrPacketIdNotUnique on replay, got: %v", err)
}
})
}
}
func TestMultiUserManager(t *testing.T) {
masterKey := generateRandomKey(16)
userKey1 := generateRandomKey(16)
userKey2 := generateRandomKey(16)
config := &MultiUserServerConfig{
Method: MethodAES128GCM,
Key: masterKey,
Users: []*protocol.User{
{
Email: "user1@example.com",
Account: serial.ToTypedMessage(&Account{Key: userKey1}),
},
},
}
inbound, err := NewMultiServer(newTestContext(), config)
common.Must(err)
if inbound.GetUsersCount(context.Background()) != 1 {
t.Fatalf("expected 1 user, got %d", inbound.GetUsersCount(context.Background()))
}
u1 := inbound.GetUser(context.Background(), "user1@example.com")
if u1 == nil || u1.Email != "user1@example.com" {
t.Fatal("user1 not found")
}
// Add User 2
rawKey2, _ := base64.StdEncoding.DecodeString(userKey2)
u2 := &protocol.MemoryUser{
Email: "user2@example.com",
Account: &MemoryAccount{
Key: rawKey2,
},
}
err = inbound.AddUser(context.Background(), u2)
common.Must(err)
if inbound.GetUsersCount(context.Background()) != 2 {
t.Fatalf("expected 2 users, got %d", inbound.GetUsersCount(context.Background()))
}
// Remove User 1
err = inbound.RemoveUser(context.Background(), "user1@example.com")
common.Must(err)
if inbound.GetUsersCount(context.Background()) != 1 {
t.Fatalf("expected 1 user, got %d", inbound.GetUsersCount(context.Background()))
}
if inbound.GetUser(context.Background(), "user1@example.com") != nil {
t.Fatal("user1 should have been removed")
}
}
func TestLargeStreamTransfer(t *testing.T) {
method, err := GetCipherMethod(MethodAES128GCM)
common.Must(err)
sessionKey := make([]byte, 16)
_, _ = rand.Read(sessionKey)
clientAead, err := method.NewAEAD(sessionKey)
common.Must(err)
serverAead, err := method.NewAEAD(sessionKey)
common.Must(err)
r, w := io.Pipe()
defer r.Close()
defer w.Close()
writer := NewStreamWriter(w, clientAead)
reader := NewStreamReader(r, serverAead)
const totalSize = 100 * 1024 // 100 KB
data := make([]byte, totalSize)
_, _ = rand.Read(data)
errCh := make(chan error, 1)
go func() {
// Write using Write (which splits by MaxPacketSize = 65535)
_, werr := writer.Write(data)
if werr != nil {
errCh <- werr
return
}
_ = w.Close()
errCh <- nil
}()
var received []byte
for {
mb, rerr := reader.ReadMultiBuffer()
if !mb.IsEmpty() {
for _, b := range mb {
received = append(received, b.Bytes()...)
}
buf.ReleaseMulti(mb)
}
if rerr != nil {
if rerr == io.EOF {
break
}
t.Fatalf("ReadMultiBuffer error: %v", rerr)
}
}
if werr := <-errCh; werr != nil {
t.Fatalf("writer error: %v", werr)
}
if len(received) != totalSize {
t.Fatalf("received size mismatch: got %d, want %d", len(received), totalSize)
}
if !bytes.Equal(received, data) {
t.Fatal("received data does not match sent data")
}
}
func TestClientUDPSessionMultiDestination(t *testing.T) {
for _, methodName := range []string{MethodAES128GCM, MethodAES256GCM, MethodChaCha20Poly1305} {
t.Run(methodName, func(t *testing.T) {
method, err := GetCipherMethod(methodName)
common.Must(err)
rawKey := make([]byte, method.KeySaltLength)
_, _ = rand.Read(rawKey)
clientCodec, err := NewUDPPacketCodec(method, [][]byte{rawKey})
common.Must(err)
serverCodec, err := NewUDPServerCodec(method, rawKey, time.Minute)
common.Must(err)
session, err := clientCodec.NewClientSession()
common.Must(err)
dest1 := net.UDPDestination(net.LocalHostIP, net.Port(53))
dest2 := net.UDPDestination(net.IPAddress([]byte{127, 0, 0, 2}), net.Port(53))
payload1 := []byte("query-google-dns")
payload2 := []byte("query-cloudflare-dns")
// Client sends to dest1 and dest2 using SAME session
pkt1, err := session.EncodePacket(dest1, payload1)
common.Must(err)
defer pkt1.Release()
pkt2, err := session.EncodePacket(dest2, payload2)
common.Must(err)
defer pkt2.Release()
// Server decodes both
dec1, err := serverCodec.DecodePacket(pkt1.Bytes())
common.Must(err)
dec2, err := serverCodec.DecodePacket(pkt2.Bytes())
common.Must(err)
if dec1.SessionID != session.ClientSessionID() || dec2.SessionID != session.ClientSessionID() {
t.Fatalf("both packets must share client session ID %d, got %d and %d", session.ClientSessionID(), dec1.SessionID, dec2.SessionID)
}
if dec1.Destination.String() != dest1.String() {
t.Fatalf("expected dest1 %s, got %s", dest1, dec1.Destination)
}
if dec2.Destination.String() != dest2.String() {
t.Fatalf("expected dest2 %s, got %s", dest2, dec2.Destination)
}
if !bytes.Equal(dec1.Payload, payload1) || !bytes.Equal(dec2.Payload, payload2) {
t.Fatal("payload mismatch")
}
// Server replies to dest1 and dest2
respPayload1 := []byte("reply-google-dns")
respPayload2 := []byte("reply-cloudflare-dns")
respPkt1, err := serverCodec.EncodeServerPacket(dec1.SessionID, dest1, respPayload1)
common.Must(err)
respPkt2, err := serverCodec.EncodeServerPacket(dec2.SessionID, dest2, respPayload2)
common.Must(err)
// Client decodes replies
clientDec1, err := session.DecodePacket(respPkt1)
common.Must(err)
if clientDec1.Destination.String() != dest1.String() {
t.Fatalf("expected client dec1 dest %s, got %s", dest1, clientDec1.Destination)
}
if !bytes.Equal(clientDec1.Payload, respPayload1) {
t.Fatal("reply payload 1 mismatch")
}
clientDec2, err := session.DecodePacket(respPkt2)
common.Must(err)
if clientDec2.Destination.String() != dest2.String() {
t.Fatalf("expected client dec2 dest %s, got %s", dest2, clientDec2.Destination)
}
if !bytes.Equal(clientDec2.Payload, respPayload2) {
t.Fatal("reply payload 2 mismatch")
}
})
}
}
-649
View File
@@ -1,649 +0,0 @@
package shadowsocks_2022
import (
"context"
"crypto/cipher"
"crypto/rand"
"encoding/binary"
"io"
"math"
mrand "math/rand/v2"
"sync"
"time"
"github.com/xtls/xray-core/common/antireplay"
"github.com/xtls/xray-core/common/buf"
"github.com/xtls/xray-core/common/errors"
"github.com/xtls/xray-core/common/net"
"github.com/xtls/xray-core/common/protocol"
"github.com/xtls/xray-core/common/signal"
"github.com/xtls/xray-core/common/task"
"github.com/xtls/xray-core/features/policy"
"github.com/xtls/xray-core/transport"
)
var addrParser = protocol.NewAddressParser(
protocol.AddressFamilyByte(0x01, net.AddressFamilyIPv4),
protocol.AddressFamilyByte(0x04, net.AddressFamilyIPv6),
protocol.AddressFamilyByte(0x03, net.AddressFamilyDomain),
protocol.WithAddressTypeParser(func(b byte) byte {
return b & 0x0F
}),
)
func IncreaseNonce(nonce []byte) {
for i := range nonce {
nonce[i]++
if nonce[i] != 0 {
return
}
}
}
// WriteAddressPort writes a destination address and port in SOCKS5 format
func WriteAddressPort(w io.Writer, dest net.Destination) error {
return addrParser.WriteAddressPort(w, dest.Address, dest.Port)
}
// AddrPortLength returns the serialized length of a destination in SOCKS5 format
func AddrPortLength(dest net.Destination) int {
switch dest.Address.Family() {
case net.AddressFamilyIPv4:
return 1 + 4 + 2
case net.AddressFamilyDomain:
return 1 + 1 + len(dest.Address.Domain()) + 2
case net.AddressFamilyIPv6:
return 1 + 16 + 2
default:
return 0
}
}
type StreamWriter struct {
writer io.Writer
cipher cipher.AEAD
nonce [StreamNonceSize]byte
lenBuf [2]byte
buf []byte
}
func NewStreamWriter(w io.Writer, c cipher.AEAD) *StreamWriter {
return &StreamWriter{
writer: w,
cipher: c,
buf: make([]byte, 0, MaxPacketSize+2+2*AEADTagSize),
}
}
func (w *StreamWriter) Nonce() []byte {
return w.nonce[:]
}
func (w *StreamWriter) WriteChunk(payload []byte) error {
payloadLen := len(payload)
if payloadLen == 0 {
return nil
}
if payloadLen > MaxPacketSize {
return errors.New("payload exceeds MaxPacketSize")
}
binary.BigEndian.PutUint16(w.lenBuf[:], uint16(payloadLen))
w.buf = w.cipher.Seal(w.buf[:0], w.nonce[:], w.lenBuf[:], nil)
IncreaseNonce(w.nonce[:])
w.buf = w.cipher.Seal(w.buf, w.nonce[:], payload, nil)
IncreaseNonce(w.nonce[:])
_, err := w.writer.Write(w.buf)
return err
}
func (w *StreamWriter) Write(p []byte) (int, error) {
n := len(p)
for len(p) > 0 {
chunkSize := len(p)
if chunkSize > MaxPacketSize {
chunkSize = MaxPacketSize
}
if err := w.WriteChunk(p[:chunkSize]); err != nil {
return 0, err
}
p = p[chunkSize:]
}
return n, nil
}
func (w *StreamWriter) WriteMultiBuffer(mb buf.MultiBuffer) error {
defer buf.ReleaseMulti(mb)
for _, b := range mb {
p := b.Bytes()
for len(p) > 0 {
chunkSize := len(p)
if chunkSize > MaxPacketSize {
chunkSize = MaxPacketSize
}
if err := w.WriteChunk(p[:chunkSize]); err != nil {
return err
}
p = p[chunkSize:]
}
}
return nil
}
type StreamReader struct {
reader io.Reader
cipher cipher.AEAD
nonce [StreamNonceSize]byte
lenBuf [2 + AEADTagSize]byte
buffer []byte
cached int
offset int
}
func NewStreamReader(r io.Reader, c cipher.AEAD) *StreamReader {
return &StreamReader{
reader: r,
cipher: c,
buffer: make([]byte, MaxPacketSize+AEADTagSize),
}
}
func (r *StreamReader) Nonce() []byte {
return r.nonce[:]
}
func (r *StreamReader) Read(p []byte) (int, error) {
if r.cached > 0 {
n := copy(p, r.buffer[r.offset:r.offset+r.cached])
r.cached -= n
r.offset += n
return n, nil
}
// Read 2-byte length + AEAD tag (18 bytes)
if _, err := io.ReadFull(r.reader, r.lenBuf[:]); err != nil {
return 0, err
}
decryptedLen, err := r.cipher.Open(r.lenBuf[:0], r.nonce[:], r.lenBuf[:], nil)
if err != nil {
return 0, errors.New("failed to decrypt chunk length").Base(err)
}
IncreaseNonce(r.nonce[:])
payloadLen := int(binary.BigEndian.Uint16(decryptedLen))
if payloadLen == 0 || payloadLen > MaxPacketSize {
return 0, ErrInvalidRequest
}
chunkEnd := payloadLen + AEADTagSize
if _, err := io.ReadFull(r.reader, r.buffer[:chunkEnd]); err != nil {
return 0, err
}
decryptedPayload, err := r.cipher.Open(r.buffer[:0], r.nonce[:], r.buffer[:chunkEnd], nil)
if err != nil {
return 0, errors.New("failed to decrypt chunk payload").Base(err)
}
IncreaseNonce(r.nonce[:])
r.cached = len(decryptedPayload)
r.offset = 0
n := copy(p, r.buffer[r.offset:r.offset+r.cached])
r.cached -= n
r.offset += n
return n, nil
}
func (r *StreamReader) ReadMultiBuffer() (buf.MultiBuffer, error) {
if r.cached > 0 {
mb := buf.MergeBytes(nil, r.buffer[r.offset:r.offset+r.cached])
r.cached = 0
r.offset = 0
return mb, nil
}
if _, err := io.ReadFull(r.reader, r.lenBuf[:]); err != nil {
return nil, err
}
decryptedLen, err := r.cipher.Open(r.lenBuf[:0], r.nonce[:], r.lenBuf[:], nil)
if err != nil {
return nil, errors.New("failed to decrypt chunk length").Base(err)
}
IncreaseNonce(r.nonce[:])
payloadLen := int(binary.BigEndian.Uint16(decryptedLen))
if payloadLen == 0 || payloadLen > MaxPacketSize {
return nil, ErrInvalidRequest
}
chunkEnd := payloadLen + AEADTagSize
if _, err := io.ReadFull(r.reader, r.buffer[:chunkEnd]); err != nil {
return nil, err
}
decryptedPayload, err := r.cipher.Open(r.buffer[:0], r.nonce[:], r.buffer[:chunkEnd], nil)
if err != nil {
return nil, errors.New("failed to decrypt chunk payload").Base(err)
}
IncreaseNonce(r.nonce[:])
mb := buf.MergeBytes(nil, decryptedPayload)
return mb, nil
}
type ClientRequestHeader struct {
Destination net.Destination
EarlyData []byte
}
func ReadClientRequestHeaderWithFixed(reader *StreamReader, fixedChunk []byte) (*ClientRequestHeader, error) {
plainFixed, err := reader.cipher.Open(fixedChunk[:0], reader.Nonce(), fixedChunk, nil)
if err != nil {
return nil, errors.New("failed to decrypt client request header").Base(err)
}
IncreaseNonce(reader.Nonce())
if plainFixed[0] != HeaderTypeClient {
return nil, ErrBadHeaderType
}
epoch := binary.BigEndian.Uint64(plainFixed[1:9])
diff := int(math.Abs(float64(time.Now().Unix() - int64(epoch))))
if diff > 30 {
return nil, ErrBadTimestamp
}
varHeaderLen := int(binary.BigEndian.Uint16(plainFixed[9:11]))
if varHeaderLen == 0 {
return nil, ErrInvalidRequest
}
var stackVarChunk [512]byte
var varChunkCipher []byte
needed := varHeaderLen + AEADTagSize
if needed <= len(stackVarChunk) {
varChunkCipher = stackVarChunk[:needed]
} else {
varChunkCipher = make([]byte, needed)
}
if _, err := io.ReadFull(reader.reader, varChunkCipher); err != nil {
return nil, err
}
plainVar, err := reader.cipher.Open(varChunkCipher[:0], reader.Nonce(), varChunkCipher, nil)
if err != nil {
return nil, errors.New("failed to decrypt variable request header").Base(err)
}
IncreaseNonce(reader.Nonce())
dest, addrLen, err := ParseAddressPort(plainVar)
if err != nil {
return nil, err
}
dest.Network = net.Network_TCP
offset := addrLen
if len(plainVar) < offset+2 {
return nil, ErrPacketTooShort
}
paddingLen := int(binary.BigEndian.Uint16(plainVar[offset : offset+2]))
offset += 2
if len(plainVar) < offset+paddingLen {
return nil, ErrNoPadding
}
offset += paddingLen
var earlyData []byte
var payloadLen int
if len(plainVar) > offset {
earlyData = plainVar[offset:]
payloadLen = len(earlyData)
}
// SIP022 §3.1.4: Servers MUST reject the request if the variable-length header chunk does not contain payload and the padding length is 0.
if paddingLen == 0 && payloadLen == 0 {
return nil, errors.New("request without payload and padding is not allowed")
}
return &ClientRequestHeader{
Destination: dest,
EarlyData: earlyData,
}, nil
}
// WriteTCPRequest writes the Shadowsocks 2022 request header into w and returns a body writer.
func WriteTCPRequest(w io.Writer, method *CipherMethod, pskList [][]byte, dest net.Destination, clientSalt []byte, payload []byte) (buf.Writer, error) {
finalPSK := pskList[len(pskList)-1]
sessionKey := DeriveSessionSubKey(finalPSK, clientSalt, method.KeySaltLength)
aead, err := method.NewAEAD(sessionKey)
if err != nil {
return nil, err
}
writer := NewStreamWriter(w, aead)
payloadLen := len(payload)
var paddingLen int
if payloadLen < MaxPaddingLength {
paddingLen = mrand.IntN(MaxPaddingLength) + 1
}
addrPortLen := AddrPortLength(dest)
varHeaderLen := addrPortLen + 2 + paddingLen + payloadLen
totalHandshakeLen := int32(method.KeySaltLength + len(pskList)*AESBlockSize + RequestHeaderFixedChunkLength + AEADTagSize + varHeaderLen + AEADTagSize)
handshakeBuf := buf.NewWithSize(totalHandshakeLen)
defer handshakeBuf.Release()
handshakeBuf.Write(clientSalt)
for i, currPSK := range pskList[:len(pskList)-1] {
identitySubkey := DeriveIdentitySubKey(currPSK, clientSalt, method.KeySaltLength)
block, err := method.NewBlock(identitySubkey)
if err != nil {
return nil, err
}
nextPSK := pskList[i+1]
pskHash := DeriveUserPSKHash(nextPSK)
var encryptedEIH [AESBlockSize]byte
block.Encrypt(encryptedEIH[:], pskHash[:])
handshakeBuf.Write(encryptedEIH[:])
}
var fixedHeaderPlaintext [RequestHeaderFixedChunkLength]byte
fixedHeaderPlaintext[0] = HeaderTypeClient
binary.BigEndian.PutUint64(fixedHeaderPlaintext[1:9], uint64(time.Now().Unix()))
binary.BigEndian.PutUint16(fixedHeaderPlaintext[9:11], uint16(varHeaderLen))
fixedChunk := writer.cipher.Seal(nil, writer.nonce[:], fixedHeaderPlaintext[:], nil)
IncreaseNonce(writer.nonce[:])
handshakeBuf.Write(fixedChunk)
varHeaderBuf := buf.NewWithSize(int32(varHeaderLen))
defer varHeaderBuf.Release()
if err := WriteAddressPort(varHeaderBuf, dest); err != nil {
return nil, err
}
var padLenBytes [2]byte
binary.BigEndian.PutUint16(padLenBytes[:], uint16(paddingLen))
varHeaderBuf.Write(padLenBytes[:])
if paddingLen > 0 {
varHeaderBuf.Write(zeroPadding[:paddingLen])
}
if payloadLen > 0 {
varHeaderBuf.Write(payload)
}
varChunk := writer.cipher.Seal(nil, writer.nonce[:], varHeaderBuf.Bytes(), nil)
IncreaseNonce(writer.nonce[:])
handshakeBuf.Write(varChunk)
if _, err := w.Write(handshakeBuf.Bytes()); err != nil {
return nil, err
}
return writer, nil
}
// ReadTCPResponse reads and verifies the server's handshake response and returns a reader for the stream.
func ReadTCPResponse(r io.Reader, method *CipherMethod, psk []byte, clientSalt []byte) (buf.Reader, error) {
fixedPlainLen := 1 + 8 + method.KeySaltLength + 2
chunkCipherLen := fixedPlainLen + AEADTagSize
headerLen := method.KeySaltLength + chunkCipherLen
// Single read call for Salt + Fixed-length response header chunk per SIP022 §3.1.4
var headerBuf [128]byte
headerSlice := headerBuf[:headerLen]
n, err := r.Read(headerSlice)
if err != nil || n < headerLen {
return nil, errors.New("failed to read complete server response header")
}
serverSaltSlice := headerSlice[:method.KeySaltLength]
chunkSlice := headerSlice[method.KeySaltLength:headerLen]
sessionKey := DeriveSessionSubKey(psk, serverSaltSlice, method.KeySaltLength)
aead, err := method.NewAEAD(sessionKey)
if err != nil {
return nil, err
}
reader := NewStreamReader(r, aead)
decryptedFixed, err := reader.cipher.Open(chunkSlice[:0], reader.nonce[:], chunkSlice, nil)
if err != nil {
return nil, errors.New("failed to decrypt server response header").Base(err)
}
IncreaseNonce(reader.nonce[:])
if decryptedFixed[0] != HeaderTypeServer {
return nil, ErrBadHeaderType
}
serverEpoch := binary.BigEndian.Uint64(decryptedFixed[1:9])
diff := int(math.Abs(float64(time.Now().Unix() - int64(serverEpoch))))
if diff > 30 {
return nil, ErrBadTimestamp
}
echoedSalt := decryptedFixed[9 : 9+method.KeySaltLength]
for i := 0; i < method.KeySaltLength; i++ {
if echoedSalt[i] != clientSalt[i] {
return nil, errors.New("bad request salt")
}
}
initialPayloadLen := int(binary.BigEndian.Uint16(decryptedFixed[9+method.KeySaltLength : 11+method.KeySaltLength]))
if initialPayloadLen > 0 {
initialCipherLen := initialPayloadLen + AEADTagSize
if _, err := io.ReadFull(r, reader.buffer[:initialCipherLen]); err != nil {
return nil, err
}
decryptedInitial, err := reader.cipher.Open(reader.buffer[:0], reader.nonce[:], reader.buffer[:initialCipherLen], nil)
if err != nil {
return nil, errors.New("failed to decrypt initial response payload").Base(err)
}
IncreaseNonce(reader.nonce[:])
reader.cached = len(decryptedInitial)
reader.offset = 0
}
return reader, nil
}
// ServerStreamWriter lazily sends the response header along with the first payload chunk per SIP022 §3.1.2 & §3.1.4.
type ServerStreamWriter struct {
mu sync.Mutex
w io.Writer
method *CipherMethod
psk []byte
clientSalt []byte
streamWriter *StreamWriter
}
func NewServerStreamWriter(w io.Writer, method *CipherMethod, psk []byte, clientSalt []byte) *ServerStreamWriter {
return &ServerStreamWriter{
w: w,
method: method,
psk: psk,
clientSalt: clientSalt,
}
}
func (s *ServerStreamWriter) sendHeaderWithFirstPayload(payload []byte) (*StreamWriter, error) {
var serverSalt [32]byte
serverSaltSlice := serverSalt[:s.method.KeySaltLength]
if _, err := io.ReadFull(rand.Reader, serverSaltSlice); err != nil {
return nil, err
}
respKey := DeriveSessionSubKey(s.psk, serverSaltSlice, s.method.KeySaltLength)
respAead, err := s.method.NewAEAD(respKey)
if err != nil {
return nil, err
}
sw := NewStreamWriter(s.w, respAead)
totalHeaderLen := int32(s.method.KeySaltLength + 1 + 8 + s.method.KeySaltLength + 2 + AEADTagSize + len(payload) + AEADTagSize)
outBuf := buf.NewWithSize(totalHeaderLen)
defer outBuf.Release()
outBuf.Write(serverSaltSlice)
var fixedRespPlain [1 + 8 + 32 + 2]byte
fixedRespSlice := fixedRespPlain[:1+8+s.method.KeySaltLength+2]
fixedRespSlice[0] = HeaderTypeServer
binary.BigEndian.PutUint64(fixedRespSlice[1:9], uint64(time.Now().Unix()))
copy(fixedRespSlice[9:9+s.method.KeySaltLength], s.clientSalt)
binary.BigEndian.PutUint16(fixedRespSlice[9+s.method.KeySaltLength:11+s.method.KeySaltLength], uint16(len(payload)))
fixedRespChunk := sw.cipher.Seal(nil, sw.nonce[:], fixedRespSlice, nil)
IncreaseNonce(sw.nonce[:])
outBuf.Write(fixedRespChunk)
if len(payload) > 0 {
payloadChunk := sw.cipher.Seal(nil, sw.nonce[:], payload, nil)
IncreaseNonce(sw.nonce[:])
outBuf.Write(payloadChunk)
}
if _, err := s.w.Write(outBuf.Bytes()); err != nil {
return nil, err
}
return sw, nil
}
func (s *ServerStreamWriter) WriteMultiBuffer(mb buf.MultiBuffer) error {
if mb.IsEmpty() {
return nil
}
if s.streamWriter == nil {
s.mu.Lock()
if s.streamWriter == nil {
firstBuf := mb[0]
firstBytes := firstBuf.Bytes()
chunkSize := len(firstBytes)
if chunkSize > MaxPacketSize {
chunkSize = MaxPacketSize
}
firstPayload := firstBytes[:chunkSize]
sw, err := s.sendHeaderWithFirstPayload(firstPayload)
if err != nil {
s.mu.Unlock()
buf.ReleaseMulti(mb)
return err
}
s.streamWriter = sw
firstBuf.Advance(int32(chunkSize))
if firstBuf.IsEmpty() {
firstBuf.Release()
mb = mb[1:]
}
}
s.mu.Unlock()
if len(mb) == 0 {
return nil
}
}
return s.streamWriter.WriteMultiBuffer(mb)
}
func (s *ServerStreamWriter) Write(p []byte) (int, error) {
n := len(p)
if s.streamWriter == nil {
s.mu.Lock()
if s.streamWriter == nil {
chunkSize := len(p)
if chunkSize > MaxPacketSize {
chunkSize = MaxPacketSize
}
firstPayload := p[:chunkSize]
sw, err := s.sendHeaderWithFirstPayload(firstPayload)
if err != nil {
s.mu.Unlock()
return 0, err
}
s.streamWriter = sw
p = p[chunkSize:]
}
s.mu.Unlock()
if len(p) == 0 {
return n, nil
}
}
_, err := s.streamWriter.Write(p)
return n, err
}
func (s *ServerStreamWriter) Close() error {
if s.streamWriter == nil {
s.mu.Lock()
defer s.mu.Unlock()
if s.streamWriter == nil {
sw, err := s.sendHeaderWithFirstPayload(nil)
if err != nil {
return err
}
s.streamWriter = sw
}
}
return nil
}
// InitServerStream decrypts the client request header, verifies the timestamp and replay filter,
// and returns a StreamReader for subsequent stream chunks.
func InitServerStream(conn net.Conn, method *CipherMethod, psk, saltSlice []byte, salt [32]byte, fixedChunk []byte, saltFilter *antireplay.ReplayFilter[[32]byte]) (*StreamReader, *ClientRequestHeader, error) {
sessionKey := DeriveSessionSubKey(psk, saltSlice, method.KeySaltLength)
aead, err := method.NewAEAD(sessionKey)
if err != nil {
return nil, nil, err
}
reader := NewStreamReader(conn, aead)
reqHeader, err := ReadClientRequestHeaderWithFixed(reader, fixedChunk)
if err != nil {
return nil, nil, err
}
_ = conn.SetReadDeadline(time.Time{})
if !saltFilter.Check(salt) {
return nil, nil, ErrSaltNotUnique
}
return reader, reqHeader, nil
}
func TransportTCP(ctx context.Context, sessionPolicy policy.Session, reader buf.Reader, writer buf.Writer, link *transport.Link) error {
ctx, cancel := context.WithCancel(ctx)
timer := signal.CancelAfterInactivity(ctx, cancel, sessionPolicy.Timeouts.ConnectionIdle)
ctx = policy.ContextWithBufferPolicy(ctx, sessionPolicy.Buffer)
requestDone := func() error {
defer timer.SetTimeout(sessionPolicy.Timeouts.DownlinkOnly)
return buf.Copy(reader, link.Writer, buf.UpdateActivity(timer))
}
responseDone := func() error {
defer timer.SetTimeout(sessionPolicy.Timeouts.UplinkOnly)
if c, ok := writer.(io.Closer); ok {
defer c.Close()
}
return buf.Copy(link.Reader, writer, buf.UpdateActivity(timer))
}
responseDoneAndCloseWriter := task.OnSuccess(responseDone, task.Close(link.Writer))
return task.Run(ctx, requestDone, responseDoneAndCloseWriter)
}
+2 -60
View File
@@ -15,57 +15,13 @@ Plainly enabling it in the config probably will result nothing, or lock your rou
## DETAILS
By default, enabling the feature will only bring the tun interface up. \
When configured explicitly, Windows and Linux can apply interface addresses from `gateway`, while macOS and FreeBSD use the first IPv4 prefix from `gateway` for the point-to-point address. \
Without `gateway`, the systems differ: Xray assigns no address on Linux, Windows gives the interface link-local addresses itself (an IPv6 one at once, an IPv4 one from `169.254.0.0/16` after a few seconds), and macOS and FreeBSD use `169.254.10.1/30`. \
When configured explicitly, Windows and Linux can apply interface addresses from `gateway`, while macOS uses the first IPv4 prefix from `gateway` to configure the utun point-to-point address. \
Windows, Linux and macOS can also apply system routes from `autoSystemRoutingTable`.
macOS does not configure system DNS from the `dns` field, and neither does Linux by default; system DNS remains managed by the OS or distribution-specific network services. \
Linux and macOS do not configure system DNS from the `dns` field; system DNS remains managed by the OS or distribution-specific network services. \
For more advanced routing policies or rules, OS level configuration can still manage the named interface (e.g. xray0) when it appears.
This keeps complex system level routing and rules in a single place of responsibility - the OS itself. \
Examples of how to achieve this on a simple Linux system (Ubuntu with systemd-networkd) can be found at the end of this README.
### SYSTEM DNS ON LINUX (`autoSystemDnsToGateway`)
On Linux, setting `autoSystemDnsToGateway` to `true` lets the inbound point the system resolver at the tun interface, so name lookups resolve through Xray instead of going out over the physical link. It is off by default, and it is Linux-only.
It uses `resolvectl`, which means it only works when all of these hold. Where Xray can tell that one does not, it does not start:
- the system runs systemd and `resolvectl` is on `PATH`
- `systemd-resolved` is enabled and actually managing DNS (installed but not running is not enough)
- systemd-resolved is version 240 or newer, where `default-route` exists
- no `dns` upstream resolves through the system resolver, directly or through its own bootstrap (see below)
The address handed over is the first IPv4 `gateway`, or without one the first IPv6 `gateway`, incremented by one (e.g. `192.168.100.1/30` -> `192.168.100.2`, `fc00::1/64` -> `fc00::2`). Without any `gateway`, the config is rejected. It is not taken from `dns`: handing `1.1.1.1` to `resolvectl dns` would make systemd-resolved query that server directly over the physical link, which is the leak this option exists to close.
Because that address has to actually answer, the takeover is checked before it happens. A query from the interface address to that address is routed through the configured rules, and host-wide DNS is only changed when the result is a DNS-capable outbound. Otherwise DNS is left alone and Xray does not start. In practice this means you also need a routing rule sending the interface's port 53 to a `dns` outbound, for example:
```json
"routing": {
"rules": [
{ "type": "field", "inboundTag": ["tun"], "port": 53, "outboundTag": "dns" }
]
}
```
The check is a preflight, not a proof for arbitrary rules. It sends its query from the interface address and from a representative ephemeral source port, so a rule that matches on the source port cannot be predicted ahead of time: if the interface's port 53 reaches the `dns` outbound only from some source ports, the takeover is accepted and queries from the other ports fail. Supported configurations are those where the DNS path does not depend on the source port, that is, where the interface's port 53 reaches a `dns` outbound whatever its source.
It is also a check for the dependencies it knows about, not a proof that no indirect one exists. A hostname-based upstream that bootstraps through system DNS is the case in point: `https+local://dns.google/dns-query` resolves its own hostname with `DialSystem`, so once the takeover is in place that bootstrap goes `resolved -> TUN -> DNS outbound -> bootstrap -> resolved` and the query times out. The preflight does not see it, because the dependency sits in the upstream's bootstrap rather than in the clients it inspects. Upstream resolution, bootstrap included, therefore has to stay independent of the resolver path being redirected; configuring the address instead of the hostname, or resolving the hostname beforehand, avoids it.
The upstream requirement in the list above matters as much as the routing rule. With no name servers configured, Core resolves through a client that forwards to the system resolver; pointing the system resolver at the TUN would then close a loop through the DNS outbound, `resolved -> TUN -> DNS outbound -> system resolver -> resolved`, and resolution stops. The takeover is refused in that case, and Xray does not start.
The same applies to a name server pointed at `localhost`, and to a `dns` section that is present but lists no name servers. One such upstream is enough to refuse the takeover even when independent upstreams are configured alongside it: name servers are selected per domain, so a domain-specific rule can still choose the local one, and the loop then affects whichever domains reach it. The check is deliberately broader than the loop it observed, because the alternative would be to drop a name server the user configured.
Where it cannot apply, Xray does not start, rather than run with the leak described in XTLS/Xray-core#6454, so leave the option off there:
| Environment | Behaviour |
|---|---|
| systemd distribution with systemd-resolved enabled | applies |
| Alpine, Void, Devuan, OpenRC-based, OpenWrt | no `resolvectl`, does not start |
| DNS managed by dnsmasq / unbound / BIND / static `resolv.conf` | unreachable by `resolvectl`, does not start |
| Containers without a systemd-resolved daemon | does not start |
| systemd older than 240 | `default-route` unavailable, does not start |
On `Close()` the setting is reverted. It is **not** reverted if the process is killed with `SIGKILL`, since a process cannot handle that signal; run `resolvectl revert <iface>` to clean up by hand. An application that brings its own DNS endpoint is unaffected either way — this only covers the system resolver.
Due to this inbound not actually being a proxy, the configuration ignore required listen and port options, and never listen on any port. \
Here is simple Xray config snippet to enable the inbound:
```
@@ -199,20 +155,6 @@ To make it start, wintun.dll specific for your Windows/arch must be present next
After the start network adapter with the name you chose in the config will be created in the system, and exist while Xray is running.
When `dns` is set, those servers are applied to the adapter. Windows is kept from registering the TUN's addresses in DNS, and its DNS cache is flushed when the TUN starts and stops.
With `autoSystemWfpBlockLeak`, which needs `autoSystemRoutingTable` (the config is rejected otherwise), Xray also adds Windows Filtering Platform filters that keep two kinds of traffic of every program but Xray itself from leaving outside the TUN, each chosen by a value in the list, e.g. `"autoSystemWfpBlockLeak": ["dns", "misconfigtun"]`:
- `"dns"` (needs `dns`, the config is rejected otherwise): DNS (port 53) only goes through the TUN. Windows keeps sending name queries to the DNS servers of the other interfaces as well, out through those interfaces whatever the routes say, and other programs reach a resolver on the local network (e.g. `192.168.1.1` handed out by DHCP) through its more specific LAN route instead of the TUN. On Windows 11 and Server 2022 and later, where those queries may also go over HTTPS or TLS, Windows' DNS Client service cannot connect outside the TUN at all, except for name resolution on the local network (LLMNR, mDNS). The `dns` servers therefore have to lie within `gateway` or `autoSystemRoutingTable` (a warning is logged otherwise), and DNS servers that should be reached directly belong in Xray's own `dns` settings.
- `"misconfigtun"`: an IP version without routes in `autoSystemRoutingTable`, IPv4 or IPv6, is blocked entirely, in both directions, as it would bypass the TUN. Only loopback and what Windows itself needs on the local link (DHCP, and for IPv6 neighbor and multicast listener discovery) remain allowed. An address of that version in `gateway` is not needed: without one, Windows gives the TUN link-local addresses itself, an IPv6 one at once and an IPv4 one from `169.254.0.0/16` after some seconds (until then, IPv4 routed to the TUN is unreachable), and what is routed to the TUN goes through it with those.
With the filters in place, Xray's own connections out also get past Windows Firewall's block rules (other firewalls may still block them), while connections to Xray's inbounds stay subject to them.
Names that Xray resolves through the system resolver, such as an outbound's server address given as a domain with the default `AsIs` domain strategy, would be looked up by Windows on Xray's behalf, and those queries would then go into the TUN too. While DNS is restricted this way and `autoOutboundsInterface` is in use (the default with `autoSystemRoutingTable`), Xray therefore resolves them itself, with its own queries to the DNS servers of the other interfaces. That bypasses Windows' DNS cache, and its name resolution on the local network (LLMNR, mDNS): a server address given as a domain is looked up again for every connection, and a DNS server that does not answer delays each lookup. Having Xray's own `dns` resolve it, through the outbound's `sockopt.domainStrategy`, avoids that. The `localhost` DNS server queries the same servers whenever `autoOutboundsInterface` is in use. Both skip the TUN's own DNS servers, unless another interface uses them as well: queried from Xray itself, they would lead back into it, or nowhere.
If the filters cannot be added, Xray does not start. They are removed when Xray exits. Not covered is name resolution on the local network (LLMNR, mDNS, NetBIOS), except over an IP version that is blocked.
`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"]`.
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.
+2 -21
View File
@@ -32,8 +32,6 @@ type Config struct {
AutoSystemRoutingTable []string `protobuf:"bytes,6,rep,name=auto_system_routing_table,json=autoSystemRoutingTable,proto3" json:"auto_system_routing_table,omitempty"`
AutoOutboundsInterface string `protobuf:"bytes,7,opt,name=auto_outbounds_interface,json=autoOutboundsInterface,proto3" json:"auto_outbounds_interface,omitempty"`
Desc string `protobuf:"bytes,8,opt,name=desc,proto3" json:"desc,omitempty"`
AutoSystemDnsToGateway bool `protobuf:"varint,9,opt,name=auto_system_dns_to_gateway,json=autoSystemDnsToGateway,proto3" json:"auto_system_dns_to_gateway,omitempty"`
AutoSystemWfpBlockLeak []string `protobuf:"bytes,10,rep,name=auto_system_wfp_block_leak,json=autoSystemWfpBlockLeak,proto3" json:"auto_system_wfp_block_leak,omitempty"`
unknownFields protoimpl.UnknownFields
sizeCache protoimpl.SizeCache
}
@@ -124,25 +122,11 @@ func (x *Config) GetDesc() string {
return ""
}
func (x *Config) GetAutoSystemDnsToGateway() bool {
if x != nil {
return x.AutoSystemDnsToGateway
}
return false
}
func (x *Config) GetAutoSystemWfpBlockLeak() []string {
if x != nil {
return x.AutoSystemWfpBlockLeak
}
return nil
}
var File_proxy_tun_config_proto protoreflect.FileDescriptor
const file_proxy_tun_config_proto_rawDesc = "" +
"\n" +
"\x16proxy/tun/config.proto\x12\x0exray.proxy.tun\"\xfa\x02\n" +
"\x16proxy/tun/config.proto\x12\x0exray.proxy.tun\"\x82\x02\n" +
"\x06Config\x12\x12\n" +
"\x04name\x18\x01 \x01(\tR\x04name\x12\x10\n" +
"\x03MTU\x18\x02 \x01(\rR\x03MTU\x12\x18\n" +
@@ -152,10 +136,7 @@ const file_proxy_tun_config_proto_rawDesc = "" +
"user_level\x18\x05 \x01(\rR\tuserLevel\x129\n" +
"\x19auto_system_routing_table\x18\x06 \x03(\tR\x16autoSystemRoutingTable\x128\n" +
"\x18auto_outbounds_interface\x18\a \x01(\tR\x16autoOutboundsInterface\x12\x12\n" +
"\x04desc\x18\b \x01(\tR\x04desc\x12:\n" +
"\x1aauto_system_dns_to_gateway\x18\t \x01(\bR\x16autoSystemDnsToGateway\x12:\n" +
"\x1aauto_system_wfp_block_leak\x18\n" +
" \x03(\tR\x16autoSystemWfpBlockLeakBL\n" +
"\x04desc\x18\b \x01(\tR\x04descBL\n" +
"\x12com.xray.proxy.tunP\x01Z#github.com/xtls/xray-core/proxy/tun\xaa\x02\x0eXray.Proxy.Tunb\x06proto3"
var (
-2
View File
@@ -15,6 +15,4 @@ message Config {
repeated string auto_system_routing_table = 6;
string auto_outbounds_interface = 7;
string desc = 8;
bool auto_system_dns_to_gateway = 9;
repeated string auto_system_wfp_block_leak = 10;
}
+4 -45
View File
@@ -37,25 +37,6 @@ type Handler struct {
downlinkCounter stats.Counter
}
type tunUDPStatsWriter struct {
writer buf.Writer
counter stats.Counter
}
func (w *tunUDPStatsWriter) WriteMultiBuffer(mb buf.MultiBuffer) error {
for len(mb) > 0 {
remaining, packet := buf.SplitFirst(mb)
packetSize := packet.Len()
if err := w.writer.WriteMultiBuffer(buf.MultiBuffer{packet}); err != nil {
buf.ReleaseMulti(remaining)
return err
}
w.counter.Add(int64(packetSize))
mb = remaining
}
return nil
}
// ConnectionHandler interface with the only method that stack is going to push new connections to
type ConnectionHandler interface {
HandleConnection(conn net.Conn, destination net.Destination)
@@ -123,7 +104,7 @@ func (t *Handler) Start() error {
iface := updater.Get()
if iface == nil {
errors.LogInfo(context.Background(), "[tun] falied to set interface > iface == nil")
return errors.New("iface not found")
return nil
}
return c.Control(func(fd uintptr) {
addrPort, _ := netip.ParseAddrPort(address)
@@ -165,18 +146,6 @@ func (t *Handler) Start() error {
return err
}
// Platform-specific system DNS takeover, where the platform implements it.
// Rather no TUN than one that the system DNS bypasses.
if c, ok := tunInterface.(interface {
ConfigureSystemDNS(context.Context, string) error
}); ok {
if err := c.ConfigureSystemDNS(t.ctx, t.tag); err != nil {
_ = tunStack.Close()
_ = tunInterface.Close()
return errors.New("unable to set the system DNS (remove autoSystemDnsToGateway to run without)").Base(err)
}
}
t.stack = tunStack
t.tun = tunInterface
@@ -202,8 +171,7 @@ func (t *Handler) HandleConnection(conn net.Conn, destination net.Destination) {
return
}
source := net.DestinationFromAddr(remote)
isUDP := destination.Network == net.Network_UDP
if !isUDP && (t.uplinkCounter != nil || t.downlinkCounter != nil) {
if t.uplinkCounter != nil || t.downlinkCounter != nil {
conn = &stat.CounterConnection{
Connection: conn,
ReadCounter: t.uplinkCounter,
@@ -235,18 +203,9 @@ func (t *Handler) HandleConnection(conn net.Conn, destination net.Destination) {
})
errors.LogInfo(ctx, "processing from ", source, " to ", destination)
reader := &buf.TimeoutWrapperReader{Reader: buf.NewReader(conn)}
writer := buf.NewWriter(conn)
if isUDP {
reader.Counter = t.uplinkCounter
if t.downlinkCounter != nil {
writer = &tunUDPStatsWriter{writer: writer, counter: t.downlinkCounter}
}
}
link := &transport.Link{
Reader: reader,
Writer: writer,
Reader: &buf.TimeoutWrapperReader{Reader: buf.NewReader(conn)},
Writer: buf.NewWriter(conn),
}
if err := t.dispatcher.DispatchLink(ctx, destination, link); err != nil {
errors.LogError(ctx, errors.New("connection closed").Base(err))
-251
View File
@@ -6,24 +6,12 @@ import (
"context"
"net"
"net/netip"
"os/exec"
"strconv"
"sync"
"github.com/vishvananda/netlink"
appdns "github.com/xtls/xray-core/app/dns"
"github.com/xtls/xray-core/common/errors"
xnet "github.com/xtls/xray-core/common/net"
"github.com/xtls/xray-core/common/platform"
"github.com/xtls/xray-core/common/serial"
"github.com/xtls/xray-core/common/session"
"github.com/xtls/xray-core/core"
feature_dns "github.com/xtls/xray-core/features/dns"
"github.com/xtls/xray-core/features/dns/localdns"
"github.com/xtls/xray-core/features/outbound"
"github.com/xtls/xray-core/features/routing"
routingsession "github.com/xtls/xray-core/features/routing/session"
"github.com/xtls/xray-core/proxy/dns"
"golang.org/x/sys/unix"
"gvisor.dev/gvisor/pkg/tcpip/link/fdbased"
"gvisor.dev/gvisor/pkg/tcpip/stack"
@@ -42,244 +30,6 @@ type LinuxTun struct {
systemRoutes []netlink.Route
routeMonitorStop chan struct{}
routeMonitorOnce sync.Once
systemDNSSet bool
systemDNSDirty bool
}
// resolvectlRunner runs a resolvectl command. Overridable for tests.
var resolvectlRunner = func(name string, args ...string) ([]byte, error) {
return exec.Command(name, args...).CombinedOutput()
}
// systemDNSAddrs derives the addresses used for the system DNS takeover from the
// first IPv4 gateway, or without one, the first IPv6 gateway: the gateway
// address itself is what a query from this interface appears to come from, and
// the next address is what the resolver is pointed at. The latter belongs to
// the TUN and is answered inside Xray; handing the configured public resolvers
// to resolvectl instead would leave the system querying them directly over the
// physical link, defeating the point of the TUN.
func systemDNSAddrs(gateway []string) (source, dns netip.Addr, ok bool) {
var first6 netip.Addr
for _, address := range gateway {
prefix, err := netip.ParsePrefix(address)
if err != nil {
continue
}
addr := prefix.Addr()
if addr.Is4() {
return addr, addr.Next(), true
}
if !first6.IsValid() {
first6 = addr
}
}
if first6.IsValid() {
return first6, first6.Next(), true
}
return netip.Addr{}, netip.Addr{}, false
}
func buildResolvectlArgs(action, iface string, extra ...string) []string {
args := make([]string, 0, 2+len(extra))
args = append(args, action, iface)
args = append(args, extra...)
return args
}
func runResolvectl(action, iface string, extra ...string) error {
args := buildResolvectlArgs(action, iface, extra...)
if _, err := resolvectlRunner("resolvectl", args...); err != nil {
return errors.New("resolvectl ", action, " failed").Base(err)
}
return nil
}
// ifaceName returns the TUN interface name, or empty when the link is not
// available. Callers must treat empty as "nothing to configure".
func (t *LinuxTun) ifaceName() string {
if t.tunLink == nil {
return ""
}
attrs := t.tunLink.Attrs()
if attrs == nil {
return ""
}
return attrs.Name
}
// probeSourcePort is a representative client port for the routing probe. A real
// query arrives from an ephemeral port that cannot be known in advance, so this
// only matters for a rule that matches on a source port.
const probeSourcePort = 49152
// verifyDNSRouting reports whether a DNS query to address would actually be
// handled. Redirecting the system resolver at an address nothing answers would
// break name resolution outright, so the takeover only proceeds when routing
// hands such a query to a DNS-capable outbound.
//
// Overridable for tests.
var verifyDNSRouting = func(ctx context.Context, inboundTag, source, address string) error {
ip, err := netip.ParseAddr(address)
if err != nil {
return errors.New("invalid DNS address ", address).Base(err)
}
src, err := netip.ParseAddr(source)
if err != nil || src.Is4() != ip.Is4() {
return errors.New("invalid source address ", source).Base(err)
}
instance := core.MustFromContext(ctx)
// Any resolution path that could still reach the system resolver has to be
// refused, because pointing the system resolver at the TUN would close a
// loop through the DNS outbound. With no `dns` section Core installs such a
// client; with a `dns` section that has no name servers app/dns falls back
// to one; and a name server pointed at "localhost" is one even when
// independent upstreams are configured alongside it, because name servers
// are selected per domain.
switch dnsFeature := instance.GetFeature(feature_dns.ClientType()).(type) {
case *localdns.Client:
return errors.New("DNS feature is the system resolver, takeover would loop")
case *appdns.DNS:
if dnsFeature.MayUseSystemResolver() {
return errors.New("DNS configuration may resolve through the system resolver, takeover would loop")
}
}
router, ok := instance.GetFeature(routing.RouterType()).(routing.Router)
if !ok {
return errors.New("router feature unavailable")
}
// A real query from this interface carries a source address, and rules may
// match on it, so the probe has to carry one too.
queryCtx := session.ContextWithInbound(ctx, &session.Inbound{
Name: "tun",
Tag: inboundTag,
Source: xnet.UDPDestination(xnet.IPAddress(src.AsSlice()), probeSourcePort),
})
queryCtx = session.ContextWithOutbounds(queryCtx, []*session.Outbound{{
Target: xnet.UDPDestination(xnet.IPAddress(ip.AsSlice()), 53),
}})
route, err := router.PickRoute(routingsession.AsRoutingContext(queryCtx))
if err != nil {
return errors.New("no route for ", address, ":53").Base(err)
}
manager, ok := instance.GetFeature(outbound.ManagerType()).(outbound.Manager)
if !ok {
return errors.New("outbound manager unavailable")
}
handler := manager.GetHandler(route.GetOutboundTag())
if handler == nil {
return errors.New("outbound ", route.GetOutboundTag(), " does not exist")
}
if settings := handler.ProxySettings(); settings == nil || settings.Type != serial.GetMessageType(&dns.Config{}) {
return errors.New("outbound ", route.GetOutboundTag(), " does not handle DNS")
}
return nil
}
// ConfigureSystemDNS points systemd-resolved at this interface so name lookups
// resolve through Xray instead of leaking to the physical link.
//
// It acts only when the config opts in, and it verifies the data path first:
// unless a query to the advertised address would actually be handled, host-wide
// resolution is left to the OS and an error returned. The caller does not start
// the TUN on an error, as the system DNS would bypass it.
func (t *LinuxTun) ConfigureSystemDNS(ctx context.Context, inboundTag string) error {
if !t.options.AutoSystemDnsToGateway {
return nil
}
if t.systemDNSSet {
return nil
}
// A previous revert may have failed. Retry before applying anything, so a
// dirty resolver does not silently outlive the attempt to clean it up.
if t.systemDNSDirty {
if err := t.revertSystemDNS(); err != nil {
return errors.New("previous system DNS revert still failing").Base(err)
}
}
source, address, ok := systemDNSAddrs(t.options.Gateway)
if !ok {
return errors.New("no gateway, cannot derive a system DNS address")
}
iface := t.ifaceName()
if iface == "" {
return errors.New("interface not available")
}
if err := verifyDNSRouting(ctx, inboundTag, source.String(), address.String()); err != nil {
return errors.New("no DNS path at ", address.String(), ":53").Base(err)
}
// Applied as a sequence with rollback: a half-configured resolver would be
// worse than none at all.
if err := runResolvectl("dns", iface, address.String()); err != nil {
return errors.New("resolvectl dns failed").Base(err)
}
if err := runResolvectl("domain", iface, "~."); err != nil {
return t.rollbackSystemDNS(iface, errors.New("resolvectl domain failed").Base(err))
}
if err := runResolvectl("default-route", iface, "true"); err != nil {
return t.rollbackSystemDNS(iface, errors.New("resolvectl default-route failed").Base(err))
}
t.systemDNSSet = true
errors.LogInfo(ctx, "[tun] system DNS set to ", address.String(), " on ", iface)
return nil
}
// rollbackSystemDNS undoes a partially applied takeover. A failed revert is
// recorded so the next attempt retries it, and is reported rather than
// swallowed.
func (t *LinuxTun) rollbackSystemDNS(iface string, cause error) error {
if err := runResolvectl("revert", iface); err != nil {
t.systemDNSDirty = true
// Combine, because Base overwrites: reporting only the cause would hide
// the revert failure, and reporting only the revert failure would hide
// why the revert was attempted.
return errors.New("revert failed, per-link DNS settings may remain").Base(errors.Combine(err, cause))
}
return cause
}
// revertSystemDNS issues the revert and keeps the dirty flag in step with the
// outcome.
func (t *LinuxTun) revertSystemDNS() error {
err := runResolvectl("revert", t.ifaceName())
t.systemDNSDirty = err != nil
if err != nil {
return err
}
t.systemDNSSet = false
return nil
}
// unsetSystemDNS hands DNS back to the OS. Only meaningful when
// ConfigureSystemDNS applied something, or a previous revert failed.
func (t *LinuxTun) unsetSystemDNS() {
if !t.systemDNSSet && !t.systemDNSDirty {
return
}
if t.ifaceName() == "" {
// The link is gone, and its per-link settings went with it.
t.systemDNSSet = false
t.systemDNSDirty = false
return
}
if err := t.revertSystemDNS(); err != nil {
errors.LogInfoInner(context.Background(), err, "[tun] failed to revert system DNS; per-link settings may remain until revert succeeds")
}
}
// LinuxTun implements Tun
@@ -450,7 +200,6 @@ func (t *LinuxTun) Close() error {
}
})
t.unsetSystemDNS()
_ = t.unsetSystemRoutes()
_ = t.unsetInterfaceAddresses()
-205
View File
@@ -1,205 +0,0 @@
//go:build linux && !android
package tun
import (
"context"
"strings"
"testing"
"github.com/xtls/xray-core/app/dispatcher"
appdns "github.com/xtls/xray-core/app/dns"
"github.com/xtls/xray-core/app/proxyman"
_ "github.com/xtls/xray-core/app/proxyman/inbound"
_ "github.com/xtls/xray-core/app/proxyman/outbound"
"github.com/xtls/xray-core/app/router"
"github.com/xtls/xray-core/common/geodata"
"github.com/xtls/xray-core/common/net"
"github.com/xtls/xray-core/common/serial"
"github.com/xtls/xray-core/core"
"github.com/xtls/xray-core/proxy/blackhole"
proxydns "github.com/xtls/xray-core/proxy/dns"
"github.com/xtls/xray-core/proxy/freedom"
)
const (
routeTestInboundTag = "tun"
routeTestSource = "192.168.100.1"
routeTestDNSAddress = "192.168.100.2"
)
// port53Rule sends DNS queries arriving from the interface to the dns outbound.
func port53Rule() *router.RoutingRule {
return &router.RoutingRule{
InboundTag: []string{routeTestInboundTag},
PortList: &net.PortList{Range: []*net.PortRange{net.SinglePortRange(53)}},
TargetTag: &router.RoutingRule_Tag{Tag: "dns"},
}
}
// sourceBlockRule diverts traffic from one address, which is the shape of a rule
// that only matches because the real request carries a source.
func sourceBlockRule(ip []byte) *router.RoutingRule {
return &router.RoutingRule{
SourceIp: []*geodata.IPRule{{
Value: &geodata.IPRule_Custom{
Custom: &geodata.CIDRRule{
Cidr: &geodata.CIDR{Ip: ip, Prefix: 32},
},
},
}},
TargetTag: &router.RoutingRule_Tag{Tag: "block"},
}
}
// newRouteTestContext builds a real but unstarted instance: no TUN device, no
// running resolver. The instance is placed in the context through the key core
// exports for tests.
func newRouteTestContext(t *testing.T, withDNSApp bool, nameServers []*appdns.NameServer, rules []*router.RoutingRule) context.Context {
t.Helper()
apps := []*serial.TypedMessage{
serial.ToTypedMessage(&dispatcher.Config{}),
serial.ToTypedMessage(&proxyman.InboundConfig{}),
serial.ToTypedMessage(&proxyman.OutboundConfig{}),
serial.ToTypedMessage(&router.Config{Rule: rules}),
}
if withDNSApp {
apps = append(apps, serial.ToTypedMessage(&appdns.Config{NameServer: nameServers}))
}
instance, err := core.New(&core.Config{
App: apps,
Outbound: []*core.OutboundHandlerConfig{
{Tag: "direct", ProxySettings: serial.ToTypedMessage(&freedom.Config{})},
{Tag: "dns", ProxySettings: serial.ToTypedMessage(&proxydns.Config{})},
{Tag: "block", ProxySettings: serial.ToTypedMessage(&blackhole.Config{})},
},
})
if err != nil {
t.Fatalf("core.New: %v", err)
}
t.Cleanup(func() { _ = instance.Close() })
return context.WithValue(context.Background(), core.XrayKey(1), instance)
}
func udpNameServer(ip []byte) []*appdns.NameServer {
return []*appdns.NameServer{{
Address: &net.Endpoint{
Network: net.Network_UDP,
Address: &net.IPOrDomain{Address: &net.IPOrDomain_Ip{Ip: ip}},
Port: 53,
},
}}
}
// localNameServer is a name server pointed at "localhost", which app/dns
// resolves through the system resolver.
func localNameServer() *appdns.NameServer {
return &appdns.NameServer{
Address: &net.Endpoint{
Network: net.Network_UDP,
Address: &net.IPOrDomain{Address: &net.IPOrDomain_Domain{Domain: "localhost"}},
Port: 53,
},
}
}
// These drive the real feature lookup and the real router. verifyDNSRouting is
// the same function ConfigureSystemDNS calls, so a false positive here is a
// false positive in the takeover decision itself, which is what assertions on
// the resolvectl arguments could never catch.
func TestVerifyDNSRoutingDecisions(t *testing.T) {
tests := []struct {
name string
withDNSApp bool
nameServers []*appdns.NameServer
rules []*router.RoutingRule
wantErr string
}{
{
name: "independent upstream reaches the dns outbound",
withDNSApp: true,
nameServers: udpNameServer([]byte{9, 9, 9, 9}),
rules: []*router.RoutingRule{port53Rule()},
wantErr: "",
},
{
name: "no dns section falls back to the system resolver",
rules: []*router.RoutingRule{port53Rule()},
wantErr: "system resolver",
},
{
name: "dns section without name servers falls back too",
withDNSApp: true,
rules: []*router.RoutingRule{port53Rule()},
wantErr: "system resolver",
},
{
// An independent upstream is not enough on its own: name servers are
// selected per domain, so a local one can still be the one chosen.
// The refusal is deliberately domain-agnostic for that reason.
name: "a local name server alongside an independent one",
withDNSApp: true,
nameServers: append(udpNameServer([]byte{9, 9, 9, 9}), localNameServer()),
rules: []*router.RoutingRule{port53Rule()},
wantErr: "system resolver",
},
{
name: "a rule on the interface address diverts the real query",
withDNSApp: true, nameServers: udpNameServer([]byte{9, 9, 9, 9}),
rules: []*router.RoutingRule{
sourceBlockRule([]byte{192, 168, 100, 1}),
port53Rule(),
},
wantErr: "does not handle DNS",
},
{
name: "a rule on another address does not match it",
withDNSApp: true, nameServers: udpNameServer([]byte{9, 9, 9, 9}),
rules: []*router.RoutingRule{
sourceBlockRule([]byte{10, 0, 0, 1}),
port53Rule(),
},
wantErr: "",
},
{
name: "no rule matches the query",
withDNSApp: true, nameServers: udpNameServer([]byte{9, 9, 9, 9}),
wantErr: "no route",
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
ctx := newRouteTestContext(t, tt.withDNSApp, tt.nameServers, tt.rules)
err := verifyDNSRouting(ctx, routeTestInboundTag, routeTestSource, routeTestDNSAddress)
if tt.wantErr == "" {
if err != nil {
t.Fatalf("expected the takeover to be accepted, got: %v", err)
}
return
}
if err == nil {
t.Fatalf("expected the takeover to be refused with %q, got nil", tt.wantErr)
}
if !strings.Contains(err.Error(), tt.wantErr) {
t.Errorf("error = %q, want it to contain %q", err.Error(), tt.wantErr)
}
})
}
}
// Without an IPv4 gateway, the takeover uses the first IPv6 one, and the probe
// carries IPv6 addresses.
func TestVerifyDNSRoutingIPv6(t *testing.T) {
ctx := newRouteTestContext(t, true, udpNameServer([]byte{9, 9, 9, 9}), []*router.RoutingRule{port53Rule()})
if err := verifyDNSRouting(ctx, routeTestInboundTag, "fc00::1", "fc00::2"); err != nil {
t.Fatalf("expected the takeover to be accepted, got: %v", err)
}
if err := verifyDNSRouting(ctx, routeTestInboundTag, routeTestSource, "fc00::2"); err == nil {
t.Fatal("expected mixed IPv4 and IPv6 addresses to be refused")
}
}
-445
View File
@@ -1,445 +0,0 @@
//go:build linux && !android
package tun
import (
"context"
"errors"
"strings"
"testing"
"github.com/vishvananda/netlink"
)
// testLink returns a minimal netlink.Link whose Attrs().Name is name, so the
// DNS helpers can be exercised without a real TUN device.
func testLink(name string) netlink.Link {
return &netlink.Dummy{LinkAttrs: netlink.LinkAttrs{Name: name}}
}
type probeCall struct {
inboundTag string
source string
address string
}
// stubDNSRouting replaces the routing probe for the duration of a test and
// records how it was called, so tests can assert the probe is representative.
func stubDNSRouting(t *testing.T, err error) *[]probeCall {
t.Helper()
original := verifyDNSRouting
calls := []probeCall{}
verifyDNSRouting = func(_ context.Context, inboundTag, source, address string) error {
calls = append(calls, probeCall{inboundTag, source, address})
return err
}
t.Cleanup(func() { verifyDNSRouting = original })
return &calls
}
// recorder installs a resolvectl stub for the duration of a test and returns the
// captured invocations. An empty failOn succeeds every call; otherwise the named
// subcommand fails.
func recorder(t *testing.T, failOn string) *[][]string {
t.Helper()
original := resolvectlRunner
calls := [][]string{}
resolvectlRunner = func(name string, args ...string) ([]byte, error) {
calls = append(calls, append([]string{name}, args...))
if failOn != "" && len(args) > 0 && args[0] == failOn {
return nil, errors.New("boom")
}
return nil, nil
}
t.Cleanup(func() { resolvectlRunner = original })
return &calls
}
func optedInTun() *LinuxTun {
return &LinuxTun{
options: &Config{
Name: "xray_tun",
Gateway: []string{"192.168.100.1/30"},
AutoSystemDnsToGateway: true,
},
tunLink: testLink("xray_tun"),
}
}
func joined(calls [][]string) string {
parts := make([]string, 0, len(calls))
for _, call := range calls {
parts = append(parts, strings.Join(call, " "))
}
return strings.Join(parts, " | ")
}
func TestConfigureSystemDNSDisabledByDefault(t *testing.T) {
probes := stubDNSRouting(t, nil)
calls := recorder(t, "")
t1 := optedInTun()
t1.options.AutoSystemDnsToGateway = false
if err := t1.ConfigureSystemDNS(context.Background(), "tun"); err != nil {
t.Fatalf("unexpected error: %v", err)
}
if len(*probes) != 0 {
t.Errorf("routing probe must not run when disabled, got %d calls", len(*probes))
}
if len(*calls) != 0 {
t.Errorf("resolvectl must not run when disabled, got %v", *calls)
}
if t1.systemDNSSet {
t.Error("systemDNSSet should stay false when disabled")
}
}
func TestConfigureSystemDNSNoGateway(t *testing.T) {
probes := stubDNSRouting(t, nil)
calls := recorder(t, "")
t1 := optedInTun()
t1.options.Gateway = nil
if err := t1.ConfigureSystemDNS(context.Background(), "tun"); err == nil {
t.Fatal("expected an error when no gateway is configured")
}
if len(*probes) != 0 {
t.Errorf("routing probe must not run without a gateway, got %d calls", len(*probes))
}
if len(*calls) != 0 {
t.Errorf("resolvectl must not run without a gateway, got %v", *calls)
}
}
// This is the case the reviewer flagged: without a routed DNS path, pointing the
// system resolver at the derived address would break resolution outright.
func TestConfigureSystemDNSLeavesOSDNSWhenNoRoute(t *testing.T) {
probes := stubDNSRouting(t, errors.New("no route"))
calls := recorder(t, "")
t1 := optedInTun()
if err := t1.ConfigureSystemDNS(context.Background(), "tun"); err == nil {
t.Fatal("expected an error when the DNS path is unverified")
}
if len(*probes) != 1 {
t.Errorf("routing probe should run once, got %d", len(*probes))
}
if len(*calls) != 0 {
t.Errorf("system DNS must be left untouched, got %v", *calls)
}
if t1.systemDNSSet {
t.Error("systemDNSSet should stay false when the path is unverified")
}
}
// A real query from the interface carries a source address, and rules may match
// on it, so the probe must not be source-less.
func TestConfigureSystemDNSProbeCarriesSource(t *testing.T) {
probes := stubDNSRouting(t, nil)
recorder(t, "")
t1 := optedInTun()
if err := t1.ConfigureSystemDNS(context.Background(), "tun"); err != nil {
t.Fatalf("unexpected error: %v", err)
}
if len(*probes) != 1 {
t.Fatalf("expected one probe call, got %d", len(*probes))
}
got := (*probes)[0]
if got.source != "192.168.100.1" {
t.Errorf("probe source = %q, want the interface address %q", got.source, "192.168.100.1")
}
if got.address != "192.168.100.2" {
t.Errorf("probe address = %q, want %q", got.address, "192.168.100.2")
}
if got.inboundTag != "tun" {
t.Errorf("probe inbound tag = %q, want %q", got.inboundTag, "tun")
}
}
func TestConfigureSystemDNSAppliesResolvectl(t *testing.T) {
stubDNSRouting(t, nil)
calls := recorder(t, "")
t1 := optedInTun()
if err := t1.ConfigureSystemDNS(context.Background(), "tun"); err != nil {
t.Fatalf("unexpected error: %v", err)
}
if !t1.systemDNSSet {
t.Fatal("systemDNSSet should be true after a successful takeover")
}
want := "resolvectl dns xray_tun 192.168.100.2 | " +
"resolvectl domain xray_tun ~. | " +
"resolvectl default-route xray_tun true"
if got := joined(*calls); got != want {
t.Errorf("resolvectl calls = %q, want %q", got, want)
}
}
func TestConfigureSystemDNSIdempotent(t *testing.T) {
stubDNSRouting(t, nil)
calls := recorder(t, "")
t1 := optedInTun()
if err := t1.ConfigureSystemDNS(context.Background(), "tun"); err != nil {
t.Fatalf("unexpected error: %v", err)
}
before := len(*calls)
if err := t1.ConfigureSystemDNS(context.Background(), "tun"); err != nil {
t.Fatalf("unexpected error: %v", err)
}
if len(*calls) != before {
t.Errorf("second call must be a no-op, calls went %d -> %d", before, len(*calls))
}
}
// A half-applied resolver is worse than none, so a failure mid-sequence reverts.
func TestConfigureSystemDNSRollsBackOnPartialFailure(t *testing.T) {
stubDNSRouting(t, nil)
calls := recorder(t, "domain")
t1 := optedInTun()
if err := t1.ConfigureSystemDNS(context.Background(), "tun"); err == nil {
t.Fatal("expected an error when a resolvectl step fails")
}
if t1.systemDNSSet {
t.Error("systemDNSSet should stay false after a failed takeover")
}
if t1.systemDNSDirty {
t.Error("a successful revert should not leave the resolver dirty")
}
if !strings.Contains(joined(*calls), "resolvectl revert xray_tun") {
t.Errorf("expected a revert after partial failure, got %q", joined(*calls))
}
}
// If the revert itself fails the settings may still be installed, so the state
// has to be remembered rather than silently dropped.
func TestConfigureSystemDNSRollbackFailureKeepsDirty(t *testing.T) {
stubDNSRouting(t, nil)
calls := recorder(t, "revert")
t1 := optedInTun()
t1.options.Gateway = []string{"192.168.100.1/30"}
// Make only the rollback path fail: "dns" succeeds, "domain" fails, "revert" fails.
*calls = nil
original := resolvectlRunner
defer func() { resolvectlRunner = original }()
resolvectlRunner = func(name string, args ...string) ([]byte, error) {
*calls = append(*calls, append([]string{name}, args...))
if len(args) > 0 && (args[0] == "domain" || args[0] == "revert") {
return nil, errors.New("boom")
}
return nil, nil
}
if err := t1.ConfigureSystemDNS(context.Background(), "tun"); err == nil {
t.Fatal("expected an error when domain fails")
}
if !t1.systemDNSDirty {
t.Error("a failed revert must leave the resolver marked dirty")
}
if t1.systemDNSSet {
t.Error("systemDNSSet must stay false when the takeover did not complete")
}
}
// A dirty resolver is retried before anything new is applied.
func TestConfigureSystemDNSRetriesDirtyBeforeApplying(t *testing.T) {
stubDNSRouting(t, nil)
calls := recorder(t, "")
t1 := optedInTun()
t1.systemDNSDirty = true
if err := t1.ConfigureSystemDNS(context.Background(), "tun"); err != nil {
t.Fatalf("unexpected error: %v", err)
}
got := joined(*calls)
if !strings.HasPrefix(got, "resolvectl revert xray_tun") {
t.Errorf("expected the stale revert first, got %q", got)
}
if t1.systemDNSDirty {
t.Error("a successful retry should clear the dirty flag")
}
if !t1.systemDNSSet {
t.Error("the takeover should proceed once the retry succeeds")
}
}
func TestUnsetSystemDNSReverts(t *testing.T) {
stubDNSRouting(t, nil)
calls := recorder(t, "")
t1 := optedInTun()
if err := t1.ConfigureSystemDNS(context.Background(), "tun"); err != nil {
t.Fatalf("setup failed: %v", err)
}
*calls = nil
t1.unsetSystemDNS()
if t1.systemDNSSet {
t.Error("systemDNSSet should be false after unset")
}
if got := joined(*calls); got != "resolvectl revert xray_tun" {
t.Errorf("unset calls = %q, want %q", got, "resolvectl revert xray_tun")
}
t1.unsetSystemDNS()
if len(*calls) != 1 {
t.Errorf("unsetSystemDNS must be idempotent, got %q", joined(*calls))
}
}
func TestUnsetSystemDNSKeepsDirtyWhenRevertFails(t *testing.T) {
stubDNSRouting(t, nil)
calls := recorder(t, "revert")
t1 := optedInTun()
t1.systemDNSSet = true
t1.unsetSystemDNS()
if !t1.systemDNSDirty {
t.Error("a failed revert during unset must be remembered")
}
if got := joined(*calls); !strings.Contains(got, "resolvectl revert xray_tun") {
t.Errorf("expected a revert attempt, got %q", got)
}
}
func TestSystemDNSAddrs(t *testing.T) {
tests := []struct {
name string
gateway []string
wantSource string
wantDNS string
wantOK bool
}{
{
name: "ipv4 /30",
gateway: []string{"192.168.100.1/30"},
wantSource: "192.168.100.1",
wantDNS: "192.168.100.2",
wantOK: true,
},
{
name: "ipv4 /16",
gateway: []string{"10.0.0.1/16"},
wantSource: "10.0.0.1",
wantDNS: "10.0.0.2",
wantOK: true,
},
{
name: "first ipv4 wins",
gateway: []string{"fc00::1/64", "172.18.0.1/30"},
wantSource: "172.18.0.1",
wantDNS: "172.18.0.2",
wantOK: true,
},
{
name: "no gateway",
gateway: nil,
wantOK: false,
},
{
name: "ipv6 only",
gateway: []string{"fc00::1/64"},
wantSource: "fc00::1",
wantDNS: "fc00::2",
wantOK: true,
},
{
name: "first ipv6 without ipv4",
gateway: []string{"fc00::1/64", "fd00::1/64"},
wantSource: "fc00::1",
wantDNS: "fc00::2",
wantOK: true,
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
source, dnsAddr, ok := systemDNSAddrs(tt.gateway)
if ok != tt.wantOK {
t.Fatalf("ok = %v, want %v", ok, tt.wantOK)
}
if !tt.wantOK {
return
}
if source.String() != tt.wantSource {
t.Errorf("source = %q, want %q", source.String(), tt.wantSource)
}
if dnsAddr.String() != tt.wantDNS {
t.Errorf("dns = %q, want %q", dnsAddr.String(), tt.wantDNS)
}
})
}
}
func TestBuildResolvectlArgs(t *testing.T) {
tests := []struct {
name string
action string
iface string
extra []string
want []string
}{
{
name: "revert",
action: "revert",
iface: "xray_tun",
want: []string{"revert", "xray_tun"},
},
{
name: "dns single",
action: "dns",
iface: "xray_tun",
extra: []string{"192.168.100.2"},
want: []string{"dns", "xray_tun", "192.168.100.2"},
},
{
name: "dns multiple",
action: "dns",
iface: "xray_tun",
extra: []string{"192.168.100.2", "fc00::2"},
want: []string{"dns", "xray_tun", "192.168.100.2", "fc00::2"},
},
{
name: "domain wildcard",
action: "domain",
iface: "xray_tun",
extra: []string{"~."},
want: []string{"domain", "xray_tun", "~."},
},
{
name: "default-route",
action: "default-route",
iface: "xray_tun",
extra: []string{"true"},
want: []string{"default-route", "xray_tun", "true"},
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
got := buildResolvectlArgs(tt.action, tt.iface, tt.extra...)
if len(got) != len(tt.want) {
t.Fatalf("args = %v, want %v", got, tt.want)
}
for i := range got {
if got[i] != tt.want[i] {
t.Errorf("args[%d] = %q, want %q (full: %v)", i, got[i], tt.want[i], got)
}
}
})
}
}
+3 -235
View File
@@ -3,25 +3,17 @@
package tun
import (
"bytes"
"context"
"crypto/md5"
"encoding/binary"
go_errors "errors"
"net"
"net/netip"
"os/exec"
"path/filepath"
"slices"
"strconv"
"strings"
"sync"
"syscall"
"time"
"unsafe"
"github.com/xtls/xray-core/common/errors"
"github.com/xtls/xray-core/transport/internet"
"golang.org/x/sys/windows"
"golang.zx2c4.com/wintun"
"golang.zx2c4.com/wireguard/windows/tunnel/winipcfg"
@@ -46,10 +38,6 @@ type WindowsTun struct {
luid winipcfg.LUID
cbr winipcfg.ChangeCallback
cbi winipcfg.ChangeCallback
wfp windows.Handle
resolver *savedResolver
skipStop chan struct{}
skipDone chan struct{}
closed bool
}
@@ -90,13 +78,8 @@ func open(name, desc string) (*wintun.Adapter, error) {
// generate a deterministic GUID from the adapter name
id := md5.Sum([]byte(name))
guid := (*windows.GUID)(unsafe.Pointer(&id[0]))
// try to open existing adapter by name
adapter, err := wintun.OpenAdapter(name)
if err == nil {
return adapter, nil
}
// try to create adapter anew
adapter, err = wintun.CreateAdapter(name, desc, guid)
adapter, err := wintun.CreateAdapter(name, desc, guid)
if err == nil {
return adapter, nil
}
@@ -209,105 +192,19 @@ startOver:
}
}
// Windows lists the TUN's DNS servers among the system's ones, which Go's
// resolver queries for Xray's own lookups past the TUN, where they lead
// nowhere or back into Xray. Not skipped are those another interface uses
// as well, as that could leave no server at all. As those can change at
// any time, they are looked at again as often as Go rereads its servers.
if len(dns) > 0 {
skipped, err := tunOnlyDNS(t.luid, dns)
if err != nil {
skipped = dns
}
internet.SkipDNSServers(skipped)
t.skipStop, t.skipDone = make(chan struct{}), make(chan struct{})
go func() {
defer close(t.skipDone)
ticker := time.NewTicker(5 * time.Second)
defer ticker.Stop()
for {
select {
case <-ticker.C:
if skipped, err := tunOnlyDNS(t.luid, dns); err == nil {
internet.SkipDNSServers(skipped)
}
case <-t.skipStop:
return
}
}
}()
}
// Keep Windows from registering the TUN's addresses, and the host name
// with them, through dynamic DNS updates. Best effort.
if address4 || address6 {
if err := disableDNSRegistration(t.luid, dns); err != nil {
errors.LogDebugInner(context.Background(), err, "[tun] unable to disable DNS registration")
}
}
// With autoSystemWfpBlockLeak, once the system routes lead to the TUN,
// keep DNS ("dns", if dns is set), and an IP version no route of which
// leads to the TUN ("misconfigtun"), from leaving through the other
// interfaces. Addresses do not matter: without one of a version in
// gateway, Windows gives the TUN a link-local one.
leaks := t.options.AutoSystemWfpBlockLeak
blockDNS := slices.Contains(leaks, "dns") && len(dns) > 0
blockIPv4 := slices.Contains(leaks, "misconfigtun") && !route4
blockIPv6 := slices.Contains(leaks, "misconfigtun") && !route6
if (route4 || route6) && (blockDNS || blockIPv4 || blockIPv6) {
if t.wfp, err = blockLeaks(t.luid, blockDNS, blockIPv4, blockIPv6); err != nil {
var blocked []string
for _, b := range []struct {
on bool
what string
}{{blockDNS, "DNS"}, {blockIPv4, "IPv4"}, {blockIPv6, "IPv6"}} {
if b.on {
blocked = append(blocked, b.what)
}
}
// Rather no TUN than a leaking one.
return errors.New("unable to block ", strings.Join(blocked, " and "), " outside the TUN (remove autoSystemWfpBlockLeak to run without)").Base(err)
}
errors.LogInfo(context.Background(), "[tun] outside the TUN, blocked DNS: ", blockDNS, ", blocked IPv4: ", blockIPv4, ", blocked IPv6: ", blockIPv6)
if blockDNS {
covered := slices.Clone(addresses)
for _, route := range routesData {
covered = append(covered, route.Destination)
}
for _, server := range dnsOutsideTUN(dns, covered) {
errors.LogWarning(context.Background(), "[tun] DNS server ", server, " is in neither gateway nor autoSystemRoutingTable, so queries to it cannot go through the TUN and are blocked")
}
// With updater, the dialer controllers bind Xray's own sockets
// to the physical interface.
if updater != nil {
t.resolver = resolveOnOwn()
}
}
}
if len(dns) > 0 || route4 || route6 {
if err := flushDNSCache(); err != nil {
errors.LogInfoInner(context.Background(), err, "[tun] unable to flush DNS cache")
}
}
if updater != nil {
// 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) {
t.cbr, err = winipcfg.RegisterRouteChangeCallback(func(notificationType winipcfg.MibNotificationType, route *winipcfg.MibIPforwardRow2) {
updater.Update()
})
if err != nil {
return err
}
t.cbr = cbr
cbi, err := winipcfg.RegisterInterfaceChangeCallback(func(notificationType winipcfg.MibNotificationType, iface *winipcfg.MibIPInterfaceRow) {
t.cbi, err = winipcfg.RegisterInterfaceChangeCallback(func(notificationType winipcfg.MibNotificationType, iface *winipcfg.MibIPInterfaceRow) {
updater.Update()
})
if err != nil {
return err
}
t.cbi = cbi
}
return nil
}
@@ -334,20 +231,6 @@ func (t *WindowsTun) Close() error {
t.luid.FlushIPAddresses(windows.AF_INET6)
t.luid.FlushDNS(windows.AF_INET6)
}
if t.wfp != 0 {
closeWFPEngine(t.wfp)
}
if t.resolver != nil {
t.resolver.restore()
}
if t.skipStop != nil {
close(t.skipStop)
<-t.skipDone
}
internet.SkipDNSServers(nil)
if len(t.options.DNS) > 0 || len(t.options.AutoSystemRoutingTable) > 0 {
flushDNSCache()
}
if t.session != (wintun.Session{}) {
t.session.End()
}
@@ -357,121 +240,6 @@ func (t *WindowsTun) Close() error {
return nil
}
type savedResolver struct {
preferGo bool
dial func(ctx context.Context, network, address string) (net.Conn, error)
}
// resolveOnOwn has Go resolve the names Xray would otherwise ask Windows for,
// on Xray's own sockets, which the dialer controllers bind to the physical
// interface, and skipping the TUN's DNS servers, as localdns does. Windows'
// resolver runs in the DNS Client service, whose queries the DNS filter lets
// through the TUN only, so Xray's own lookups, like of an outbound's server
// domain, would go into Xray again and could end up waiting on themselves.
//
// It changes net.DefaultResolver for the whole process, which covers every
// lookup that would reach Windows' resolver; restore undoes it.
func resolveOnOwn() *savedResolver {
saved := &savedResolver{net.DefaultResolver.PreferGo, net.DefaultResolver.Dial}
dialer := &net.Dialer{Control: func(network, address string, c syscall.RawConn) error {
for _, ctl := range internet.Controllers {
if err := ctl(network, address, c); err != nil {
return err
}
}
return nil
}}
// Go's resolver moves on to the next server right away when a dial fails.
net.DefaultResolver.Dial = func(ctx context.Context, network, address string) (net.Conn, error) {
if internet.IsSkippedDNSServer(address) {
return nil, errors.New("skipped DNS server ", address)
}
return dialer.DialContext(ctx, network, address)
}
net.DefaultResolver.PreferGo = true
return saved
}
func (s *savedResolver) restore() {
net.DefaultResolver.PreferGo = s.preferGo
net.DefaultResolver.Dial = s.dial
}
// tunOnlyDNS returns those of servers, the TUN's DNS servers, that Go's
// resolver does not also get from another interface: one that is up and has
// a gateway, as it reads them.
func tunOnlyDNS(tun winipcfg.LUID, servers []netip.Addr) ([]netip.Addr, error) {
adapters, err := winipcfg.GetAdaptersAddresses(windows.AF_UNSPEC, winipcfg.GAAFlagIncludeGateways)
if err != nil {
return nil, err
}
var others []netip.Addr
for _, adapter := range adapters {
if adapter.LUID == tun || adapter.OperStatus != winipcfg.IfOperStatusUp || adapter.FirstGatewayAddress == nil {
continue
}
for server := adapter.FirstDNSServerAddress; server != nil; server = server.Next {
if addr, ok := netip.AddrFromSlice(server.Address.IP()); ok {
others = append(others, addr.Unmap())
}
}
}
return slices.DeleteFunc(slices.Clone(servers), func(server netip.Addr) bool {
return slices.Contains(others, server.Unmap())
}), nil
}
// disableDNSRegistration turns off the dynamic DNS registration of the
// interface's addresses. dns are its DNS servers.
func disableDNSRegistration(luid winipcfg.LUID, dns []netip.Addr) error {
guid, err := luid.GUID()
if err != nil {
return err
}
err = winipcfg.SetInterfaceDnsSettings(*guid, &winipcfg.DnsInterfaceSettings{
Version: winipcfg.DnsInterfaceSettingsVersion1,
Flags: winipcfg.DnsInterfaceSettingsFlagRegistrationEnabled,
})
if err == nil || !go_errors.Is(err, windows.ERROR_PROC_NOT_FOUND) {
return err
}
return disableDNSRegistrationByNetsh(luid, dns)
}
// disableDNSRegistrationByNetsh does it for Windows before 10 1809, which
// lacks SetInterfaceDnsSettings. The setting is the interface's, not the
// address family's, but netsh only applies it along with a DNS server, which
// replaces the IPv4 ones, so they are set again afterwards.
func disableDNSRegistrationByNetsh(luid winipcfg.LUID, dns []netip.Addr) error {
row, err := luid.Interface()
if err != nil {
return err
}
server := "127.0.0.1" // any will do when there is no IPv4 one
if i := slices.IndexFunc(dns, netip.Addr.Is4); i >= 0 {
server = dns[i].String()
}
err = runNetsh("interface", "ipv4", "set", "dnsservers", "name="+strconv.FormatUint(uint64(row.InterfaceIndex), 10), "source=static", "address="+server, "register=none", "validate=no")
return errors.Combine(err, luid.SetDNS(windows.AF_INET, dns, nil))
}
// runNetsh runs netsh.exe from the system directory. netsh reports some
// failures, like a syntax error, only in its output, even with exit code 0,
// so any output counts as a failure.
func runNetsh(args ...string) error {
system32, err := windows.GetSystemDirectory()
if err != nil {
return err
}
cmd := exec.Command(filepath.Join(system32, "netsh.exe"), args...)
cmd.SysProcAttr = &syscall.SysProcAttr{HideWindow: true}
output, err := cmd.CombinedOutput()
if output = bytes.TrimSpace(output); err != nil || len(output) > 0 {
return errors.New("netsh ", strings.Join(args, " "), ": ", string(output)).Base(err)
}
return nil
}
func (t *WindowsTun) Name() (string, error) {
row, err := t.luid.Interface()
if err != nil {
-471
View File
@@ -1,471 +0,0 @@
//go:build windows
package tun
import (
"net/netip"
"os"
"runtime"
"slices"
"unsafe"
"github.com/xtls/xray-core/common/errors"
"golang.org/x/sys/windows"
"golang.zx2c4.com/wireguard/windows/tunnel/winipcfg"
)
var (
modfwpuclnt = windows.NewLazySystemDLL("fwpuclnt.dll")
moddnsapi = windows.NewLazySystemDLL("dnsapi.dll")
procFwpmEngineOpen0 = modfwpuclnt.NewProc("FwpmEngineOpen0")
procFwpmEngineClose0 = modfwpuclnt.NewProc("FwpmEngineClose0")
procFwpmTransactionBegin0 = modfwpuclnt.NewProc("FwpmTransactionBegin0")
procFwpmTransactionCommit0 = modfwpuclnt.NewProc("FwpmTransactionCommit0")
procFwpmTransactionAbort0 = modfwpuclnt.NewProc("FwpmTransactionAbort0")
procFwpmSubLayerAdd0 = modfwpuclnt.NewProc("FwpmSubLayerAdd0")
procFwpmFilterAdd0 = modfwpuclnt.NewProc("FwpmFilterAdd0")
procFwpmGetAppIdFromFileName0 = modfwpuclnt.NewProc("FwpmGetAppIdFromFileName0")
procFwpmFreeMemory0 = modfwpuclnt.NewProc("FwpmFreeMemory0")
procDnsFlushResolverCache = moddnsapi.NewProc("DnsFlushResolverCache")
)
// fwptypes.h and fwpmtypes.h
const (
rpcCAuthnWinNT = 10 // RPC_C_AUTHN_WINNT
fwpmSessionFlagDynamic = 1 // FWPM_SESSION_FLAG_DYNAMIC
fwpmFilterFlagClearActionRight = 8 // FWPM_FILTER_FLAG_CLEAR_ACTION_RIGHT
fwpUint8 = 1 // FWP_UINT8
fwpUint16 = 2 // FWP_UINT16
fwpUint32 = 3 // FWP_UINT32
fwpUint64 = 4 // FWP_UINT64
fwpByteArray16Type = 11 // FWP_BYTE_ARRAY16_TYPE
fwpByteBlobType = 12 // FWP_BYTE_BLOB_TYPE
fwpSecurityDescriptorType = 14 // FWP_SECURITY_DESCRIPTOR_TYPE
fwpMatchEqual = 0 // FWP_MATCH_EQUAL
fwpMatchFlagsAllSet = 6 // FWP_MATCH_FLAGS_ALL_SET
fwpConditionFlagIsLoopback = 1 // FWP_CONDITION_FLAG_IS_LOOPBACK
fwpActionBlock = 0x1001 // FWP_ACTION_BLOCK
fwpActionPermit = 0x1002 // FWP_ACTION_PERMIT
)
// fwpmu.h
var (
fwpmLayerALEAuthConnectV4 = windows.GUID{Data1: 0xc38d57d1, Data2: 0x05a7, Data3: 0x4c33, Data4: [8]byte{0x90, 0x4f, 0x7f, 0xbc, 0xee, 0xe6, 0x0e, 0x82}}
fwpmLayerALEAuthConnectV6 = windows.GUID{Data1: 0x4a72393b, Data2: 0x319f, Data3: 0x44bc, Data4: [8]byte{0x84, 0xc3, 0xba, 0x54, 0xdc, 0xb3, 0xb6, 0xb4}}
fwpmLayerALEAuthRecvAcceptV4 = windows.GUID{Data1: 0xe1cd9fe7, Data2: 0xf4b5, Data3: 0x4273, Data4: [8]byte{0x96, 0xc0, 0x59, 0x2e, 0x48, 0x7b, 0x86, 0x50}}
fwpmLayerALEAuthRecvAcceptV6 = windows.GUID{Data1: 0xa3b42c97, Data2: 0x9f04, Data3: 0x4672, Data4: [8]byte{0xb8, 0x7e, 0xce, 0xe9, 0xc4, 0x83, 0x25, 0x7f}}
fwpmConditionFlags = windows.GUID{Data1: 0x632ce23b, Data2: 0x5167, Data3: 0x435c, Data4: [8]byte{0x86, 0xd7, 0xe9, 0x03, 0x68, 0x4a, 0xa8, 0x0c}}
fwpmConditionIPArrivalInterface = windows.GUID{Data1: 0x618a9b6d, Data2: 0x386b, Data3: 0x4136, Data4: [8]byte{0xad, 0x6e, 0xb5, 0x15, 0x87, 0xcf, 0xb1, 0xcd}}
fwpmConditionIPLocalInterface = windows.GUID{Data1: 0x4cd62a49, Data2: 0x59c3, Data3: 0x4969, Data4: [8]byte{0xb7, 0xf3, 0xbd, 0xa5, 0xd3, 0x28, 0x90, 0xa4}}
fwpmConditionIPLocalPort = windows.GUID{Data1: 0x0c1ba1af, Data2: 0x5765, Data3: 0x453f, Data4: [8]byte{0xaf, 0x22, 0xa8, 0xf7, 0x91, 0xac, 0x77, 0x5b}} // also FWPM_CONDITION_ICMP_TYPE
fwpmConditionIPNexthopInterface = windows.GUID{Data1: 0x93ae8f5b, Data2: 0x7f6f, Data3: 0x4719, Data4: [8]byte{0x98, 0xc8, 0x14, 0xe9, 0x74, 0x29, 0xef, 0x04}}
fwpmConditionIPProtocol = windows.GUID{Data1: 0x3971ef2b, Data2: 0x623e, Data3: 0x4f9a, Data4: [8]byte{0x8c, 0xb1, 0x6e, 0x79, 0xb8, 0x06, 0xb9, 0xa7}}
fwpmConditionIPRemoteAddress = windows.GUID{Data1: 0xb235ae9a, Data2: 0x1d64, Data3: 0x49b8, Data4: [8]byte{0xa4, 0x4c, 0x5f, 0xf3, 0xd9, 0x09, 0x50, 0x45}}
fwpmConditionIPRemotePort = windows.GUID{Data1: 0xc35a604d, Data2: 0xd22b, Data3: 0x4e1a, Data4: [8]byte{0x91, 0xb4, 0x68, 0xf6, 0x74, 0xee, 0x67, 0x4b}} // also FWPM_CONDITION_ICMP_CODE
fwpmConditionALEAppID = windows.GUID{Data1: 0xd78e1e87, Data2: 0x8644, Data3: 0x4ea5, Data4: [8]byte{0x94, 0x37, 0xd8, 0x09, 0xec, 0xef, 0xc9, 0x71}}
fwpmConditionALEUserID = windows.GUID{Data1: 0xaf043a0a, Data2: 0xb34d, Data3: 0x4f86, Data4: [8]byte{0x97, 0x9c, 0xc9, 0x03, 0x71, 0xaf, 0x6e, 0x66}}
)
// dnsClientSID is the SID of Windows' DNS Client service, NT SERVICE\Dnscache.
// Service SIDs derive from the service name, so it is the same everywhere (sc
// showsid dnscache).
const dnsClientSID = "S-1-5-80-859482183-879914841-863379149-1145462774-2388618682"
// ff02::1:2, where DHCPv6 clients send to. A package-level variable never
// moves, so conditions may refer to it through uintptr.
var ipv6AllDHCPv6Servers = [16]byte{0xff, 0x02, 13: 0x01, 15: 0x02}
type fwpByteBlob struct {
size uint32
data *byte
}
// fwpValue0 is FWP_VALUE0 as well as FWP_CONDITION_VALUE0. Their union holds
// a scalar of at most 32 bits, or a pointer for the larger types.
type fwpValue0 struct {
typ uint32
value uintptr
}
type fwpmDisplayData0 struct {
name *uint16
description *uint16
}
type fwpmSession0 struct {
sessionKey windows.GUID
displayData fwpmDisplayData0
flags uint32
txnWaitTimeoutInMSec uint32
processID uint32
sid *windows.SID
username *uint16
kernelMode int32
}
type fwpmSublayer0 struct {
subLayerKey windows.GUID
displayData fwpmDisplayData0
flags uint32
providerKey *windows.GUID
providerData fwpByteBlob
weight uint16
}
type fwpmFilterCondition0 struct {
fieldKey windows.GUID
matchType uint32
conditionValue fwpValue0
}
type fwpmAction0 struct {
typ uint32
filterType windows.GUID
}
type fwpmFilter0 struct {
filterKey windows.GUID
displayData fwpmDisplayData0
flags uint32
providerKey *windows.GUID
providerData fwpByteBlob
layerKey windows.GUID
subLayerKey windows.GUID
weight fwpValue0
numFilterConditions uint32
filterCondition *fwpmFilterCondition0
action fwpmAction0
_ uint32 // C aligns the following union to 8 bytes, as it holds a UINT64
providerContextKey windows.GUID
reserved *windows.GUID
_ [8 - unsafe.Sizeof(uintptr(0))]byte // and filterId as well, also on 32-bit
filterID uint64
effectiveWeight fwpValue0
}
// fwpmResult converts the DWORD status the Fwpm functions return.
func fwpmResult(r1, _ uintptr, _ error) error {
if r1 != 0 {
return windows.Errno(r1)
}
return nil
}
func utf16Ptr(s string) *uint16 {
p, _ := windows.UTF16PtrFromString(s)
return p
}
func condition(field *windows.GUID, typ uint32, value uintptr) fwpmFilterCondition0 {
return fwpmFilterCondition0{
fieldKey: *field,
matchType: fwpMatchEqual,
conditionValue: fwpValue0{typ: typ, value: value},
}
}
// blockLeaks keeps traffic from leaving through interfaces other than tun,
// for every program but Xray itself, whose outbounds (DNS included) use the
// other interfaces on purpose:
//
// - dns: DNS (port 53) may only go through the TUN. Windows sends a name
// query to the DNS servers of all interfaces, not only to those of the TUN:
// to the first server of each interface, then to all of them when no answer
// arrives within a second or two. It sends the queries for the servers of
// an interface out through that interface, whatever the routes say, and
// other programs reach an on-link resolver, like 192.168.1.1 from DHCP,
// through its LAN route, which is more specific than the TUN's default
// route. Since Windows 11 and Server 2022, Windows may also send its
// queries over HTTPS or TLS, so there its DNS Client service may not
// connect outside the TUN at all, except for name resolution on the local
// link (mDNS, LLMNR).
// - ipv4, ipv6: no IPv4, or no IPv6, at all, in either direction, for a TUN
// that no route of it leads to, except loopback and what Windows itself
// needs on the local link (DHCP, and for IPv6 neighbor and multicast
// listener discovery), none of which can leave it. The TUN carries what
// is routed to it even without an address of that IP version in gateway:
// Windows gives it link-local ones itself, an IPv6 one at once, an IPv4
// one from 169.254.0.0/16 after some seconds (until then, IPv4 routed to
// the TUN is unreachable).
//
// The filters live in a dynamic WFP session: closing the returned engine handle
// with closeWFPEngine deletes them, and so does Windows when the process dies.
func blockLeaks(tun winipcfg.LUID, dns, ipv4, ipv6 bool) (windows.Handle, error) {
engine, err := openWFPEngine()
if err != nil {
return 0, err
}
if err := fwpmResult(procFwpmTransactionBegin0.Call(uintptr(engine), 0)); err != nil {
closeWFPEngine(engine)
return 0, errors.New("FwpmTransactionBegin0 failed").Base(err)
}
err = addLeakFilters(engine, tun, dns, ipv4, ipv6)
if err == nil {
if err = fwpmResult(procFwpmTransactionCommit0.Call(uintptr(engine))); err != nil {
err = errors.New("FwpmTransactionCommit0 failed").Base(err)
}
}
if err != nil {
procFwpmTransactionAbort0.Call(uintptr(engine))
closeWFPEngine(engine)
return 0, err
}
return engine, nil
}
func openWFPEngine() (windows.Handle, error) {
if err := modfwpuclnt.Load(); err != nil {
return 0, err
}
// txnWaitTimeoutInMSec stays 0 for BFE's default, so that a transaction
// held by another program cannot hang the start forever.
session := fwpmSession0{
displayData: fwpmDisplayData0{name: utf16Ptr("Xray TUN")},
flags: fwpmSessionFlagDynamic,
}
var engine windows.Handle
if err := fwpmResult(procFwpmEngineOpen0.Call(0, rpcCAuthnWinNT, 0, uintptr(unsafe.Pointer(&session)), uintptr(unsafe.Pointer(&engine)))); err != nil {
return 0, errors.New("FwpmEngineOpen0 failed").Base(err)
}
return engine, nil
}
func closeWFPEngine(engine windows.Handle) {
procFwpmEngineClose0.Call(uintptr(engine))
}
// addLeakFilters adds the filters of blockLeaks in a sublayer of their own.
// blockLeaks runs it in a transaction, so that they take effect all at once.
func addLeakFilters(engine windows.Handle, tun winipcfg.LUID, dns, ipv4, ipv6 bool) error {
exe, err := os.Executable()
if err != nil {
return err
}
exePath, err := windows.UTF16PtrFromString(exe)
if err != nil {
return err
}
var appID *fwpByteBlob
if err := fwpmResult(procFwpmGetAppIdFromFileName0.Call(uintptr(unsafe.Pointer(exePath)), uintptr(unsafe.Pointer(&appID)))); err != nil {
return errors.New("FwpmGetAppIdFromFileName0 failed for ", exe).Base(err)
}
defer func() { procFwpmFreeMemory0.Call(uintptr(unsafe.Pointer(&appID))) }()
sublayer := fwpmSublayer0{
displayData: fwpmDisplayData0{name: utf16Ptr("Xray TUN")},
weight: 0xffff,
}
if sublayer.subLayerKey, err = windows.GenerateGUID(); err != nil {
return err
}
if err := fwpmResult(procFwpmSubLayerAdd0.Call(uintptr(engine), uintptr(unsafe.Pointer(&sublayer)), 0)); err != nil {
return errors.New("FwpmSubLayerAdd0 failed").Base(err)
}
add := func(layer *windows.GUID, name string, flags, action uint32, weight uint8, conditions ...fwpmFilterCondition0) error {
return addFilter(engine, &sublayer.subLayerKey, layer, "Xray TUN: "+name, flags, action, weight, conditions...)
}
var pinner runtime.Pinner
defer pinner.Unpin()
tunLUID := new(uint64)
*tunLUID = uint64(tun)
pinner.Pin(tunLUID) // the condition only holds it as uintptr
// The heaviest matching filter of a sublayer decides. All sublayers have
// their say, though, and a block in any of them beats a permit, unless
// the permit is hard: it clears the action right, and then the blocks of
// lower sublayers, Windows Firewall rules among them, no longer override
// it, only a callout's veto does. Xray's own connections out get such a
// hard permit. Connections from outside to Xray get an ordinary one, so
// that firewalls keep guarding its inbounds.
self := condition(&fwpmConditionALEAppID, fwpByteBlobType, uintptr(unsafe.Pointer(appID)))
dns53 := condition(&fwpmConditionIPRemotePort, fwpUint16, 53)
// DNS goes through the TUN when its local address is the TUN's, and it
// also leaves, or arrives, through the TUN. The local address alone
// decides by default, but with weak host sending or receiving enabled,
// packets of the TUN's address can use other interfaces. (The next hop,
// the interface replies would leave by, is not known for arriving ones.)
onTUN := func(field *windows.GUID) fwpmFilterCondition0 {
return condition(field, fwpUint64, uintptr(unsafe.Pointer(tunLUID)))
}
out := []fwpmFilterCondition0{dns53, onTUN(&fwpmConditionIPLocalInterface), onTUN(&fwpmConditionIPNexthopInterface)}
in := []fwpmFilterCondition0{dns53, onTUN(&fwpmConditionIPLocalInterface), onTUN(&fwpmConditionIPArrivalInterface)}
for _, layer := range []struct {
key *windows.GUID
selfFlags uint32
throughTUN []fwpmFilterCondition0
}{
{&fwpmLayerALEAuthConnectV4, fwpmFilterFlagClearActionRight, out},
{&fwpmLayerALEAuthRecvAcceptV4, 0, in},
{&fwpmLayerALEAuthConnectV6, fwpmFilterFlagClearActionRight, out},
{&fwpmLayerALEAuthRecvAcceptV6, 0, in},
} {
if err := add(layer.key, "permit Xray", layer.selfFlags, fwpActionPermit, 4, self); err != nil {
return err
}
if dns {
if err := add(layer.key, "permit DNS through the TUN", 0, fwpActionPermit, 3, layer.throughTUN...); err != nil {
return err
}
if err := add(layer.key, "block DNS", 0, fwpActionBlock, 2, dns53); err != nil {
return err
}
}
}
// Since Windows 11 and Server 2022 (build 20348), the DNS Client service
// may also send the queries for an interface's servers over HTTPS or TLS,
// out through that interface and to any port. So there it may only
// connect through the TUN, except for mDNS and LLMNR, which stay on the
// local link (over an IP version only while it is not blocked altogether).
// Earlier versions only query port 53, and may run the service in one
// process with others, which the filters would catch as well. Like
// Windows Firewall's rules for it, they recognize the service by its SID,
// which Windows puts in the token of its process: the security descriptor
// grants that SID the right to match (FWP_ACTRL_MATCH_FILTER, CC in SDDL).
if _, _, build := windows.RtlGetNtVersionNumbers(); dns && build >= 20348 {
sd, err := windows.SecurityDescriptorFromString("O:SYG:SYD:(A;;CCRC;;;" + dnsClientSID + ")")
if err != nil {
return err
}
sdBlob := &fwpByteBlob{size: sd.Length(), data: (*byte)(unsafe.Pointer(sd))}
pinner.Pin(sdBlob) // the condition only holds it as uintptr
dnsClient := condition(&fwpmConditionALEUserID, fwpSecurityDescriptorType, uintptr(unsafe.Pointer(sdBlob)))
// Conditions on the same field match when any of them does.
mdnsLLMNR := []fwpmFilterCondition0{dnsClient, condition(&fwpmConditionIPRemotePort, fwpUint16, 5353), condition(&fwpmConditionIPRemotePort, fwpUint16, 5355)}
for _, layer := range []struct {
key *windows.GUID
localLink bool
}{
{&fwpmLayerALEAuthConnectV4, !ipv4},
{&fwpmLayerALEAuthConnectV6, !ipv6},
} {
if err := add(layer.key, "permit the DNS Client service through the TUN", 0, fwpActionPermit, 3, dnsClient, onTUN(&fwpmConditionIPLocalInterface), onTUN(&fwpmConditionIPNexthopInterface)); err != nil {
return err
}
if layer.localLink {
if err := add(layer.key, "permit the DNS Client service's mDNS and LLMNR", 0, fwpActionPermit, 3, mdnsLLMNR...); err != nil {
return err
}
}
if err := add(layer.key, "block the DNS Client service", 0, fwpActionBlock, 2, dnsClient); err != nil {
return err
}
}
}
// Both directions: replies to a connection accepted from outside would
// leave through the physical link as well.
loopback := fwpmFilterCondition0{
fieldKey: fwpmConditionFlags,
matchType: fwpMatchFlagsAllSet,
conditionValue: fwpValue0{typ: fwpUint32, value: fwpConditionFlagIsLoopback},
}
if ipv4 {
// DHCP keeps the addresses of the other interfaces, which Xray's own
// connections use.
dhcp := []fwpmFilterCondition0{
condition(&fwpmConditionIPProtocol, fwpUint8, windows.IPPROTO_UDP),
condition(&fwpmConditionIPLocalPort, fwpUint16, 68),
condition(&fwpmConditionIPRemotePort, fwpUint16, 67),
}
for _, layer := range []*windows.GUID{&fwpmLayerALEAuthConnectV4, &fwpmLayerALEAuthRecvAcceptV4} {
if err := add(layer, "permit IPv4 loopback", 0, fwpActionPermit, 1, loopback); err != nil {
return err
}
if err := add(layer, "permit DHCP", 0, fwpActionPermit, 1, dhcp...); err != nil {
return err
}
if err := add(layer, "block IPv4", 0, fwpActionBlock, 0); err != nil {
return err
}
}
}
if ipv6 {
// Neighbor and multicast listener discovery, ICMPv6 130-137 and 143,
// whose type and code sit where the local and remote port are.
discovery := []fwpmFilterCondition0{condition(&fwpmConditionIPProtocol, fwpUint8, windows.IPPROTO_ICMPV6)}
for _, typ := range []uintptr{130, 131, 132, 133, 134, 135, 136, 137, 143} {
discovery = append(discovery, condition(&fwpmConditionIPLocalPort, fwpUint16, typ))
}
discovery = append(discovery, condition(&fwpmConditionIPRemotePort, fwpUint16, 0))
dhcpv6 := []fwpmFilterCondition0{
condition(&fwpmConditionIPProtocol, fwpUint8, windows.IPPROTO_UDP),
condition(&fwpmConditionIPLocalPort, fwpUint16, 546),
condition(&fwpmConditionIPRemotePort, fwpUint16, 547),
}
for _, direction := range []struct {
layer *windows.GUID
dhcpv6 []fwpmFilterCondition0
}{
// The client sends to the servers' multicast address, and they
// answer from their own.
{&fwpmLayerALEAuthConnectV6, slices.Concat(dhcpv6, []fwpmFilterCondition0{condition(&fwpmConditionIPRemoteAddress, fwpByteArray16Type, uintptr(unsafe.Pointer(&ipv6AllDHCPv6Servers)))})},
{&fwpmLayerALEAuthRecvAcceptV6, dhcpv6},
} {
if err := add(direction.layer, "permit IPv6 loopback", 0, fwpActionPermit, 1, loopback); err != nil {
return err
}
if err := add(direction.layer, "permit IPv6 neighbor and multicast listener discovery", 0, fwpActionPermit, 1, discovery...); err != nil {
return err
}
if err := add(direction.layer, "permit DHCPv6", 0, fwpActionPermit, 1, direction.dhcpv6...); err != nil {
return err
}
if err := add(direction.layer, "block IPv6", 0, fwpActionBlock, 0); err != nil {
return err
}
}
}
return nil
}
func addFilter(engine windows.Handle, sublayer, layer *windows.GUID, name string, flags, action uint32, weight uint8, conditions ...fwpmFilterCondition0) error {
filter := fwpmFilter0{
displayData: fwpmDisplayData0{name: utf16Ptr(name)},
flags: flags,
layerKey: *layer,
subLayerKey: *sublayer,
weight: fwpValue0{typ: fwpUint8, value: uintptr(weight)},
numFilterConditions: uint32(len(conditions)),
action: fwpmAction0{typ: action},
}
if len(conditions) > 0 {
filter.filterCondition = &conditions[0]
}
if err := fwpmResult(procFwpmFilterAdd0.Call(uintptr(engine), uintptr(unsafe.Pointer(&filter)), 0, 0)); err != nil {
return errors.New("FwpmFilterAdd0 failed for ", name).Base(err)
}
return nil
}
// dnsOutsideTUN returns the servers outside all of prefixes, the TUN's own
// subnets and routes: queries to them cannot go through the TUN.
func dnsOutsideTUN(servers []netip.Addr, prefixes []netip.Prefix) []netip.Addr {
var outside []netip.Addr
for _, server := range servers {
server = server.Unmap()
if !slices.ContainsFunc(prefixes, func(p netip.Prefix) bool { return p.Contains(server) }) {
outside = append(outside, server)
}
}
return outside
}
// flushDNSCache drops the answers Windows cached so far, like ipconfig
// /flushdns, so that names get resolved again with the current DNS setup.
func flushDNSCache() error {
if err := procDnsFlushResolverCache.Find(); err != nil {
return err
}
if r, _, err := procDnsFlushResolverCache.Call(); r == 0 {
return err
}
return nil
}
-206
View File
@@ -1,206 +0,0 @@
//go:build windows
package tun
import (
"context"
go_errors "errors"
"net"
"net/netip"
"slices"
"testing"
"unsafe"
"github.com/xtls/xray-core/transport/internet"
"golang.org/x/sys/windows"
"golang.zx2c4.com/wireguard/windows/tunnel/winipcfg"
)
// The WFP structures are handed to fwpuclnt.dll as they are, so their layout
// has to match what MSVC produces for 64-bit and for 32-bit Windows.
func TestWFPStructLayout(t *testing.T) {
check := func(name string, got, want64, want32 []uintptr) {
t.Helper()
want := want32
if unsafe.Sizeof(uintptr(0)) == 8 {
want = want64
}
if !slices.Equal(got, want) {
t.Errorf("%s: size and offsets are %v, want %v", name, got, want)
}
}
var blob fwpByteBlob
check("FWP_BYTE_BLOB",
[]uintptr{unsafe.Sizeof(blob), unsafe.Offsetof(blob.data)},
[]uintptr{16, 8}, []uintptr{8, 4})
var value fwpValue0
check("FWP_VALUE0",
[]uintptr{unsafe.Sizeof(value), unsafe.Offsetof(value.value)},
[]uintptr{16, 8}, []uintptr{8, 4})
var display fwpmDisplayData0
check("FWPM_DISPLAY_DATA0",
[]uintptr{unsafe.Sizeof(display), unsafe.Offsetof(display.description)},
[]uintptr{16, 8}, []uintptr{8, 4})
var action fwpmAction0
check("FWPM_ACTION0",
[]uintptr{unsafe.Sizeof(action), unsafe.Offsetof(action.filterType)},
[]uintptr{20, 4}, []uintptr{20, 4})
var cond fwpmFilterCondition0
check("FWPM_FILTER_CONDITION0",
[]uintptr{unsafe.Sizeof(cond), unsafe.Offsetof(cond.matchType), unsafe.Offsetof(cond.conditionValue)},
[]uintptr{40, 16, 24}, []uintptr{28, 16, 20})
var session fwpmSession0
check("FWPM_SESSION0",
[]uintptr{
unsafe.Sizeof(session), unsafe.Offsetof(session.displayData), unsafe.Offsetof(session.flags),
unsafe.Offsetof(session.txnWaitTimeoutInMSec), unsafe.Offsetof(session.processID), unsafe.Offsetof(session.sid),
unsafe.Offsetof(session.username), unsafe.Offsetof(session.kernelMode),
},
[]uintptr{72, 16, 32, 36, 40, 48, 56, 64},
[]uintptr{48, 16, 24, 28, 32, 36, 40, 44})
var sublayer fwpmSublayer0
check("FWPM_SUBLAYER0",
[]uintptr{
unsafe.Sizeof(sublayer), unsafe.Offsetof(sublayer.displayData), unsafe.Offsetof(sublayer.flags),
unsafe.Offsetof(sublayer.providerKey), unsafe.Offsetof(sublayer.providerData), unsafe.Offsetof(sublayer.weight),
},
[]uintptr{72, 16, 32, 40, 48, 64},
[]uintptr{44, 16, 24, 28, 32, 40})
var filter fwpmFilter0
check("FWPM_FILTER0",
[]uintptr{
unsafe.Sizeof(filter), unsafe.Offsetof(filter.displayData), unsafe.Offsetof(filter.flags),
unsafe.Offsetof(filter.providerKey), unsafe.Offsetof(filter.providerData), unsafe.Offsetof(filter.layerKey),
unsafe.Offsetof(filter.subLayerKey), unsafe.Offsetof(filter.weight), unsafe.Offsetof(filter.numFilterConditions),
unsafe.Offsetof(filter.filterCondition), unsafe.Offsetof(filter.action), unsafe.Offsetof(filter.providerContextKey),
unsafe.Offsetof(filter.reserved), unsafe.Offsetof(filter.filterID), unsafe.Offsetof(filter.effectiveWeight),
},
[]uintptr{200, 16, 32, 40, 48, 64, 80, 96, 112, 120, 128, 152, 168, 176, 184},
[]uintptr{152, 16, 24, 28, 32, 40, 56, 72, 80, 84, 88, 112, 128, 136, 144})
}
// TestLeakFiltersAccepted has WFP validate the filters by adding them inside a
// transaction that is then aborted, which leaves the system untouched. Adding
// filters requires an elevated process.
func TestLeakFiltersAccepted(t *testing.T) {
skipUnlessElevated := func(err error) {
t.Helper()
if go_errors.Is(err, windows.ERROR_ACCESS_DENIED) {
t.Skipf("WFP filters can only be added by an elevated process: %v", err)
}
t.Fatal(err)
}
engine, err := openWFPEngine()
if err != nil {
skipUnlessElevated(err)
}
defer closeWFPEngine(engine)
if err := fwpmResult(procFwpmTransactionBegin0.Call(uintptr(engine), 0)); err != nil {
skipUnlessElevated(err)
}
defer procFwpmTransactionAbort0.Call(uintptr(engine))
// Any interface stands in for the TUN; the loopback one always exists.
loopback, err := winipcfg.LUIDFromIndex(1)
if err != nil {
t.Fatal(err)
}
if err := addLeakFilters(engine, loopback, true, true, true); err != nil {
skipUnlessElevated(err)
}
}
func TestDNSClientSID(t *testing.T) {
sid, _, _, err := windows.LookupSID("", `NT SERVICE\Dnscache`)
if err != nil {
t.Fatal(err)
}
if sid.String() != dnsClientSID {
t.Errorf(`NT SERVICE\Dnscache is %v, not %v`, sid, dnsClientSID)
}
}
func TestDNSOutsideTUN(t *testing.T) {
prefixes := []netip.Prefix{
netip.MustParsePrefix("198.51.100.1/30"), // gateway, not masked
netip.MustParsePrefix("203.0.113.0/24"), // route
}
servers := []netip.Addr{
netip.MustParseAddr("198.51.100.2"),
netip.MustParseAddr("203.0.113.53"),
netip.MustParseAddr("::ffff:203.0.113.54"),
netip.MustParseAddr("8.8.8.8"),
netip.MustParseAddr("2001:db8::53"),
}
want := []netip.Addr{netip.MustParseAddr("8.8.8.8"), netip.MustParseAddr("2001:db8::53")}
if got := dnsOutsideTUN(servers, prefixes); !slices.Equal(got, want) {
t.Errorf("got %v, want %v", got, want)
}
}
func TestResolveOnOwn(t *testing.T) {
internet.SkipDNSServers([]netip.Addr{netip.MustParseAddr("::ffff:203.0.113.53")})
t.Cleanup(func() { internet.SkipDNSServers(nil) })
preferGo, dial := net.DefaultResolver.PreferGo, net.DefaultResolver.Dial
saved := resolveOnOwn()
t.Cleanup(saved.restore)
if !net.DefaultResolver.PreferGo || net.DefaultResolver.Dial == nil {
t.Fatal("net.DefaultResolver is unchanged")
}
if _, err := net.DefaultResolver.Dial(context.Background(), "udp", "203.0.113.53:53"); err == nil {
t.Error("the TUN's DNS server was not skipped")
}
conn, err := net.DefaultResolver.Dial(context.Background(), "udp", "127.0.0.1:53")
if err != nil {
t.Fatal(err)
}
conn.Close()
saved.restore()
if net.DefaultResolver.PreferGo != preferGo || (net.DefaultResolver.Dial == nil) != (dial == nil) {
t.Error("net.DefaultResolver is not restored")
}
}
// TestTunOnlyDNS checks that a DNS server another interface uses as well is
// not skipped, while one of the TUN alone is.
func TestTunOnlyDNS(t *testing.T) {
adapters, err := winipcfg.GetAdaptersAddresses(windows.AF_UNSPEC, winipcfg.GAAFlagIncludeGateways)
if err != nil {
t.Fatal(err)
}
var other netip.Addr
for _, adapter := range adapters {
if adapter.OperStatus == winipcfg.IfOperStatusUp && adapter.FirstGatewayAddress != nil && adapter.FirstDNSServerAddress != nil {
other, _ = netip.AddrFromSlice(adapter.FirstDNSServerAddress.Address.IP())
other = other.Unmap()
break
}
}
if !other.IsValid() {
t.Skip("no interface with a gateway and a DNS server")
}
tunOnly := netip.MustParseAddr("203.0.113.53")
// LUID 0 is no interface, so every one counts as another.
got, err := tunOnlyDNS(0, []netip.Addr{other, tunOnly})
if err != nil {
t.Fatal(err)
}
if !slices.Equal(got, []netip.Addr{tunOnly}) {
t.Errorf("got %v, want [%v]", got, tunOnly)
}
}
func TestFlushDNSCache(t *testing.T) {
if err := flushDNSCache(); err != nil {
t.Fatal(err)
}
}
+5 -5
View File
@@ -78,11 +78,11 @@ func (u *udpConnectionHandler) HandlePacket(src net.Destination, dst net.Destina
}
}
func (u *udpConnectionHandler) connectionFinished(conn *udpConn) {
func (u *udpConnectionHandler) connectionFinished(src net.Destination) {
u.Lock()
// Close runs twice per flow; a newer conn may already own this src.
if u.udpConns[conn.src] == conn {
delete(u.udpConns, conn.src)
conn, found := u.udpConns[src]
if found {
delete(u.udpConns, src)
close(conn.egress)
}
u.Unlock()
@@ -161,7 +161,7 @@ func (c *udpConn) Write(p []byte) (int, error) {
}
func (c *udpConn) Close() error {
c.handler.connectionFinished(c)
c.handler.connectionFinished(c.src)
return nil
}
+2 -12
View File
@@ -52,12 +52,9 @@ func (b *bind) Open(port uint16) (fns []conn.ReceiveFunc, actualPort uint16, err
case <-ch:
default:
errors.LogErrorInner(context.Background(), err, "unexpected closed")
b.mu.Lock()
downFunc := b.downFunc
b.mu.Unlock()
if downFunc != nil {
if b.downFunc != nil {
go func() {
common.Must(downFunc())
common.Must(b.downFunc())
}()
}
}
@@ -79,13 +76,6 @@ func (b *bind) Open(port uint16) (fns []conn.ReceiveFunc, actualPort uint16, err
}, uint16(c.LocalAddr().(*net.UDPAddr).Port), nil
}
// setDownFunc sets downFunc after the device is created, since the device may already be using the bind.
func (b *bind) setDownFunc(f func() error) {
b.mu.Lock()
defer b.mu.Unlock()
b.downFunc = f
}
func (b *bind) Close() error {
b.mu.Lock()
defer b.mu.Unlock()
+145 -132
View File
@@ -3,6 +3,7 @@ package wireguard
import (
"context"
"fmt"
gonet "net"
"net/netip"
"reflect"
"strings"
@@ -30,6 +31,11 @@ import (
"golang.zx2c4.com/wireguard/device"
)
type entry struct {
got []net.IP
time time.Time
}
type Handler struct {
conf *DeviceConfig
policyManager policy.Manager
@@ -43,6 +49,11 @@ type Handler struct {
tnet *Net
dev *device.Device
mu sync.Mutex
// TODO: cache cleanup loop
local bool
cache map[string]entry
cacheMu sync.Mutex
}
func NewClient(ctx context.Context, conf *DeviceConfig) (*Handler, error) {
@@ -98,10 +109,15 @@ func NewClient(ctx context.Context, conf *DeviceConfig) (*Handler, error) {
return nil, err
}
local := false
dns := conf.DNS
if len(dns) == 0 {
dns = []string{"1.1.1.1", "1.0.0.1", "2606:4700:4700::1111", "2606:4700:4700::1001"}
}
if len(dns) == 1 && dns[0] == "local" {
local = true
dns = nil
}
dnses := make([]netip.Addr, 0, len(dns))
for _, dns := range dns {
dnses = append(dnses, netip.MustParseAddr(dns))
@@ -135,6 +151,9 @@ func NewClient(ctx context.Context, conf *DeviceConfig) (*Handler, error) {
tun: tun,
tnet: tnet,
local: local,
cache: make(map[string]entry),
}, nil
}
@@ -153,6 +172,22 @@ func (h *Handler) Process(ctx context.Context, link *transport.Link, dialer inte
return err
}
var addr netip.Addr
if ob.Target.Address.Family().IsDomain() {
ip, err := h.resolveRemote(ob.Target.Address.String())
if err != nil {
return errors.New("failed to resolve domain").Base(err)
}
addr, _ = netip.AddrFromSlice(ip)
} else {
addr, _ = netip.AddrFromSlice(ob.Target.Address.IP())
}
addrPort := netip.AddrPortFrom(addr, ob.Target.Port.Value())
if !addrPort.IsValid() {
return errors.New("invalid target ", ob.Target)
}
var newCtx context.Context
var newCancel context.CancelFunc
if session.TimeoutOnlyFromContext(ctx) {
@@ -181,10 +216,10 @@ func (h *Handler) Process(ctx context.Context, link *transport.Link, dialer inte
var err error
if sessionPolicy.Timeouts.Handshake != 0 {
timeoutCtx, timeoutCancel := context.WithTimeout(ctx, sessionPolicy.Timeouts.Handshake)
conn, err = h.tnet.DialContext(timeoutCtx, "tcp", ob.Target.NetAddr())
conn, err = h.tnet.DialContextTCPAddrPort(timeoutCtx, addrPort)
timeoutCancel()
} else {
conn, err = h.tnet.Dial("tcp", ob.Target.NetAddr())
conn, err = h.tnet.DialContextTCPAddrPort(ctx, addrPort)
}
if err != nil {
return errors.New("failed to create TCP connection").Base(err)
@@ -193,14 +228,15 @@ func (h *Handler) Process(ctx context.Context, link *transport.Link, dialer inte
reader = buf.NewReader(conn)
writer = buf.NewWriter(conn)
case net.Network_UDP:
conn, err := h.tnet.Dial("udp", ob.Target.NetAddr())
conn, err := h.tnet.DialUDPAddrPort(netip.AddrPort{}, addrPort)
if err != nil {
return errors.New("failed to create UDP connection").Base(err)
}
defer conn.Close()
c := &UDPConnClient{
PacketConn: conn.(*net.PacketConnWrapper).PacketConn,
Dest: conn.RemoteAddr().(*net.UDPAddr),
c := &udpConnClient{
PacketConn: conn.(*internet.PacketConnWrapper).PacketConn,
resolveFunc: h.resolveRemote,
dest: gonet.UDPAddrFromAddrPort(addrPort),
}
reader = c
writer = c
@@ -257,26 +293,26 @@ func (h *Handler) init(ctx context.Context) error {
if err != nil {
return nil, err
}
conn, err := internet.DialSystem(ctx, dest, h.streamSettings.SocketSettings)
if err != nil {
return nil, err
}
var pktConn net.PacketConn
if h.streamSettings.FinalMask != nil {
conn, err := h.streamSettings.FinalMask.DialUDP(ctx, dest)
switch c := conn.(type) {
case *internet.PacketConnWrapper:
pktConn = c.PacketConn
case *cnc.Connection:
pktConn = &internet.FakePacketConn{Conn: c}
default:
panic(reflect.TypeOf(c))
}
if h.streamSettings.UdpmaskManager != nil {
newConn, err := h.streamSettings.UdpmaskManager.WrapPacketConnClient(pktConn)
if err != nil {
return nil, errors.New("failed to dial to dest").Base(err)
}
pktConn = conn.(*net.PacketConnWrapper).PacketConn
} else {
conn, err := internet.DialSystem(ctx, dest, h.streamSettings.SocketSettings)
if err != nil {
return nil, errors.New("failed to dial to dest").Base(err)
}
switch c := conn.(type) {
case *net.PacketConnWrapper:
pktConn = c.PacketConn
case *cnc.Connection:
pktConn = &internet.FakePacketConn{Conn: c}
default:
panic(reflect.TypeOf(c))
pktConn.Close()
return nil, errors.New("mask err").Base(err)
}
pktConn = newConn
}
if h.uplinkCounter != nil || h.downlinkCounter != nil {
pktConn = &PacketCounterConnection{
@@ -287,13 +323,7 @@ func (h *Handler) init(ctx context.Context) error {
}
return pktConn, nil
}
// device.NewDevice may use the bind right away (Up -> BindUpdate -> Open),
// so everything it reads must be set before creating the device.
bind := &bind{
resolveFunc: resolveFunc,
listenFunc: listenFunc,
reserved: h.conf.Reserved,
}
bind := &bind{}
logger := &device.Logger{
Verbosef: func(format string, args ...any) {
log.Record(&log.GeneralMessage{
@@ -309,7 +339,10 @@ func (h *Handler) init(ctx context.Context) error {
},
}
dev := device.NewDevice(h.tun, bind, logger)
bind.setDownFunc(dev.Down)
bind.resolveFunc = resolveFunc
bind.listenFunc = listenFunc
bind.downFunc = dev.Down
bind.reserved = h.conf.Reserved
var cfg strings.Builder
cfg.WriteString("private_key=" + h.conf.SecretKey + "\n")
for _, peer := range h.conf.Peers {
@@ -338,54 +371,90 @@ func (h *Handler) init(ctx context.Context) error {
}
func (h *Handler) resolveLocal(host string) (net.IP, error) {
return h.resolveDomain(host, h.conf.DomainStrategy, func(host string) ([]net.IP, uint32, error) {
return h.dns.LookupIP(host, dns.IPOption{IPv4Enable: true, IPv6Enable: true})
})
}
func (h *Handler) resolveRemote(host string) (net.IP, error) {
return h.resolveDomain(host, h.conf.DomainStrategy, func(host string) ([]net.IP, uint32, error) {
if h.local {
return h.dns.LookupIP(host, dns.IPOption{IPv4Enable: true, IPv6Enable: true})
}
return h.tnet.LookupHost(host)
})
}
func (h *Handler) resolveDomain(host string, strategy DeviceConfig_DomainStrategy, lookupIP func(host string) ([]net.IP, uint32, error)) (net.IP, error) {
if ip := net.ParseIP(host); ip != nil {
return ip, nil
}
ips, _, err := h.dns.LookupIP(host, dns.IPOption{IPv4Enable: true, IPv6Enable: true})
h.cacheMu.Lock()
if entry, ok := h.cache[host]; ok {
if time.Now().Before(entry.time) {
h.cacheMu.Unlock()
return entry.got[dice.Roll(len(entry.got))], nil
}
delete(h.cache, host)
}
h.cacheMu.Unlock()
ips, ttl, err := lookupIP(host)
if err != nil {
return nil, err
}
got := ips
if h.streamSettings.SocketSettings != nil {
var got4, got6 []net.IP
for _, ip := range ips {
if ip.To4() != nil {
got4 = append(got4, ip)
} else {
got6 = append(got6, ip)
}
}
switch h.streamSettings.SocketSettings.DomainStrategy {
case internet.DomainStrategy_AS_IS, internet.DomainStrategy_USE_IP, internet.DomainStrategy_FORCE_IP:
got = ips
case internet.DomainStrategy_USE_IP4, internet.DomainStrategy_FORCE_IP4:
got = got4
case internet.DomainStrategy_USE_IP6, internet.DomainStrategy_FORCE_IP6:
got = got6
case internet.DomainStrategy_USE_IP46, internet.DomainStrategy_FORCE_IP46:
got = got4
if len(got) == 0 {
got = got6
}
case internet.DomainStrategy_USE_IP64, internet.DomainStrategy_FORCE_IP64:
got = got6
if len(got) == 0 {
got = got4
}
}
if len(got) == 0 {
return nil, dns.ErrEmptyResponse
if len(ips) == 0 {
return nil, dns.ErrEmptyResponse
}
var got4, got6 []net.IP
for _, ip := range ips {
if ip.To4() != nil {
got4 = append(got4, ip)
} else {
got6 = append(got6, ip)
}
}
var got []net.IP
switch strategy {
case DeviceConfig_FORCE_IP:
got = ips
return ips[dice.Roll(len(ips))], nil
case DeviceConfig_FORCE_IP4:
got = got4
case DeviceConfig_FORCE_IP6:
got = got6
case DeviceConfig_FORCE_IP46:
got = got4
if len(got) == 0 {
got = got6
}
case DeviceConfig_FORCE_IP64:
got = got6
if len(got) == 0 {
got = got4
}
default:
panic(strategy)
}
if len(got) == 0 {
return nil, dns.ErrEmptyResponse
}
entry := entry{
got: got,
time: time.Now().Add(time.Duration(ttl) * time.Second),
}
h.cacheMu.Lock()
h.cache[host] = entry
h.cacheMu.Unlock()
return got[dice.Roll(len(got))], nil
}
type UDPConnClient struct {
type udpConnClient struct {
net.PacketConn
Dest *net.UDPAddr
resolveFunc func(host string) (net.IP, error)
dest *net.UDPAddr
}
func (c *UDPConnClient) ReadMultiBuffer() (buf.MultiBuffer, error) {
func (c *udpConnClient) ReadMultiBuffer() (buf.MultiBuffer, error) {
b := buf.New()
b.Resize(0, buf.Size)
n, addr, err := c.PacketConn.ReadFrom(b.Bytes())
@@ -404,13 +473,20 @@ func (c *UDPConnClient) ReadMultiBuffer() (buf.MultiBuffer, error) {
return buf.MultiBuffer{b}, nil
}
func (c *UDPConnClient) WriteMultiBuffer(mb buf.MultiBuffer) error {
func (c *udpConnClient) WriteMultiBuffer(mb buf.MultiBuffer) error {
for i, b := range mb {
dst := c.Dest
dst := c.dest
if b.UDP != nil {
if b.UDP.Address.Family().IsDomain() {
if b.UDP.Port != net.Port(dst.Port) {
dst = &net.UDPAddr{IP: dst.IP, Port: int(b.UDP.Port)}
ip, err := c.resolveFunc(b.UDP.Address.String())
if err != nil {
errors.LogErrorInner(context.Background(), err, "drop packet to ", b.UDP, " with size ", len(b.Bytes()))
b.Release()
continue
}
dst = &net.UDPAddr{
IP: ip,
Port: int(b.UDP.Port),
}
} else {
dst = b.UDP.RawNetAddr().(*net.UDPAddr)
@@ -447,66 +523,3 @@ func (c *PacketCounterConnection) WriteTo(p []byte, addr net.Addr) (n int, err e
}
return
}
type entry struct {
saddr []string
deadline time.Time
}
type cache struct {
running bool
m map[string]entry
mu sync.Mutex
}
func (c *cache) run() {
if c.running {
return
}
c.running = true
if c.m == nil {
c.m = make(map[string]entry)
}
go c.gc()
}
func (c *cache) gc() {
ticker := time.NewTicker(time.Minute)
defer ticker.Stop()
for now := range ticker.C {
c.mu.Lock()
for key, entry := range c.m {
if now.After(entry.deadline) {
delete(c.m, key)
}
}
if len(c.m) == 0 {
c.running = false
c.mu.Unlock()
return
}
c.mu.Unlock()
}
}
func (c *cache) LookupHost(host string) []string {
c.mu.Lock()
defer c.mu.Unlock()
c.run()
if entry, ok := c.m[host]; ok {
if time.Now().Before(entry.deadline) {
return entry.saddr
}
delete(c.m, host)
}
return nil
}
func (c *cache) Cache(host string, saddr []string, ttl uint32) {
c.mu.Lock()
defer c.mu.Unlock()
c.m[host] = entry{
saddr: saddr,
deadline: time.Now().Add(time.Second * time.Duration(ttl)),
}
}
+102 -26
View File
@@ -22,6 +22,61 @@ const (
_ = protoimpl.EnforceVersion(protoimpl.MaxVersion - 20)
)
type DeviceConfig_DomainStrategy int32
const (
DeviceConfig_FORCE_IP DeviceConfig_DomainStrategy = 0
DeviceConfig_FORCE_IP4 DeviceConfig_DomainStrategy = 1
DeviceConfig_FORCE_IP6 DeviceConfig_DomainStrategy = 2
DeviceConfig_FORCE_IP46 DeviceConfig_DomainStrategy = 3
DeviceConfig_FORCE_IP64 DeviceConfig_DomainStrategy = 4
)
// Enum value maps for DeviceConfig_DomainStrategy.
var (
DeviceConfig_DomainStrategy_name = map[int32]string{
0: "FORCE_IP",
1: "FORCE_IP4",
2: "FORCE_IP6",
3: "FORCE_IP46",
4: "FORCE_IP64",
}
DeviceConfig_DomainStrategy_value = map[string]int32{
"FORCE_IP": 0,
"FORCE_IP4": 1,
"FORCE_IP6": 2,
"FORCE_IP46": 3,
"FORCE_IP64": 4,
}
)
func (x DeviceConfig_DomainStrategy) Enum() *DeviceConfig_DomainStrategy {
p := new(DeviceConfig_DomainStrategy)
*p = x
return p
}
func (x DeviceConfig_DomainStrategy) String() string {
return protoimpl.X.EnumStringOf(x.Descriptor(), protoreflect.EnumNumber(x))
}
func (DeviceConfig_DomainStrategy) Descriptor() protoreflect.EnumDescriptor {
return file_proxy_wireguard_config_proto_enumTypes[0].Descriptor()
}
func (DeviceConfig_DomainStrategy) Type() protoreflect.EnumType {
return &file_proxy_wireguard_config_proto_enumTypes[0]
}
func (x DeviceConfig_DomainStrategy) Number() protoreflect.EnumNumber {
return protoreflect.EnumNumber(x)
}
// Deprecated: Use DeviceConfig_DomainStrategy.Descriptor instead.
func (DeviceConfig_DomainStrategy) EnumDescriptor() ([]byte, []int) {
return file_proxy_wireguard_config_proto_rawDescGZIP(), []int{1, 0}
}
type PeerConfig struct {
state protoimpl.MessageState `protogen:"open.v1"`
PublicKey string `protobuf:"bytes,1,opt,name=public_key,json=publicKey,proto3" json:"public_key,omitempty"`
@@ -99,18 +154,19 @@ func (x *PeerConfig) GetAllowedIps() []string {
}
type DeviceConfig struct {
state protoimpl.MessageState `protogen:"open.v1"`
SecretKey string `protobuf:"bytes,1,opt,name=secret_key,json=secretKey,proto3" json:"secret_key,omitempty"`
Endpoint []string `protobuf:"bytes,2,rep,name=endpoint,proto3" json:"endpoint,omitempty"`
Peers []*PeerConfig `protobuf:"bytes,3,rep,name=peers,proto3" json:"peers,omitempty"`
Users []*protocol.User `protobuf:"bytes,5,rep,name=users,proto3" json:"users,omitempty"`
Mtu int32 `protobuf:"varint,4,opt,name=mtu,proto3" json:"mtu,omitempty"`
Reserved []byte `protobuf:"bytes,6,opt,name=reserved,proto3" json:"reserved,omitempty"`
IsClient bool `protobuf:"varint,8,opt,name=is_client,json=isClient,proto3" json:"is_client,omitempty"`
NoKernelTun bool `protobuf:"varint,9,opt,name=no_kernel_tun,json=noKernelTun,proto3" json:"no_kernel_tun,omitempty"`
DNS []string `protobuf:"bytes,10,rep,name=DNS,proto3" json:"DNS,omitempty"`
unknownFields protoimpl.UnknownFields
sizeCache protoimpl.SizeCache
state protoimpl.MessageState `protogen:"open.v1"`
SecretKey string `protobuf:"bytes,1,opt,name=secret_key,json=secretKey,proto3" json:"secret_key,omitempty"`
Endpoint []string `protobuf:"bytes,2,rep,name=endpoint,proto3" json:"endpoint,omitempty"`
Peers []*PeerConfig `protobuf:"bytes,3,rep,name=peers,proto3" json:"peers,omitempty"`
Users []*protocol.User `protobuf:"bytes,5,rep,name=users,proto3" json:"users,omitempty"`
Mtu int32 `protobuf:"varint,4,opt,name=mtu,proto3" json:"mtu,omitempty"`
Reserved []byte `protobuf:"bytes,6,opt,name=reserved,proto3" json:"reserved,omitempty"`
DomainStrategy DeviceConfig_DomainStrategy `protobuf:"varint,7,opt,name=domain_strategy,json=domainStrategy,proto3,enum=xray.proxy.wireguard.DeviceConfig_DomainStrategy" json:"domain_strategy,omitempty"`
IsClient bool `protobuf:"varint,8,opt,name=is_client,json=isClient,proto3" json:"is_client,omitempty"`
NoKernelTun bool `protobuf:"varint,9,opt,name=no_kernel_tun,json=noKernelTun,proto3" json:"no_kernel_tun,omitempty"`
DNS []string `protobuf:"bytes,10,rep,name=DNS,proto3" json:"DNS,omitempty"`
unknownFields protoimpl.UnknownFields
sizeCache protoimpl.SizeCache
}
func (x *DeviceConfig) Reset() {
@@ -185,6 +241,13 @@ func (x *DeviceConfig) GetReserved() []byte {
return nil
}
func (x *DeviceConfig) GetDomainStrategy() DeviceConfig_DomainStrategy {
if x != nil {
return x.DomainStrategy
}
return DeviceConfig_FORCE_IP
}
func (x *DeviceConfig) GetIsClient() bool {
if x != nil {
return x.IsClient
@@ -220,7 +283,7 @@ const file_proxy_wireguard_config_proto_rawDesc = "" +
"\n" +
"keep_alive\x18\x04 \x01(\tR\tkeepAlive\x12\x1f\n" +
"\vallowed_ips\x18\x05 \x03(\tR\n" +
"allowedIps\"\xb4\x02\n" +
"allowedIps\"\xee\x03\n" +
"\fDeviceConfig\x12\x1d\n" +
"\n" +
"secret_key\x18\x01 \x01(\tR\tsecretKey\x12\x1a\n" +
@@ -228,11 +291,20 @@ const file_proxy_wireguard_config_proto_rawDesc = "" +
"\x05peers\x18\x03 \x03(\v2 .xray.proxy.wireguard.PeerConfigR\x05peers\x120\n" +
"\x05users\x18\x05 \x03(\v2\x1a.xray.common.protocol.UserR\x05users\x12\x10\n" +
"\x03mtu\x18\x04 \x01(\x05R\x03mtu\x12\x1a\n" +
"\breserved\x18\x06 \x01(\fR\breserved\x12\x1b\n" +
"\breserved\x18\x06 \x01(\fR\breserved\x12Z\n" +
"\x0fdomain_strategy\x18\a \x01(\x0e21.xray.proxy.wireguard.DeviceConfig.DomainStrategyR\x0edomainStrategy\x12\x1b\n" +
"\tis_client\x18\b \x01(\bR\bisClient\x12\"\n" +
"\rno_kernel_tun\x18\t \x01(\bR\vnoKernelTun\x12\x10\n" +
"\x03DNS\x18\n" +
" \x03(\tR\x03DNSB^\n" +
" \x03(\tR\x03DNS\"\\\n" +
"\x0eDomainStrategy\x12\f\n" +
"\bFORCE_IP\x10\x00\x12\r\n" +
"\tFORCE_IP4\x10\x01\x12\r\n" +
"\tFORCE_IP6\x10\x02\x12\x0e\n" +
"\n" +
"FORCE_IP46\x10\x03\x12\x0e\n" +
"\n" +
"FORCE_IP64\x10\x04B^\n" +
"\x18com.xray.proxy.wireguardP\x01Z)github.com/xtls/xray-core/proxy/wireguard\xaa\x02\x14Xray.Proxy.WireGuardb\x06proto3"
var (
@@ -247,20 +319,23 @@ func file_proxy_wireguard_config_proto_rawDescGZIP() []byte {
return file_proxy_wireguard_config_proto_rawDescData
}
var file_proxy_wireguard_config_proto_enumTypes = make([]protoimpl.EnumInfo, 1)
var file_proxy_wireguard_config_proto_msgTypes = make([]protoimpl.MessageInfo, 2)
var file_proxy_wireguard_config_proto_goTypes = []any{
(*PeerConfig)(nil), // 0: xray.proxy.wireguard.PeerConfig
(*DeviceConfig)(nil), // 1: xray.proxy.wireguard.DeviceConfig
(*protocol.User)(nil), // 2: xray.common.protocol.User
(DeviceConfig_DomainStrategy)(0), // 0: xray.proxy.wireguard.DeviceConfig.DomainStrategy
(*PeerConfig)(nil), // 1: xray.proxy.wireguard.PeerConfig
(*DeviceConfig)(nil), // 2: xray.proxy.wireguard.DeviceConfig
(*protocol.User)(nil), // 3: xray.common.protocol.User
}
var file_proxy_wireguard_config_proto_depIdxs = []int32{
0, // 0: xray.proxy.wireguard.DeviceConfig.peers:type_name -> xray.proxy.wireguard.PeerConfig
2, // 1: xray.proxy.wireguard.DeviceConfig.users:type_name -> xray.common.protocol.User
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
1, // 0: xray.proxy.wireguard.DeviceConfig.peers:type_name -> xray.proxy.wireguard.PeerConfig
3, // 1: xray.proxy.wireguard.DeviceConfig.users:type_name -> xray.common.protocol.User
0, // 2: xray.proxy.wireguard.DeviceConfig.domain_strategy:type_name -> xray.proxy.wireguard.DeviceConfig.DomainStrategy
3, // [3:3] is the sub-list for method output_type
3, // [3:3] is the sub-list for method input_type
3, // [3:3] is the sub-list for extension type_name
3, // [3:3] is the sub-list for extension extendee
0, // [0:3] is the sub-list for field type_name
}
func init() { file_proxy_wireguard_config_proto_init() }
@@ -273,13 +348,14 @@ func file_proxy_wireguard_config_proto_init() {
File: protoimpl.DescBuilder{
GoPackagePath: reflect.TypeOf(x{}).PkgPath(),
RawDescriptor: unsafe.Slice(unsafe.StringData(file_proxy_wireguard_config_proto_rawDesc), len(file_proxy_wireguard_config_proto_rawDesc)),
NumEnums: 0,
NumEnums: 1,
NumMessages: 2,
NumExtensions: 0,
NumServices: 0,
},
GoTypes: file_proxy_wireguard_config_proto_goTypes,
DependencyIndexes: file_proxy_wireguard_config_proto_depIdxs,
EnumInfos: file_proxy_wireguard_config_proto_enumTypes,
MessageInfos: file_proxy_wireguard_config_proto_msgTypes,
}.Build()
File_proxy_wireguard_config_proto = out.File
+8
View File
@@ -17,6 +17,13 @@ message PeerConfig {
}
message DeviceConfig {
enum DomainStrategy {
FORCE_IP = 0;
FORCE_IP4 = 1;
FORCE_IP6 = 2;
FORCE_IP46 = 3;
FORCE_IP64 = 4;
}
string secret_key = 1;
repeated string endpoint = 2;
repeated PeerConfig peers = 3;
@@ -24,6 +31,7 @@ message DeviceConfig {
int32 mtu = 4;
bytes reserved = 6;
DomainStrategy domain_strategy = 7;
bool is_client = 8;
bool no_kernel_tun = 9;
repeated string DNS = 10;
+16 -161
View File
@@ -15,13 +15,11 @@ import (
"net"
"net/netip"
"os"
"regexp"
"strconv"
"strings"
"syscall"
"time"
xnet "github.com/xtls/xray-core/common/net"
"github.com/xtls/xray-core/transport/internet"
"golang.zx2c4.com/wireguard/tun"
"golang.org/x/net/dns/dnsmessage"
@@ -44,7 +42,6 @@ type netTun struct {
events chan tun.Event
notifyHandle *channel.NotificationHandle
incomingPacket chan *buffer.View
closed chan struct{}
mtu int
dnsServers []netip.Addr
hasV4, hasV6 bool
@@ -61,7 +58,6 @@ func CreateNetTUN(localAddresses, dnsServers []netip.Addr, mtu int, handleLocal
stack: stack.New(opts),
events: make(chan tun.Event, 10),
incomingPacket: make(chan *buffer.View),
closed: make(chan struct{}),
dnsServers: dnsServers,
mtu: mtu,
}
@@ -128,15 +124,12 @@ func (tun *netTun) Events() <-chan tun.Event {
}
func (tun *netTun) Read(buf [][]byte, sizes []int, offset int) (int, error) {
var view *buffer.View
select {
case view = <-tun.incomingPacket:
case <-tun.closed:
view, ok := <-tun.incomingPacket
if !ok {
return 0, os.ErrClosed
}
n, err := view.Read(buf[0][offset:])
view.Release()
if err != nil {
return 0, err
}
@@ -173,11 +166,7 @@ func (tun *netTun) WriteNotify() {
view := pkt.ToView()
pkt.DecRef()
select {
case tun.incomingPacket <- view:
case <-tun.closed:
view.Release()
}
tun.incomingPacket <- view
}
func (tun *netTun) Close() error {
@@ -190,9 +179,8 @@ func (tun *netTun) Close() error {
close(tun.events)
}
// we don't close incomingPacket, because WriteNotify may be mid-send on it (DNS lookup) and would panic.
if tun.closed != nil {
close(tun.closed)
if tun.incomingPacket != nil {
close(tun.incomingPacket)
}
return nil
@@ -220,7 +208,7 @@ func (tun *netTun) DialUDPAddrPort(laddr, raddr netip.AddrPort) (net.Conn, error
if err != nil {
return nil, err
}
return &xnet.PacketConnWrapper{
return &internet.PacketConnWrapper{
PacketConn: conn,
Dest: net.UDPAddrFromAddrPort(raddr),
}, nil
@@ -231,7 +219,6 @@ type Net struct {
DialUDPAddrPort func(laddr, raddr netip.AddrPort) (net.Conn, error)
dnsServers []netip.Addr
hasV4, hasV6 bool
cache cache
}
func convertToFullAddr(endpoint netip.AddrPort) (tcpip.FullAddress, tcpip.NetworkProtocolNumber) {
@@ -259,12 +246,9 @@ var (
errServerTemporarilyMisbehaving = errors.New("server misbehaving")
errCanceled = errors.New("operation was canceled")
errTimeout = errors.New("i/o timeout")
errNumericPort = errors.New("port must be numeric")
errNoSuitableAddress = errors.New("no suitable address found")
errMissingAddress = errors.New("missing address")
)
func (net *Net) LookupHost(host string) (addrs []string, err error) {
func (net *Net) LookupHost(host string) (addrs []net.IP, ttl uint32, err error) {
return net.LookupContextHost(context.Background(), host)
}
@@ -583,12 +567,9 @@ func (tnet *Net) tryOneName(ctx context.Context, name string, qtype dnsmessage.T
return dnsmessage.Parser{}, "", lastErr
}
func (tnet *Net) LookupContextHost(ctx context.Context, host string) ([]string, error) {
if saddr := tnet.cache.LookupHost(host); saddr != nil {
return saddr, nil
}
func (tnet *Net) LookupContextHost(ctx context.Context, host string) ([]net.IP, uint32, error) {
if host == "" || (!tnet.hasV6 && !tnet.hasV4) {
return nil, &net.DNSError{Err: errNoSuchHost.Error(), Name: host, IsNotFound: true}
return nil, 0, &net.DNSError{Err: errNoSuchHost.Error(), Name: host, IsNotFound: true}
}
zlen := len(host)
if strings.IndexByte(host, ':') != -1 {
@@ -597,11 +578,11 @@ func (tnet *Net) LookupContextHost(ctx context.Context, host string) ([]string,
}
}
if ip, err := netip.ParseAddr(host[:zlen]); err == nil {
return []string{ip.String()}, nil
return []net.IP{ip.AsSlice()}, 0, nil
}
if !isDomainName(host) {
return nil, &net.DNSError{Err: errNoSuchHost.Error(), Name: host, IsNotFound: true}
return nil, 0, &net.DNSError{Err: errNoSuchHost.Error(), Name: host, IsNotFound: true}
}
type result struct {
p dnsmessage.Parser
@@ -702,137 +683,11 @@ func (tnet *Net) LookupContextHost(ctx context.Context, host string) ([]string,
}
if len(addrs) == 0 && lastErr != nil {
return nil, lastErr
return nil, 0, lastErr
}
saddrs := make([]string, 0, len(addrs))
ips := make([]net.IP, 0, len(addrs))
for _, ip := range addrs {
saddrs = append(saddrs, ip.String())
ips = append(ips, ip.AsSlice())
}
tnet.cache.Cache(host, saddrs, ttl)
return saddrs, nil
}
func partialDeadline(now, deadline time.Time, addrsRemaining int) (time.Time, error) {
if deadline.IsZero() {
return deadline, nil
}
timeRemaining := deadline.Sub(now)
if timeRemaining <= 0 {
return time.Time{}, errTimeout
}
timeout := timeRemaining / time.Duration(addrsRemaining)
const saneMinimum = 2 * time.Second
if timeout < saneMinimum {
if timeRemaining < saneMinimum {
timeout = timeRemaining
} else {
timeout = saneMinimum
}
}
return now.Add(timeout), nil
}
var protoSplitter = regexp.MustCompile(`^(tcp|udp|ping)(4|6)?$`)
func (tnet *Net) DialContext(ctx context.Context, network, address string) (net.Conn, error) {
if ctx == nil {
panic("nil context")
}
var acceptV4, acceptV6 bool
matches := protoSplitter.FindStringSubmatch(network)
if matches == nil {
return nil, &net.OpError{Op: "dial", Err: net.UnknownNetworkError(network)}
} else if len(matches[2]) == 0 {
acceptV4 = true
acceptV6 = true
} else {
acceptV4 = matches[2][0] == '4'
acceptV6 = !acceptV4
}
var host string
var port int
if matches[1] == "ping" {
host = address
} else {
var sport string
var err error
host, sport, err = net.SplitHostPort(address)
if err != nil {
return nil, &net.OpError{Op: "dial", Err: err}
}
port, err = strconv.Atoi(sport)
if err != nil || port < 0 || port > 65535 {
return nil, &net.OpError{Op: "dial", Err: errNumericPort}
}
}
allAddr, err := tnet.LookupContextHost(ctx, host)
if err != nil {
return nil, &net.OpError{Op: "dial", Err: err}
}
var addrs []netip.AddrPort
for _, addr := range allAddr {
ip, err := netip.ParseAddr(addr)
if err == nil && ((ip.Is4() && acceptV4) || (ip.Is6() && acceptV6)) {
addrs = append(addrs, netip.AddrPortFrom(ip, uint16(port)))
}
}
if len(addrs) == 0 && len(allAddr) != 0 {
return nil, &net.OpError{Op: "dial", Err: errNoSuitableAddress}
}
var firstErr error
for i, addr := range addrs {
select {
case <-ctx.Done():
err := ctx.Err()
if err == context.Canceled {
err = errCanceled
} else if err == context.DeadlineExceeded {
err = errTimeout
}
return nil, &net.OpError{Op: "dial", Err: err}
default:
}
dialCtx := ctx
if deadline, hasDeadline := ctx.Deadline(); hasDeadline {
partialDeadline, err := partialDeadline(time.Now(), deadline, len(addrs)-i)
if err != nil {
if firstErr == nil {
firstErr = &net.OpError{Op: "dial", Err: err}
}
break
}
if partialDeadline.Before(deadline) {
var cancel context.CancelFunc
dialCtx, cancel = context.WithDeadline(ctx, partialDeadline)
defer cancel()
}
}
var c net.Conn
switch matches[1] {
case "tcp":
c, err = tnet.DialContextTCPAddrPort(dialCtx, addr)
case "udp":
c, err = tnet.DialUDPAddrPort(netip.AddrPort{}, addr)
case "ping":
err = errors.New("not support")
// c, err = tnet.DialPingAddr(netip.Addr{}, addr.Addr())
}
if err == nil {
return c, nil
}
if firstErr == nil {
firstErr = err
}
}
if firstErr == nil {
firstErr = &net.OpError{Op: "dial", Err: errMissingAddress}
}
return nil, firstErr
}
func (tnet *Net) Dial(network, address string) (net.Conn, error) {
return tnet.DialContext(context.Background(), network, address)
return ips, ttl, nil
}
+12 -12
View File
@@ -113,7 +113,7 @@ func NewServer(ctx context.Context, conf *DeviceConfig) (*Server, error) {
users.Store(user.Account.(*MemoryAccount).Pub, user)
}
s := &Server{
return &Server{
conf: conf,
ctx: core.ToBackgroundDetachedContext(ctx),
policyManager: p,
@@ -131,10 +131,7 @@ func NewServer(ctx context.Context, conf *DeviceConfig) (*Server, error) {
pub: pub,
users: users,
}
// Install the stack's protocol handlers before the device can deliver packets to it (Start -> dev.Up).
CreateForwarder(stack, s.HandleConnection)
return s, nil
}, nil
}
func (s *Server) AddUser(ctx context.Context, user *protocol.MemoryUser) error {
@@ -261,16 +258,18 @@ func (s *Server) Start() error {
return errors.New("address is domain")
}
listenFunc := func() (net.PacketConn, error) {
var pktConn net.PacketConn
var err error
if s.streamSettings.FinalMask != nil {
pktConn, err = s.streamSettings.FinalMask.ListenPacket(context.Background(), &net.UDPAddr{IP: s.src.Address.IP(), Port: int(s.src.Port)})
} else {
pktConn, err = internet.ListenSystemPacket(context.Background(), &net.UDPAddr{IP: s.src.Address.IP(), Port: int(s.src.Port)}, s.streamSettings.SocketSettings)
}
pktConn, err := internet.ListenSystemPacket(context.Background(), &net.UDPAddr{IP: s.src.Address.IP(), Port: int(s.src.Port)}, s.streamSettings.SocketSettings)
if err != nil {
return nil, err
}
if s.streamSettings.UdpmaskManager != nil {
newConn, err := s.streamSettings.UdpmaskManager.WrapPacketConnServer(pktConn)
if err != nil {
pktConn.Close()
return nil, errors.New("mask err").Base(err)
}
pktConn = newConn
}
if s.uplinkCounter != nil || s.downlinkCounter != nil {
pktConn = &PacketCounterConnection{
PacketConn: pktConn,
@@ -323,6 +322,7 @@ func (s *Server) Start() error {
return err
}
s.dev = dev
createForwarder(s.stack, s.HandleConnection)
return nil
}
+1 -1
View File
@@ -49,7 +49,7 @@ func CalculateInterfaceName(name string) (tunName string) {
return
}
func CreateForwarder(gstack *stack.Stack, handler func(conn net.Conn, dest net.Destination)) {
func createForwarder(gstack *stack.Stack, handler func(conn net.Conn, dest net.Destination)) {
gstack.SetPromiscuousMode(1, true)
gstack.SetSpoofing(1, true)
+2 -2
View File
@@ -16,7 +16,7 @@ import (
"github.com/vishvananda/netlink"
"github.com/xtls/xray-core/common/errors"
xnet "github.com/xtls/xray-core/common/net"
"github.com/xtls/xray-core/transport/internet"
"golang.zx2c4.com/wireguard/tun"
)
@@ -263,7 +263,7 @@ func (tun *kernelTun) DialUDPAddrPort(laddr, raddr netip.AddrPort) (net.Conn, er
if err != nil {
return nil, err
}
return &xnet.PacketConnWrapper{
return &internet.PacketConnWrapper{
PacketConn: conn,
Dest: net.UDPAddrFromAddrPort(raddr),
}, nil
-689
View File
@@ -1,689 +0,0 @@
package scenarios
import (
"bufio"
"bytes"
"context"
"crypto/rand"
gotls "crypto/tls"
"crypto/x509"
"encoding/binary"
go_errors "errors"
"io"
"net/http"
"net/netip"
"net/url"
"strconv"
"strings"
"sync"
"sync/atomic"
"testing"
"time"
"github.com/apernet/quic-go"
"github.com/apernet/quic-go/http3"
"golang.org/x/net/http2"
"golang.org/x/net/http2/hpack"
"golang.org/x/sync/errgroup"
"gvisor.dev/gvisor/pkg/tcpip"
"gvisor.dev/gvisor/pkg/tcpip/adapters/gonet"
"gvisor.dev/gvisor/pkg/tcpip/network/ipv4"
"gvisor.dev/gvisor/pkg/tcpip/network/ipv6"
"github.com/xtls/xray-core/app/log"
"github.com/xtls/xray-core/app/proxyman"
"github.com/xtls/xray-core/app/router"
"github.com/xtls/xray-core/common"
clog "github.com/xtls/xray-core/common/log"
"github.com/xtls/xray-core/common/net"
"github.com/xtls/xray-core/common/protocol"
"github.com/xtls/xray-core/common/protocol/tls/cert"
"github.com/xtls/xray-core/common/serial"
core "github.com/xtls/xray-core/core"
"github.com/xtls/xray-core/proxy/dokodemo"
"github.com/xtls/xray-core/proxy/freedom"
"github.com/xtls/xray-core/proxy/masque"
"github.com/xtls/xray-core/proxy/wireguard"
"github.com/xtls/xray-core/testing/servers/tcp"
"github.com/xtls/xray-core/testing/servers/udp"
"github.com/xtls/xray-core/transport/internet"
transmasque "github.com/xtls/xray-core/transport/internet/masque"
"github.com/xtls/xray-core/transport/internet/masque/connectip"
"github.com/xtls/xray-core/transport/internet/tls"
)
var (
masqueServerV4 = netip.MustParseAddr("10.13.0.1")
masqueServerV6 = netip.MustParseAddr("fd13::1")
masqueClientV4 = netip.MustParsePrefix("10.13.0.2/32")
masqueClientV6 = netip.MustParsePrefix("fd13::2/128")
)
const (
masqueEchoPort = 7
masqueAuthorization = "Basic dUBleGFtcGxlLmNvbTpw"
)
func startMasqueServer(t *testing.T, h2 bool) (net.Port, [32]byte) {
dev, _, gstack, err := wireguard.CreateNetTUN([]netip.Addr{masqueServerV4, masqueServerV6}, nil, transmasque.MinPacketSize, false)
common.Must(err)
t.Cleanup(func() { dev.Close() })
for _, addr := range []netip.Addr{masqueServerV4, masqueServerV6} {
proto := ipv4.ProtocolNumber
if addr.Is6() {
proto = ipv6.ProtocolNumber
}
local := tcpip.FullAddress{NIC: 1, Addr: tcpip.AddrFromSlice(addr.AsSlice()), Port: masqueEchoPort}
l, err := gonet.ListenTCP(gstack, local, proto)
common.Must(err)
go func() {
for {
c, err := l.Accept()
if err != nil {
return
}
go func() {
defer c.Close()
b := make([]byte, 2048)
for {
n, err := c.Read(b)
if err != nil {
return
}
if _, err := c.Write(xor(b[:n])); err != nil {
return
}
}
}()
}
}()
u, err := gonet.DialUDP(gstack, &local, nil, proto)
common.Must(err)
go func() {
b := make([]byte, 2048)
for {
n, addr, err := u.ReadFrom(b)
if err != nil {
return
}
u.WriteTo(xor(b[:n]), addr)
}
}()
}
var current atomic.Pointer[connectip.Conn]
go func() {
bufs := [][]byte{make([]byte, transmasque.MinPacketSize)}
sizes := []int{0}
for {
if _, err := dev.Read(bufs, sizes, 0); err != nil {
return
}
if conn := current.Load(); conn != nil {
if icmp, _ := conn.WritePacket(bufs[0][:sizes[0]]); len(icmp) > 0 {
go dev.Write([][]byte{icmp}, 0)
}
}
}
}()
handler := func(w http.ResponseWriter, r *http.Request) {
if r.URL.Path != transmasque.DefaultPath {
w.WriteHeader(http.StatusNotFound)
return
}
if r.Header.Get("Authorization") != masqueAuthorization {
w.WriteHeader(http.StatusUnauthorized)
return
}
req, err := connectip.ParseProxyRequest(r)
if err != nil {
var perr *connectip.ProxyRequestParseError
if go_errors.As(err, &perr) {
w.WriteHeader(perr.HTTPStatus)
}
return
}
conn, err := (&connectip.Proxy{}).Proxy(w, req)
if err != nil {
return
}
defer conn.Close()
common.Must(conn.AssignAddresses([]netip.Prefix{masqueClientV4, masqueClientV6}))
common.Must(conn.AdvertiseRoute([]connectip.IPRoute{
{StartIP: netip.IPv4Unspecified(), EndIP: netip.AddrFrom4([4]byte{255, 255, 255, 255})},
{StartIP: netip.IPv6Unspecified(), EndIP: netip.AddrFrom16([16]byte{0: 0xff, 1: 0xff, 2: 0xff, 3: 0xff, 4: 0xff, 5: 0xff, 6: 0xff, 7: 0xff, 8: 0xff, 9: 0xff, 10: 0xff, 11: 0xff, 12: 0xff, 13: 0xff, 14: 0xff, 15: 0xff})},
}))
go func() {
for {
ar, err := conn.ReceiveAddressRequest(context.Background())
if err != nil {
return
}
assigned := make([]netip.Prefix, len(ar.Prefixes))
for i, p := range ar.Prefixes {
if p.Addr().Is4() {
assigned[i] = masqueClientV4
} else {
assigned[i] = masqueClientV6
}
}
ar.Respond(assigned, nil)
}
}()
current.Store(conn)
b := make([]byte, 2048)
for {
n, err := conn.ReadPacket(b)
if err != nil {
if go_errors.Is(err, io.ErrShortBuffer) {
continue
}
return
}
dev.Write([][]byte{b[:n]}, 0)
}
}
certificate, certHash := cert.MustGenerate(nil, cert.CommonName("localhost"))
key := common.Must2(x509.ParsePKCS8PrivateKey(certificate.PrivateKey))
tlsConfig := &gotls.Config{
Certificates: []gotls.Certificate{{Certificate: [][]byte{certificate.Certificate}, PrivateKey: key}},
NextProtos: []string{http3.NextProtoH3},
}
if h2 {
tlsConfig.NextProtos = []string{http2.NextProtoTLS}
ln := common.Must2(gotls.Listen("tcp", "127.0.0.1:0", tlsConfig))
t.Cleanup(func() { ln.Close() })
go serveHTTP2(ln, http.HandlerFunc(handler))
return net.Port(ln.Addr().(*net.TCPAddr).Port), certHash
}
pktConn := common.Must2(net.ListenUDP("udp", &net.UDPAddr{IP: net.LocalHostIP.IP()}))
tr := &quic.Transport{Conn: pktConn}
ln := common.Must2(tr.ListenEarly(tlsConfig, &quic.Config{EnableDatagrams: true, InitialPacketSize: 1350}))
server := &http3.Server{Handler: http.HandlerFunc(handler), EnableDatagrams: true}
go server.ServeListener(ln)
t.Cleanup(func() {
server.Close()
ln.Close()
tr.Close()
pktConn.Close()
})
return net.Port(pktConn.LocalAddr().(*net.UDPAddr).Port), certHash
}
func serveHTTP2(ln net.Listener, handler http.Handler) {
for {
conn, err := ln.Accept()
if err != nil {
return
}
go serveHTTP2Conn(conn, handler)
}
}
type http2ServerConn struct {
mu sync.Mutex
fr *http2.Framer
hbuf bytes.Buffer
henc *hpack.Encoder
}
func (c *http2ServerConn) write(f func(*http2.Framer) error) error {
c.mu.Lock()
defer c.mu.Unlock()
return f(c.fr)
}
func (c *http2ServerConn) writeHeaders(streamID uint32, status int, header http.Header) error {
c.mu.Lock()
defer c.mu.Unlock()
c.hbuf.Reset()
c.henc.WriteField(hpack.HeaderField{Name: ":status", Value: strconv.Itoa(status)})
for k, vv := range header {
for _, v := range vv {
c.henc.WriteField(hpack.HeaderField{Name: strings.ToLower(k), Value: v})
}
}
return c.fr.WriteHeaders(http2.HeadersFrameParam{StreamID: streamID, BlockFragment: c.hbuf.Bytes(), EndHeaders: true})
}
func (c *http2ServerConn) writeData(streamID uint32, endStream bool, data []byte) error {
c.mu.Lock()
defer c.mu.Unlock()
for {
n := min(len(data), 16384)
if err := c.fr.WriteData(streamID, endStream && n == len(data), data[:n]); err != nil {
return err
}
if data = data[n:]; len(data) == 0 {
return nil
}
}
}
func serveHTTP2Conn(conn net.Conn, handler http.Handler) {
defer conn.Close()
br := bufio.NewReader(conn)
preface := make([]byte, len(http2.ClientPreface))
if _, err := io.ReadFull(br, preface); err != nil || string(preface) != http2.ClientPreface {
return
}
sc := &http2ServerConn{fr: http2.NewFramer(conn, br)}
sc.henc = hpack.NewEncoder(&sc.hbuf)
sc.fr.ReadMetaHeaders = hpack.NewDecoder(4096, nil)
if err := sc.write(func(fr *http2.Framer) error {
if err := fr.WriteSettings(
http2.Setting{ID: http2.SettingEnableConnectProtocol, Val: 1},
http2.Setting{ID: http2.SettingInitialWindowSize, Val: 1 << 30},
); err != nil {
return err
}
return fr.WriteWindowUpdate(0, 1<<30)
}); err != nil {
return
}
bodies := make(map[uint32]*io.PipeWriter)
defer func() {
for _, body := range bodies {
body.Close()
}
}()
for {
f, err := sc.fr.ReadFrame()
if err != nil {
return
}
switch f := f.(type) {
case *http2.SettingsFrame:
if !f.IsAck() {
err = sc.write((*http2.Framer).WriteSettingsAck)
}
case *http2.PingFrame:
if !f.IsAck() {
err = sc.write(func(fr *http2.Framer) error { return fr.WritePing(true, f.Data) })
}
case *http2.MetaHeadersFrame:
u, err := url.ParseRequestURI(f.PseudoValue("path"))
if err != nil {
return
}
pr, pw := io.Pipe()
bodies[f.StreamID] = pw
req := &http.Request{
Method: f.PseudoValue("method"),
URL: u,
Proto: "HTTP/2.0",
ProtoMajor: 2,
Header: http.Header{},
Host: f.PseudoValue("authority"),
Body: pr,
}
for _, hf := range f.RegularFields() {
req.Header.Add(hf.Name, hf.Value)
}
if protocol := f.PseudoValue("protocol"); protocol != "" {
req.Header.Set(":protocol", protocol)
}
streamID := f.StreamID
w := &http2ResponseWriter{conn: sc, streamID: streamID, header: http.Header{}}
go func() {
handler.ServeHTTP(w, req)
w.WriteHeader(http.StatusOK)
sc.writeData(streamID, true, nil)
}()
case *http2.DataFrame:
if body := bodies[f.StreamID]; body != nil {
if _, err := body.Write(f.Data()); err != nil || f.StreamEnded() {
body.Close()
delete(bodies, f.StreamID)
}
}
case *http2.RSTStreamFrame:
if body := bodies[f.StreamID]; body != nil {
body.CloseWithError(http2.StreamError{StreamID: f.StreamID, Code: f.ErrCode})
delete(bodies, f.StreamID)
}
}
if err != nil {
return
}
}
}
type http2ResponseWriter struct {
conn *http2ServerConn
streamID uint32
header http.Header
wroteHeader bool
}
func (w *http2ResponseWriter) Header() http.Header { return w.header }
func (w *http2ResponseWriter) WriteHeader(code int) {
if !w.wroteHeader {
w.wroteHeader = true
w.conn.writeHeaders(w.streamID, code, w.header)
}
}
func (w *http2ResponseWriter) Write(b []byte) (int, error) {
w.WriteHeader(http.StatusOK)
if err := w.conn.writeData(w.streamID, false, b); err != nil {
return 0, err
}
return len(b), nil
}
func (w *http2ResponseWriter) Flush() {}
func TestMasque(t *testing.T) {
testMasque(t, false)
}
func TestMasqueHTTP2(t *testing.T) {
testMasque(t, true)
}
func masqueDokodemo(port net.Port, addr netip.Addr, network net.Network) *core.InboundHandlerConfig {
return &core.InboundHandlerConfig{
ReceiverSettings: serial.ToTypedMessage(&proxyman.ReceiverConfig{
PortList: &net.PortList{Range: []*net.PortRange{net.SinglePortRange(port)}},
Listen: net.NewIPOrDomain(net.LocalHostIP),
}),
ProxySettings: serial.ToTypedMessage(&dokodemo.Config{
RewriteAddress: net.NewIPOrDomain(net.IPAddress(addr.AsSlice())),
RewritePort: masqueEchoPort,
AllowedNetworks: []net.Network{network},
}),
}
}
func masqueStreamSettings(tlsConfig *tls.Config, config *transmasque.Config) *internet.StreamConfig {
return &internet.StreamConfig{
ProtocolName: "masque",
TransportSettings: []*internet.TransportConfig{
{
ProtocolName: "masque",
Settings: serial.ToTypedMessage(config),
},
},
SecurityType: serial.GetMessageType(&tls.Config{}),
SecuritySettings: []*serial.TypedMessage{serial.ToTypedMessage(tlsConfig)},
}
}
func masqueClientTLS(certHash [32]byte, alpn ...string) *tls.Config {
return &tls.Config{
ServerName: "localhost",
PinnedPeerCertSha256: [][]byte{certHash[:]},
NextProtocol: alpn,
}
}
func masqueOutbound(serverPort net.Port, certHash [32]byte, h2 bool, authorization string) *core.OutboundHandlerConfig {
tlsConfig := masqueClientTLS(certHash)
if h2 {
tlsConfig.NextProtocol = []string{http2.NextProtoTLS}
}
return &core.OutboundHandlerConfig{
ProxySettings: serial.ToTypedMessage(&masque.ClientConfig{
Server: &protocol.ServerEndpoint{
Address: net.NewIPOrDomain(net.LocalHostIP),
Port: uint32(serverPort),
},
}),
SenderSettings: serial.ToTypedMessage(&proxyman.SenderConfig{
StreamSettings: masqueStreamSettings(tlsConfig, &transmasque.Config{
Path: transmasque.DefaultPath,
Headers: map[string]string{"Authorization": authorization},
}),
}),
}
}
func masqueClientConfig(serverPort net.Port, certHash [32]byte, h2 bool, authorization string, tcpPort, tcp6Port, udpPort net.Port, v4, v6 netip.Addr) *core.Config {
return &core.Config{
App: []*serial.TypedMessage{
serial.ToTypedMessage(&log.Config{
ErrorLogLevel: clog.Severity_Debug,
ErrorLogType: log.LogType_Console,
}),
},
Inbound: []*core.InboundHandlerConfig{
masqueDokodemo(tcpPort, v4, net.Network_TCP),
masqueDokodemo(tcp6Port, v6, net.Network_TCP),
masqueDokodemo(udpPort, v4, net.Network_UDP),
},
Outbound: []*core.OutboundHandlerConfig{
masqueOutbound(serverPort, certHash, h2, authorization),
},
}
}
func testMasqueTraffic(t *testing.T, tcpPort, tcp6Port, udpPort net.Port) {
var errg errgroup.Group
for range 3 {
errg.Go(testTCPConn(tcpPort, 1024*1024, time.Second*20))
}
errg.Go(testTCPConn(tcp6Port, 1024*1024, time.Second*20))
errg.Go(testUDPConn(udpPort, 1024, time.Second*5))
if err := errg.Wait(); err != nil {
t.Error(err)
}
}
func testMasque(t *testing.T, h2 bool) {
serverPort, certHash := startMasqueServer(t, h2)
tcpPort := tcp.PickPort()
tcp6Port := tcp.PickPort()
udpPort := udp.PickPort()
clientConfig := masqueClientConfig(serverPort, certHash, h2, masqueAuthorization, tcpPort, tcp6Port, udpPort, masqueServerV4, masqueServerV6)
servers, err := InitializeServerConfigs(clientConfig)
common.Must(err)
defer CloseAllServers(servers)
testMasqueTraffic(t, tcpPort, tcp6Port, udpPort)
}
func masqueServerInbound(serverPort net.Port, certificate *tls.Certificate, alpn ...string) *core.InboundHandlerConfig {
return &core.InboundHandlerConfig{
ReceiverSettings: serial.ToTypedMessage(&proxyman.ReceiverConfig{
PortList: &net.PortList{Range: []*net.PortRange{net.SinglePortRange(serverPort)}},
Listen: net.NewIPOrDomain(net.LocalHostIP),
StreamSettings: masqueStreamSettings(&tls.Config{
Certificate: []*tls.Certificate{certificate},
NextProtocol: alpn,
}, &transmasque.Config{Path: transmasque.DefaultPath}),
}),
ProxySettings: serial.ToTypedMessage(&masque.ServerConfig{
Users: []*protocol.User{{
Email: "u@example.com",
Account: serial.ToTypedMessage(&masque.Account{Password: "p"}),
}},
Address: []string{"10.14.0.1/24", "fd14::1/64"},
}),
}
}
func masqueServerConfig(serverPort net.Port, certificate *tls.Certificate, h2 bool, tcpDest, udpDest net.Destination) *core.Config {
var alpn []string
if h2 {
alpn = []string{http2.NextProtoTLS}
}
redirect := func(tag string, dest net.Destination) *core.OutboundHandlerConfig {
return &core.OutboundHandlerConfig{
Tag: tag,
ProxySettings: serial.ToTypedMessage(&freedom.Config{
DestinationOverride: &freedom.DestinationOverride{
Server: &protocol.ServerEndpoint{
Address: net.NewIPOrDomain(dest.Address),
Port: uint32(dest.Port),
},
},
FinalRules: []*freedom.FinalRuleConfig{{Action: freedom.RuleAction_Allow}},
}),
}
}
return &core.Config{
App: []*serial.TypedMessage{
serial.ToTypedMessage(&log.Config{
ErrorLogLevel: clog.Severity_Debug,
ErrorLogType: log.LogType_Console,
}),
serial.ToTypedMessage(&router.Config{
Rule: []*router.RoutingRule{
{Networks: []net.Network{net.Network_TCP}, TargetTag: &router.RoutingRule_Tag{Tag: "tcp"}},
{Networks: []net.Network{net.Network_UDP}, TargetTag: &router.RoutingRule_Tag{Tag: "udp"}},
},
}),
},
Inbound: []*core.InboundHandlerConfig{
masqueServerInbound(serverPort, certificate, alpn...),
},
Outbound: []*core.OutboundHandlerConfig{
redirect("tcp", tcpDest),
redirect("udp", udpDest),
},
}
}
func testMasqueServer(t *testing.T, h2 bool, authorization string) error {
tcpServer := tcp.Server{MsgProcessor: xor}
tcpDest, err := tcpServer.Start()
common.Must(err)
defer tcpServer.Close()
udpServer := udp.Server{MsgProcessor: xor}
udpDest, err := udpServer.Start()
common.Must(err)
defer udpServer.Close()
ct, ctHash := cert.MustGenerate(nil, cert.CommonName("localhost"))
serverPort := udp.PickPort()
if h2 {
serverPort = tcp.PickPort()
}
tcpPort := tcp.PickPort()
tcp6Port := tcp.PickPort()
udpPort := udp.PickPort()
servers, err := InitializeServerConfigs(
masqueServerConfig(serverPort, tls.ParseCertificate(ct), h2, tcpDest, udpDest),
masqueClientConfig(serverPort, ctHash, h2, authorization, tcpPort, tcp6Port, udpPort, netip.MustParseAddr("192.0.2.1"), netip.MustParseAddr("2001:db8::1")),
)
common.Must(err)
defer CloseAllServers(servers)
if authorization != masqueAuthorization {
return testTCPConn(tcpPort, 1024, time.Second*5)()
}
testMasqueTraffic(t, tcpPort, tcp6Port, udpPort)
return nil
}
func TestMasqueServer(t *testing.T) {
testMasqueServer(t, false, masqueAuthorization)
}
func TestMasqueServerHTTP2(t *testing.T) {
testMasqueServer(t, true, masqueAuthorization)
}
func TestMasqueServerRejectsWrongPassword(t *testing.T) {
for _, h2 := range []bool{false, true} {
if err := testMasqueServer(t, h2, "Basic dUBleGFtcGxlLmNvbTp3cm9uZw=="); err == nil {
t.Errorf("a wrong password got through (h2: %v)", h2)
}
}
}
func masqueIPPacket(src, dst netip.Addr, payload []byte) []byte {
if src.Is4() {
p := make([]byte, 20, 20+len(payload))
p[0] = 0x45
binary.BigEndian.PutUint16(p[2:], uint16(20+len(payload)))
p[8] = 64
p[9] = 253
copy(p[12:], src.AsSlice())
copy(p[16:], dst.AsSlice())
return append(p, payload...)
}
p := make([]byte, 40, 40+len(payload))
p[0] = 0x60
binary.BigEndian.PutUint16(p[4:], uint16(len(payload)))
p[6] = 253
p[7] = 64
copy(p[8:], src.AsSlice())
copy(p[24:], dst.AsSlice())
return append(p, payload...)
}
func masqueIPAddrs(p []byte) (src, dst netip.Addr) {
if p[0]>>4 == 4 {
return netip.AddrFrom4([4]byte(p[12:16])), netip.AddrFrom4([4]byte(p[16:20]))
}
return netip.AddrFrom16([16]byte(p[8:24])), netip.AddrFrom16([16]byte(p[24:40]))
}
func TestMasqueServerClientToClient(t *testing.T) {
ct, ctHash := cert.MustGenerate(nil, cert.CommonName("localhost"))
serverPort := udp.PickPort()
servers, err := InitializeServerConfigs(&core.Config{
Inbound: []*core.InboundHandlerConfig{
masqueServerInbound(serverPort, tls.ParseCertificate(ct), http3.NextProtoH3, http2.NextProtoTLS),
},
Outbound: []*core.OutboundHandlerConfig{
{ProxySettings: serial.ToTypedMessage(&freedom.Config{})},
},
})
common.Must(err)
defer CloseAllServers(servers)
dial := func(alpn ...string) *transmasque.Conn {
streamSettings, err := internet.ToMemoryStreamConfig(masqueStreamSettings(masqueClientTLS(ctHash, alpn...), &transmasque.Config{
Path: transmasque.DefaultPath,
Headers: map[string]string{"Authorization": masqueAuthorization},
}))
common.Must(err)
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
defer cancel()
conn, err := transmasque.Dial(ctx, net.TCPDestination(net.LocalHostIP, serverPort), streamSettings)
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { conn.Close() })
return conn.(*transmasque.Conn)
}
h3 := dial()
h2 := dial(http2.NextProtoTLS)
for _, c := range []struct{ from, to *transmasque.Conn }{{h3, h2}, {h2, h3}} {
for i := range c.from.LocalAddrs() {
src, dst := c.from.LocalAddrs()[i], c.to.LocalAddrs()[i]
payload := make([]byte, 1000)
rand.Read(payload)
if _, err := c.from.Write(masqueIPPacket(src, dst, payload)); err != nil {
t.Fatal(err)
}
received := make(chan []byte, 1)
go func() {
b := make([]byte, 2048)
n, _ := c.to.Read(b)
received <- b[:n]
}()
select {
case p := <-received:
gotSrc, gotDst := masqueIPAddrs(p)
if gotSrc != src || gotDst != dst || !bytes.HasSuffix(p, payload) {
t.Fatalf("unexpected packet from %s to %s: %x", gotSrc, gotDst, p)
}
case <-time.After(5 * time.Second):
t.Fatalf("no packet from %s to %s", src, dst)
}
}
}
}
+7 -16
View File
@@ -6,6 +6,7 @@ import (
"testing"
"time"
"github.com/sagernet/sing-shadowsocks/shadowaead_2022"
"github.com/xtls/xray-core/app/log"
"github.com/xtls/xray-core/app/proxyman"
"github.com/xtls/xray-core/common"
@@ -21,19 +22,9 @@ import (
"golang.org/x/sync/errgroup"
)
var ss2022Methods = []string{
shadowsocks_2022.MethodAES128GCM,
shadowsocks_2022.MethodAES256GCM,
shadowsocks_2022.MethodChaCha20Poly1305,
}
func TestShadowsocks2022Tcp(t *testing.T) {
for _, method := range ss2022Methods {
keySize := 32
if method == shadowsocks_2022.MethodAES128GCM {
keySize = 16
}
password := make([]byte, keySize)
for _, method := range shadowaead_2022.List {
password := make([]byte, 32)
rand.Read(password)
t.Run(method, func(t *testing.T) {
testShadowsocks2022Tcp(t, method, base64.StdEncoding.EncodeToString(password))
@@ -42,21 +33,21 @@ func TestShadowsocks2022Tcp(t *testing.T) {
}
func TestShadowsocks2022UdpAES128(t *testing.T) {
password := make([]byte, 16)
password := make([]byte, 32)
rand.Read(password)
testShadowsocks2022Udp(t, shadowsocks_2022.MethodAES128GCM, base64.StdEncoding.EncodeToString(password))
testShadowsocks2022Udp(t, shadowaead_2022.List[0], base64.StdEncoding.EncodeToString(password))
}
func TestShadowsocks2022UdpAES256(t *testing.T) {
password := make([]byte, 32)
rand.Read(password)
testShadowsocks2022Udp(t, shadowsocks_2022.MethodAES256GCM, base64.StdEncoding.EncodeToString(password))
testShadowsocks2022Udp(t, shadowaead_2022.List[1], base64.StdEncoding.EncodeToString(password))
}
func TestShadowsocks2022UdpChacha(t *testing.T) {
password := make([]byte, 32)
rand.Read(password)
testShadowsocks2022Udp(t, shadowsocks_2022.MethodChaCha20Poly1305, base64.StdEncoding.EncodeToString(password))
testShadowsocks2022Udp(t, shadowaead_2022.List[2], base64.StdEncoding.EncodeToString(password))
}
func testShadowsocks2022Tcp(t *testing.T, method string, password string) {
-2
View File
@@ -65,7 +65,6 @@ func TestWireguard(t *testing.T) {
ProxySettings: serial.ToTypedMessage(&freedom.Config{
FinalRules: []*freedom.FinalRuleConfig{{Action: freedom.RuleAction_Allow}},
}),
SenderSettings: serial.ToTypedMessage(&proxyman.SenderConfig{}),
},
},
}
@@ -105,7 +104,6 @@ func TestWireguard(t *testing.T) {
AllowedIps: []string{"0.0.0.0/0", "::0/0"},
}},
}),
SenderSettings: serial.ToTypedMessage(&proxyman.SenderConfig{}),
},
},
}
-32
View File
@@ -1,32 +0,0 @@
package internet
import (
"net/netip"
"slices"
"sync/atomic"
)
var skippedDNSServers atomic.Pointer[[]netip.Addr]
// SkipDNSServers has the queries Xray sends to the system's DNS servers on its
// own, like those of localdns, skip servers until it is called again. The DNS
// servers of a TUN are only meant for what goes through it: queried by Xray
// itself they lead back into it, or nowhere.
func SkipDNSServers(servers []netip.Addr) {
skipped := make([]netip.Addr, len(servers))
for i, server := range servers {
skipped[i] = server.Unmap()
}
skippedDNSServers.Store(&skipped)
}
// IsSkippedDNSServer reports whether address, a DNS server as host:port, is to
// be skipped, see SkipDNSServers.
func IsSkippedDNSServer(address string) bool {
skipped := skippedDNSServers.Load()
if skipped == nil {
return false
}
server, err := netip.ParseAddrPort(address)
return err == nil && slices.Contains(*skipped, server.Addr().Unmap())
}
-27
View File
@@ -1,27 +0,0 @@
package internet_test
import (
"net/netip"
"testing"
"github.com/xtls/xray-core/transport/internet"
)
func TestSkipDNSServers(t *testing.T) {
internet.SkipDNSServers([]netip.Addr{netip.MustParseAddr("::ffff:203.0.113.53"), netip.MustParseAddr("2001:db8::53")})
t.Cleanup(func() { internet.SkipDNSServers(nil) })
for address, want := range map[string]bool{
"203.0.113.53:53": true,
"[2001:db8::53]:53": true,
"198.51.100.53:53": false,
"localhost:53": false,
} {
if got := internet.IsSkippedDNSServer(address); got != want {
t.Errorf("IsSkippedDNSServer(%q) = %v, want %v", address, got, want)
}
}
internet.SkipDNSServers(nil)
if internet.IsSkippedDNSServer("203.0.113.53:53") {
t.Error("still skipped after SkipDNSServers(nil)")
}
}
+99 -225
View File
@@ -2,276 +2,105 @@ package finalmask
import (
"context"
"fmt"
"net"
"slices"
"github.com/xtls/xray-core/common"
"github.com/xtls/xray-core/common/buf"
"github.com/xtls/xray-core/common/errors"
"github.com/xtls/xray-core/common/net"
)
type Dialer struct {
DialTCP func(net.Destination) (net.Conn, error)
DialUDP func(net.Destination) (net.Conn, error)
type Udpmask interface {
WrapPacketConnClient(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error)
WrapPacketConnServer(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error)
}
type ListenConfig struct {
Listen func(net.Addr) (net.Listener, error)
ListenPacket func(net.Addr) (net.PacketConn, error)
type UdpmaskManager struct {
udpmasks []Udpmask
}
type TCPMask interface {
WrapConnClient(net.Conn, *net.Destination, *Dialer) (net.Conn, error)
WrapConnServer(net.Conn) (net.Conn, error)
// Listen(net.Listener) (net.Listener, error)
func NewUdpmaskManager(udpmasks []Udpmask) *UdpmaskManager {
slices.Reverse(udpmasks)
return &UdpmaskManager{udpmasks: udpmasks}
}
type UDPMask interface {
WrapPacketConnClient(net.PacketConn, *net.Destination, *Dialer) (net.PacketConn, error)
WrapPacketConnServer(net.PacketConn, net.Addr, *ListenConfig) (net.PacketConn, error)
}
type FinalMask struct {
tcpMasks []TCPMask
udpMasks []UDPMask
dialTCP func(context.Context, net.Destination) (net.Conn, error)
listen func(context.Context, net.Addr) (net.Listener, error)
dialUDP func(context.Context, net.Destination) (net.PacketConn, net.Addr, error)
listenPacket func(context.Context, net.Addr) (net.PacketConn, error)
}
func NewFinalMask(tcpMasks []TCPMask, udpMasks []UDPMask, dialTCP func(context.Context, net.Destination) (net.Conn, error), listen func(context.Context, net.Addr) (net.Listener, error), dialUDP func(context.Context, net.Destination) (net.PacketConn, net.Addr, error), listenPacket func(context.Context, net.Addr) (net.PacketConn, error)) *FinalMask {
slices.Reverse(tcpMasks)
slices.Reverse(udpMasks)
return &FinalMask{
tcpMasks: tcpMasks,
udpMasks: udpMasks,
dialTCP: dialTCP,
dialUDP: dialUDP,
listen: listen,
listenPacket: listenPacket,
}
}
func (fm *FinalMask) DialTCP(ctx context.Context, dest net.Destination) (net.Conn, error) {
if len(fm.tcpMasks) == 0 {
return fm.dialTCP(ctx, dest)
}
for i := range fm.tcpMasks {
if i > 0 {
if _, ok := fm.tcpMasks[i].(interface{ HandleDial() }); ok {
return nil, fmt.Errorf("incorrect index: %d %T", i, fm.tcpMasks[i])
}
}
}
var conn net.Conn
var err error
if _, ok := fm.tcpMasks[0].(interface{ HandleDial() }); !ok {
conn, err = fm.dialTCP(ctx, dest)
if err != nil {
return nil, err
}
}
dialer := &Dialer{
DialTCP: func(dest net.Destination) (net.Conn, error) {
return fm.dialTCP(ctx, dest)
},
DialUDP: func(dest net.Destination) (net.Conn, error) {
conn, addr, err := fm.dialUDP(ctx, dest)
if err != nil {
return nil, err
}
return &net.PacketConnWrapper{PacketConn: conn, Dest: addr}, err
},
}
for i := range fm.tcpMasks {
var newConn net.Conn
newConn, err = fm.tcpMasks[i].WrapConnClient(conn, &dest, dialer)
if err != nil {
common.CloseIfExists(conn)
return nil, err
}
conn = newConn
}
return conn, nil
}
func (fm *FinalMask) Listen(ctx context.Context, addr net.Addr) (net.Listener, error) {
if len(fm.tcpMasks) == 0 {
return fm.listen(ctx, addr)
}
off := 0
listener, err := fm.listen(ctx, addr)
if err != nil {
return nil, err
}
for i := range fm.tcpMasks {
if _, ok := fm.tcpMasks[i].(interface {
Listen(net.Listener) (net.Listener, error)
}); ok {
if i-off == 0 {
l, err := fm.tcpMasks[i].(interface {
Listen(net.Listener) (net.Listener, error)
}).Listen(listener)
if err != nil {
listener.Close()
return nil, err
}
listener = l
} else {
l, err := fm.tcpMasks[i].(interface {
Listen(net.Listener) (net.Listener, error)
}).Listen(&TCPListener{Listener: listener, tcpMasks: fm.tcpMasks[off:i]})
if err != nil {
listener.Close()
return nil, err
}
listener = l
}
off = i + 1
}
}
if off < len(fm.tcpMasks) {
return &TCPListener{Listener: listener, tcpMasks: fm.tcpMasks[off:]}, nil
}
return listener, nil
}
func (fm *FinalMask) DialUDP(ctx context.Context, dest net.Destination) (net.Conn, error) {
if len(fm.udpMasks) == 0 {
conn, addr, err := fm.dialUDP(ctx, dest)
if err != nil {
return nil, err
}
return &net.PacketConnWrapper{PacketConn: conn, Dest: addr}, nil
}
for i := range fm.udpMasks {
if i > 0 {
if _, ok := fm.udpMasks[i].(interface{ HandleDial() }); ok {
return nil, fmt.Errorf("incorrect index: %d %T", i, fm.udpMasks[i])
}
}
}
var conn net.PacketConn
var addr net.Addr
var err error
if _, ok := fm.udpMasks[0].(interface{ HandleDial() }); !ok {
conn, addr, err = fm.dialUDP(ctx, dest)
if err != nil {
return nil, err
}
}
dialer := &Dialer{
DialTCP: func(dest net.Destination) (net.Conn, error) {
return fm.dialTCP(ctx, dest)
},
DialUDP: func(dest net.Destination) (net.Conn, error) {
conn, addr, err := fm.dialUDP(ctx, dest)
if err != nil {
return nil, err
}
return &net.PacketConnWrapper{PacketConn: conn, Dest: addr}, err
},
}
func (m *UdpmaskManager) WrapPacketConnClient(raw net.PacketConn) (net.PacketConn, error) {
var sizes []int
var conns []net.PacketConn
for i := range fm.udpMasks {
var newConn net.PacketConn
if _, ok := fm.udpMasks[i].(interface{ HeaderConn() }); ok {
newConn, err = fm.udpMasks[i].WrapPacketConnClient(nil, nil, nil)
for i, mask := range m.udpmasks {
if _, ok := mask.(headerConn); ok {
conn, err := mask.WrapPacketConnClient(nil, i, len(m.udpmasks)-1)
if err != nil {
_ = conn.Close()
return nil, err
}
sizes = append(sizes, newConn.(interface{ Size() int }).Size())
conns = append(conns, newConn)
sizes = append(sizes, conn.(headerSize).Size())
conns = append(conns, conn)
} else {
if len(conns) > 0 {
conn = &headerManagerConn{PacketConn: conn, sizes: sizes, conns: conns}
raw = &headerManagerConn{sizes: sizes, conns: conns, PacketConn: raw}
sizes = nil
conns = nil
}
newConn, err = fm.udpMasks[i].WrapPacketConnClient(conn, &dest, dialer)
var err error
raw, err = mask.WrapPacketConnClient(raw, i, len(m.udpmasks)-1)
if err != nil {
common.CloseIfExists(conn)
return nil, err
}
conn = newConn
}
}
if len(conns) > 0 {
conn = &headerManagerConn{PacketConn: conn, sizes: sizes, conns: conns}
raw = &headerManagerConn{sizes: sizes, conns: conns, PacketConn: raw}
sizes = nil
conns = nil
}
if addr == nil {
addr = &net.UDPAddr{IP: []byte{0, 0, 0, 0}}
}
return &net.PacketConnWrapper{PacketConn: conn, Dest: addr}, nil
return raw, nil
}
func (fm *FinalMask) ListenPacket(ctx context.Context, addr net.Addr) (net.PacketConn, error) {
if len(fm.udpMasks) == 0 {
return fm.listenPacket(ctx, addr)
}
for i := range fm.udpMasks {
if i > 0 {
if _, ok := fm.udpMasks[i].(interface{ HandleListen() }); ok {
return nil, fmt.Errorf("incorrect index: %d %T", i, fm.udpMasks[i])
}
}
}
var conn net.PacketConn
var err error
if _, ok := fm.udpMasks[0].(interface{ HandleListen() }); !ok {
conn, err = fm.listenPacket(ctx, addr)
if err != nil {
return nil, err
}
}
lc := &ListenConfig{
Listen: func(addr net.Addr) (net.Listener, error) { return fm.listen(ctx, addr) },
ListenPacket: func(addr net.Addr) (net.PacketConn, error) { return fm.listenPacket(ctx, addr) },
}
func (m *UdpmaskManager) WrapPacketConnServer(raw net.PacketConn) (net.PacketConn, error) {
var sizes []int
var conns []net.PacketConn
for i := range fm.udpMasks {
var newConn net.PacketConn
if _, ok := fm.udpMasks[i].(interface{ HeaderConn() }); ok {
newConn, err = fm.udpMasks[i].WrapPacketConnServer(nil, nil, nil)
for i, mask := range m.udpmasks {
if _, ok := mask.(headerConn); ok {
conn, err := mask.WrapPacketConnServer(nil, i, len(m.udpmasks)-1)
if err != nil {
common.CloseIfExists(conn)
return nil, err
}
sizes = append(sizes, newConn.(interface{ Size() int }).Size())
conns = append(conns, newConn)
sizes = append(sizes, conn.(headerSize).Size())
conns = append(conns, conn)
} else {
if len(conns) > 0 {
conn = &headerManagerConn{PacketConn: conn, sizes: sizes, conns: conns}
raw = &headerManagerConn{sizes: sizes, conns: conns, PacketConn: raw}
sizes = nil
conns = nil
}
newConn, err = fm.udpMasks[i].WrapPacketConnServer(conn, addr, lc)
var err error
raw, err = mask.WrapPacketConnServer(raw, i, len(m.udpmasks)-1)
if err != nil {
common.CloseIfExists(conn)
return nil, err
}
conn = newConn
}
}
if len(conns) > 0 {
conn = &headerManagerConn{PacketConn: conn, sizes: sizes, conns: conns}
raw = &headerManagerConn{sizes: sizes, conns: conns, PacketConn: raw}
sizes = nil
conns = nil
}
return conn, nil
return raw, nil
}
const (
UDPSize = 4096
)
type headerConn interface {
HeaderConn()
}
type headerSize interface {
Size() int
}
type headerManagerConn struct {
net.PacketConn
@@ -362,27 +191,72 @@ func (c *headerManagerConn) WriteTo(p []byte, addr net.Addr) (n int, err error)
return len(p), nil
}
type TCPListener struct {
net.Listener
tcpMasks []TCPMask
type Tcpmask interface {
WrapConnClient(net.Conn) (net.Conn, error)
WrapConnServer(net.Conn) (net.Conn, error)
}
func (l *TCPListener) Accept() (net.Conn, error) {
type TcpmaskManager struct {
tcpmasks []Tcpmask
}
func NewTcpmaskManager(tcpmasks []Tcpmask) *TcpmaskManager {
slices.Reverse(tcpmasks)
return &TcpmaskManager{tcpmasks: tcpmasks}
}
func (m *TcpmaskManager) WrapConnClient(raw net.Conn) (net.Conn, error) {
var err error
for _, mask := range m.tcpmasks {
raw, err = mask.WrapConnClient(raw)
if err != nil {
return nil, err
}
}
return raw, nil
}
func (m *TcpmaskManager) WrapConnServer(raw net.Conn) (net.Conn, error) {
var err error
for _, mask := range m.tcpmasks {
raw, err = mask.WrapConnServer(raw)
if err != nil {
return nil, err
}
}
return raw, nil
}
func (m *TcpmaskManager) WrapListener(l net.Listener) (net.Listener, error) {
return NewTcpListener(m, l)
}
type tcpListener struct {
m *TcpmaskManager
net.Listener
}
func NewTcpListener(m *TcpmaskManager, l net.Listener) (net.Listener, error) {
return &tcpListener{
m: m,
Listener: l,
}, nil
}
func (l *tcpListener) Accept() (net.Conn, error) {
conn, err := l.Listener.Accept()
if err != nil {
return conn, err
}
for i := range l.tcpMasks {
var newConn net.Conn
newConn, err = l.tcpMasks[i].WrapConnServer(conn)
if err != nil {
_ = conn.Close()
return nil, err
}
conn = newConn
newConn, err := l.m.WrapConnServer(conn)
if err != nil {
errors.LogDebugInner(context.Background(), err, "mask err")
_ = conn.Close()
return nil, err
}
return conn, nil
return newConn, nil
}
type TcpMaskConn interface {
@@ -1,14 +1,11 @@
package fragment
import (
"github.com/xtls/xray-core/common/net"
"github.com/xtls/xray-core/transport/internet/finalmask"
)
import "net"
func (c *Config) WrapConnClient(conn net.Conn, dest *net.Destination, dialer *finalmask.Dialer) (net.Conn, error) {
return NewConnClient(c, conn, false)
func (c *Config) WrapConnClient(raw net.Conn) (net.Conn, error) {
return NewConnClient(c, raw, false)
}
func (c *Config) WrapConnServer(conn net.Conn) (net.Conn, error) {
return NewConnServer(c, conn, true)
func (c *Config) WrapConnServer(raw net.Conn) (net.Conn, error) {
return NewConnServer(c, raw, true)
}
@@ -1,30 +1,29 @@
package custom
import (
"github.com/xtls/xray-core/common/net"
"github.com/xtls/xray-core/transport/internet/finalmask"
"net"
)
func (c *TCPConfig) WrapConnClient(conn net.Conn, dest *net.Destination, dialer *finalmask.Dialer) (net.Conn, error) {
return NewConnClientTCP(c, conn)
func (c *TCPConfig) WrapConnClient(raw net.Conn) (net.Conn, error) {
return NewConnClientTCP(c, raw)
}
func (c *TCPConfig) WrapConnServer(conn net.Conn) (net.Conn, error) {
return NewConnServerTCP(c, conn)
func (c *TCPConfig) WrapConnServer(raw net.Conn) (net.Conn, error) {
return NewConnServerTCP(c, raw)
}
func (c *UDPConfig) WrapPacketConnClient(conn net.PacketConn, dest *net.Destination, dialer *finalmask.Dialer) (net.PacketConn, error) {
return NewConnClientUDP(c, conn)
func (c *UDPConfig) WrapPacketConnClient(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error) {
return NewConnClientUDP(c, raw)
}
func (c *UDPConfig) WrapPacketConnServer(conn net.PacketConn, addr net.Addr, lc *finalmask.ListenConfig) (net.PacketConn, error) {
return NewConnServerUDP(c, conn)
func (c *UDPConfig) WrapPacketConnServer(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error) {
return NewConnServerUDP(c, raw)
}
func (c *UDPStandaloneConfig) WrapPacketConnClient(conn net.PacketConn, dest *net.Destination, dialer *finalmask.Dialer) (net.PacketConn, error) {
return NewConnClientUDPStandalone(c, conn)
func (c *UDPStandaloneConfig) WrapPacketConnClient(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error) {
return NewConnClientUDPStandalone(c, raw)
}
func (c *UDPStandaloneConfig) WrapPacketConnServer(conn net.PacketConn, addr net.Addr, lc *finalmask.ListenConfig) (net.PacketConn, error) {
return NewConnServerUDPStandalone(c, conn)
func (c *UDPStandaloneConfig) WrapPacketConnServer(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error) {
return NewConnServerUDPStandalone(c, raw)
}
@@ -9,6 +9,8 @@ import (
"strings"
"testing"
"time"
"github.com/xtls/xray-core/transport/internet/finalmask"
)
func TestMetadataEvaluatorRejectsUnknownName(t *testing.T) {
@@ -154,7 +156,7 @@ func TestMetadataUDPStandaloneWriteUsesRemotePort(t *testing.T) {
}
defer serverRaw.Close()
client, err := cfg.WrapPacketConnClient(clientRaw, nil, nil)
client, err := finalmask.NewUdpmaskManager([]finalmask.Udpmask{cfg}).WrapPacketConnClient(clientRaw)
if err != nil {
t.Fatal(err)
}
@@ -299,7 +301,7 @@ func TestMetadataTCPHandshakeUsesEndpointPorts(t *testing.T) {
}
defer serverRaw.Close()
client, err := clientCfg.WrapConnClient(clientRaw, nil, nil)
client, err := clientCfg.WrapConnClient(clientRaw)
if err != nil {
t.Fatal(err)
}
@@ -5,6 +5,8 @@ import (
"net"
"testing"
"time"
"github.com/xtls/xray-core/transport/internet/finalmask"
)
func mustSendRecvUDP(t *testing.T, from net.PacketConn, to net.PacketConn, msg []byte) {
@@ -46,6 +48,7 @@ func TestStateUDPResponseReusesPriorCapturedValues(t *testing.T) {
},
},
}
maskManager := finalmask.NewUdpmaskManager([]finalmask.Udpmask{cfg})
clientRaw, err := net.ListenPacket("udp", "127.0.0.1:0")
if err != nil {
@@ -59,11 +62,11 @@ func TestStateUDPResponseReusesPriorCapturedValues(t *testing.T) {
}
defer serverRaw.Close()
client, err := cfg.WrapPacketConnClient(clientRaw, nil, nil)
client, err := maskManager.WrapPacketConnClient(clientRaw)
if err != nil {
t.Fatal(err)
}
server, err := cfg.WrapPacketConnServer(serverRaw, nil, nil)
server, err := maskManager.WrapPacketConnServer(serverRaw)
if err != nil {
t.Fatal(err)
}
@@ -37,7 +37,7 @@ func TestDSLTCPHandshakeReusesCapturedValue(t *testing.T) {
defer clientRaw.Close()
defer serverRaw.Close()
client, err := cfg.WrapConnClient(clientRaw, nil, nil)
client, err := cfg.WrapConnClient(clientRaw)
if err != nil {
t.Fatal(err)
}
@@ -117,7 +117,7 @@ func TestDSLTCPClientRejectsMismatchedResponseSequence(t *testing.T) {
defer clientRaw.Close()
defer serverRaw.Close()
client, err := clientCfg.WrapConnClient(clientRaw, nil, nil)
client, err := clientCfg.WrapConnClient(clientRaw)
if err != nil {
t.Fatal(err)
}

Some files were not shown because too many files have changed in this diff Show More