Compare commits

..
10 Commits
Author SHA1 Message Date
RPRXandGitHub b26a91de4f Xray-core v26.9.30
Sponsor & Donation & NFTs: https://github.com/XTLS/Xray-core/issues/3668
Project X Channel: https://t.me/projectXtls

Announcement of NFTs by Project X: https://github.com/XTLS/Xray-core/discussions/3633
Project X NFT: https://opensea.io/assets/ethereum/0x5ee362866001613093361eb8569d59c4141b76d1/1

VLESS Post-Quantum Encryption: https://github.com/XTLS/Xray-core/pull/5067
VLESS NFT: https://opensea.io/collection/vless

XHTTP: Beyond REALITY: https://github.com/XTLS/Xray-core/discussions/4113
REALITY NFT: https://opensea.io/assets/ethereum/0x5ee362866001613093361eb8569d59c4141b76d1/2
2026-09-30 07:40:04 +00:00
1f304916bd TUN inbound: Add autoSystemWfpBlockLeak on Windows (blocks "dns" and "misconfigtun" IPv4/IPv6 traffic leaks outside the TUN); Rename autoSystemDNS to autoSystemDnsToGateway on Linux (and change some behaviors) (#6853)
https://github.com/XTLS/Xray-core/pull/6853#issuecomment-5899791359
https://github.com/XTLS/Xray-core/pull/6853#issuecomment-5901287980
https://github.com/XTLS/Xray-core/pull/6853#issuecomment-5903680113
https://github.com/XTLS/Xray-core/pull/6853#issuecomment-5904123488
https://github.com/XTLS/Xray-core/pull/6853#issuecomment-5904647772
https://github.com/XTLS/Xray-core/pull/6853#issuecomment-5905047424

Fixes https://github.com/XTLS/Xray-core/issues/6454#issuecomment-5863800676

---------

Co-authored-by: Claude Opus 5.5 <noreply@anthropic.com>
2026-09-30 06:26:16 +00:00
风扇滑翔翼 0086362663 SS2022 proxy: Refactor to fix some issues (#6866)
https://github.com/XTLS/Xray-core/pull/6866#issuecomment-5904374529

Fixes https://github.com/XTLS/Xray-core/pull/6831#issuecomment-5884408501 and https://github.com/XTLS/Xray-core/pull/6866#issuecomment-5889568266
2026-09-30 13:16:27 +08:00
CluvexandGitHub e51b3c3621 Noise finalmask: type supports "exp" (#6862)
https://github.com/XTLS/Xray-core/pull/6844#issuecomment-5859092964
https://github.com/XTLS/Xray-core/pull/6844#issuecomment-5859855674
https://github.com/XTLS/Xray-core/pull/6862#issuecomment-5897471145
2026-09-30 03:11:12 +00:00
Artem LytkinandGitHub 6243d2a26e Geodata: Reduce matcher memory on mobile and desktop, skip regexes that cannot match (#6867)
https://github.com/XTLS/Xray-core/pull/6867#issuecomment-5895171934
2026-09-29 19:25:53 +00:00
风扇滑翔翼 35e616d3b9 Xray-core: Move PacketConnWrapper to common/net (#6854)
https://github.com/XTLS/Xray-core/pull/6854#issuecomment-5895541650

Fixes https://github.com/XTLS/Xray-core/issues/6849
2026-09-30 02:33:53 +08:00
v2rayandRPRX 08cb6e6bca WireGuard proxy: Fix potential startup races (#6852)
https://github.com/XTLS/Xray-core/pull/6852#issuecomment-5871816550

Fixes https://github.com/XTLS/Xray-core/issues/6850
2026-09-29 18:13:43 +00:00
wewon-backdownandRPRX 48ad0300ea API: xray api adu supports Hysteria (#6847)
https://github.com/XTLS/Xray-core/pull/6847#issuecomment-5895325181

---------

Co-authored-by: RPRX <63339210+RPRX@users.noreply.github.com>
2026-09-29 18:11:26 +00:00
0fc379203f HTTPUpgrade transport: Send some Sec-WebSocket-* headers (#6835)
https://github.com/XTLS/Xray-core/pull/6835#issuecomment-5853986428

---------

Co-authored-by: 风扇滑翔翼 <Fangliding.fshxy@outlook.com>
2026-09-29 17:25:26 +00:00
LjhAUMEMandGitHub fc8f8a451d XDNS finalmask: Refactor and new parameters (#6718)
https://github.com/XTLS/Xray-core/pull/6718#issuecomment-5894987590

Fixes https://github.com/XTLS/Xray-core/issues/6692
2026-09-29 17:16:59 +00:00
72 changed files with 7111 additions and 3868 deletions
+33 -50
View File
@@ -82,19 +82,10 @@ func (f *MphDomainMatcherFactory) BuildMatcher(rules []*DomainRule) (DomainMatch
} }
g.Add(m, uint32(i)) g.Add(m, uint32(i))
case *DomainRule_Geosite: case *DomainRule_Geosite:
domains, err := loadSiteWithAttrs(v.Geosite.File, v.Geosite.Code, v.Geosite.Attrs) err := loadSiteMatchers(v.Geosite, func(m strmatcher.Matcher) { g.Add(m, uint32(i)) })
if err != nil { if err != nil {
return nil, err 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: default:
panic("unknown domain rule type") panic("unknown domain rule type")
} }
@@ -108,12 +99,12 @@ func (f *MphDomainMatcherFactory) BuildMatcher(rules []*DomainRule) (DomainMatch
return g, nil return g, nil
} }
type CompactDomainMatcherFactory struct { type CompactMphDomainMatcherFactory struct {
sync.Mutex sync.Mutex
shared *utils.WeakCacheMap[string, strmatcher.LinearAnyMatcher] shared *utils.WeakCacheMap[string, strmatcher.MphValueMatcher]
} }
func (f *CompactDomainMatcherFactory) getOrCreateFrom(rule *GeoSiteRule) (strmatcher.MatcherSet, error) { func (f *CompactMphDomainMatcherFactory) getOrCreateFrom(rule *GeoSiteRule) (*strmatcher.MphValueMatcher, error) {
key := rule.File + ":" + rule.Code + "@" + rule.Attrs key := rule.File + ":" + rule.Code + "@" + rule.Attrs
f.Lock() f.Lock()
@@ -125,33 +116,23 @@ func (f *CompactDomainMatcherFactory) getOrCreateFrom(rule *GeoSiteRule) (strmat
} }
errors.LogDebug(context.Background(), "geodata geosite matcher cache MISS ", key) errors.LogDebug(context.Background(), "geodata geosite matcher cache MISS ", key)
s := strmatcher.NewLinearAnyMatcher() s := strmatcher.NewMphValueMatcher()
domains, err := loadSiteWithAttrs(rule.File, rule.Code, rule.Attrs) if err := loadSiteMatchers(rule, func(m strmatcher.Matcher) { s.Add(m, 0) }); err != nil {
if err != nil {
return nil, err return nil, err
} }
for i, d := range domains { if err := s.Build(); err != nil {
domains[i] = nil // peak mem return nil, err
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) f.shared.Store(key, s)
return s, err return s, nil
} }
// BuildMatcher implements DomainMatcherFactory. // BuildMatcher implements DomainMatcherFactory.
func (f *CompactDomainMatcherFactory) BuildMatcher(rules []*DomainRule) (DomainMatcher, error) { func (f *CompactMphDomainMatcherFactory) BuildMatcher(rules []*DomainRule) (DomainMatcher, error) {
if len(rules) == 0 { if len(rules) == 0 {
return nil, errors.New("empty domain rule list") return nil, errors.New("empty domain rule list")
} }
compact := &CompactDomainMatcher{ compact := new(CompactMphDomainMatcher)
matchers: make([]strmatcher.MatcherSet, 0, len(rules)),
values: make([]uint32, 0, len(rules)),
}
for i, r := range rules { for i, r := range rules {
switch v := r.Value.(type) { switch v := r.Value.(type) {
case *DomainRule_Custom: case *DomainRule_Custom:
@@ -168,8 +149,7 @@ func (f *CompactDomainMatcherFactory) BuildMatcher(rules []*DomainRule) (DomainM
if err != nil { if err != nil {
return nil, err return nil, err
} }
compact.matchers = append(compact.matchers, m) compact.combiner.Add(m, uint32(i))
compact.values = append(compact.values, uint32(i))
default: default:
panic("unknown domain rule type") panic("unknown domain rule type")
} }
@@ -177,37 +157,40 @@ func (f *CompactDomainMatcherFactory) BuildMatcher(rules []*DomainRule) (DomainM
return compact, nil return compact, nil
} }
type CompactDomainMatcher struct { type CompactMphDomainMatcher struct {
custom strmatcher.ValueMatcher custom strmatcher.ValueMatcher
matchers []strmatcher.MatcherSet combiner strmatcher.MphValueMatcherCombiner
values []uint32
} }
// Match implements DomainMatcher. // Match implements DomainMatcher.
func (c *CompactDomainMatcher) Match(input string) []uint32 { func (c *CompactMphDomainMatcher) Match(input string) []uint32 {
var result []uint32 result := c.combiner.Match(input)
if c.custom != nil { if c.custom != nil {
result = append(result, c.custom.Match(input)...) result = append(c.custom.Match(input), result...)
}
for i, m := range c.matchers {
if m.MatchAny(input) {
result = append(result, c.values[i])
}
} }
return result return result
} }
// MatchAny implements DomainMatcher. // MatchAny implements DomainMatcher.
func (c *CompactDomainMatcher) MatchAny(input string) bool { func (c *CompactMphDomainMatcher) MatchAny(input string) bool {
if c.custom != nil && c.custom.MatchAny(input) { if c.custom != nil && c.custom.MatchAny(input) {
return true return true
} }
for _, m := range c.matchers { return c.combiner.MatchAny(input)
if m.MatchAny(input) { }
return true
// 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)
} }
} i++
return false })
} }
func parseDomain(d *Domain) (strmatcher.Matcher, error) { func parseDomain(d *Domain) (strmatcher.Matcher, error) {
@@ -231,7 +214,7 @@ func parseDomain(d *Domain) (strmatcher.Matcher, error) {
func newDomainMatcherFactory() DomainMatcherFactory { func newDomainMatcherFactory() DomainMatcherFactory {
switch runtime.GOOS { switch runtime.GOOS {
case "ios", "android": case "ios", "android":
return &CompactDomainMatcherFactory{shared: utils.NewWeakCacheMap[string, strmatcher.LinearAnyMatcher]()} return &CompactMphDomainMatcherFactory{shared: utils.NewWeakCacheMap[string, strmatcher.MphValueMatcher]()}
default: default:
return &MphDomainMatcherFactory{shared: utils.NewWeakCacheMap[string, strmatcher.MphValueMatcher]()} return &MphDomainMatcherFactory{shared: utils.NewWeakCacheMap[string, strmatcher.MphValueMatcher]()}
} }
+76 -2
View File
@@ -4,6 +4,7 @@ import (
"path/filepath" "path/filepath"
"reflect" "reflect"
"slices" "slices"
"sync"
"testing" "testing"
"github.com/xtls/xray-core/common/geodata/strmatcher" "github.com/xtls/xray-core/common/geodata/strmatcher"
@@ -11,7 +12,7 @@ import (
) )
func TestCompactDomainMatcher_PreservesCustomRuleIndices(t *testing.T) { func TestCompactDomainMatcher_PreservesCustomRuleIndices(t *testing.T) {
factory := &CompactDomainMatcherFactory{shared: utils.NewWeakCacheMap[string, strmatcher.LinearAnyMatcher]()} factory := &CompactMphDomainMatcherFactory{shared: utils.NewWeakCacheMap[string, strmatcher.MphValueMatcher]()}
matcher, err := factory.BuildMatcher([]*DomainRule{ matcher, err := factory.BuildMatcher([]*DomainRule{
{Value: &DomainRule_Custom{Custom: &Domain{Type: Domain_Full, Value: "example.com"}}}, {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_Domain, Value: "example.com"}}},
@@ -32,7 +33,7 @@ func TestCompactDomainMatcher_PreservesCustomRuleIndices(t *testing.T) {
func TestCompactDomainMatcher_PreservesMixedRuleIndices(t *testing.T) { func TestCompactDomainMatcher_PreservesMixedRuleIndices(t *testing.T) {
t.Setenv("xray.location.asset", filepath.Join("..", "..", "resources")) t.Setenv("xray.location.asset", filepath.Join("..", "..", "resources"))
factory := &CompactDomainMatcherFactory{shared: utils.NewWeakCacheMap[string, strmatcher.LinearAnyMatcher]()} factory := &CompactMphDomainMatcherFactory{shared: utils.NewWeakCacheMap[string, strmatcher.MphValueMatcher]()}
matcher, err := factory.BuildMatcher([]*DomainRule{ matcher, err := factory.BuildMatcher([]*DomainRule{
{Value: &DomainRule_Geosite{Geosite: &GeoSiteRule{File: DefaultGeoSiteDat, Code: "CN"}}}, {Value: &DomainRule_Geosite{Geosite: &GeoSiteRule{File: DefaultGeoSiteDat, Code: "CN"}}},
{Value: &DomainRule_Custom{Custom: &Domain{Type: Domain_Full, Value: "163.com"}}}, {Value: &DomainRule_Custom{Custom: &Domain{Type: Domain_Full, Value: "163.com"}}},
@@ -72,3 +73,76 @@ func TestMphDomainMatcher_MatchReturnsDetachedSlice(t *testing.T) {
t.Fatalf("Match() after caller mutation = %v, want %v", gotAgain, []uint32{0, 1}) 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()
})
}
}
+213 -62
View File
@@ -5,11 +5,14 @@ import (
"bytes" "bytes"
"io" "io"
"runtime" "runtime"
"slices"
"strings" "strings"
"unicode/utf8"
"github.com/xtls/xray-core/common/errors" "github.com/xtls/xray-core/common/errors"
"github.com/xtls/xray-core/common/platform/filesystem" "github.com/xtls/xray-core/common/platform/filesystem"
"google.golang.org/protobuf/encoding/protowire"
"google.golang.org/protobuf/proto" "google.golang.org/protobuf/proto"
) )
@@ -52,17 +55,56 @@ func loadIP(file, code string) ([]*CIDR, error) {
return geoip.Cidr, nil return geoip.Cidr, nil
} }
func loadSite(file, code string) ([]*Domain, error) { // loadSite calls fn, in file order, with the type and value of every domain of the geosite code
bs, err := loadFile(file, 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)
if err != nil { if err != nil {
return nil, err return errors.New("failed to open ", file).Base(err)
} }
defer runtime.GC() // peak mem defer r.Close()
var geosite GeoSite br := bufio.NewReaderSize(r, 64*1024)
if err := proto.Unmarshal(bs, &geosite); err != nil { n, err := seek(br, []byte(code))
return nil, errors.New("error unmarshal Site in ", file, ":", code).Base(err) if err != nil {
return errors.New("failed to load code ", code, " from ", file).Base(err)
} }
return geosite.Domain, nil 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
} }
func decodeVarint(br *bufio.Reader) (uint64, error) { func decodeVarint(br *bufio.Reader) (uint64, error) {
@@ -82,68 +124,63 @@ func decodeVarint(br *bufio.Reader) (uint64, error) {
} }
func find(r io.Reader, code []byte, readBody bool) ([]byte, 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) codeL := len(code)
if codeL == 0 { if codeL == 0 {
return nil, errors.New("empty code") return 0, errors.New("empty code")
} }
br := bufio.NewReaderSize(r, 64*1024)
need := 2 + codeL // TODO: if code too long need := 2 + codeL // TODO: if code too long
prefixBuf := make([]byte, need)
for { for {
if _, err := br.ReadByte(); err != nil { if _, err := br.ReadByte(); err != nil {
return nil, err return 0, err
} }
x, err := decodeVarint(br) x, err := decodeVarint(br)
if err != nil { if err != nil {
return nil, err return 0, err
} }
bodyL := int(x) bodyL := int(x)
if bodyL <= 0 { if bodyL <= 0 {
return nil, errors.New("invalid body length: ", bodyL) return 0, errors.New("invalid body length: ", bodyL)
} }
prefixL := bodyL // Peek no more than the buffer holds: a code longer than the buffer cannot match a single
if prefixL > need { // length byte anyway, so a short peek only skips it, as base find (io.ReadFull) does.
prefixL = need prefix, err := br.Peek(min(bodyL, need, br.Size()))
} if err != nil {
prefix := prefixBuf[:prefixL] if err == io.EOF && len(prefix) > 0 {
if _, err := io.ReadFull(br, prefix); err != nil { err = io.ErrUnexpectedEOF // as io.ReadFull
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) {
remain := bodyL - prefixL return bodyL, nil
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 {
if remain > 0 { return 0, err
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 { type AttributeMatcher interface {
Match(*Domain) bool Match(*Domain) bool
} }
@@ -185,23 +222,137 @@ func NewAllAttrsMatcher(attrs string) AttributeMatcher {
return m return m
} }
func loadSiteWithAttrs(file, code, attrs string) ([]*Domain, error) { var errInvalidUTF8 = errors.New("string field contains invalid UTF-8")
domains, err := loadSite(file, code)
if err != nil {
return nil, err
}
matcher := NewAllAttrsMatcher(attrs) type siteDecoder struct {
if matcher == nil { want []string
return domains, nil has []bool
} fn func(Domain_Type, []byte)
}
filtered := make([]*Domain, 0, len(domains)) func newSiteDecoder(attrs string, fn func(Domain_Type, []byte)) *siteDecoder {
for _, d := range domains { d := &siteDecoder{fn: fn}
if matcher.Match(d) { if attrs != "" {
filtered = append(filtered, d) d.want = strings.Split(attrs, "@")
d.has = make([]bool, len(d.want))
}
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
return filtered, 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
} }
+283
View File
@@ -0,0 +1,283 @@
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")
}
}
+8 -12
View File
@@ -52,7 +52,9 @@ func (g *MphIndexMatcher) Add(matcher Matcher) uint32 {
func (g *MphIndexMatcher) Build() error { func (g *MphIndexMatcher) Build() error {
if g.mph != nil { if g.mph != nil {
runtime.GC() // peak mem runtime.GC() // peak mem
g.mph.Build() if err := g.mph.Build(); err != nil {
return err
}
} }
runtime.GC() // peak mem runtime.GC() // peak mem
if g.ac != nil { if g.ac != nil {
@@ -64,23 +66,17 @@ func (g *MphIndexMatcher) Build() error {
// Match implements IndexMatcher.Match. // Match implements IndexMatcher.Match.
func (g *MphIndexMatcher) Match(input string) []uint32 { func (g *MphIndexMatcher) Match(input string) []uint32 {
result := make([][]uint32, 0, 5) var result []uint32
if g.mph != nil { if g.mph != nil {
if matches := g.mph.Match(input); len(matches) > 0 { result = g.mph.Match(input) // a new slice, returned without another copy
result = append(result, matches)
}
} }
if g.ac != nil { if g.ac != nil {
if matches := g.ac.Match(input); len(matches) > 0 { result = append(result, g.ac.Match(input)...)
result = append(result, matches)
}
} }
if g.regex != nil { if g.regex != nil {
if matches := g.regex.Match(input); len(matches) > 0 { result = append(result, g.regex.Match(input)...)
result = append(result, matches)
}
} }
return CompositeMatches(result) return result
} }
// MatchAny implements IndexMatcher.MatchAny. // MatchAny implements IndexMatcher.MatchAny.
@@ -78,6 +78,10 @@ func TestMphIndexMatcher(t *testing.T) {
Input: "example.com", Input: "example.com",
Output: []uint32{10, 4}, Output: []uint32{10, 4},
}, },
{
Input: "apis.org",
Output: []uint32{2, 6},
},
} }
matcherGroup := NewMphIndexMatcher() matcherGroup := NewMphIndexMatcher()
for _, rule := range rules { for _, rule := range rules {
@@ -87,8 +91,13 @@ func TestMphIndexMatcher(t *testing.T) {
} }
matcherGroup.Build() matcherGroup.Build()
for _, test := range cases { for _, test := range cases {
if m := matcherGroup.Match(test.Input); !reflect.DeepEqual(m, test.Output) { m := matcherGroup.Match(test.Input)
if !reflect.DeepEqual(m, test.Output) {
t.Error("unexpected output: ", m, " for test case ", test) 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)
}
} }
} }
+375 -166
View File
@@ -1,231 +1,440 @@
package strmatcher package strmatcher
import ( import (
"bytes"
"cmp"
"encoding/binary"
"errors" "errors"
"math" "math"
"math/bits" "slices"
"runtime"
"sort"
"strings" "strings"
"unsafe" "unsafe"
) )
// PrimeRK is the prime base used in Rabin-Karp algorithm. // Flags of a level1 slot, stored above the record offset.
const PrimeRK = 16777619
// 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
}
// 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
}
const ( const (
mphMatchTypeCount = 2 // Full and Domain 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
) )
type mphRuleInfo struct { // Kinds of an added pattern, indexes of mphKinds.
rollingHash uint32 const (
matchers [mphMatchTypeCount][]uint32 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
} }
// MphMatcherGroup is an implementation of MatcherGroup. // MphMatcherGroup is an implementation of MatcherGroup for Full and Domain matchers.
// It implements Rabin-Karp algorithm and minimal perfect hash table for Full and Domain matcher. // 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 { type MphMatcherGroup struct {
patterns string // All rule patterns concatenated arena string
patternOffs []uint32 // RuleIdx -> patterns[patternOffs[i]:patternOffs[i+1]], index 0 reserved for failed lookup level0 []uint16 // bucket -> seed
values []uint32 // All registered matcher values concatenated level1 []uint32 // slot -> flags | record offset
valueOffs []uint32 // RuleIdx -> values[valueOffs[i]:valueOffs[i+1]] (Full Matcher takes precedence) fp []uint8 // slot -> low byte of its pattern's hash, rejects most misses without reading arena
level0 []uint32 // RollingHash & Mask -> seed for Memhash n0, n1 uint32
level0Mask uint32 // Mask restricting RollingHash to 0 ~ len(level0) mul uint64 // multiplier of the suffix hash
level1 []uint32 // Memhash<seed> & Mask -> stored index for rules single uint32 // the only value if !multi
level1Mask uint32 // Mask for restricting Memhash<seed> to 0 ~ len(level1) multi bool
rules []string // RuleIdx -> pattern string, only used for building
ruleInfos *map[string]mphRuleInfo buf []byte // build only, patterns in Add order
entries []mphEntry
} }
func NewMphMatcherGroup() *MphMatcherGroup { func NewMphMatcherGroup() *MphMatcherGroup {
return &MphMatcherGroup{ return new(MphMatcherGroup)
rules: []string{""},
level0: nil,
level0Mask: 0,
level1: nil,
level1Mask: 0,
ruleInfos: &map[string]mphRuleInfo{}, // Only used for building, destroyed after build complete
}
} }
// AddFullMatcher implements MatcherGroupForFull. // AddFullMatcher implements MatcherGroupForFull.
func (g *MphMatcherGroup) AddFullMatcher(matcher FullMatcher, value uint32) { func (g *MphMatcherGroup) AddFullMatcher(matcher FullMatcher, value uint32) {
pattern := strings.ToLower(matcher.Pattern()) g.add(matcher.Pattern(), mphKindFull, value)
g.addPattern(0, "", pattern, matcher.Type(), value)
} }
// AddDomainMatcher implements MatcherGroupForDomain. // AddDomainMatcher implements MatcherGroupForDomain.
func (g *MphMatcherGroup) AddDomainMatcher(matcher DomainMatcher, value uint32) { func (g *MphMatcherGroup) AddDomainMatcher(matcher DomainMatcher, value uint32) {
pattern := strings.ToLower(matcher.Pattern()) g.add(matcher.Pattern(), mphKindDomain, value)
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) addPattern(suffixHash uint32, suffixPattern string, pattern string, matcherType Type, value uint32) uint32 { func (g *MphMatcherGroup) add(pattern string, kind uint8, value uint32) {
fullPattern := pattern + suffixPattern if g.arena != "" {
info, found := (*g.ruleInfos)[fullPattern] panic(errMphBuilt)
if !found { }
info = mphRuleInfo{rollingHash: RollingHash(suffixHash, pattern)} pattern = strings.ToLower(pattern)
g.rules = append(g.rules, fullPattern) 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})
} }
info.matchers[matcherType] = append(info.matchers[matcherType], value)
(*g.ruleInfos)[fullPattern] = info
return info.rollingHash
} }
// Build builds a minimal perfect hash table for insert rules. func (g *MphMatcherGroup) key(i uint32) []byte {
// Algorithm used: Hash, displace, and compress. See http://cmph.sourceforge.net/papers/esa09.pdf 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.
func (g *MphMatcherGroup) Build() error { func (g *MphMatcherGroup) Build() error {
ruleCount := len(*g.ruleInfos) if g.arena != "" {
g.level0 = make([]uint32, nextPow2(ruleCount/4)) return errMphBuilt
g.level0Mask = uint32(len(g.level0) - 1)
g.level1 = make([]uint32, nextPow2(ruleCount))
g.level1Mask = uint32(len(g.level1) - 1)
// Flatten patterns and values so the built group has no per-rule objects
valueCount := 0
for _, ruleInfo := range *g.ruleInfos {
valueCount += len(ruleInfo.matchers[Full]) + len(ruleInfo.matchers[Domain])
} }
g.patterns = strings.Join(g.rules, "") if uint64(len(g.buf)) > math.MaxUint32 {
if uint64(len(g.patterns)) > math.MaxUint32 || uint64(valueCount) > math.MaxUint32 {
return errors.New("too many rules for MphMatcherGroup") return errors.New("too many rules for MphMatcherGroup")
} }
g.patternOffs = make([]uint32, len(g.rules)+1) recs := g.writeRecords()
g.values = make([]uint32, 0, valueCount) if len(g.arena) > mphOffMask {
g.valueOffs = make([]uint32, len(g.rules)+1) return errors.New("too many rules for MphMatcherGroup")
// 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.patternOffs[ruleIdx+1] = g.patternOffs[ruleIdx] + uint32(len(g.rules[ruleIdx]))
g.values = append(append(g.values, ruleInfo.matchers[Full]...), ruleInfo.matchers[Domain]...)
g.valueOffs[ruleIdx+1] = uint32(len(g.values))
} }
g.rules = nil hashes := make([]uint64, len(recs))
g.ruleInfos = nil // Set ruleInfos nil to release memory for _, mul := range mphMultipliers {
runtime.GC() // peak mem for i, rec := range recs {
hashes[i] = mphMix(mphHash(mul, g.recKey(rec)))
// Sort buckets in descending order with respect to each bucket's size }
bucketIdxs := make([]int, len(buckets)) g.mul = mul
for bucketIdx := range buckets { if err := g.place(recs, hashes); err != errMphCollision {
bucketIdxs[bucketIdx] = bucketIdx return err
}
} }
sort.Slice(bucketIdxs, func(i, j int) bool { return len(buckets[bucketIdxs[i]]) > len(buckets[bucketIdxs[j]]) }) return errMphCollision
}
// Exercise Hash, Displace, and Compress algorithm to construct minimal perfect hash table // writeRecords writes one record per distinct pattern to arena and returns flags | offset of each.
occupied := make([]bool, len(g.level1)) // Whether a second-level hash has been already used func (g *MphMatcherGroup) writeRecords() []uint32 {
hashedBucket := make([]uint32, 0, 4) // Second-level hashes for each rule in a specific bucket g.multi = false
for _, bucketIdx := range bucketIdxs { if len(g.entries) > 0 {
bucket := buckets[bucketIdx] g.single = g.entries[0].value
hashedBucket = hashedBucket[:0] for _, e := range g.entries {
seed := uint32(0) if e.value != g.single {
for len(hashedBucket) != len(bucket) { g.multi = true
for _, ruleIdx := range bucket { break
memHash := MemHash(seed, g.pattern(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)
} }
} }
g.level0[bucketIdx] = seed // Displacement value for this bucket }
// 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))
})
size := len(g.buf) + len(g.entries) + 2
if g.multi {
size += 3 * len(g.entries)
}
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))
}
}
}
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")
} }
return nil return nil
} }
func (g *MphMatcherGroup) pattern(ruleIdx uint32) string { // mphHash is the suffix hash of s, taken from the right: the hash of s[i:] is the state after reading s[i].
return g.patterns[g.patternOffs[ruleIdx]:g.patternOffs[ruleIdx+1]] 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
} }
// valuesOf caps the capacity, so appending to a Match result can't overwrite the next rule's values. // mphMix spreads the weak low bits of a suffix hash.
func (g *MphMatcherGroup) valuesOf(ruleIdx uint32) []uint32 { func mphMix(h uint64) uint64 {
start, end := g.valueOffs[ruleIdx], g.valueOffs[ruleIdx+1] h ^= h >> 32
return g.values[start:end:end] h *= 0xd6e8feb86659fd93
return h ^ h>>32
} }
// Lookup searches for input in minimal perfect hash table and returns its index. 0 indicates not found. func (g *MphMatcherGroup) bucket(f uint64) uint32 {
func (g *MphMatcherGroup) Lookup(rollingHash uint32, input string) uint32 { return uint32(((f >> 32) * uint64(g.n0)) >> 32)
i0 := rollingHash & g.level0Mask }
seed := g.level0[i0]
i1 := MemHash(seed, input) & g.level1Mask func (g *MphMatcherGroup) slot(f uint64, seed uint16) uint32 {
n := g.level1[i1] x := ((f ^ uint64(seed)*0x9e3779b97f4a7c15) * 0xc4ceb9fe1a85ec53) >> 32
// Build only puts valid rule indices in level1, so n+1 < len(patternOffs) and the span is inside patterns. return uint32((x * uint64(g.n1)) >> 32)
// Skip the bounds checks, they made this hot path measurably slower than indexing a []string }
offs := (*[2]uint32)(unsafe.Add(unsafe.Pointer(unsafe.SliceData(g.patternOffs)), uintptr(n)*4))
if start := offs[0]; int(offs[1]-start) == len(input) && unsafe.String((*byte)(unsafe.Add(unsafe.Pointer(unsafe.StringData(g.patterns)), start)), len(input)) == input { func (g *MphMatcherGroup) uvarint(p uint32) (x, next uint32) {
return n 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
} }
return 0 return 0
} }
// Match implements MatcherGroup.Match. // appendValues appends the values of record e for the flags in want, in mphKinds order.
func (g *MphMatcherGroup) Match(input string) []uint32 { func (g *MphMatcherGroup) appendValues(dst []uint32, e, want uint32) []uint32 {
matches := make([][]uint32, 0, 5) if !g.multi {
hash := uint32(0) for _, flag := range mphKinds {
for i := len(input) - 1; i >= 0; i-- { if e&want&flag != 0 {
hash = hash*PrimeRK + uint32(input[i]) dst = append(dst, g.single)
if input[i] == '.' { }
if mphIdx := g.Lookup(hash, input[i:]); mphIdx != 0 { }
matches = append(matches, g.valuesOf(mphIdx)) 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)
} }
} }
} }
if mphIdx := g.Lookup(hash, input); mphIdx != 0 { return dst
matches = append(matches, g.valuesOf(mphIdx)) }
// Match implements MatcherGroup.Match. Values of an exact match come first (Full, then Domain), then those of
// the parent domains, nearest first.
func (g *MphMatcherGroup) Match(input string) []uint32 {
var stack [8]uint32
parents := stack[:0] // TLD side first
h, mul := uint64(0), g.mul
for i := len(input) - 1; i >= 0; i-- {
if input[i] == '.' {
if e := g.lookup(h, input[i+1:]); e&(mphDomain|mphParent) != 0 {
parents = append(parents, e)
}
}
h = h*mul + uint64(input[i])
} }
return CompositeMatchesReverse(matches) exact := g.lookup(h, input)
if exact&(mphFull|mphDomain) == 0 && len(parents) == 0 {
return nil
}
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
} }
// MatchAny implements MatcherGroup.MatchAny. // MatchAny implements MatcherGroup.MatchAny.
func (g *MphMatcherGroup) MatchAny(input string) bool { func (g *MphMatcherGroup) MatchAny(input string) bool {
hash := uint32(0) 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)
for i := len(input) - 1; i >= 0; i-- { for i := len(input) - 1; i >= 0; i-- {
hash = hash*PrimeRK + uint32(input[i])
if input[i] == '.' { if input[i] == '.' {
if g.Lookup(hash, input[i:]) != 0 { dst = append(dst, mphSuffix{h, i + 1})
return true }
} h = h*mul + uint64(input[i])
}
return dst, h
}
// 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
}
for _, p := range parents {
if g.lookup(p.h, input[p.off:])&(mphDomain|mphParent) != 0 {
return true
} }
} }
return g.Lookup(hash, input) != 0 return g.lookup(h, input)&(mphFull|mphDomain) != 0
} }
func nextPow2(v int) int {
if v <= 1 {
return 1
}
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
@@ -0,0 +1,108 @@
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)
}
}
@@ -4,6 +4,7 @@ import (
"math/rand" "math/rand"
"reflect" "reflect"
"slices" "slices"
"strings"
"testing" "testing"
"github.com/xtls/xray-core/common" "github.com/xtls/xray-core/common"
@@ -304,7 +305,7 @@ func TestMphMatcherGroupRandom(t *testing.T) {
domain["."+p] = append(domain["."+p], value) domain["."+p] = append(domain["."+p], value)
} }
} }
g.Build() common.Must(g.Build())
for _, input := range inputs { for _, input := range inputs {
keys := []string{input} // Whole input first, then "." suffixes from longest to shortest keys := []string{input} // Whole input first, then "." suffixes from longest to shortest
for i := range len(input) { for i := range len(input) {
@@ -316,7 +317,10 @@ func TestMphMatcherGroupRandom(t *testing.T) {
for _, k := range keys { for _, k := range keys {
want = append(append(want, full[k]...), domain[k]...) want = append(append(want, full[k]...), domain[k]...)
} }
if m := g.Match(input); !slices.Equal(m, want) { // 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) t.Fatalf("seed %d: Match(%q) = %v, want %v", seed, input, m, want)
} }
if m := g.MatchAny(input); m != (len(want) > 0) { if m := g.MatchAny(input); m != (len(want) > 0) {
@@ -338,3 +342,79 @@ func TestMphMatcherGroupAppend(t *testing.T) {
t.Error("expect [2], but ", m) 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)
}
+237 -1
View File
@@ -2,10 +2,12 @@ package strmatcher
import ( import (
"errors" "errors"
"math/bits"
"regexp" "regexp"
"regexp/syntax" "regexp/syntax"
"slices" "slices"
"strings" "strings"
"unicode"
"unicode/utf8" "unicode/utf8"
"golang.org/x/net/idna" "golang.org/x/net/idna"
@@ -75,7 +77,9 @@ func (m SubstrMatcher) Match(s string) bool {
// RegexMatcher is an implementation of Matcher. // RegexMatcher is an implementation of Matcher.
type RegexMatcher struct { type RegexMatcher struct {
pattern *regexp.Regexp pattern *regexp.Regexp
literals []string // every match contains all of them, longest first 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) { func newRegexMatcher(pattern string) (Matcher, error) {
@@ -87,10 +91,239 @@ func newRegexMatcher(pattern string) (Matcher, error) {
if re, err := syntax.Parse(pattern, syntax.Perl); err == nil { // same flags as regexp.Compile if re, err := syntax.Parse(pattern, syntax.Perl); err == nil { // same flags as regexp.Compile
m.literals = requiredLiterals(re, nil) m.literals = requiredLiterals(re, nil)
slices.SortStableFunc(m.literals, func(a, b string) int { return len(b) - len(a) }) slices.SortStableFunc(m.literals, func(a, b string) int { return len(b) - len(a) })
m.tail, m.rest = tailGuard(re)
} }
return m, nil 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. // requiredLiterals appends to dst the case-sensitive strings that every match of re contains.
func requiredLiterals(re *syntax.Regexp, dst []string) []string { func requiredLiterals(re *syntax.Regexp, dst []string) []string {
switch re.Op { switch re.Op {
@@ -126,6 +359,9 @@ func (m *RegexMatcher) String() string {
} }
func (m *RegexMatcher) Match(s string) bool { func (m *RegexMatcher) Match(s string) bool {
if !m.mayMatch(s) {
return false
}
for _, l := range m.literals { for _, l := range m.literals {
if !strings.Contains(s, l) { if !strings.Contains(s, l) {
return false return false
@@ -1,9 +1,16 @@
package strmatcher package strmatcher
import ( import (
"hash/fnv"
"math/rand/v2"
"regexp" "regexp"
"regexp/syntax"
"slices" "slices"
"strconv"
"strings"
"testing" "testing"
"unicode"
"unicode/utf8"
) )
var regexLiteralCases = []struct { var regexLiteralCases = []struct {
@@ -37,6 +44,147 @@ func TestRegexRequiredLiterals(t *testing.T) {
} }
} }
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) { func FuzzRegexMatcher(f *testing.F) {
inputs := []string{ inputs := []string{
"", "x", "yy", "abd", "ccc", "ABC", "abCDef", "abcdef", "abababcc", "a.b", "a\xffb", "a\uFFFDb", "", "x", "yy", "abd", "ccc", "ABC", "abCDef", "abcdef", "abababcc", "a.b", "a\xffb", "a\uFFFDb",
@@ -47,14 +195,39 @@ func FuzzRegexMatcher(f *testing.F) {
f.Add(test.pattern, s) 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) { f.Fuzz(func(t *testing.T, pattern, s string) {
re, err := regexp.Compile(pattern) re, err := regexp.Compile(pattern)
if err != nil { if err != nil {
return return
} }
m, _ := newRegexMatcher(pattern) m, _ := newRegexMatcher(pattern)
if got, want := m.Match(s), re.MatchString(s); got != want { check := func(s string) {
t.Errorf("pattern %q, input %q: got %v, want %v", pattern, s, got, want) 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:])
}
} }
}) })
} }
+67 -12
View File
@@ -46,7 +46,9 @@ func (g *MphValueMatcher) Add(matcher Matcher, value uint32) {
func (g *MphValueMatcher) Build() error { func (g *MphValueMatcher) Build() error {
if g.mph != nil { if g.mph != nil {
runtime.GC() // peak mem runtime.GC() // peak mem
g.mph.Build() if err := g.mph.Build(); err != nil {
return err
}
} }
runtime.GC() // peak mem runtime.GC() // peak mem
if g.ac != nil { if g.ac != nil {
@@ -58,23 +60,17 @@ func (g *MphValueMatcher) Build() error {
// Match implements ValueMatcher.Match. // Match implements ValueMatcher.Match.
func (g *MphValueMatcher) Match(input string) []uint32 { func (g *MphValueMatcher) Match(input string) []uint32 {
result := make([][]uint32, 0, 5) var result []uint32
if g.mph != nil { if g.mph != nil {
if matches := g.mph.Match(input); len(matches) > 0 { result = g.mph.Match(input) // a new slice, returned without another copy
result = append(result, matches)
}
} }
if g.ac != nil { if g.ac != nil {
if matches := g.ac.Match(input); len(matches) > 0 { result = append(result, g.ac.Match(input)...)
result = append(result, matches)
}
} }
if g.regex != nil { if g.regex != nil {
if matches := g.regex.Match(input); len(matches) > 0 { result = append(result, g.regex.Match(input)...)
result = append(result, matches)
}
} }
return CompositeMatches(result) return result
} }
// MatchAny implements ValueMatcher.MatchAny. // MatchAny implements ValueMatcher.MatchAny.
@@ -87,3 +83,62 @@ func (g *MphValueMatcher) MatchAny(input string) bool {
} }
return g.regex != nil && g.regex.MatchAny(input) 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
}
+1 -1
View File
@@ -20,7 +20,7 @@ import (
var ( var (
Version_x byte = 26 Version_x byte = 26
Version_y byte = 9 Version_y byte = 9
Version_z byte = 9 Version_z byte = 30
) )
var ( var (
+3
View File
@@ -97,6 +97,9 @@ func New() *Client {
r := &net.Resolver{ r := &net.Resolver{
PreferGo: true, PreferGo: true,
Dial: func(ctx context.Context, network, address string) (net.Conn, error) { 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) return d.DialContext(ctx, network, address)
}, },
} }
+23
View File
@@ -0,0 +1,23 @@
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()
}
+178 -25
View File
@@ -1,6 +1,7 @@
package conf package conf
import ( import (
"context"
"crypto/x509" "crypto/x509"
"encoding/base64" "encoding/base64"
"encoding/hex" "encoding/hex"
@@ -14,6 +15,7 @@ import (
googleuuid "github.com/google/uuid" googleuuid "github.com/google/uuid"
"github.com/xtls/xray-core/common/errors" "github.com/xtls/xray-core/common/errors"
"github.com/xtls/xray-core/common/net" "github.com/xtls/xray-core/common/net"
"github.com/xtls/xray-core/common/serial"
"github.com/xtls/xray-core/transport/internet/finalmask/fragment" "github.com/xtls/xray-core/transport/internet/finalmask/fragment"
"github.com/xtls/xray-core/transport/internet/finalmask/header/custom" "github.com/xtls/xray-core/transport/internet/finalmask/header/custom"
"github.com/xtls/xray-core/transport/internet/finalmask/mkcp/aes128gcm" "github.com/xtls/xray-core/transport/internet/finalmask/mkcp/aes128gcm"
@@ -81,7 +83,7 @@ var (
"noise": func() interface{} { return new(NoiseMask) }, "noise": func() interface{} { return new(NoiseMask) },
"salamander": func() interface{} { return new(Salamander) }, "salamander": func() interface{} { return new(Salamander) },
"sudoku": func() interface{} { return new(Sudoku) }, "sudoku": func() interface{} { return new(Sudoku) },
"xdns": func() interface{} { return new(Xdns) }, "xdns": func() interface{} { return new(XDNS) },
"xicmp": func() interface{} { return new(Xicmp) }, "xicmp": func() interface{} { return new(Xicmp) },
"realm": func() interface{} { return new(Realm) }, "realm": func() interface{} { return new(Realm) },
"udphop": func() interface{} { return new(UDPHop) }, "udphop": func() interface{} { return new(UDPHop) },
@@ -308,14 +310,27 @@ type NoiseMask struct {
} }
func (c *NoiseMask) Build() (proto.Message, error) { func (c *NoiseMask) Build() (proto.Message, error) {
noiseSlice := make([]*noise.Item, 0, len(c.Noise))
for _, item := range c.Noise { for _, item := range c.Noise {
if len(item.Packet) > 0 && item.Rand.To > 0 { if len(item.Packet) > 0 && item.Rand.To > 0 {
return nil, errors.New("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
noiseSlice := make([]*noise.Item, 0, len(c.Noise)) if err := json.Unmarshal(item.Packet, &exp); err != nil {
for _, item := range c.Noise { 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
}
if item.RandRange == nil { if item.RandRange == nil {
item.RandRange = &Int32Range{From: 0, To: 255} item.RandRange = &Int32Range{From: 0, To: 255}
} }
@@ -344,6 +359,88 @@ func (c *NoiseMask) Build() (proto.Message, error) {
}, nil }, 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 { type UDPItem struct {
Rand int32 `json:"rand"` Rand int32 `json:"rand"`
RandRange *Int32Range `json:"randRange"` RandRange *Int32Range `json:"randRange"`
@@ -694,32 +791,88 @@ func (c *Sudoku) Build() (proto.Message, error) {
}, nil }, nil
} }
type Xdns struct { type XDNSDomain struct {
Domain json.RawMessage `json:"domain"` Name string `json:"name"`
LenLimit int32 `json:"lenLimit"`
Domains []string `json:"domains"` LabelLimit int32 `json:"labelLimit"`
Resolvers []string `json:"resolvers"` Types []int32 `json:"types"`
Edns0 int32 `json:"edns0"`
} }
func (c *Xdns) Build() (proto.Message, error) { type XDNSResolverTCP struct {
if c.Domain != nil { Addr string `json:"addr"`
return nil, errors.PrintRemovedFeatureError("domain", "domains(server) & resolvers(client)") }
}
if len(c.Domains) == 0 && len(c.Resolvers) == 0 { func (c *XDNSResolverTCP) Build() (proto.Message, error) {
return nil, errors.New("empty domains & empty resolvers") return &xdns.TCPResolverProto{Addr: c.Addr}, nil
} }
for _, r := range c.Resolvers { type XDNSResolverUDP struct {
if !strings.Contains(r, "+udp://") { Addr string `json:"addr"`
return nil, errors.New("invalid resolver ", r) }
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,
})
} }
for i := range c.Resolvers {
return &xdns.Config{ config, err := xdnsLoader.LoadWithID(c.Resolvers[i].Settings, c.Resolvers[i].Type)
Domains: c.Domains, if err != nil {
Resolvers: c.Resolvers, return nil, err
}, nil }
pm, err := config.(interface{ Build() (proto.Message, error) }).Build()
if err != nil {
return nil, err
}
resolvers = append(resolvers, serial.ToTypedMessage(pm))
}
if c.ExtraPoll < 0 || c.ExtraPoll > 3 {
return nil, errors.New("c.ExtraPoll < 0 || c.ExtraPoll > 3")
}
return &xdns.Config{Domains: domains, Resolvers: resolvers, ExtraPoll: c.ExtraPoll}, nil
} }
type XMC struct { type XMC struct {
@@ -0,0 +1,136 @@
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])
}
}
+32 -2
View File
@@ -5,8 +5,12 @@ import (
"fmt" "fmt"
"math/big" "math/big"
"net" "net"
"runtime"
"slices"
"strconv" "strconv"
"strings"
"github.com/xtls/xray-core/common/errors"
"github.com/xtls/xray-core/proxy/tun" "github.com/xtls/xray-core/proxy/tun"
"google.golang.org/protobuf/proto" "google.golang.org/protobuf/proto"
) )
@@ -20,7 +24,8 @@ type TunConfig struct {
UserLevel uint32 `json:"userLevel"` UserLevel uint32 `json:"userLevel"`
AutoSystemRoutingTable []string `json:"autoSystemRoutingTable"` AutoSystemRoutingTable []string `json:"autoSystemRoutingTable"`
AutoOutboundsInterface *string `json:"autoOutboundsInterface"` AutoOutboundsInterface *string `json:"autoOutboundsInterface"`
AutoSystemDNS bool `json:"autoSystemDNS"` AutoSystemDnsToGateway bool `json:"autoSystemDnsToGateway"`
AutoSystemWfpBlockLeak []string `json:"autoSystemWfpBlockLeak"`
} }
func (v *TunConfig) Build() (proto.Message, error) { func (v *TunConfig) Build() (proto.Message, error) {
@@ -32,7 +37,32 @@ func (v *TunConfig) Build() (proto.Message, error) {
DNS: v.DNS, DNS: v.DNS,
UserLevel: v.UserLevel, UserLevel: v.UserLevel,
AutoSystemRoutingTable: v.AutoSystemRoutingTable, AutoSystemRoutingTable: v.AutoSystemRoutingTable,
AutoSystemDns: v.AutoSystemDNS, 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 { if v.AutoOutboundsInterface != nil {
config.AutoOutboundsInterface = *v.AutoOutboundsInterface config.AutoOutboundsInterface = *v.AutoOutboundsInterface
+71
View File
@@ -0,0 +1,71 @@
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")
}
}
@@ -12,6 +12,7 @@ import (
"github.com/xtls/xray-core/core" "github.com/xtls/xray-core/core"
"github.com/xtls/xray-core/infra/conf" "github.com/xtls/xray-core/infra/conf"
"github.com/xtls/xray-core/infra/conf/serial" "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/masque"
"github.com/xtls/xray-core/proxy/shadowsocks" "github.com/xtls/xray-core/proxy/shadowsocks"
"github.com/xtls/xray-core/proxy/shadowsocks_2022" "github.com/xtls/xray-core/proxy/shadowsocks_2022"
@@ -91,6 +92,8 @@ func extractInboundUsers(inb *core.InboundHandlerConfig) []*protocol.User {
return ty.Users return ty.Users
case *masque.ServerConfig: case *masque.ServerConfig:
return ty.Users return ty.Users
case *hysteria.ServerConfig:
return ty.Users
default: default:
fmt.Println("unsupported inbound type") fmt.Println("unsupported inbound type")
} }
+34 -114
View File
@@ -2,7 +2,6 @@ package shadowsocks_2022
import ( import (
"context" "context"
"io"
"time" "time"
"github.com/xtls/xray-core/common" "github.com/xtls/xray-core/common"
@@ -13,9 +12,6 @@ import (
"github.com/xtls/xray-core/common/net" "github.com/xtls/xray-core/common/net"
"github.com/xtls/xray-core/common/protocol" "github.com/xtls/xray-core/common/protocol"
"github.com/xtls/xray-core/common/session" "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/common/utils"
"github.com/xtls/xray-core/core" "github.com/xtls/xray-core/core"
"github.com/xtls/xray-core/features/policy" "github.com/xtls/xray-core/features/policy"
"github.com/xtls/xray-core/features/routing" "github.com/xtls/xray-core/features/routing"
@@ -101,35 +97,29 @@ func (i *Inbound) processTCP(ctx context.Context, conn net.Conn, dispatcher rout
return errors.New("unable to set read deadline").Base(err) 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 var salt [32]byte
copy(salt[:i.method.KeySaltLength], headerBuf[:i.method.KeySaltLength])
saltSlice := salt[:i.method.KeySaltLength] saltSlice := salt[:i.method.KeySaltLength]
if _, err := io.ReadFull(conn, saltSlice); err != nil { fixedChunk := headerBuf[i.method.KeySaltLength:]
return err
}
if !i.saltFilter.Check(salt) { reader, reqHeader, err := InitServerStream(conn, i.method, i.psk, saltSlice, salt, fixedChunk, i.saltFilter)
return ErrSaltNotUnique
}
sessionKey := DeriveSessionSubKey(i.psk, saltSlice, i.method.KeySaltLength)
aead, err := i.method.NewAEAD(sessionKey)
if err != nil { if err != nil {
ResetTCPConn(conn)
return err return err
} }
reader := NewStreamReader(conn, aead)
reqHeader, err := ReadClientRequestHeader(conn, reader)
if err != nil {
return err
}
conn.SetReadDeadline(time.Time{})
dest := reqHeader.Destination dest := reqHeader.Destination
writer, err := WriteTCPResponse(conn, i.method, i.psk, saltSlice, nil) writer := NewServerStreamWriter(conn, i.method, i.psk, saltSlice)
if err != nil {
return err
}
ctx = log.ContextWithAccessMessage(ctx, &log.AccessMessage{ ctx = log.ContextWithAccessMessage(ctx, &log.AccessMessage{
From: conn.RemoteAddr(), From: conn.RemoteAddr(),
@@ -146,42 +136,17 @@ func (i *Inbound) processTCP(ctx context.Context, conn net.Conn, dispatcher rout
} }
if len(reqHeader.EarlyData) > 0 { if len(reqHeader.EarlyData) > 0 {
earlyBuf := buf.New() mb := buf.MergeBytes(nil, reqHeader.EarlyData)
earlyBuf.Write(reqHeader.EarlyData) if err := link.Writer.WriteMultiBuffer(mb); err != nil {
if err := link.Writer.WriteMultiBuffer(buf.MultiBuffer{earlyBuf}); err != nil {
return err return err
} }
} }
sessionPolicy = i.policyManager.ForLevel(uint32(i.user.Level)) return TransportTCP(ctx, i.policyManager.ForLevel(uint32(i.user.Level)), reader, writer, link)
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)
return buf.Copy(link.Reader, writer, buf.UpdateActivity(timer))
}
responseDoneAndCloseWriter := task.OnSuccess(responseDone, task.Close(link.Writer))
return task.Run(ctx, requestDone, responseDoneAndCloseWriter)
} }
func (i *Inbound) processUDP(ctx context.Context, conn stat.Connection, dispatcher routing.Dispatcher) error { func (i *Inbound) processUDP(ctx context.Context, conn stat.Connection, dispatcher routing.Dispatcher) error {
udpConns := utils.NewTypedSyncMap[uint64, *udpConnEntry]() reader := buf.NewPacketReader(conn)
defer func() {
udpConns.Range(func(key uint64, entry *udpConnEntry) bool {
entry.timer.SetTimeout(0)
return true
})
}()
reader := buf.NewReader(conn)
for { for {
mb, err := reader.ReadMultiBuffer() mb, err := reader.ReadMultiBuffer()
if err != nil { if err != nil {
@@ -191,75 +156,30 @@ func (i *Inbound) processUDP(ctx context.Context, conn stat.Connection, dispatch
for _, b := range mb { for _, b := range mb {
decoded, err := i.udpCodec.DecodePacket(b.Bytes()) decoded, err := i.udpCodec.DecodePacket(b.Bytes())
if err != nil { b.Release()
b.Release() if err != nil || decoded.HeaderType != HeaderTypeClient {
continue continue
} }
entry, ok := udpConns.Load(decoded.SessionID) sessionItem := i.udpCodec.GetSession(decoded.SessionID)
if !ok { if sessionItem.User == nil {
sessCtx, cancel := context.WithCancel(ctx) sessionItem.Lock()
sessCtx = log.ContextWithAccessMessage(sessCtx, &log.AccessMessage{ if sessionItem.User == nil {
From: conn.RemoteAddr(), sessionItem.User = i.user
To: decoded.Destination,
Status: log.AccessAccepted,
Email: i.user.Email,
})
link, err := dispatcher.Dispatch(sessCtx, decoded.Destination)
if err != nil {
cancel()
b.Release()
continue
}
newEntry := &udpConnEntry{
link: link,
cancel: cancel,
}
sessionPolicy := i.policyManager.ForLevel(uint32(i.user.Level))
newEntry.timer = signal.CancelAfterInactivity(sessCtx, func() {
udpConns.Delete(decoded.SessionID)
common.Interrupt(link.Reader)
common.Interrupt(link.Writer)
cancel()
}, sessionPolicy.Timeouts.ConnectionIdle)
actual, loaded := udpConns.LoadOrStore(decoded.SessionID, newEntry)
if loaded {
// Another goroutine/packet beat us to storing, terminate our redundant link
newEntry.timer.SetTimeout(0)
entry = actual
} else {
entry = newEntry
go func(sessID uint64, dest net.Destination, cEntry *udpConnEntry) {
defer func() {
cEntry.timer.SetTimeout(0)
}()
for {
resMb, err := cEntry.link.Reader.ReadMultiBuffer()
if err != nil {
return
}
cEntry.timer.Update()
for _, rb := range resMb {
encPacket, err := i.udpCodec.EncodeServerPacket(sessID, dest, rb.Bytes())
rb.Release()
if err != nil {
continue
}
_, _ = conn.Write(encPacket)
}
}
}(decoded.SessionID, decoded.Destination, entry)
} }
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)
})
if err != nil {
continue
} }
entry.timer.Update()
payloadBuf := buf.New() payloadBuf := buf.New()
payloadBuf.Write(decoded.Payload) payloadBuf.Write(decoded.Payload)
b.Release() payloadBuf.UDP = &decoded.Destination
_ = entry.link.Writer.WriteMultiBuffer(buf.MultiBuffer{payloadBuf}) _ = link.Writer.WriteMultiBuffer(buf.MultiBuffer{payloadBuf})
} }
} }
} }
+49 -202
View File
@@ -4,7 +4,6 @@ import (
"context" "context"
"crypto/cipher" "crypto/cipher"
"encoding/binary" "encoding/binary"
"io"
"strconv" "strconv"
"strings" "strings"
"sync" "sync"
@@ -19,8 +18,6 @@ import (
"github.com/xtls/xray-core/common/net" "github.com/xtls/xray-core/common/net"
"github.com/xtls/xray-core/common/protocol" "github.com/xtls/xray-core/common/protocol"
"github.com/xtls/xray-core/common/session" "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/common/utils" "github.com/xtls/xray-core/common/utils"
"github.com/xtls/xray-core/common/uuid" "github.com/xtls/xray-core/common/uuid"
"github.com/xtls/xray-core/core" "github.com/xtls/xray-core/core"
@@ -207,64 +204,46 @@ func (i *MultiUserInbound) processTCP(ctx context.Context, conn net.Conn, dispat
return errors.New("unable to set read deadline").Base(err) return errors.New("unable to set read deadline").Base(err)
} }
// 1. Read Request Salt (16 or 32 bytes) // 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 var salt [32]byte
copy(salt[:i.method.KeySaltLength], headerBuf[:i.method.KeySaltLength])
saltSlice := salt[:i.method.KeySaltLength] saltSlice := salt[:i.method.KeySaltLength]
if _, err := io.ReadFull(conn, saltSlice); err != nil { eih := headerBuf[i.method.KeySaltLength : i.method.KeySaltLength+AESBlockSize]
return err fixedChunk := headerBuf[i.method.KeySaltLength+AESBlockSize:]
}
if !i.saltFilter.Check(salt) { decryptedHash, err := DecryptEIH(i.method, i.masterPSK, saltSlice, eih)
return ErrSaltNotUnique
}
// 2. Read Extended Identity Header (16 bytes)
var eih [AESBlockSize]byte
if _, err := io.ReadFull(conn, eih[:]); err != nil {
return err
}
// Decrypt EIH with IdentitySubKey derived from masterPSK and salt
identitySubkey := DeriveIdentitySubKey(i.masterPSK, saltSlice, i.method.KeySaltLength)
block, err := i.method.NewBlock(identitySubkey)
if err != nil { if err != nil {
ResetTCPConn(conn)
return err return err
} }
var decryptedHash [AESBlockSize]byte
block.Decrypt(decryptedHash[:], eih[:])
// Lookup user // Lookup user
user, ok := i.usersByHash.Load(decryptedHash) user, ok := i.usersByHash.Load(decryptedHash)
if !ok || user == nil { if !ok {
ResetTCPConn(conn)
return ErrInvalidRequest return ErrInvalidRequest
} }
userPSK := user.Account.(*MemoryAccount).Key userPSK := user.Account.(*MemoryAccount).Key
// 3. Derive Session Subkey using matched user's PSK reader, reqHeader, err := InitServerStream(conn, i.method, userPSK, saltSlice, salt, fixedChunk, i.saltFilter)
sessionKey := DeriveSessionSubKey(userPSK, saltSlice, i.method.KeySaltLength)
aead, err := i.method.NewAEAD(sessionKey)
if err != nil { if err != nil {
ResetTCPConn(conn)
return err return err
} }
reader := NewStreamReader(conn, aead)
// 4 & 5. Read Client Request Header
reqHeader, err := ReadClientRequestHeader(conn, reader)
if err != nil {
return err
}
conn.SetReadDeadline(time.Time{})
dest := reqHeader.Destination dest := reqHeader.Destination
// 6. Send Server Response Handshake writer := NewServerStreamWriter(conn, i.method, userPSK, saltSlice)
writer, err := WriteTCPResponse(conn, i.method, userPSK, saltSlice, nil)
if err != nil {
return err
}
// 7. Dispatch Connection to Xray routing with matched User // Dispatch Connection to Xray routing with matched User
inbound := session.InboundFromContext(ctx) inbound := session.InboundFromContext(ctx)
inbound.User = user inbound.User = user
@@ -283,42 +262,17 @@ func (i *MultiUserInbound) processTCP(ctx context.Context, conn net.Conn, dispat
} }
if len(reqHeader.EarlyData) > 0 { if len(reqHeader.EarlyData) > 0 {
earlyBuf := buf.New() mb := buf.MergeBytes(nil, reqHeader.EarlyData)
earlyBuf.Write(reqHeader.EarlyData) if err := link.Writer.WriteMultiBuffer(mb); err != nil {
if err := link.Writer.WriteMultiBuffer(buf.MultiBuffer{earlyBuf}); err != nil {
return err return err
} }
} }
sessionPolicy = i.policyManager.ForLevel(user.Level) return TransportTCP(ctx, i.policyManager.ForLevel(user.Level), reader, writer, link)
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)
return buf.Copy(link.Reader, writer, buf.UpdateActivity(timer))
}
responseDoneAndCloseWriter := task.OnSuccess(responseDone, task.Close(link.Writer))
return task.Run(ctx, requestDone, responseDoneAndCloseWriter)
} }
func (i *MultiUserInbound) processUDP(ctx context.Context, conn stat.Connection, dispatcher routing.Dispatcher) error { func (i *MultiUserInbound) processUDP(ctx context.Context, conn stat.Connection, dispatcher routing.Dispatcher) error {
udpConns := utils.NewTypedSyncMap[uint64, *udpConnEntry]() reader := buf.NewPacketReader(conn)
defer func() {
udpConns.Range(func(key uint64, entry *udpConnEntry) bool {
entry.timer.SetTimeout(0)
return true
})
}()
reader := buf.NewReader(conn)
for { for {
mb, err := reader.ReadMultiBuffer() mb, err := reader.ReadMultiBuffer()
if err != nil { if err != nil {
@@ -342,168 +296,61 @@ func (i *MultiUserInbound) processUDP(ctx context.Context, conn stat.Connection,
sessionID := binary.BigEndian.Uint64(rawHeader[:8]) sessionID := binary.BigEndian.Uint64(rawHeader[:8])
packetID := binary.BigEndian.Uint64(rawHeader[8:16]) packetID := binary.BigEndian.Uint64(rawHeader[8:16])
// Replay protection & session lookup
sessionItem := i.udpSessions.GetOrCreate(sessionID) sessionItem := i.udpSessions.GetOrCreate(sessionID)
sessionItem.Lock() if !sessionItem.CheckPacketID(packetID) {
if !sessionItem.Window.Check(packetID) {
sessionItem.Unlock()
b.Release() b.Release()
continue continue
} }
var userPSK []byte var userPSK []byte
var currentUser *protocol.MemoryUser var currentUser *protocol.MemoryUser
if sessionItem.User != nil { sessionItem.Lock()
currentUser = sessionItem.User currentUser = sessionItem.User
userPSK = sessionItem.UserPSK userPSK = sessionItem.UserPSK
sessionItem.Unlock() sessionItem.Unlock()
} else {
sessionItem.Unlock()
// Decrypt EIH
identitySubkey := DeriveIdentitySubKey(i.masterPSK, rawHeader[:8], i.method.KeySaltLength)
idBlock, err := i.method.NewBlock(identitySubkey)
if err != nil {
b.Release()
continue
}
var decryptedHash [16]byte if currentUser == nil {
idBlock.Decrypt(decryptedHash[:], packetBytes[16:32]) // Decrypt EIH
decryptedHash := DecryptUDPEIH(i.udpMasterCipher, rawHeader[:], packetBytes[16:32])
user, ok := i.usersByHash.Load(decryptedHash) user, ok := i.usersByHash.Load(decryptedHash)
if !ok || user == nil { if !ok {
b.Release() b.Release()
continue continue
} }
currentUser = user currentUser = user
userPSK = user.Account.(*MemoryAccount).Key userPSK = user.Account.(*MemoryAccount).Key
sessionItem.Lock()
sessionItem.User = user
sessionItem.UserPSK = userPSK
sessionItem.Unlock()
} }
// Decrypt Body (with AEAD caching per session) decoded, err := sessionItem.DecryptAESPayload(i.method, userPSK, sessionID, packetID, rawHeader[:], packetBytes[32:])
bodyAead := sessionItem.GetRemoteCipher()
if bodyAead == nil {
bodyKey := DeriveSessionSubKey(userPSK, rawHeader[:8], i.method.KeySaltLength)
var err error
bodyAead, err = i.method.NewAEAD(bodyKey)
if err != nil {
b.Release()
continue
}
sessionItem.SetRemoteCipher(bodyAead)
}
bodyNonce := rawHeader[4:16]
bodyCipher := packetBytes[32:]
bodyPlain, err := bodyAead.Open(nil, bodyNonce, bodyCipher, nil)
b.Release() b.Release()
if err != nil || len(bodyPlain) < 1+8+2 {
continue
}
sessionItem.Lock()
sessionItem.Window.Add(packetID)
sessionItem.Unlock()
if bodyPlain[0] != HeaderTypeClient {
continue
}
epoch := binary.BigEndian.Uint64(bodyPlain[1:9])
diff := time.Now().Unix() - int64(epoch)
if diff < -30 || diff > 30 {
continue
}
paddingLen := int(binary.BigEndian.Uint16(bodyPlain[9:11]))
offset := 11 + paddingLen
if len(bodyPlain) < offset {
continue
}
dest, addrLen, err := parseAddressPort(bodyPlain[offset:])
if err != nil { if err != nil {
continue continue
} }
payload := bodyPlain[offset+addrLen:] sessionItem.Lock()
if sessionItem.User == nil {
sessionItem.User = currentUser
sessionItem.UserPSK = userPSK
}
sessionItem.Unlock()
entry, ok := udpConns.Load(sessionID) link, err := sessionItem.EnsureLink(ctx, conn, decoded.Destination, dispatcher, i.policyManager, func(replyDest net.Destination, payload []byte) ([]byte, error) {
if !ok { return i.encodeServerUDPPacket(sessionID, userPSK, replyDest, payload)
sessCtx, cancel := context.WithCancel(ctx) })
inbound := session.InboundFromContext(sessCtx) if err != nil {
inbound.User = currentUser continue
sessCtx = log.ContextWithAccessMessage(sessCtx, &log.AccessMessage{
From: conn.RemoteAddr(),
To: dest,
Status: log.AccessAccepted,
Email: currentUser.Email,
})
link, err := dispatcher.Dispatch(sessCtx, dest)
if err != nil {
cancel()
continue
}
newEntry := &udpConnEntry{
link: link,
cancel: cancel,
}
sessionPolicy := i.policyManager.ForLevel(currentUser.Level)
newEntry.timer = signal.CancelAfterInactivity(sessCtx, func() {
udpConns.Delete(sessionID)
common.Interrupt(link.Reader)
common.Interrupt(link.Writer)
cancel()
}, sessionPolicy.Timeouts.ConnectionIdle)
actual, loaded := udpConns.LoadOrStore(sessionID, newEntry)
if loaded {
newEntry.timer.SetTimeout(0)
entry = actual
} else {
entry = newEntry
go func(sessID uint64, uPSK []byte, d net.Destination, cEntry *udpConnEntry) {
defer func() {
cEntry.timer.SetTimeout(0)
}()
for {
resMb, err := cEntry.link.Reader.ReadMultiBuffer()
if err != nil {
return
}
cEntry.timer.Update()
for _, rb := range resMb {
encPacket, err := i.encodeServerUDPPacket(sessID, uPSK, d, rb.Bytes())
rb.Release()
if err != nil {
continue
}
_, _ = conn.Write(encPacket)
}
}
}(sessionID, userPSK, dest, entry)
}
} }
entry.timer.Update()
pBuf := buf.New() pBuf := buf.New()
pBuf.Write(payload) pBuf.Write(decoded.Payload)
_ = entry.link.Writer.WriteMultiBuffer(buf.MultiBuffer{pBuf}) pBuf.UDP = &decoded.Destination
_ = link.Writer.WriteMultiBuffer(buf.MultiBuffer{pBuf})
} }
} }
} }
func (i *MultiUserInbound) encodeServerUDPPacket(clientSessionID uint64, userPSK []byte, dest net.Destination, payload []byte) ([]byte, error) { func (i *MultiUserInbound) encodeServerUDPPacket(clientSessionID uint64, userPSK []byte, dest net.Destination, payload []byte) ([]byte, error) {
sessionItem := i.udpSessions.GetOrCreate(clientSessionID) return i.udpSessions.EncodeServerPacket(i.method, userPSK, clientSessionID, dest, payload)
if err := sessionItem.EnsureServerState(i.method, i.udpMasterCipher, nil, userPSK); err != nil {
return nil, err
}
return sessionItem.EncodeServerPacket(i.method, clientSessionID, dest, payload)
} }
+63 -128
View File
@@ -4,7 +4,6 @@ import (
"context" "context"
"crypto/cipher" "crypto/cipher"
"encoding/binary" "encoding/binary"
"io"
"strconv" "strconv"
"time" "time"
@@ -15,9 +14,6 @@ import (
"github.com/xtls/xray-core/common/net" "github.com/xtls/xray-core/common/net"
"github.com/xtls/xray-core/common/protocol" "github.com/xtls/xray-core/common/protocol"
"github.com/xtls/xray-core/common/session" "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/common/utils"
"github.com/xtls/xray-core/common/uuid" "github.com/xtls/xray-core/common/uuid"
"github.com/xtls/xray-core/core" "github.com/xtls/xray-core/core"
"github.com/xtls/xray-core/features/policy" "github.com/xtls/xray-core/features/policy"
@@ -35,18 +31,17 @@ type relayDest struct {
destination net.Destination destination net.Destination
email string email string
level uint32 level uint32
key []byte
blockCipher cipher.Block blockCipher cipher.Block
} }
type RelayInbound struct { type RelayInbound struct {
networks []net.Network networks []net.Network
method *CipherMethod method *CipherMethod
relayPSK []byte relayPSK []byte
relayBlock cipher.Block relayBlock cipher.Block
destinations map[[AESBlockSize]byte]*relayDest destinations map[[AESBlockSize]byte]*relayDest
rawDestinations []*RelayDestination udpSessions *UDPSessionManager
policyManager policy.Manager policyManager policy.Manager
} }
func NewRelayServer(ctx context.Context, config *RelayServerConfig) (*RelayInbound, error) { func NewRelayServer(ctx context.Context, config *RelayServerConfig) (*RelayInbound, error) {
@@ -78,13 +73,13 @@ func NewRelayServer(ctx context.Context, config *RelayServerConfig) (*RelayInbou
v := core.MustFromContext(ctx) v := core.MustFromContext(ctx)
i := &RelayInbound{ i := &RelayInbound{
networks: networks, networks: networks,
method: method, method: method,
relayPSK: relayPSK, relayPSK: relayPSK,
relayBlock: relayBlock, relayBlock: relayBlock,
destinations: make(map[[AESBlockSize]byte]*relayDest), destinations: make(map[[AESBlockSize]byte]*relayDest),
rawDestinations: config.Destinations, udpSessions: NewUDPSessionManager(500 * time.Second),
policyManager: v.GetFeature(policy.ManagerType()).(policy.Manager), policyManager: v.GetFeature(policy.ManagerType()).(policy.Manager),
} }
for idx, d := range config.Destinations { for idx, d := range config.Destinations {
@@ -108,7 +103,6 @@ func NewRelayServer(ctx context.Context, config *RelayServerConfig) (*RelayInbou
destination: net.TCPDestination(d.Address.AsAddress(), net.Port(d.Port)), destination: net.TCPDestination(d.Address.AsAddress(), net.Port(d.Port)),
email: d.Email, email: d.Email,
level: uint32(d.Level), level: uint32(d.Level),
key: destKey,
blockCipher: destBlock, blockCipher: destBlock,
} }
} }
@@ -139,28 +133,36 @@ func (i *RelayInbound) processTCP(ctx context.Context, conn net.Conn, dispatcher
return errors.New("unable to set read deadline").Base(err) return errors.New("unable to set read deadline").Base(err)
} }
// Read Salt + Outer EIH // Read initial handshake in a single read call per SIP022 §3.1.3 & §3.1.4
needed := i.method.KeySaltLength + AESBlockSize needed := i.method.KeySaltLength + AESBlockSize
var headerBuf [48]byte requestHeader := buf.New()
headerSlice := headerBuf[:needed] n, err := requestHeader.ReadFrom(conn)
if _, err := io.ReadFull(conn, headerSlice); err != nil {
return err
}
salt := headerSlice[:i.method.KeySaltLength]
eih := headerSlice[i.method.KeySaltLength:]
identitySubkey := DeriveIdentitySubKey(i.relayPSK, salt, i.method.KeySaltLength)
block, err := i.method.NewBlock(identitySubkey)
if err != nil { if err != nil {
requestHeader.Release()
ResetTCPConn(conn)
return err return err
} }
if int(n) < needed {
requestHeader.Release()
ResetTCPConn(conn)
return ErrInvalidRequest
}
var decryptedHash [AESBlockSize]byte headerSlice := requestHeader.Bytes()
block.Decrypt(decryptedHash[:], eih) 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] targetDest, ok := i.destinations[decryptedHash]
if !ok { if !ok {
requestHeader.Release()
ResetTCPConn(conn)
return ErrInvalidRequest return ErrInvalidRequest
} }
conn.SetReadDeadline(time.Time{}) conn.SetReadDeadline(time.Time{})
@@ -182,45 +184,26 @@ func (i *RelayInbound) processTCP(ctx context.Context, conn net.Conn, dispatcher
link, err := dispatcher.Dispatch(ctx, targetDest.destination) link, err := dispatcher.Dispatch(ctx, targetDest.destination)
if err != nil { if err != nil {
requestHeader.Release()
return err return err
} }
// Unwrap outer EIH: send client salt to next hop, stripping this hop's EIH // Unwrap outer EIH: send client salt and remaining handshake bytes to next hop
saltBuf := buf.New() // in a single write call, satisfying downstream server's single-read handshake expectation (SIP022 §3.1.3).
saltBuf.Write(salt) var saltCopy [32]byte
if err := link.Writer.WriteMultiBuffer(buf.MultiBuffer{saltBuf}); err != nil { 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 err
} }
sessionPolicy = i.policyManager.ForLevel(targetDest.level) return TransportTCP(ctx, i.policyManager.ForLevel(targetDest.level), buf.NewReader(conn), buf.NewWriter(conn), link)
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(buf.NewReader(conn), link.Writer, buf.UpdateActivity(timer))
}
responseDone := func() error {
defer timer.SetTimeout(sessionPolicy.Timeouts.UplinkOnly)
return buf.Copy(link.Reader, buf.NewWriter(conn), buf.UpdateActivity(timer))
}
responseDoneAndCloseWriter := task.OnSuccess(responseDone, task.Close(link.Writer))
return task.Run(ctx, requestDone, responseDoneAndCloseWriter)
} }
func (i *RelayInbound) processUDP(ctx context.Context, conn stat.Connection, dispatcher routing.Dispatcher) error { func (i *RelayInbound) processUDP(ctx context.Context, conn stat.Connection, dispatcher routing.Dispatcher) error {
udpConns := utils.NewTypedSyncMap[uint64, *udpConnEntry]() reader := buf.NewPacketReader(conn)
defer func() {
udpConns.Range(func(key uint64, entry *udpConnEntry) bool {
entry.timer.SetTimeout(0)
return true
})
}()
reader := buf.NewReader(conn)
for { for {
mb, err := reader.ReadMultiBuffer() mb, err := reader.ReadMultiBuffer()
if err != nil { if err != nil {
@@ -238,11 +221,7 @@ func (i *RelayInbound) processUDP(ctx context.Context, conn stat.Connection, dis
var packetHeader [AESBlockSize]byte var packetHeader [AESBlockSize]byte
i.relayBlock.Decrypt(packetHeader[:], data[:AESBlockSize]) i.relayBlock.Decrypt(packetHeader[:], data[:AESBlockSize])
var eiHeader [AESBlockSize]byte eiHeader := DecryptUDPEIH(i.relayBlock, packetHeader[:], data[AESBlockSize:2*AESBlockSize])
i.relayBlock.Decrypt(eiHeader[:], data[AESBlockSize:2*AESBlockSize])
for idx := 0; idx < AESBlockSize; idx++ {
eiHeader[idx] ^= packetHeader[idx]
}
targetDest, ok := i.destinations[eiHeader] targetDest, ok := i.destinations[eiHeader]
if !ok { if !ok {
@@ -263,68 +242,24 @@ func (i *RelayInbound) processUDP(ctx context.Context, conn stat.Connection, dis
dest := targetDest.destination dest := targetDest.destination
dest.Network = net.Network_UDP dest.Network = net.Network_UDP
entry, ok := udpConns.Load(sessionID) sessionItem := i.udpSessions.GetOrCreate(sessionID)
if !ok { if sessionItem.User == nil {
sessCtx, cancel := context.WithCancel(ctx) sessionItem.Lock()
inbound := session.InboundFromContext(sessCtx) if sessionItem.User == nil {
inbound.User = &protocol.MemoryUser{ sessionItem.User = &protocol.MemoryUser{
Email: targetDest.email, Email: targetDest.email,
Level: targetDest.level, Level: targetDest.level,
} }
sessCtx = log.ContextWithAccessMessage(sessCtx, &log.AccessMessage{
From: conn.RemoteAddr(),
To: dest,
Status: log.AccessAccepted,
Email: targetDest.email,
})
link, err := dispatcher.Dispatch(sessCtx, dest)
if err != nil {
cancel()
b.Release()
continue
}
newEntry := &udpConnEntry{
link: link,
cancel: cancel,
}
sessionPolicy := i.policyManager.ForLevel(targetDest.level)
newEntry.timer = signal.CancelAfterInactivity(sessCtx, func() {
udpConns.Delete(sessionID)
common.Interrupt(link.Reader)
common.Interrupt(link.Writer)
cancel()
}, sessionPolicy.Timeouts.ConnectionIdle)
actual, loaded := udpConns.LoadOrStore(sessionID, newEntry)
if loaded {
newEntry.timer.SetTimeout(0)
entry = actual
} else {
entry = newEntry
go func(cEntry *udpConnEntry) {
defer func() {
cEntry.timer.SetTimeout(0)
}()
for {
resMb, err := cEntry.link.Reader.ReadMultiBuffer()
if err != nil {
return
}
cEntry.timer.Update()
for _, rb := range resMb {
_, _ = conn.Write(rb.Bytes())
rb.Release()
}
}
}(entry)
} }
sessionItem.Unlock()
}
link, err := sessionItem.EnsureLink(ctx, conn, dest, dispatcher, i.policyManager, nil)
if err != nil {
b.Release()
continue
} }
entry.timer.Update() _ = link.Writer.WriteMultiBuffer(buf.MultiBuffer{b})
_ = entry.link.Writer.WriteMultiBuffer(buf.MultiBuffer{b})
} }
} }
} }
+11
View File
@@ -61,3 +61,14 @@ func DeriveUserPSKHash(userPSK []byte) [AESBlockSize]byte {
copy(out[:], h[:AESBlockSize]) copy(out[:], h[:AESBlockSize])
return out 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
}
+33 -13
View File
@@ -4,7 +4,6 @@ import (
"context" "context"
"crypto/rand" "crypto/rand"
"io" "io"
"time"
"github.com/xtls/xray-core/common" "github.com/xtls/xray-core/common"
"github.com/xtls/xray-core/common/buf" "github.com/xtls/xray-core/common/buf"
@@ -46,8 +45,12 @@ func NewClient(ctx context.Context, config *ClientConfig) (*Outbound, error) {
return nil, errors.New("invalid key: ", config.Key).Base(err) 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] finalPSK := pskList[len(pskList)-1]
udpCodec, err := NewUDPPacketCodec(method, finalPSK) udpCodec, err := NewUDPPacketCodec(method, pskList)
if err != nil { if err != nil {
return nil, errors.New("failed to create udp packet codec").Base(err) return nil, errors.New("failed to create udp packet codec").Base(err)
} }
@@ -126,18 +129,30 @@ func (o *Outbound) Process(ctx context.Context, link *transport.Link, dialer int
requestDone := func() error { requestDone := func() error {
defer timer.SetTimeout(sessionPolicy.Timeouts.DownlinkOnly) defer timer.SetTimeout(sessionPolicy.Timeouts.DownlinkOnly)
bufferedWriter := buf.NewBufferedWriter(buf.NewWriter(conn))
bodyWriter, err := WriteTCPRequest(bufferedWriter, o.method, o.pskList, destination, clientSaltSlice, nil) 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()
}
}
bodyWriter, err := WriteTCPRequest(conn, o.method, o.pskList, destination, clientSaltSlice, initialPayload)
if firstBuf != nil {
firstBuf.Release()
}
if err != nil { if err != nil {
buf.ReleaseMulti(remainingMB)
return errors.New("failed to write request").Base(err) return errors.New("failed to write request").Base(err)
} }
if err = buf.CopyOnceTimeout(link.Reader, bodyWriter, time.Millisecond*100); err != nil && err != buf.ErrNotTimeoutReader && err != buf.ErrReadTimeout { if !remainingMB.IsEmpty() {
return errors.New("failed to write A request payload").Base(err) if err := bodyWriter.WriteMultiBuffer(remainingMB); err != nil {
} return err
}
if err := bufferedWriter.SetBuffered(false); err != nil {
return err
} }
return buf.Copy(link.Reader, bodyWriter, buf.UpdateActivity(timer)) return buf.Copy(link.Reader, bodyWriter, buf.UpdateActivity(timer))
@@ -163,13 +178,18 @@ func (o *Outbound) Process(ctx context.Context, link *transport.Link, dialer int
} }
if network == net.Network_UDP { 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 { requestDone := func() error {
defer timer.SetTimeout(sessionPolicy.Timeouts.DownlinkOnly) defer timer.SetTimeout(sessionPolicy.Timeouts.DownlinkOnly)
writer := &UDPWriter{ writer := &UDPWriter{
Writer: conn, Writer: conn,
Destination: destination, Destination: destination,
Codec: o.udpCodec, Session: session,
} }
if err := buf.Copy(link.Reader, writer, buf.UpdateActivity(timer)); err != nil { if err := buf.Copy(link.Reader, writer, buf.UpdateActivity(timer)); err != nil {
@@ -182,8 +202,8 @@ func (o *Outbound) Process(ctx context.Context, link *transport.Link, dialer int
defer timer.SetTimeout(sessionPolicy.Timeouts.UplinkOnly) defer timer.SetTimeout(sessionPolicy.Timeouts.UplinkOnly)
reader := &UDPReader{ reader := &UDPReader{
Reader: conn, Reader: conn,
Codec: o.udpCodec, Session: session,
} }
if err := buf.Copy(reader, link.Writer, buf.UpdateActivity(timer)); err != nil { if err := buf.Copy(reader, link.Writer, buf.UpdateActivity(timer)); err != nil {
+441 -187
View File
@@ -16,14 +16,13 @@ import (
) )
type UDPCodec struct { type UDPCodec struct {
method *CipherMethod method *CipherMethod
psk []byte pskList [][]byte
blockCipher cipher.Block psk []byte
chachaCipher cipher.AEAD blockCipher cipher.Block
clientBodyCipher cipher.AEAD blockCiphers []cipher.Block
clientSessionID uint64 chachaCipher cipher.AEAD
nextPacketID atomic.Uint64 sessions *UDPSessionManager
sessions *UDPSessionManager
} }
type ( type (
@@ -48,22 +47,23 @@ func newUDPCodec(method *CipherMethod, psk []byte) (*UDPCodec, error) {
return c, nil return c, nil
} }
func NewUDPPacketCodec(method *CipherMethod, psk []byte) (*UDPCodec, error) { func NewUDPPacketCodec(method *CipherMethod, pskList [][]byte) (*UDPCodec, error) {
c, err := newUDPCodec(method, psk) 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 { if err != nil {
return nil, err return nil, err
} }
var sessID [8]byte c.pskList = pskList
if _, err := io.ReadFull(rand.Reader, sessID[:]); err != nil { if len(pskList) > 1 {
return nil, err c.blockCiphers = make([]cipher.Block, len(pskList))
} for i, psk := range pskList {
c.clientSessionID = binary.BigEndian.Uint64(sessID[:]) c.blockCiphers[i], err = method.NewBlock(psk)
if err != nil {
if !method.IsChaCha { return nil, err
clientBodyKey := DeriveSessionSubKey(psk, sessID[:], method.KeySaltLength) }
c.clientBodyCipher, err = method.NewAEAD(clientBodyKey)
if err != nil {
return nil, err
} }
} }
return c, nil return c, nil
@@ -78,108 +78,37 @@ func NewUDPServerCodec(method *CipherMethod, psk []byte, sessionTimeout time.Dur
return c, nil return c, nil
} }
func (c *UDPCodec) EncodeClientPacket(dest net.Destination, payload []byte) (*buf.Buffer, error) { func (c *UDPCodec) Sessions() *UDPSessionManager {
packetID := c.nextPacketID.Add(1) return c.sessions
sessID := c.clientSessionID }
// Padding determination (e.g. DNS port 53 disguise) func (c *UDPCodec) GetSession(sessionID uint64) *ServerUDPSession {
var paddingLen int if c.sessions == nil {
if dest.Port == 53 && len(payload) < MaxPaddingLength { return nil
paddingLen = mrand.IntN(MaxPaddingLength-len(payload)) + 1
} }
return c.sessions.GetOrCreate(sessionID)
addrPortLen := AddrPortLength(dest)
if c.method.IsChaCha {
// ChaCha20 mode: 24-byte nonce + plaintext header (27B) + padding + dest + payload + AEAD tag (16B)
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(c.chachaCipher.Overhead()))
c.chachaCipher.Seal(plainBytes[:0], nonce[:], plainBytes, nil)
return outBuf, nil
}
// AES mode:
// 16B Encrypted Header + (11B header + padding + dest + payload + 16B AEAD tag)
totalLen := 16 + 11 + paddingLen + addrPortLen + len(payload) + AEADTagSize
if totalLen > buf.Size {
return nil, ErrPacketTooLarge
}
outBuf := buf.New()
var rawHeader [16]byte
binary.BigEndian.PutUint64(rawHeader[:8], sessID)
binary.BigEndian.PutUint64(rawHeader[8:16], packetID)
var encryptedHeader [16]byte
c.blockCipher.Encrypt(encryptedHeader[:], rawHeader[:])
outBuf.Write(encryptedHeader[:])
bodyAead := c.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)
plainBytes := outBuf.Bytes()[16:]
bodyNonce := rawHeader[4:16]
outBuf.Extend(int32(bodyAead.Overhead()))
bodyAead.Seal(plainBytes[:0], bodyNonce, plainBytes, nil)
return outBuf, nil
} }
type DecodedUDPPacket struct { type DecodedUDPPacket struct {
SessionID uint64 SessionID uint64
PacketID uint64 PacketID uint64
HeaderType byte HeaderType byte
Timestamp uint64 Timestamp uint64
Destination net.Destination ClientSessionID uint64
Payload []byte Destination net.Destination
Payload []byte
} }
func parseAddressPort(data []byte) (net.Destination, int, error) { 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 { if len(data) < 1 {
return net.Destination{}, 0, ErrPacketTooShort return net.Destination{}, 0, ErrPacketTooShort
} }
@@ -220,6 +149,9 @@ func parsePlainUDPPacket(sessionID, packetID uint64, bodyPlain []byte) (DecodedU
} }
headerType := bodyPlain[0] headerType := bodyPlain[0]
if headerType != HeaderTypeClient && headerType != HeaderTypeServer {
return DecodedUDPPacket{}, ErrBadHeaderType
}
epoch := binary.BigEndian.Uint64(bodyPlain[1:9]) epoch := binary.BigEndian.Uint64(bodyPlain[1:9])
diff := int(math.Abs(float64(time.Now().Unix() - int64(epoch)))) diff := int(math.Abs(float64(time.Now().Unix() - int64(epoch))))
if diff > 30 { if diff > 30 {
@@ -227,11 +159,13 @@ func parsePlainUDPPacket(sessionID, packetID uint64, bodyPlain []byte) (DecodedU
} }
offset := 9 offset := 9
var clientSessionID uint64
if headerType == HeaderTypeServer { if headerType == HeaderTypeServer {
if len(bodyPlain) < offset+8+2 { if len(bodyPlain) < offset+8+2 {
return DecodedUDPPacket{}, ErrPacketTooShort return DecodedUDPPacket{}, ErrPacketTooShort
} }
offset += 8 // skip clientSessionID clientSessionID = binary.BigEndian.Uint64(bodyPlain[offset : offset+8])
offset += 8
} }
paddingLen := int(binary.BigEndian.Uint16(bodyPlain[offset : offset+2])) paddingLen := int(binary.BigEndian.Uint16(bodyPlain[offset : offset+2]))
@@ -242,19 +176,20 @@ func parsePlainUDPPacket(sessionID, packetID uint64, bodyPlain []byte) (DecodedU
} }
offset += paddingLen offset += paddingLen
dest, addrLen, err := parseAddressPort(bodyPlain[offset:]) dest, addrLen, err := ParseAddressPort(bodyPlain[offset:])
if err != nil { if err != nil {
return DecodedUDPPacket{}, err return DecodedUDPPacket{}, err
} }
payload := bodyPlain[offset+addrLen:] payload := bodyPlain[offset+addrLen:]
return DecodedUDPPacket{ return DecodedUDPPacket{
SessionID: sessionID, SessionID: sessionID,
PacketID: packetID, PacketID: packetID,
HeaderType: headerType, HeaderType: headerType,
Timestamp: epoch, Timestamp: epoch,
Destination: dest, ClientSessionID: clientSessionID,
Payload: payload, Destination: dest,
Payload: payload,
}, nil }, nil
} }
@@ -269,7 +204,7 @@ func (c *UDPCodec) DecodePacket(data []byte) (DecodedUDPPacket, error) {
} }
nonce := data[:PacketNonceSize] nonce := data[:PacketNonceSize]
ciphertext := data[PacketNonceSize:] ciphertext := data[PacketNonceSize:]
plain, err := c.chachaCipher.Open(ciphertext[:0], nonce, ciphertext, nil) plain, err := c.chachaCipher.Open(nil, nonce, ciphertext, nil)
if err != nil { if err != nil {
return DecodedUDPPacket{}, errors.New("failed to decrypt chacha udp packet").Base(err) return DecodedUDPPacket{}, errors.New("failed to decrypt chacha udp packet").Base(err)
} }
@@ -280,17 +215,22 @@ func (c *UDPCodec) DecodePacket(data []byte) (DecodedUDPPacket, error) {
sessionID := binary.BigEndian.Uint64(plain[:8]) sessionID := binary.BigEndian.Uint64(plain[:8])
packetID := binary.BigEndian.Uint64(plain[8:16]) packetID := binary.BigEndian.Uint64(plain[8:16])
if c.sessions != nil { sessionItem := c.sessions.GetOrCreate(sessionID)
sessionItem := c.sessions.GetOrCreate(sessionID) if !sessionItem.CheckPacketID(packetID) {
sessionItem.Lock() return DecodedUDPPacket{}, ErrPacketIdNotUnique
if !sessionItem.Window.CheckAndAdd(packetID) {
sessionItem.Unlock()
return DecodedUDPPacket{}, ErrPacketIdNotUnique
}
sessionItem.Unlock()
} }
return parsePlainUDPPacket(sessionID, packetID, plain[16:]) 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 // AES mode
@@ -299,54 +239,52 @@ func (c *UDPCodec) DecodePacket(data []byte) (DecodedUDPPacket, error) {
sessionID := binary.BigEndian.Uint64(rawHeader[:8]) sessionID := binary.BigEndian.Uint64(rawHeader[:8])
packetID := binary.BigEndian.Uint64(rawHeader[8:16]) packetID := binary.BigEndian.Uint64(rawHeader[8:16])
var bodyAead cipher.AEAD sessionItem := c.sessions.GetOrCreate(sessionID)
var sessionItem *ServerUDPSession if !sessionItem.CheckPacketID(packetID) {
return DecodedUDPPacket{}, ErrPacketIdNotUnique
}
if c.sessions != nil { return sessionItem.DecryptAESPayload(c.method, c.psk, sessionID, packetID, rawHeader[:], data[16:])
sessionItem = c.sessions.GetOrCreate(sessionID) }
sessionItem.Lock()
if !sessionItem.Window.Check(packetID) {
sessionItem.Unlock()
return DecodedUDPPacket{}, ErrPacketIdNotUnique
}
sessionItem.Unlock()
bodyAead = sessionItem.GetRemoteCipher() func (s *ServerUDPSession) DecryptAESPayload(method *CipherMethod, psk []byte, sessionID, packetID uint64, rawHeader, bodyCipher []byte) (DecodedUDPPacket, error) {
if bodyAead == nil { bodyAead := s.clientBodyCipher
bodyKey := DeriveSessionSubKey(c.psk, rawHeader[:8], c.method.KeySaltLength) isNewCipher := false
var err error if bodyAead == nil {
bodyAead, err = c.method.NewAEAD(bodyKey) bodyKey := DeriveSessionSubKey(psk, rawHeader[:8], method.KeySaltLength)
if err != nil {
return DecodedUDPPacket{}, err
}
sessionItem.SetRemoteCipher(bodyAead)
}
} else {
bodyKey := DeriveSessionSubKey(c.psk, rawHeader[:8], c.method.KeySaltLength)
var err error var err error
bodyAead, err = c.method.NewAEAD(bodyKey) bodyAead, err = method.NewAEAD(bodyKey)
if err != nil { if err != nil {
return DecodedUDPPacket{}, err return DecodedUDPPacket{}, err
} }
isNewCipher = true
} }
bodyNonce := rawHeader[4:16] bodyNonce := rawHeader[4:16]
bodyCipher := data[16:] bodyPlain, err := bodyAead.Open(nil, bodyNonce, bodyCipher, nil)
bodyPlain, err := bodyAead.Open(bodyCipher[:0], bodyNonce, bodyCipher, nil)
if err != nil { if err != nil {
return DecodedUDPPacket{}, errors.New("failed to decrypt aes udp body").Base(err) return DecodedUDPPacket{}, errors.New("failed to decrypt aes udp body").Base(err)
} }
if sessionItem != nil { decoded, err := parsePlainUDPPacket(sessionID, packetID, bodyPlain)
sessionItem.Lock() if err != nil {
sessionItem.Window.Add(packetID) return DecodedUDPPacket{}, err
sessionItem.Unlock()
} }
return parsePlainUDPPacket(sessionID, packetID, bodyPlain) if decoded.HeaderType != HeaderTypeClient {
return DecodedUDPPacket{}, ErrBadHeaderType
}
s.AddPacketID(packetID)
if isNewCipher {
s.clientBodyCipher = bodyAead
}
return decoded, nil
} }
func (s *ServerUDPSession) EnsureServerState(method *CipherMethod, headerBlock cipher.Block, chachaCipher cipher.AEAD, psk []byte) error { func (s *ServerUDPSession) EnsureServerState(method *CipherMethod, psk []byte) error {
s.Lock() s.Lock()
defer s.Unlock() defer s.Unlock()
if s.ServerSessionID != 0 { if s.ServerSessionID != 0 {
@@ -363,23 +301,29 @@ func (s *ServerUDPSession) EnsureServerState(method *CipherMethod, headerBlock c
} }
} }
if method.IsChaCha { if method.IsChaCha {
s.ServerChaCha = chachaCipher var err error
} else { s.serverChaCha, err = method.NewUDPCipher(psk)
s.ServerBlockCipher = headerBlock return err
bodyKey := DeriveSessionSubKey(psk, sidBuf[:], method.KeySaltLength) }
bodyAead, err := method.NewAEAD(bodyKey)
if err != nil { var err error
s.ServerSessionID = 0 s.serverHeaderBlock, err = method.NewBlock(psk)
return err if err != nil {
} s.ServerSessionID = 0
s.ServerCipher = bodyAead return err
}
bodyKey := DeriveSessionSubKey(psk, sidBuf[:], method.KeySaltLength)
s.serverBodyCipher, err = method.NewAEAD(bodyKey)
if err != nil {
s.ServerSessionID = 0
return err
} }
return nil return nil
} }
func (s *ServerUDPSession) EncodeServerPacket(method *CipherMethod, clientSessionID uint64, dest net.Destination, payload []byte) ([]byte, error) { func (s *ServerUDPSession) EncodeServerPacket(method *CipherMethod, clientSessionID uint64, dest net.Destination, payload []byte) ([]byte, error) {
serverSessionID := s.ServerSessionID serverSessionID := s.ServerSessionID
serverPacketID := s.ServerPacketID.Add(1) serverPacketID := s.ServerPacketID.Add(1) - 1
if method.IsChaCha { if method.IsChaCha {
var nonce [PacketNonceSize]byte var nonce [PacketNonceSize]byte
@@ -404,7 +348,7 @@ func (s *ServerUDPSession) EncodeServerPacket(method *CipherMethod, clientSessio
} }
plainBuf.Write(payload) plainBuf.Write(payload)
sealed := s.ServerChaCha.Seal(nil, nonce[:], plainBuf.Bytes(), nil) sealed := s.serverChaCha.Seal(nil, nonce[:], plainBuf.Bytes(), nil)
res := make([]byte, PacketNonceSize+len(sealed)) res := make([]byte, PacketNonceSize+len(sealed))
copy(res[:PacketNonceSize], nonce[:]) copy(res[:PacketNonceSize], nonce[:])
copy(res[PacketNonceSize:], sealed) copy(res[PacketNonceSize:], sealed)
@@ -417,7 +361,7 @@ func (s *ServerUDPSession) EncodeServerPacket(method *CipherMethod, clientSessio
binary.BigEndian.PutUint64(rawHeader[8:16], serverPacketID) binary.BigEndian.PutUint64(rawHeader[8:16], serverPacketID)
var encryptedHeader [16]byte var encryptedHeader [16]byte
s.ServerBlockCipher.Encrypt(encryptedHeader[:], rawHeader[:]) s.serverHeaderBlock.Encrypt(encryptedHeader[:], rawHeader[:])
bodyBuf := buf.New() bodyBuf := buf.New()
defer bodyBuf.Release() defer bodyBuf.Release()
@@ -435,7 +379,7 @@ func (s *ServerUDPSession) EncodeServerPacket(method *CipherMethod, clientSessio
bodyBuf.Write(payload) bodyBuf.Write(payload)
bodyNonce := rawHeader[4:16] bodyNonce := rawHeader[4:16]
sealedBody := s.ServerCipher.Seal(nil, bodyNonce, bodyBuf.Bytes(), nil) sealedBody := s.serverBodyCipher.Seal(nil, bodyNonce, bodyBuf.Bytes(), nil)
res := make([]byte, 16+len(sealedBody)) res := make([]byte, 16+len(sealedBody))
copy(res[:16], encryptedHeader[:]) copy(res[:16], encryptedHeader[:])
@@ -444,17 +388,327 @@ func (s *ServerUDPSession) EncodeServerPacket(method *CipherMethod, clientSessio
} }
func (c *UDPCodec) EncodeServerPacket(clientSessionID uint64, dest net.Destination, payload []byte) ([]byte, error) { func (c *UDPCodec) EncodeServerPacket(clientSessionID uint64, dest net.Destination, payload []byte) ([]byte, error) {
sessionItem := c.sessions.GetOrCreate(clientSessionID) return c.sessions.EncodeServerPacket(c.method, c.psk, clientSessionID, dest, payload)
if err := sessionItem.EnsureServerState(c.method, c.blockCipher, c.chachaCipher, c.psk); err != nil { }
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 return nil, err
} }
return sessionItem.EncodeServerPacket(c.method, clientSessionID, dest, payload) 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 { type UDPWriter struct {
Writer io.Writer Writer io.Writer
Destination net.Destination Destination net.Destination
Codec *UDPPacketCodec Session *ClientUDPSession
} }
func (w *UDPWriter) WriteMultiBuffer(mb buf.MultiBuffer) error { func (w *UDPWriter) WriteMultiBuffer(mb buf.MultiBuffer) error {
@@ -468,7 +722,7 @@ func (w *UDPWriter) WriteMultiBuffer(mb buf.MultiBuffer) error {
if b.UDP != nil { if b.UDP != nil {
dest = *b.UDP dest = *b.UDP
} }
pktBuf, err := w.Codec.EncodeClientPacket(dest, b.Bytes()) pktBuf, err := w.Session.EncodePacket(dest, b.Bytes())
b.Release() b.Release()
if err != nil { if err != nil {
buf.ReleaseMulti(mb) buf.ReleaseMulti(mb)
@@ -485,8 +739,8 @@ func (w *UDPWriter) WriteMultiBuffer(mb buf.MultiBuffer) error {
} }
type UDPReader struct { type UDPReader struct {
Reader io.Reader Reader io.Reader
Codec *UDPPacketCodec Session *ClientUDPSession
} }
func (r *UDPReader) ReadMultiBuffer() (buf.MultiBuffer, error) { func (r *UDPReader) ReadMultiBuffer() (buf.MultiBuffer, error) {
@@ -498,7 +752,7 @@ func (r *UDPReader) ReadMultiBuffer() (buf.MultiBuffer, error) {
return nil, err return nil, err
} }
decoded, err := r.Codec.DecodePacket(buffer.Bytes()) decoded, err := r.Session.DecodePacket(buffer.Bytes())
if err != nil { if err != nil {
buffer.Release() buffer.Release()
continue continue
+105
View File
@@ -2,9 +2,11 @@ package shadowsocks_2022_test
import ( import (
"context" "context"
"crypto/rand"
"encoding/base64" "encoding/base64"
"encoding/binary" "encoding/binary"
"errors" "errors"
"io"
gonet "net" gonet "net"
"sync" "sync"
"sync/atomic" "sync/atomic"
@@ -269,3 +271,106 @@ func (c *dummyStatConn) WriteMultiBuffer(mb buf.MultiBuffer) error {
} }
return nil 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))
}
})
}
}
+41 -16
View File
@@ -6,8 +6,11 @@ import (
"sync/atomic" "sync/atomic"
"time" "time"
"github.com/xtls/xray-core/common/net"
"github.com/xtls/xray-core/common/protocol" "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/common/utils"
"github.com/xtls/xray-core/transport"
) )
const ( const (
@@ -74,30 +77,42 @@ func (f *SlidingWindow) CheckAndAdd(counter uint64) bool {
type ServerUDPSession struct { type ServerUDPSession struct {
sync.Mutex sync.Mutex
SessionID uint64 SessionID uint64
RemoteCipher atomic.Pointer[cipher.AEAD] Window *SlidingWindow
Window SlidingWindow User *protocol.MemoryUser
User *protocol.MemoryUser UserPSK []byte
UserPSK []byte LastActive atomic.Int64 // Unix timestamp in seconds
LastActive atomic.Int64 // Unix timestamp in seconds
clientBodyCipher cipher.AEAD
ServerSessionID uint64 ServerSessionID uint64
ServerPacketID atomic.Uint64 ServerPacketID atomic.Uint64
ServerCipher cipher.AEAD serverBodyCipher cipher.AEAD
ServerBlockCipher cipher.Block serverHeaderBlock cipher.Block
ServerChaCha cipher.AEAD serverChaCha cipher.AEAD
manager *UDPSessionManager
link atomic.Pointer[transport.Link]
timer *signal.ActivityTimer
currentConn atomic.Value // stores stat.Connection
} }
func (s *ServerUDPSession) GetRemoteCipher() cipher.AEAD { func (s *ServerUDPSession) CheckPacketID(packetID uint64) bool {
ptr := s.RemoteCipher.Load() s.Lock()
if ptr == nil { defer s.Unlock()
return nil if s.Window == nil {
s.Window = new(SlidingWindow)
} }
return *ptr return s.Window.Check(packetID)
} }
func (s *ServerUDPSession) SetRemoteCipher(c cipher.AEAD) { func (s *ServerUDPSession) AddPacketID(packetID uint64) {
s.RemoteCipher.Store(&c) s.Lock()
defer s.Unlock()
if s.Window == nil {
s.Window = new(SlidingWindow)
}
s.Window.Add(packetID)
} }
type UDPSessionManager struct { type UDPSessionManager struct {
@@ -122,6 +137,7 @@ func (m *UDPSessionManager) GetOrCreate(sessionID uint64) *ServerUDPSession {
s := &ServerUDPSession{ s := &ServerUDPSession{
SessionID: sessionID, SessionID: sessionID,
manager: m,
} }
s.LastActive.Store(now) s.LastActive.Store(now)
@@ -148,6 +164,7 @@ func (m *UDPSessionManager) cleanup(now int64) {
m.sessions.Range(func(k uint64, v *ServerUDPSession) bool { m.sessions.Range(func(k uint64, v *ServerUDPSession) bool {
if now-v.LastActive.Load() > timeoutSec { if now-v.LastActive.Load() > timeoutSec {
m.sessions.Delete(k) m.sessions.Delete(k)
v.Close()
} }
return true return true
}) })
@@ -156,3 +173,11 @@ func (m *UDPSessionManager) cleanup(now int64) {
func (m *UDPSessionManager) Delete(sessionID uint64) { func (m *UDPSessionManager) Delete(sessionID uint64) {
m.sessions.Delete(sessionID) 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)
}
+149 -6
View File
@@ -2,18 +2,161 @@ package shadowsocks_2022
import ( import (
"context" "context"
"sync"
"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/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/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"
"github.com/xtls/xray-core/transport/internet/stat"
) )
type udpConnEntry struct { func (s *ServerUDPSession) UpdateConn(conn stat.Connection) {
sync.Mutex if s.currentConn.Load() == nil {
link *transport.Link s.currentConn.Store(conn)
timer *signal.ActivityTimer }
cancel context.CancelFunc 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 ( const (
+172 -36
View File
@@ -182,57 +182,48 @@ func TestTCPStream(t *testing.T) {
common.Must(err) common.Must(err)
IncreaseNonce(reader.Nonce()) IncreaseNonce(reader.Nonce())
vBuf := buf.New() dest, addrLen, err := ParseAddressPort(plainVar)
vBuf.Write(plainVar)
receivedDest, err = ReadAddressPort(vBuf)
common.Must(err) common.Must(err)
receivedDest = net.TCPDestination(dest.Address, dest.Port)
plainVar = plainVar[addrLen:]
padLen := int(binary.BigEndian.Uint16(plainVar[:2]))
receivedPayload = plainVar[2+padLen:]
// Skip padding // Server sends response stream with receivedPayload as first payload
var padBytes [2]byte writer := NewServerStreamWriter(serverConn, method, rawKey, salt)
_, _ = vBuf.Read(padBytes[:]) pBuf := buf.New()
padLen := int(padBytes[0])<<8 | int(padBytes[1]) pBuf.Write(receivedPayload)
vBuf.Advance(int32(padLen)) _ = writer.WriteMultiBuffer(buf.MultiBuffer{pBuf})
receivedPayload = make([]byte, vBuf.Len()) // Read and echo additional stream data
copy(receivedPayload, vBuf.Bytes())
vBuf.Release()
// Server sends response handshake
serverSalt := make([]byte, method.KeySaltLength)
_, _ = rand.Read(serverSalt)
respKey := DeriveSessionSubKey(rawKey, serverSalt, method.KeySaltLength)
respAead, err := method.NewAEAD(respKey)
writer := NewStreamWriter(serverConn, respAead)
_, _ = serverConn.Write(serverSalt)
fixedResp := make([]byte, 1+8+method.KeySaltLength+2)
fixedResp[0] = HeaderTypeServer
binary.BigEndian.PutUint64(fixedResp[1:9], uint64(time.Now().Unix()))
copy(fixedResp[9:9+method.KeySaltLength], salt)
binary.BigEndian.PutUint16(fixedResp[9+method.KeySaltLength:11+method.KeySaltLength], 0)
fixedChunk := respAead.Seal(nil, writer.Nonce(), fixedResp, nil)
IncreaseNonce(writer.Nonce())
_, _ = serverConn.Write(fixedChunk)
// Echo stream data
mb, err := reader.ReadMultiBuffer() mb, err := reader.ReadMultiBuffer()
common.Must(err) common.Must(err)
_ = writer.WriteMultiBuffer(mb) _ = writer.WriteMultiBuffer(mb)
_ = writer.Close()
}() }()
// Client goroutine // Client goroutine
go func() { go func() {
defer wg.Done() defer wg.Done()
clientSalt, writer, err := ClientHandshake(clientConn, method, [][]byte{rawKey}, dest, testPayload) 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) common.Must(err)
reader, _, err := ClientVerifyServerResponse(clientConn, method, rawKey, clientSalt) reader, err := ReadTCPResponse(clientConn, method, rawKey, clientSalt)
common.Must(err) 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 // Send additional stream data
streamData := []byte("stream chunk test") streamData := []byte("stream chunk test")
_ = writer.WriteChunk(streamData) _ = writer.WriteMultiBuffer(buf.MultiBuffer{buf.FromBytes(streamData)})
mb, err := reader.ReadMultiBuffer() mb, err := reader.ReadMultiBuffer()
common.Must(err) common.Must(err)
@@ -272,12 +263,14 @@ func TestUDPCodec(t *testing.T) {
psk := make([]byte, method.KeySaltLength) psk := make([]byte, method.KeySaltLength)
_, _ = rand.Read(psk) _, _ = rand.Read(psk)
clientCodec, err := NewUDPPacketCodec(method, psk) clientCodec, err := NewUDPPacketCodec(method, [][]byte{psk})
common.Must(err) common.Must(err)
serverCodec, err := NewUDPServerCodec(method, psk, time.Minute) serverCodec, err := NewUDPServerCodec(method, psk, time.Minute)
common.Must(err) common.Must(err)
pktBuf, err := clientCodec.EncodeClientPacket(dest, payload) session, err := clientCodec.NewClientSession()
common.Must(err)
pktBuf, err := session.EncodePacket(dest, payload)
common.Must(err) common.Must(err)
defer pktBuf.Release() defer pktBuf.Release()
@@ -360,3 +353,146 @@ func TestMultiUserManager(t *testing.T) {
t.Fatal("user1 should have been removed") 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")
}
})
}
}
+235 -115
View File
@@ -1,18 +1,25 @@
package shadowsocks_2022 package shadowsocks_2022
import ( import (
"context"
"crypto/cipher" "crypto/cipher"
"crypto/rand" "crypto/rand"
"encoding/binary" "encoding/binary"
"io" "io"
"math" "math"
mrand "math/rand/v2" mrand "math/rand/v2"
"sync"
"time" "time"
"github.com/xtls/xray-core/common/antireplay"
"github.com/xtls/xray-core/common/buf" "github.com/xtls/xray-core/common/buf"
"github.com/xtls/xray-core/common/errors" "github.com/xtls/xray-core/common/errors"
"github.com/xtls/xray-core/common/net" "github.com/xtls/xray-core/common/net"
"github.com/xtls/xray-core/common/protocol" "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( var addrParser = protocol.NewAddressParser(
@@ -38,15 +45,6 @@ func WriteAddressPort(w io.Writer, dest net.Destination) error {
return addrParser.WriteAddressPort(w, dest.Address, dest.Port) return addrParser.WriteAddressPort(w, dest.Address, dest.Port)
} }
// ReadAddressPort reads a destination address and port in SOCKS5 format
func ReadAddressPort(r io.Reader) (net.Destination, error) {
addr, port, err := addrParser.ReadAddressPort(nil, r)
if err != nil {
return net.Destination{}, err
}
return net.TCPDestination(addr, port), nil
}
// AddrPortLength returns the serialized length of a destination in SOCKS5 format // AddrPortLength returns the serialized length of a destination in SOCKS5 format
func AddrPortLength(dest net.Destination) int { func AddrPortLength(dest net.Destination) int {
switch dest.Address.Family() { switch dest.Address.Family() {
@@ -119,8 +117,16 @@ func (w *StreamWriter) Write(p []byte) (int, error) {
func (w *StreamWriter) WriteMultiBuffer(mb buf.MultiBuffer) error { func (w *StreamWriter) WriteMultiBuffer(mb buf.MultiBuffer) error {
defer buf.ReleaseMulti(mb) defer buf.ReleaseMulti(mb)
for _, b := range mb { for _, b := range mb {
if err := w.WriteChunk(b.Bytes()); err != nil { p := b.Bytes()
return err 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 return nil
@@ -168,7 +174,7 @@ func (r *StreamReader) Read(p []byte) (int, error) {
IncreaseNonce(r.nonce[:]) IncreaseNonce(r.nonce[:])
payloadLen := int(binary.BigEndian.Uint16(decryptedLen)) payloadLen := int(binary.BigEndian.Uint16(decryptedLen))
if payloadLen == 0 { if payloadLen == 0 || payloadLen > MaxPacketSize {
return 0, ErrInvalidRequest return 0, ErrInvalidRequest
} }
@@ -194,11 +200,10 @@ func (r *StreamReader) Read(p []byte) (int, error) {
func (r *StreamReader) ReadMultiBuffer() (buf.MultiBuffer, error) { func (r *StreamReader) ReadMultiBuffer() (buf.MultiBuffer, error) {
if r.cached > 0 { if r.cached > 0 {
b := buf.New() mb := buf.MergeBytes(nil, r.buffer[r.offset:r.offset+r.cached])
b.Write(r.buffer[r.offset : r.offset+r.cached])
r.cached = 0 r.cached = 0
r.offset = 0 r.offset = 0
return buf.MultiBuffer{b}, nil return mb, nil
} }
if _, err := io.ReadFull(r.reader, r.lenBuf[:]); err != nil { if _, err := io.ReadFull(r.reader, r.lenBuf[:]); err != nil {
@@ -212,7 +217,7 @@ func (r *StreamReader) ReadMultiBuffer() (buf.MultiBuffer, error) {
IncreaseNonce(r.nonce[:]) IncreaseNonce(r.nonce[:])
payloadLen := int(binary.BigEndian.Uint16(decryptedLen)) payloadLen := int(binary.BigEndian.Uint16(decryptedLen))
if payloadLen == 0 { if payloadLen == 0 || payloadLen > MaxPacketSize {
return nil, ErrInvalidRequest return nil, ErrInvalidRequest
} }
@@ -227,9 +232,8 @@ func (r *StreamReader) ReadMultiBuffer() (buf.MultiBuffer, error) {
} }
IncreaseNonce(r.nonce[:]) IncreaseNonce(r.nonce[:])
b := buf.New() mb := buf.MergeBytes(nil, decryptedPayload)
b.Write(decryptedPayload) return mb, nil
return buf.MultiBuffer{b}, nil
} }
type ClientRequestHeader struct { type ClientRequestHeader struct {
@@ -237,13 +241,8 @@ type ClientRequestHeader struct {
EarlyData []byte EarlyData []byte
} }
func ReadClientRequestHeader(conn io.Reader, reader *StreamReader) (*ClientRequestHeader, error) { func ReadClientRequestHeaderWithFixed(reader *StreamReader, fixedChunk []byte) (*ClientRequestHeader, error) {
var fixedBuf [RequestHeaderFixedChunkLength + AEADTagSize]byte plainFixed, err := reader.cipher.Open(fixedChunk[:0], reader.Nonce(), fixedChunk, nil)
if _, err := io.ReadFull(conn, fixedBuf[:]); err != nil {
return nil, err
}
plainFixed, err := reader.cipher.Open(fixedBuf[:0], reader.Nonce(), fixedBuf[:], nil)
if err != nil { if err != nil {
return nil, errors.New("failed to decrypt client request header").Base(err) return nil, errors.New("failed to decrypt client request header").Base(err)
} }
@@ -272,7 +271,7 @@ func ReadClientRequestHeader(conn io.Reader, reader *StreamReader) (*ClientReque
} else { } else {
varChunkCipher = make([]byte, needed) varChunkCipher = make([]byte, needed)
} }
if _, err := io.ReadFull(conn, varChunkCipher); err != nil { if _, err := io.ReadFull(reader.reader, varChunkCipher); err != nil {
return nil, err return nil, err
} }
@@ -282,31 +281,34 @@ func ReadClientRequestHeader(conn io.Reader, reader *StreamReader) (*ClientReque
} }
IncreaseNonce(reader.Nonce()) IncreaseNonce(reader.Nonce())
b := buf.New() dest, addrLen, err := ParseAddressPort(plainVar)
b.Write(plainVar)
defer b.Release()
dest, err := ReadAddressPort(b)
if err != nil { if err != nil {
return nil, err return nil, err
} }
dest.Network = net.Network_TCP
var padLenBytes [2]byte offset := addrLen
if _, err := b.Read(padLenBytes[:]); err != nil { if len(plainVar) < offset+2 {
return nil, err return nil, ErrPacketTooShort
} }
paddingLen := int(binary.BigEndian.Uint16(padLenBytes[:])) paddingLen := int(binary.BigEndian.Uint16(plainVar[offset : offset+2]))
if int(b.Len()) < paddingLen { offset += 2
if len(plainVar) < offset+paddingLen {
return nil, ErrNoPadding return nil, ErrNoPadding
} }
if paddingLen > 0 { offset += paddingLen
b.Advance(int32(paddingLen))
}
var earlyData []byte var earlyData []byte
if b.Len() > 0 { var payloadLen int
earlyData = make([]byte, b.Len()) if len(plainVar) > offset {
copy(earlyData, b.Bytes()) 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{ return &ClientRequestHeader{
@@ -315,34 +317,6 @@ func ReadClientRequestHeader(conn io.Reader, reader *StreamReader) (*ClientReque
}, nil }, nil
} }
// ClientHandshake writes the full client request header to w
func ClientHandshake(w io.Writer, method *CipherMethod, pskList [][]byte, dest net.Destination, payload []byte) ([]byte, *StreamWriter, error) {
salt := make([]byte, method.KeySaltLength)
if _, err := io.ReadFull(rand.Reader, salt); err != nil {
return nil, nil, err
}
writer, err := WriteTCPRequest(w, method, pskList, dest, salt, payload)
if err != nil {
return nil, nil, err
}
return salt, writer.(*StreamWriter), nil
}
// ClientVerifyServerResponse reads and verifies the server's handshake response
func ClientVerifyServerResponse(r io.Reader, method *CipherMethod, psk []byte, clientSalt []byte) (*StreamReader, []byte, error) {
reader, err := ReadTCPResponse(r, method, psk, clientSalt)
if err != nil {
return nil, nil, err
}
sr := reader.(*StreamReader)
var initialPayload []byte
if sr.cached > 0 {
initialPayload = make([]byte, sr.cached)
copy(initialPayload, sr.buffer[sr.offset:sr.offset+sr.cached])
}
return sr, initialPayload, nil
}
// WriteTCPRequest writes the Shadowsocks 2022 request header into w and returns a body writer. // 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) { func WriteTCPRequest(w io.Writer, method *CipherMethod, pskList [][]byte, dest net.Destination, clientSalt []byte, payload []byte) (buf.Writer, error) {
finalPSK := pskList[len(pskList)-1] finalPSK := pskList[len(pskList)-1]
@@ -354,7 +328,16 @@ func WriteTCPRequest(w io.Writer, method *CipherMethod, pskList [][]byte, dest n
writer := NewStreamWriter(w, aead) writer := NewStreamWriter(w, aead)
handshakeBuf := buf.New() 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() defer handshakeBuf.Release()
handshakeBuf.Write(clientSalt) handshakeBuf.Write(clientSalt)
@@ -372,14 +355,6 @@ func WriteTCPRequest(w io.Writer, method *CipherMethod, pskList [][]byte, dest n
handshakeBuf.Write(encryptedEIH[:]) handshakeBuf.Write(encryptedEIH[:])
} }
payloadLen := len(payload)
var paddingLen int
if payloadLen < MaxPaddingLength {
paddingLen = mrand.IntN(MaxPaddingLength-payloadLen) + 1
}
addrPortLen := AddrPortLength(dest)
varHeaderLen := addrPortLen + 2 + paddingLen + payloadLen
var fixedHeaderPlaintext [RequestHeaderFixedChunkLength]byte var fixedHeaderPlaintext [RequestHeaderFixedChunkLength]byte
fixedHeaderPlaintext[0] = HeaderTypeClient fixedHeaderPlaintext[0] = HeaderTypeClient
binary.BigEndian.PutUint64(fixedHeaderPlaintext[1:9], uint64(time.Now().Unix())) binary.BigEndian.PutUint64(fixedHeaderPlaintext[1:9], uint64(time.Now().Unix()))
@@ -389,7 +364,7 @@ func WriteTCPRequest(w io.Writer, method *CipherMethod, pskList [][]byte, dest n
IncreaseNonce(writer.nonce[:]) IncreaseNonce(writer.nonce[:])
handshakeBuf.Write(fixedChunk) handshakeBuf.Write(fixedChunk)
varHeaderBuf := buf.New() varHeaderBuf := buf.NewWithSize(int32(varHeaderLen))
defer varHeaderBuf.Release() defer varHeaderBuf.Release()
if err := WriteAddressPort(varHeaderBuf, dest); err != nil { if err := WriteAddressPort(varHeaderBuf, dest); err != nil {
@@ -421,12 +396,21 @@ func WriteTCPRequest(w io.Writer, method *CipherMethod, pskList [][]byte, dest n
// ReadTCPResponse reads and verifies the server's handshake response and returns a reader for the stream. // 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) { func ReadTCPResponse(r io.Reader, method *CipherMethod, psk []byte, clientSalt []byte) (buf.Reader, error) {
var serverSalt [32]byte fixedPlainLen := 1 + 8 + method.KeySaltLength + 2
serverSaltSlice := serverSalt[:method.KeySaltLength] chunkCipherLen := fixedPlainLen + AEADTagSize
if _, err := io.ReadFull(r, serverSaltSlice); err != nil { headerLen := method.KeySaltLength + chunkCipherLen
return nil, err
// 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) sessionKey := DeriveSessionSubKey(psk, serverSaltSlice, method.KeySaltLength)
aead, err := method.NewAEAD(sessionKey) aead, err := method.NewAEAD(sessionKey)
if err != nil { if err != nil {
@@ -435,14 +419,6 @@ func ReadTCPResponse(r io.Reader, method *CipherMethod, psk []byte, clientSalt [
reader := NewStreamReader(r, aead) reader := NewStreamReader(r, aead)
fixedPlainLen := 1 + 8 + method.KeySaltLength + 2
chunkCipherLen := fixedPlainLen + AEADTagSize
var chunkBuf [64]byte
chunkSlice := chunkBuf[:chunkCipherLen]
if _, err := io.ReadFull(r, chunkSlice); err != nil {
return nil, err
}
decryptedFixed, err := reader.cipher.Open(chunkSlice[:0], reader.nonce[:], chunkSlice, nil) decryptedFixed, err := reader.cipher.Open(chunkSlice[:0], reader.nonce[:], chunkSlice, nil)
if err != nil { if err != nil {
return nil, errors.New("failed to decrypt server response header").Base(err) return nil, errors.New("failed to decrypt server response header").Base(err)
@@ -484,46 +460,190 @@ func ReadTCPResponse(r io.Reader, method *CipherMethod, psk []byte, clientSalt [
return reader, nil return reader, nil
} }
// WriteTCPResponse writes the server handshake response and returns a body writer for server stream. // ServerStreamWriter lazily sends the response header along with the first payload chunk per SIP022 §3.1.2 & §3.1.4.
func WriteTCPResponse(w io.Writer, method *CipherMethod, psk []byte, clientSalt []byte, initialPayload []byte) (buf.Writer, error) { 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 var serverSalt [32]byte
serverSaltSlice := serverSalt[:method.KeySaltLength] serverSaltSlice := serverSalt[:s.method.KeySaltLength]
if _, err := io.ReadFull(rand.Reader, serverSaltSlice); err != nil { if _, err := io.ReadFull(rand.Reader, serverSaltSlice); err != nil {
return nil, err return nil, err
} }
respKey := DeriveSessionSubKey(psk, serverSaltSlice, method.KeySaltLength) respKey := DeriveSessionSubKey(s.psk, serverSaltSlice, s.method.KeySaltLength)
respAead, err := method.NewAEAD(respKey) respAead, err := s.method.NewAEAD(respKey)
if err != nil { if err != nil {
return nil, err return nil, err
} }
writer := NewStreamWriter(w, respAead) sw := NewStreamWriter(s.w, respAead)
respBuf := buf.New() totalHeaderLen := int32(s.method.KeySaltLength + 1 + 8 + s.method.KeySaltLength + 2 + AEADTagSize + len(payload) + AEADTagSize)
defer respBuf.Release() outBuf := buf.NewWithSize(totalHeaderLen)
defer outBuf.Release()
respBuf.Write(serverSaltSlice) outBuf.Write(serverSaltSlice)
var fixedRespPlain [1 + 8 + 32 + 2]byte var fixedRespPlain [1 + 8 + 32 + 2]byte
fixedRespSlice := fixedRespPlain[:1+8+method.KeySaltLength+2] fixedRespSlice := fixedRespPlain[:1+8+s.method.KeySaltLength+2]
fixedRespSlice[0] = HeaderTypeServer fixedRespSlice[0] = HeaderTypeServer
binary.BigEndian.PutUint64(fixedRespSlice[1:9], uint64(time.Now().Unix())) binary.BigEndian.PutUint64(fixedRespSlice[1:9], uint64(time.Now().Unix()))
copy(fixedRespSlice[9:9+method.KeySaltLength], clientSalt) copy(fixedRespSlice[9:9+s.method.KeySaltLength], s.clientSalt)
binary.BigEndian.PutUint16(fixedRespSlice[9+method.KeySaltLength:11+method.KeySaltLength], uint16(len(initialPayload))) binary.BigEndian.PutUint16(fixedRespSlice[9+s.method.KeySaltLength:11+s.method.KeySaltLength], uint16(len(payload)))
fixedRespChunk := writer.cipher.Seal(nil, writer.nonce[:], fixedRespSlice, nil) fixedRespChunk := sw.cipher.Seal(nil, sw.nonce[:], fixedRespSlice, nil)
IncreaseNonce(writer.nonce[:]) IncreaseNonce(sw.nonce[:])
respBuf.Write(fixedRespChunk) outBuf.Write(fixedRespChunk)
if len(initialPayload) > 0 { if len(payload) > 0 {
initialChunk := writer.cipher.Seal(nil, writer.nonce[:], initialPayload, nil) payloadChunk := sw.cipher.Seal(nil, sw.nonce[:], payload, nil)
IncreaseNonce(writer.nonce[:]) IncreaseNonce(sw.nonce[:])
respBuf.Write(initialChunk) outBuf.Write(payloadChunk)
} }
if _, err := w.Write(respBuf.Bytes()); err != nil { if _, err := s.w.Write(outBuf.Bytes()); err != nil {
return nil, err return nil, err
} }
return sw, nil
return writer, 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)
} }
+28 -13
View File
@@ -15,27 +15,28 @@ Plainly enabling it in the config probably will result nothing, or lock your rou
## DETAILS ## DETAILS
By default, enabling the feature will only bring the tun interface up. \ 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 uses the first IPv4 prefix from `gateway` to configure the utun point-to-point address. \ 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`. \
Windows, Linux and macOS can also apply system routes from `autoSystemRoutingTable`. 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. \ 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. \
For more advanced routing policies or rules, OS level configuration can still manage the named interface (e.g. xray0) when it appears. 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. \ 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. 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 (`autoSystemDNS`) ### SYSTEM DNS ON LINUX (`autoSystemDnsToGateway`)
On Linux, setting `autoSystemDNS` 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. 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 applies only when all of these hold: 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` - the system runs systemd and `resolvectl` is on `PATH`
- `systemd-resolved` is enabled and actually managing DNS (installed but not running has no effect) - `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 - 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) - 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` incremented by one (e.g. `192.168.100.1/30` -> `192.168.100.2`). 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. 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 the option does nothing and DNS is left to the OS. In practice this means you also need a routing rule sending the interface's port 53 to a `dns` outbound, for example: 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 ```json
"routing": { "routing": {
@@ -49,19 +50,19 @@ The check is a preflight, not a proof for arbitrary rules. It sends its query fr
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. 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. 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. 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 does not apply, DNS is left alone and the leak described in XTLS/Xray-core#6454 remains: 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 | | Environment | Behaviour |
|---|---| |---|---|
| systemd distribution with systemd-resolved enabled | applies | | systemd distribution with systemd-resolved enabled | applies |
| Alpine, Void, Devuan, OpenRC-based, OpenWrt | no `resolvectl`, skipped | | Alpine, Void, Devuan, OpenRC-based, OpenWrt | no `resolvectl`, does not start |
| DNS managed by dnsmasq / unbound / BIND / static `resolv.conf` | unreachable by `resolvectl`, skipped | | DNS managed by dnsmasq / unbound / BIND / static `resolv.conf` | unreachable by `resolvectl`, does not start |
| Containers without a systemd-resolved daemon | skipped | | Containers without a systemd-resolved daemon | does not start |
| systemd older than 240 | `default-route` unavailable, skipped | | 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. 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.
@@ -198,6 +199,20 @@ 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. 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. \ 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. 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. 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.
+16 -6
View File
@@ -32,7 +32,8 @@ type Config struct {
AutoSystemRoutingTable []string `protobuf:"bytes,6,rep,name=auto_system_routing_table,json=autoSystemRoutingTable,proto3" json:"auto_system_routing_table,omitempty"` 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"` 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"` Desc string `protobuf:"bytes,8,opt,name=desc,proto3" json:"desc,omitempty"`
AutoSystemDns bool `protobuf:"varint,9,opt,name=auto_system_dns,json=autoSystemDns,proto3" json:"auto_system_dns,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 unknownFields protoimpl.UnknownFields
sizeCache protoimpl.SizeCache sizeCache protoimpl.SizeCache
} }
@@ -123,18 +124,25 @@ func (x *Config) GetDesc() string {
return "" return ""
} }
func (x *Config) GetAutoSystemDns() bool { func (x *Config) GetAutoSystemDnsToGateway() bool {
if x != nil { if x != nil {
return x.AutoSystemDns return x.AutoSystemDnsToGateway
} }
return false return false
} }
func (x *Config) GetAutoSystemWfpBlockLeak() []string {
if x != nil {
return x.AutoSystemWfpBlockLeak
}
return nil
}
var File_proxy_tun_config_proto protoreflect.FileDescriptor var File_proxy_tun_config_proto protoreflect.FileDescriptor
const file_proxy_tun_config_proto_rawDesc = "" + const file_proxy_tun_config_proto_rawDesc = "" +
"\n" + "\n" +
"\x16proxy/tun/config.proto\x12\x0exray.proxy.tun\"\xaa\x02\n" + "\x16proxy/tun/config.proto\x12\x0exray.proxy.tun\"\xfa\x02\n" +
"\x06Config\x12\x12\n" + "\x06Config\x12\x12\n" +
"\x04name\x18\x01 \x01(\tR\x04name\x12\x10\n" + "\x04name\x18\x01 \x01(\tR\x04name\x12\x10\n" +
"\x03MTU\x18\x02 \x01(\rR\x03MTU\x12\x18\n" + "\x03MTU\x18\x02 \x01(\rR\x03MTU\x12\x18\n" +
@@ -144,8 +152,10 @@ const file_proxy_tun_config_proto_rawDesc = "" +
"user_level\x18\x05 \x01(\rR\tuserLevel\x129\n" + "user_level\x18\x05 \x01(\rR\tuserLevel\x129\n" +
"\x19auto_system_routing_table\x18\x06 \x03(\tR\x16autoSystemRoutingTable\x128\n" + "\x19auto_system_routing_table\x18\x06 \x03(\tR\x16autoSystemRoutingTable\x128\n" +
"\x18auto_outbounds_interface\x18\a \x01(\tR\x16autoOutboundsInterface\x12\x12\n" + "\x18auto_outbounds_interface\x18\a \x01(\tR\x16autoOutboundsInterface\x12\x12\n" +
"\x04desc\x18\b \x01(\tR\x04desc\x12&\n" + "\x04desc\x18\b \x01(\tR\x04desc\x12:\n" +
"\x0fauto_system_dns\x18\t \x01(\bR\rautoSystemDnsBL\n" + "\x1aauto_system_dns_to_gateway\x18\t \x01(\bR\x16autoSystemDnsToGateway\x12:\n" +
"\x1aauto_system_wfp_block_leak\x18\n" +
" \x03(\tR\x16autoSystemWfpBlockLeakBL\n" +
"\x12com.xray.proxy.tunP\x01Z#github.com/xtls/xray-core/proxy/tun\xaa\x02\x0eXray.Proxy.Tunb\x06proto3" "\x12com.xray.proxy.tunP\x01Z#github.com/xtls/xray-core/proxy/tun\xaa\x02\x0eXray.Proxy.Tunb\x06proto3"
var ( var (
+2 -1
View File
@@ -15,5 +15,6 @@ message Config {
repeated string auto_system_routing_table = 6; repeated string auto_system_routing_table = 6;
string auto_outbounds_interface = 7; string auto_outbounds_interface = 7;
string desc = 8; string desc = 8;
bool auto_system_dns = 9; bool auto_system_dns_to_gateway = 9;
repeated string auto_system_wfp_block_leak = 10;
} }
+4 -2
View File
@@ -166,12 +166,14 @@ func (t *Handler) Start() error {
} }
// Platform-specific system DNS takeover, where the platform implements it. // Platform-specific system DNS takeover, where the platform implements it.
// Non-fatal: a failure leaves DNS management with the OS. // Rather no TUN than one that the system DNS bypasses.
if c, ok := tunInterface.(interface { if c, ok := tunInterface.(interface {
ConfigureSystemDNS(context.Context, string) error ConfigureSystemDNS(context.Context, string) error
}); ok { }); ok {
if err := c.ConfigureSystemDNS(t.ctx, t.tag); err != nil { if err := c.ConfigureSystemDNS(t.ctx, t.tag); err != nil {
errors.LogInfoInner(t.ctx, err, "[tun] system DNS not configured") _ = tunStack.Close()
_ = tunInterface.Close()
return errors.New("unable to set the system DNS (remove autoSystemDnsToGateway to run without)").Base(err)
} }
} }
+21 -15
View File
@@ -53,23 +53,29 @@ var resolvectlRunner = func(name string, args ...string) ([]byte, error) {
} }
// systemDNSAddrs derives the addresses used for the system DNS takeover from the // systemDNSAddrs derives the addresses used for the system DNS takeover from the
// first IPv4 gateway: the gateway address itself is what a query from this // first IPv4 gateway, or without one, the first IPv6 gateway: the gateway
// interface appears to come from, and the next address is what the resolver is // address itself is what a query from this interface appears to come from, and
// pointed at. The latter belongs to the TUN and is answered inside Xray; // the next address is what the resolver is pointed at. The latter belongs to
// handing the configured public resolvers to resolvectl instead would leave the // the TUN and is answered inside Xray; handing the configured public resolvers
// system querying them directly over the physical link, defeating the point of // to resolvectl instead would leave the system querying them directly over the
// the TUN. // physical link, defeating the point of the TUN.
func systemDNSAddrs(gateway []string) (source, dns netip.Addr, ok bool) { func systemDNSAddrs(gateway []string) (source, dns netip.Addr, ok bool) {
var first6 netip.Addr
for _, address := range gateway { for _, address := range gateway {
prefix, err := netip.ParsePrefix(address) prefix, err := netip.ParsePrefix(address)
if err != nil { if err != nil {
continue continue
} }
addr := prefix.Addr() addr := prefix.Addr()
if !addr.Is4() { if addr.Is4() {
continue return addr, addr.Next(), true
} }
return addr, addr.Next(), true if !first6.IsValid() {
first6 = addr
}
}
if first6.IsValid() {
return first6, first6.Next(), true
} }
return netip.Addr{}, netip.Addr{}, false return netip.Addr{}, netip.Addr{}, false
} }
@@ -115,11 +121,11 @@ const probeSourcePort = 49152
// Overridable for tests. // Overridable for tests.
var verifyDNSRouting = func(ctx context.Context, inboundTag, source, address string) error { var verifyDNSRouting = func(ctx context.Context, inboundTag, source, address string) error {
ip, err := netip.ParseAddr(address) ip, err := netip.ParseAddr(address)
if err != nil || !ip.Is4() { if err != nil {
return errors.New("invalid DNS address ", address).Base(err) return errors.New("invalid DNS address ", address).Base(err)
} }
src, err := netip.ParseAddr(source) src, err := netip.ParseAddr(source)
if err != nil || !src.Is4() { if err != nil || src.Is4() != ip.Is4() {
return errors.New("invalid source address ", source).Base(err) return errors.New("invalid source address ", source).Base(err)
} }
@@ -182,10 +188,10 @@ var verifyDNSRouting = func(ctx context.Context, inboundTag, source, address str
// //
// It acts only when the config opts in, and it verifies the data path first: // 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 // unless a query to the advertised address would actually be handled, host-wide
// resolution is left to the OS, which is the documented default. Errors are // resolution is left to the OS and an error returned. The caller does not start
// returned to the caller, which treats them as non-fatal. // the TUN on an error, as the system DNS would bypass it.
func (t *LinuxTun) ConfigureSystemDNS(ctx context.Context, inboundTag string) error { func (t *LinuxTun) ConfigureSystemDNS(ctx context.Context, inboundTag string) error {
if !t.options.AutoSystemDns { if !t.options.AutoSystemDnsToGateway {
return nil return nil
} }
if t.systemDNSSet { if t.systemDNSSet {
@@ -202,7 +208,7 @@ func (t *LinuxTun) ConfigureSystemDNS(ctx context.Context, inboundTag string) er
source, address, ok := systemDNSAddrs(t.options.Gateway) source, address, ok := systemDNSAddrs(t.options.Gateway)
if !ok { if !ok {
return errors.New("no IPv4 gateway, cannot derive a system DNS address") return errors.New("no gateway, cannot derive a system DNS address")
} }
iface := t.ifaceName() iface := t.ifaceName()
+12
View File
@@ -191,3 +191,15 @@ func TestVerifyDNSRoutingDecisions(t *testing.T) {
}) })
} }
} }
// 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")
}
}
+17 -8
View File
@@ -58,9 +58,9 @@ func recorder(t *testing.T, failOn string) *[][]string {
func optedInTun() *LinuxTun { func optedInTun() *LinuxTun {
return &LinuxTun{ return &LinuxTun{
options: &Config{ options: &Config{
Name: "xray_tun", Name: "xray_tun",
Gateway: []string{"192.168.100.1/30"}, Gateway: []string{"192.168.100.1/30"},
AutoSystemDns: true, AutoSystemDnsToGateway: true,
}, },
tunLink: testLink("xray_tun"), tunLink: testLink("xray_tun"),
} }
@@ -79,7 +79,7 @@ func TestConfigureSystemDNSDisabledByDefault(t *testing.T) {
calls := recorder(t, "") calls := recorder(t, "")
t1 := optedInTun() t1 := optedInTun()
t1.options.AutoSystemDns = false t1.options.AutoSystemDnsToGateway = false
if err := t1.ConfigureSystemDNS(context.Background(), "tun"); err != nil { if err := t1.ConfigureSystemDNS(context.Background(), "tun"); err != nil {
t.Fatalf("unexpected error: %v", err) t.Fatalf("unexpected error: %v", err)
@@ -103,7 +103,7 @@ func TestConfigureSystemDNSNoGateway(t *testing.T) {
t1.options.Gateway = nil t1.options.Gateway = nil
if err := t1.ConfigureSystemDNS(context.Background(), "tun"); err == nil { if err := t1.ConfigureSystemDNS(context.Background(), "tun"); err == nil {
t.Fatal("expected an error when no IPv4 gateway is configured") t.Fatal("expected an error when no gateway is configured")
} }
if len(*probes) != 0 { if len(*probes) != 0 {
t.Errorf("routing probe must not run without a gateway, got %d calls", len(*probes)) t.Errorf("routing probe must not run without a gateway, got %d calls", len(*probes))
@@ -351,9 +351,18 @@ func TestSystemDNSAddrs(t *testing.T) {
wantOK: false, wantOK: false,
}, },
{ {
name: "ipv6 only", name: "ipv6 only",
gateway: []string{"fc00::1/64"}, gateway: []string{"fc00::1/64"},
wantOK: false, 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,
}, },
} }
+229 -2
View File
@@ -3,17 +3,25 @@
package tun package tun
import ( import (
"bytes"
"context" "context"
"crypto/md5" "crypto/md5"
"encoding/binary" "encoding/binary"
go_errors "errors" go_errors "errors"
"net" "net"
"net/netip" "net/netip"
"os/exec"
"path/filepath"
"slices"
"strconv"
"strings"
"sync" "sync"
"syscall"
"time" "time"
"unsafe" "unsafe"
"github.com/xtls/xray-core/common/errors" "github.com/xtls/xray-core/common/errors"
"github.com/xtls/xray-core/transport/internet"
"golang.org/x/sys/windows" "golang.org/x/sys/windows"
"golang.zx2c4.com/wintun" "golang.zx2c4.com/wintun"
"golang.zx2c4.com/wireguard/windows/tunnel/winipcfg" "golang.zx2c4.com/wireguard/windows/tunnel/winipcfg"
@@ -38,6 +46,10 @@ type WindowsTun struct {
luid winipcfg.LUID luid winipcfg.LUID
cbr winipcfg.ChangeCallback cbr winipcfg.ChangeCallback
cbi winipcfg.ChangeCallback cbi winipcfg.ChangeCallback
wfp windows.Handle
resolver *savedResolver
skipStop chan struct{}
skipDone chan struct{}
closed bool closed bool
} }
@@ -197,19 +209,105 @@ 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 { if updater != nil {
t.cbr, err = winipcfg.RegisterRouteChangeCallback(func(notificationType winipcfg.MibNotificationType, route *winipcfg.MibIPforwardRow2) { // Only a registered callback goes into the fields: a nil pointer in
// them would not compare equal to nil in Close.
cbr, err := winipcfg.RegisterRouteChangeCallback(func(notificationType winipcfg.MibNotificationType, route *winipcfg.MibIPforwardRow2) {
updater.Update() updater.Update()
}) })
if err != nil { if err != nil {
return err return err
} }
t.cbi, err = winipcfg.RegisterInterfaceChangeCallback(func(notificationType winipcfg.MibNotificationType, iface *winipcfg.MibIPInterfaceRow) { t.cbr = cbr
cbi, err := winipcfg.RegisterInterfaceChangeCallback(func(notificationType winipcfg.MibNotificationType, iface *winipcfg.MibIPInterfaceRow) {
updater.Update() updater.Update()
}) })
if err != nil { if err != nil {
return err return err
} }
t.cbi = cbi
} }
return nil return nil
} }
@@ -236,6 +334,20 @@ func (t *WindowsTun) Close() error {
t.luid.FlushIPAddresses(windows.AF_INET6) t.luid.FlushIPAddresses(windows.AF_INET6)
t.luid.FlushDNS(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{}) { if t.session != (wintun.Session{}) {
t.session.End() t.session.End()
} }
@@ -245,6 +357,121 @@ func (t *WindowsTun) Close() error {
return nil 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) { func (t *WindowsTun) Name() (string, error) {
row, err := t.luid.Interface() row, err := t.luid.Interface()
if err != nil { if err != nil {
+471
View File
@@ -0,0 +1,471 @@
//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
@@ -0,0 +1,206 @@
//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)
}
}
+12 -2
View File
@@ -52,9 +52,12 @@ func (b *bind) Open(port uint16) (fns []conn.ReceiveFunc, actualPort uint16, err
case <-ch: case <-ch:
default: default:
errors.LogErrorInner(context.Background(), err, "unexpected closed") errors.LogErrorInner(context.Background(), err, "unexpected closed")
if b.downFunc != nil { b.mu.Lock()
downFunc := b.downFunc
b.mu.Unlock()
if downFunc != nil {
go func() { go func() {
common.Must(b.downFunc()) common.Must(downFunc())
}() }()
} }
} }
@@ -76,6 +79,13 @@ func (b *bind) Open(port uint16) (fns []conn.ReceiveFunc, actualPort uint16, err
}, uint16(c.LocalAddr().(*net.UDPAddr).Port), nil }, 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 { func (b *bind) Close() error {
b.mu.Lock() b.mu.Lock()
defer b.mu.Unlock() defer b.mu.Unlock()
+8 -5
View File
@@ -287,7 +287,13 @@ func (h *Handler) init(ctx context.Context) error {
} }
return pktConn, nil return pktConn, nil
} }
bind := &bind{} // 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,
}
logger := &device.Logger{ logger := &device.Logger{
Verbosef: func(format string, args ...any) { Verbosef: func(format string, args ...any) {
log.Record(&log.GeneralMessage{ log.Record(&log.GeneralMessage{
@@ -303,10 +309,7 @@ func (h *Handler) init(ctx context.Context) error {
}, },
} }
dev := device.NewDevice(h.tun, bind, logger) dev := device.NewDevice(h.tun, bind, logger)
bind.resolveFunc = resolveFunc bind.setDownFunc(dev.Down)
bind.listenFunc = listenFunc
bind.downFunc = dev.Down
bind.reserved = h.conf.Reserved
var cfg strings.Builder var cfg strings.Builder
cfg.WriteString("private_key=" + h.conf.SecretKey + "\n") cfg.WriteString("private_key=" + h.conf.SecretKey + "\n")
for _, peer := range h.conf.Peers { for _, peer := range h.conf.Peers {
+5 -3
View File
@@ -113,7 +113,7 @@ func NewServer(ctx context.Context, conf *DeviceConfig) (*Server, error) {
users.Store(user.Account.(*MemoryAccount).Pub, user) users.Store(user.Account.(*MemoryAccount).Pub, user)
} }
return &Server{ s := &Server{
conf: conf, conf: conf,
ctx: core.ToBackgroundDetachedContext(ctx), ctx: core.ToBackgroundDetachedContext(ctx),
policyManager: p, policyManager: p,
@@ -131,7 +131,10 @@ func NewServer(ctx context.Context, conf *DeviceConfig) (*Server, error) {
pub: pub, pub: pub,
users: users, users: users,
}, nil }
// Install the stack's protocol handlers before the device can deliver packets to it (Start -> dev.Up).
CreateForwarder(stack, s.HandleConnection)
return s, nil
} }
func (s *Server) AddUser(ctx context.Context, user *protocol.MemoryUser) error { func (s *Server) AddUser(ctx context.Context, user *protocol.MemoryUser) error {
@@ -320,7 +323,6 @@ func (s *Server) Start() error {
return err return err
} }
s.dev = dev s.dev = dev
CreateForwarder(s.stack, s.HandleConnection)
return nil return nil
} }
+32
View File
@@ -0,0 +1,32 @@
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
@@ -0,0 +1,27 @@
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)")
}
}
+177 -19
View File
@@ -21,6 +21,135 @@ const (
_ = protoimpl.EnforceVersion(protoimpl.MaxVersion - 20) _ = protoimpl.EnforceVersion(protoimpl.MaxVersion - 20)
) )
type Segment_Kind int32
const (
Segment_BYTES Segment_Kind = 0
Segment_RANDOM Segment_Kind = 1
Segment_RANDOM_ASCII Segment_Kind = 2
Segment_RANDOM_DIGIT Segment_Kind = 3
Segment_TIMESTAMP Segment_Kind = 4
Segment_COUNTER Segment_Kind = 5
Segment_NONCE Segment_Kind = 6
)
// Enum value maps for Segment_Kind.
var (
Segment_Kind_name = map[int32]string{
0: "BYTES",
1: "RANDOM",
2: "RANDOM_ASCII",
3: "RANDOM_DIGIT",
4: "TIMESTAMP",
5: "COUNTER",
6: "NONCE",
}
Segment_Kind_value = map[string]int32{
"BYTES": 0,
"RANDOM": 1,
"RANDOM_ASCII": 2,
"RANDOM_DIGIT": 3,
"TIMESTAMP": 4,
"COUNTER": 5,
"NONCE": 6,
}
)
func (x Segment_Kind) Enum() *Segment_Kind {
p := new(Segment_Kind)
*p = x
return p
}
func (x Segment_Kind) String() string {
return protoimpl.X.EnumStringOf(x.Descriptor(), protoreflect.EnumNumber(x))
}
func (Segment_Kind) Descriptor() protoreflect.EnumDescriptor {
return file_transport_internet_finalmask_noise_config_proto_enumTypes[0].Descriptor()
}
func (Segment_Kind) Type() protoreflect.EnumType {
return &file_transport_internet_finalmask_noise_config_proto_enumTypes[0]
}
func (x Segment_Kind) Number() protoreflect.EnumNumber {
return protoreflect.EnumNumber(x)
}
// Deprecated: Use Segment_Kind.Descriptor instead.
func (Segment_Kind) EnumDescriptor() ([]byte, []int) {
return file_transport_internet_finalmask_noise_config_proto_rawDescGZIP(), []int{0, 0}
}
type Segment struct {
state protoimpl.MessageState `protogen:"open.v1"`
Kind Segment_Kind `protobuf:"varint,1,opt,name=kind,proto3,enum=xray.transport.internet.finalmask.noise.Segment_Kind" json:"kind,omitempty"`
Bytes []byte `protobuf:"bytes,2,opt,name=bytes,proto3" json:"bytes,omitempty"`
MinSize int64 `protobuf:"varint,3,opt,name=min_size,json=minSize,proto3" json:"min_size,omitempty"`
MaxSize int64 `protobuf:"varint,4,opt,name=max_size,json=maxSize,proto3" json:"max_size,omitempty"`
unknownFields protoimpl.UnknownFields
sizeCache protoimpl.SizeCache
}
func (x *Segment) Reset() {
*x = Segment{}
mi := &file_transport_internet_finalmask_noise_config_proto_msgTypes[0]
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
ms.StoreMessageInfo(mi)
}
func (x *Segment) String() string {
return protoimpl.X.MessageStringOf(x)
}
func (*Segment) ProtoMessage() {}
func (x *Segment) ProtoReflect() protoreflect.Message {
mi := &file_transport_internet_finalmask_noise_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 Segment.ProtoReflect.Descriptor instead.
func (*Segment) Descriptor() ([]byte, []int) {
return file_transport_internet_finalmask_noise_config_proto_rawDescGZIP(), []int{0}
}
func (x *Segment) GetKind() Segment_Kind {
if x != nil {
return x.Kind
}
return Segment_BYTES
}
func (x *Segment) GetBytes() []byte {
if x != nil {
return x.Bytes
}
return nil
}
func (x *Segment) GetMinSize() int64 {
if x != nil {
return x.MinSize
}
return 0
}
func (x *Segment) GetMaxSize() int64 {
if x != nil {
return x.MaxSize
}
return 0
}
type Item struct { type Item struct {
state protoimpl.MessageState `protogen:"open.v1"` state protoimpl.MessageState `protogen:"open.v1"`
RandMin int64 `protobuf:"varint,1,opt,name=rand_min,json=randMin,proto3" json:"rand_min,omitempty"` RandMin int64 `protobuf:"varint,1,opt,name=rand_min,json=randMin,proto3" json:"rand_min,omitempty"`
@@ -30,13 +159,14 @@ type Item struct {
Packet []byte `protobuf:"bytes,5,opt,name=packet,proto3" json:"packet,omitempty"` Packet []byte `protobuf:"bytes,5,opt,name=packet,proto3" json:"packet,omitempty"`
DelayMin int64 `protobuf:"varint,6,opt,name=delay_min,json=delayMin,proto3" json:"delay_min,omitempty"` DelayMin int64 `protobuf:"varint,6,opt,name=delay_min,json=delayMin,proto3" json:"delay_min,omitempty"`
DelayMax int64 `protobuf:"varint,7,opt,name=delay_max,json=delayMax,proto3" json:"delay_max,omitempty"` DelayMax int64 `protobuf:"varint,7,opt,name=delay_max,json=delayMax,proto3" json:"delay_max,omitempty"`
Segments []*Segment `protobuf:"bytes,8,rep,name=segments,proto3" json:"segments,omitempty"`
unknownFields protoimpl.UnknownFields unknownFields protoimpl.UnknownFields
sizeCache protoimpl.SizeCache sizeCache protoimpl.SizeCache
} }
func (x *Item) Reset() { func (x *Item) Reset() {
*x = Item{} *x = Item{}
mi := &file_transport_internet_finalmask_noise_config_proto_msgTypes[0] mi := &file_transport_internet_finalmask_noise_config_proto_msgTypes[1]
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
ms.StoreMessageInfo(mi) ms.StoreMessageInfo(mi)
} }
@@ -48,7 +178,7 @@ func (x *Item) String() string {
func (*Item) ProtoMessage() {} func (*Item) ProtoMessage() {}
func (x *Item) ProtoReflect() protoreflect.Message { func (x *Item) ProtoReflect() protoreflect.Message {
mi := &file_transport_internet_finalmask_noise_config_proto_msgTypes[0] mi := &file_transport_internet_finalmask_noise_config_proto_msgTypes[1]
if x != nil { if x != nil {
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
if ms.LoadMessageInfo() == nil { if ms.LoadMessageInfo() == nil {
@@ -61,7 +191,7 @@ func (x *Item) ProtoReflect() protoreflect.Message {
// Deprecated: Use Item.ProtoReflect.Descriptor instead. // Deprecated: Use Item.ProtoReflect.Descriptor instead.
func (*Item) Descriptor() ([]byte, []int) { func (*Item) Descriptor() ([]byte, []int) {
return file_transport_internet_finalmask_noise_config_proto_rawDescGZIP(), []int{0} return file_transport_internet_finalmask_noise_config_proto_rawDescGZIP(), []int{1}
} }
func (x *Item) GetRandMin() int64 { func (x *Item) GetRandMin() int64 {
@@ -113,6 +243,13 @@ func (x *Item) GetDelayMax() int64 {
return 0 return 0
} }
func (x *Item) GetSegments() []*Segment {
if x != nil {
return x.Segments
}
return nil
}
type Config struct { type Config struct {
state protoimpl.MessageState `protogen:"open.v1"` state protoimpl.MessageState `protogen:"open.v1"`
ResetMin int64 `protobuf:"varint,1,opt,name=reset_min,json=resetMin,proto3" json:"reset_min,omitempty"` ResetMin int64 `protobuf:"varint,1,opt,name=reset_min,json=resetMin,proto3" json:"reset_min,omitempty"`
@@ -124,7 +261,7 @@ type Config struct {
func (x *Config) Reset() { func (x *Config) Reset() {
*x = Config{} *x = Config{}
mi := &file_transport_internet_finalmask_noise_config_proto_msgTypes[1] mi := &file_transport_internet_finalmask_noise_config_proto_msgTypes[2]
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
ms.StoreMessageInfo(mi) ms.StoreMessageInfo(mi)
} }
@@ -136,7 +273,7 @@ func (x *Config) String() string {
func (*Config) ProtoMessage() {} func (*Config) ProtoMessage() {}
func (x *Config) ProtoReflect() protoreflect.Message { func (x *Config) ProtoReflect() protoreflect.Message {
mi := &file_transport_internet_finalmask_noise_config_proto_msgTypes[1] mi := &file_transport_internet_finalmask_noise_config_proto_msgTypes[2]
if x != nil { if x != nil {
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
if ms.LoadMessageInfo() == nil { if ms.LoadMessageInfo() == nil {
@@ -149,7 +286,7 @@ func (x *Config) ProtoReflect() protoreflect.Message {
// Deprecated: Use Config.ProtoReflect.Descriptor instead. // Deprecated: Use Config.ProtoReflect.Descriptor instead.
func (*Config) Descriptor() ([]byte, []int) { func (*Config) Descriptor() ([]byte, []int) {
return file_transport_internet_finalmask_noise_config_proto_rawDescGZIP(), []int{1} return file_transport_internet_finalmask_noise_config_proto_rawDescGZIP(), []int{2}
} }
func (x *Config) GetResetMin() int64 { func (x *Config) GetResetMin() int64 {
@@ -177,7 +314,21 @@ var File_transport_internet_finalmask_noise_config_proto protoreflect.FileDescri
const file_transport_internet_finalmask_noise_config_proto_rawDesc = "" + const file_transport_internet_finalmask_noise_config_proto_rawDesc = "" +
"\n" + "\n" +
"/transport/internet/finalmask/noise/config.proto\x12'xray.transport.internet.finalmask.noise\"\xda\x01\n" + "/transport/internet/finalmask/noise/config.proto\x12'xray.transport.internet.finalmask.noise\"\x8a\x02\n" +
"\aSegment\x12I\n" +
"\x04kind\x18\x01 \x01(\x0e25.xray.transport.internet.finalmask.noise.Segment.KindR\x04kind\x12\x14\n" +
"\x05bytes\x18\x02 \x01(\fR\x05bytes\x12\x19\n" +
"\bmin_size\x18\x03 \x01(\x03R\aminSize\x12\x19\n" +
"\bmax_size\x18\x04 \x01(\x03R\amaxSize\"h\n" +
"\x04Kind\x12\t\n" +
"\x05BYTES\x10\x00\x12\n" +
"\n" +
"\x06RANDOM\x10\x01\x12\x10\n" +
"\fRANDOM_ASCII\x10\x02\x12\x10\n" +
"\fRANDOM_DIGIT\x10\x03\x12\r\n" +
"\tTIMESTAMP\x10\x04\x12\v\n" +
"\aCOUNTER\x10\x05\x12\t\n" +
"\x05NONCE\x10\x06\"\xa8\x02\n" +
"\x04Item\x12\x19\n" + "\x04Item\x12\x19\n" +
"\brand_min\x18\x01 \x01(\x03R\arandMin\x12\x19\n" + "\brand_min\x18\x01 \x01(\x03R\arandMin\x12\x19\n" +
"\brand_max\x18\x02 \x01(\x03R\arandMax\x12$\n" + "\brand_max\x18\x02 \x01(\x03R\arandMax\x12$\n" +
@@ -185,7 +336,8 @@ const file_transport_internet_finalmask_noise_config_proto_rawDesc = "" +
"\x0erand_range_max\x18\x04 \x01(\x05R\frandRangeMax\x12\x16\n" + "\x0erand_range_max\x18\x04 \x01(\x05R\frandRangeMax\x12\x16\n" +
"\x06packet\x18\x05 \x01(\fR\x06packet\x12\x1b\n" + "\x06packet\x18\x05 \x01(\fR\x06packet\x12\x1b\n" +
"\tdelay_min\x18\x06 \x01(\x03R\bdelayMin\x12\x1b\n" + "\tdelay_min\x18\x06 \x01(\x03R\bdelayMin\x12\x1b\n" +
"\tdelay_max\x18\a \x01(\x03R\bdelayMax\"\x87\x01\n" + "\tdelay_max\x18\a \x01(\x03R\bdelayMax\x12L\n" +
"\bsegments\x18\b \x03(\v20.xray.transport.internet.finalmask.noise.SegmentR\bsegments\"\x87\x01\n" +
"\x06Config\x12\x1b\n" + "\x06Config\x12\x1b\n" +
"\treset_min\x18\x01 \x01(\x03R\bresetMin\x12\x1b\n" + "\treset_min\x18\x01 \x01(\x03R\bresetMin\x12\x1b\n" +
"\treset_max\x18\x02 \x01(\x03R\bresetMax\x12C\n" + "\treset_max\x18\x02 \x01(\x03R\bresetMax\x12C\n" +
@@ -204,18 +356,23 @@ func file_transport_internet_finalmask_noise_config_proto_rawDescGZIP() []byte {
return file_transport_internet_finalmask_noise_config_proto_rawDescData return file_transport_internet_finalmask_noise_config_proto_rawDescData
} }
var file_transport_internet_finalmask_noise_config_proto_msgTypes = make([]protoimpl.MessageInfo, 2) var file_transport_internet_finalmask_noise_config_proto_enumTypes = make([]protoimpl.EnumInfo, 1)
var file_transport_internet_finalmask_noise_config_proto_msgTypes = make([]protoimpl.MessageInfo, 3)
var file_transport_internet_finalmask_noise_config_proto_goTypes = []any{ var file_transport_internet_finalmask_noise_config_proto_goTypes = []any{
(*Item)(nil), // 0: xray.transport.internet.finalmask.noise.Item (Segment_Kind)(0), // 0: xray.transport.internet.finalmask.noise.Segment.Kind
(*Config)(nil), // 1: xray.transport.internet.finalmask.noise.Config (*Segment)(nil), // 1: xray.transport.internet.finalmask.noise.Segment
(*Item)(nil), // 2: xray.transport.internet.finalmask.noise.Item
(*Config)(nil), // 3: xray.transport.internet.finalmask.noise.Config
} }
var file_transport_internet_finalmask_noise_config_proto_depIdxs = []int32{ var file_transport_internet_finalmask_noise_config_proto_depIdxs = []int32{
0, // 0: xray.transport.internet.finalmask.noise.Config.items:type_name -> xray.transport.internet.finalmask.noise.Item 0, // 0: xray.transport.internet.finalmask.noise.Segment.kind:type_name -> xray.transport.internet.finalmask.noise.Segment.Kind
1, // [1:1] is the sub-list for method output_type 1, // 1: xray.transport.internet.finalmask.noise.Item.segments:type_name -> xray.transport.internet.finalmask.noise.Segment
1, // [1:1] is the sub-list for method input_type 2, // 2: xray.transport.internet.finalmask.noise.Config.items:type_name -> xray.transport.internet.finalmask.noise.Item
1, // [1:1] is the sub-list for extension type_name 3, // [3:3] is the sub-list for method output_type
1, // [1:1] is the sub-list for extension extendee 3, // [3:3] is the sub-list for method input_type
0, // [0:1] is the sub-list for field type_name 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_transport_internet_finalmask_noise_config_proto_init() } func init() { file_transport_internet_finalmask_noise_config_proto_init() }
@@ -228,13 +385,14 @@ func file_transport_internet_finalmask_noise_config_proto_init() {
File: protoimpl.DescBuilder{ File: protoimpl.DescBuilder{
GoPackagePath: reflect.TypeOf(x{}).PkgPath(), GoPackagePath: reflect.TypeOf(x{}).PkgPath(),
RawDescriptor: unsafe.Slice(unsafe.StringData(file_transport_internet_finalmask_noise_config_proto_rawDesc), len(file_transport_internet_finalmask_noise_config_proto_rawDesc)), RawDescriptor: unsafe.Slice(unsafe.StringData(file_transport_internet_finalmask_noise_config_proto_rawDesc), len(file_transport_internet_finalmask_noise_config_proto_rawDesc)),
NumEnums: 0, NumEnums: 1,
NumMessages: 2, NumMessages: 3,
NumExtensions: 0, NumExtensions: 0,
NumServices: 0, NumServices: 0,
}, },
GoTypes: file_transport_internet_finalmask_noise_config_proto_goTypes, GoTypes: file_transport_internet_finalmask_noise_config_proto_goTypes,
DependencyIndexes: file_transport_internet_finalmask_noise_config_proto_depIdxs, DependencyIndexes: file_transport_internet_finalmask_noise_config_proto_depIdxs,
EnumInfos: file_transport_internet_finalmask_noise_config_proto_enumTypes,
MessageInfos: file_transport_internet_finalmask_noise_config_proto_msgTypes, MessageInfos: file_transport_internet_finalmask_noise_config_proto_msgTypes,
}.Build() }.Build()
File_transport_internet_finalmask_noise_config_proto = out.File File_transport_internet_finalmask_noise_config_proto = out.File
@@ -6,6 +6,22 @@ option go_package = "github.com/xtls/xray-core/transport/internet/finalmask/nois
option java_package = "com.xray.transport.internet.finalmask.noise"; option java_package = "com.xray.transport.internet.finalmask.noise";
option java_multiple_files = true; option java_multiple_files = true;
message Segment {
enum Kind {
BYTES = 0;
RANDOM = 1;
RANDOM_ASCII = 2;
RANDOM_DIGIT = 3;
TIMESTAMP = 4;
COUNTER = 5;
NONCE = 6;
}
Kind kind = 1;
bytes bytes = 2;
int64 min_size = 3;
int64 max_size = 4;
}
message Item { message Item {
int64 rand_min = 1; int64 rand_min = 1;
int64 rand_max = 2; int64 rand_max = 2;
@@ -14,6 +30,7 @@ message Item {
bytes packet = 5; bytes packet = 5;
int64 delay_min = 6; int64 delay_min = 6;
int64 delay_max = 7; int64 delay_max = 7;
repeated Segment segments = 8;
} }
message Config { message Config {
+67 -10
View File
@@ -1,18 +1,25 @@
package noise package noise
import ( import (
"crypto/rand"
"encoding/binary"
"net" "net"
"sync" "sync"
"sync/atomic"
"time" "time"
"github.com/xtls/xray-core/common"
"github.com/xtls/xray-core/common/crypto" "github.com/xtls/xray-core/common/crypto"
) )
const asciiLetters = "abcdefghijklmnopqrstuvwxyzABCDEFGHIJKLMNOPQRSTUVWXYZ"
type noiseConn struct { type noiseConn struct {
net.PacketConn net.PacketConn
config *Config config *Config
m map[string]time.Time m map[string]time.Time
mu sync.Mutex mu sync.Mutex
counter atomic.Uint32
} }
func NewConnClient(c *Config, raw net.PacketConn) (net.PacketConn, error) { func NewConnClient(c *Config, raw net.PacketConn) (net.PacketConn, error) {
@@ -27,6 +34,62 @@ func NewConnServer(c *Config, raw net.PacketConn) (net.PacketConn, error) {
return NewConnClient(c, raw) return NewConnClient(c, raw)
} }
func (c *noiseConn) buildPacket(item *Item) []byte {
if len(item.Segments) == 0 {
if item.RandMax > 0 {
buf := make([]byte, crypto.RandBetween(item.RandMin, item.RandMax))
crypto.RandBytesBetween(buf, byte(item.RandRangeMin), byte(item.RandRangeMax))
return buf
}
return item.Packet
}
var out []byte
for _, seg := range item.Segments {
out = append(out, c.buildSegment(seg)...)
}
return out
}
func (c *noiseConn) buildSegment(seg *Segment) []byte {
switch seg.Kind {
case Segment_BYTES:
return seg.Bytes
case Segment_TIMESTAMP:
b := make([]byte, 4)
binary.BigEndian.PutUint32(b, uint32(time.Now().Unix()))
return b
case Segment_COUNTER:
b := make([]byte, 4)
binary.BigEndian.PutUint32(b, c.counter.Add(1))
return b
case Segment_NONCE:
b := make([]byte, 8)
common.Must2(rand.Read(b))
return b
default:
size := crypto.RandBetween(seg.MinSize, seg.MaxSize+1)
if size <= 0 {
return nil
}
buf := make([]byte, size)
switch seg.Kind {
case Segment_RANDOM_ASCII:
common.Must2(rand.Read(buf))
for i := range buf {
buf[i] = asciiLetters[int(buf[i])%len(asciiLetters)]
}
case Segment_RANDOM_DIGIT:
common.Must2(rand.Read(buf))
for i := range buf {
buf[i] = '0' + buf[i]%10
}
default:
common.Must2(rand.Read(buf))
}
return buf
}
}
func (c *noiseConn) WriteTo(p []byte, addr net.Addr) (n int, err error) { func (c *noiseConn) WriteTo(p []byte, addr net.Addr) (n int, err error) {
c.mu.Lock() c.mu.Lock()
defer c.mu.Unlock() defer c.mu.Unlock()
@@ -35,13 +98,7 @@ func (c *noiseConn) WriteTo(p []byte, addr net.Addr) (n int, err error) {
if t.IsZero() || (c.config.ResetMax > 0 && time.Now().After(t)) { if t.IsZero() || (c.config.ResetMax > 0 && time.Now().After(t)) {
for _, item := range c.config.Items { for _, item := range c.config.Items {
if item.RandMax > 0 { c.PacketConn.WriteTo(c.buildPacket(item), addr)
buf := make([]byte, crypto.RandBetween(item.RandMin, item.RandMax))
crypto.RandBytesBetween(buf, byte(item.RandRangeMin), byte(item.RandRangeMax))
c.PacketConn.WriteTo(buf, addr)
} else {
c.PacketConn.WriteTo(item.Packet, addr)
}
time.Sleep(time.Duration(crypto.RandBetween(item.DelayMin, item.DelayMax)) * time.Millisecond) time.Sleep(time.Duration(crypto.RandBetween(item.DelayMin, item.DelayMax)) * time.Millisecond)
} }
} }
@@ -0,0 +1,137 @@
package noise
import (
"bytes"
"encoding/binary"
"net"
"sync"
"testing"
"time"
"github.com/stretchr/testify/require"
)
type fakePacketConn struct {
mu sync.Mutex
written [][]byte
}
func (c *fakePacketConn) WriteTo(p []byte, _ net.Addr) (int, error) {
c.mu.Lock()
defer c.mu.Unlock()
c.written = append(c.written, bytes.Clone(p))
return len(p), nil
}
func (c *fakePacketConn) packets() [][]byte {
c.mu.Lock()
defer c.mu.Unlock()
return c.written
}
func (c *fakePacketConn) ReadFrom(_ []byte) (int, net.Addr, error) { return 0, nil, nil }
func (c *fakePacketConn) Close() error { return nil }
func (c *fakePacketConn) LocalAddr() net.Addr { return &net.UDPAddr{} }
func (c *fakePacketConn) SetDeadline(time.Time) error { return nil }
func (c *fakePacketConn) SetReadDeadline(time.Time) error { return nil }
func (c *fakePacketConn) SetWriteDeadline(time.Time) error { return nil }
func newConn() *noiseConn {
return &noiseConn{PacketConn: &fakePacketConn{}, config: &Config{}, m: make(map[string]time.Time)}
}
func TestBuildSegmentBytes(t *testing.T) {
c := newConn()
got := c.buildSegment(&Segment{Kind: Segment_BYTES, Bytes: []byte{0x0d, 0x0a, 0x0d, 0x0a}})
require.Equal(t, []byte{0x0d, 0x0a, 0x0d, 0x0a}, got)
}
func TestBuildSegmentTimestamp(t *testing.T) {
c := newConn()
before := time.Now().Unix()
got := c.buildSegment(&Segment{Kind: Segment_TIMESTAMP})
require.Len(t, got, 4)
ts := int64(binary.BigEndian.Uint32(got))
require.GreaterOrEqual(t, ts, before)
require.LessOrEqual(t, ts, time.Now().Unix())
}
func TestBuildSegmentCounter(t *testing.T) {
c := newConn()
first := binary.BigEndian.Uint32(c.buildSegment(&Segment{Kind: Segment_COUNTER}))
second := binary.BigEndian.Uint32(c.buildSegment(&Segment{Kind: Segment_COUNTER}))
require.Equal(t, uint32(1), first)
require.Equal(t, uint32(2), second)
}
func TestBuildSegmentNonce(t *testing.T) {
c := newConn()
a := c.buildSegment(&Segment{Kind: Segment_NONCE})
b := c.buildSegment(&Segment{Kind: Segment_NONCE})
require.Len(t, a, 8)
require.Len(t, b, 8)
require.NotEqual(t, a, b)
}
func TestBuildSegmentRandomSizes(t *testing.T) {
c := newConn()
for range 200 {
require.Len(t, c.buildSegment(&Segment{Kind: Segment_RANDOM, MinSize: 24, MaxSize: 24}), 24)
n := len(c.buildSegment(&Segment{Kind: Segment_RANDOM, MinSize: 20, MaxSize: 32}))
require.GreaterOrEqual(t, n, 20)
require.LessOrEqual(t, n, 32)
for _, b := range c.buildSegment(&Segment{Kind: Segment_RANDOM_ASCII, MinSize: 40, MaxSize: 40}) {
require.True(t, (b >= 'a' && b <= 'z') || (b >= 'A' && b <= 'Z'), "not a letter: %q", b)
}
for _, b := range c.buildSegment(&Segment{Kind: Segment_RANDOM_DIGIT, MinSize: 40, MaxSize: 40}) {
require.True(t, b >= '0' && b <= '9', "not a digit: %q", b)
}
}
}
func TestBuildPacketComposite(t *testing.T) {
c := newConn()
item := &Item{Segments: []*Segment{
{Kind: Segment_BYTES, Bytes: []byte{0x0d, 0x0a, 0x0d, 0x0a}},
{Kind: Segment_TIMESTAMP},
{Kind: Segment_RANDOM, MinSize: 24, MaxSize: 24},
}}
got := c.buildPacket(item)
require.Len(t, got, 4+4+24)
require.Equal(t, []byte{0x0d, 0x0a, 0x0d, 0x0a}, got[:4])
}
func TestBuildPacketLegacy(t *testing.T) {
c := newConn()
require.Equal(t, []byte{1, 2, 3}, c.buildPacket(&Item{Packet: []byte{1, 2, 3}}))
require.Len(t, c.buildPacket(&Item{RandMin: 16, RandMax: 17}), 16)
}
func TestWriteToSendsNoiseThenPayload(t *testing.T) {
raw := &fakePacketConn{}
c := &noiseConn{
PacketConn: raw,
m: make(map[string]time.Time),
config: &Config{Items: []*Item{
{Segments: []*Segment{{Kind: Segment_BYTES, Bytes: []byte{0x0d, 0x0a, 0x0d, 0x0a}}, {Kind: Segment_RANDOM, MinSize: 8, MaxSize: 8}}},
{RandMin: 40, RandMax: 41},
}},
}
addr := &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1), Port: 51820}
payload := []byte("real-handshake")
_, err := c.WriteTo(payload, addr)
require.NoError(t, err)
sent := raw.packets()
require.Len(t, sent, 3)
require.Len(t, sent[0], 12)
require.Equal(t, []byte{0x0d, 0x0a, 0x0d, 0x0a}, sent[0][:4])
require.Len(t, sent[1], 40)
require.Equal(t, payload, sent[2])
_, err = c.WriteTo(payload, addr)
require.NoError(t, err)
require.Len(t, raw.packets(), 4)
}
@@ -223,13 +223,6 @@ func (c *udpHopConn) Close() error {
} }
_ = c.cur.Close() _ = c.cur.Close()
c.wg.Wait() c.wg.Wait()
select {
case packet := <-c.readCh:
if packet.p != nil {
pool.Put(packet.p[:cap(packet.p)])
}
default:
}
close(c.readCh) close(c.readCh)
return nil return nil
} }
+345 -321
View File
@@ -1,417 +1,441 @@
package xdns package xdns
import ( import (
"bytes"
"context" "context"
"crypto/rand" "crypto/rand"
"encoding/base32"
"encoding/binary"
go_errors "errors"
"io" "io"
"net" mrand "math/rand"
"strconv"
"sync" "sync"
"sync/atomic" "sync/atomic"
"time" "time"
"github.com/xtls/xray-core/common" "github.com/xtls/xray-core/common"
"github.com/xtls/xray-core/common/errors" "github.com/xtls/xray-core/common/errors"
"github.com/xtls/xray-core/common/net"
"github.com/xtls/xray-core/transport/internet/finalmask" "github.com/xtls/xray-core/transport/internet/finalmask"
"golang.org/x/net/dns/dnsmessage"
) )
const ( const (
numPadding = 3
numPaddingForPoll = 8
initPollDelay = 500 * time.Millisecond initPollDelay = 500 * time.Millisecond
maxPollDelay = 10 * time.Second maxPollDelay = 10 * time.Second
pollDelayMultiplier = 2.0 pollDelayMultiplier = 2.0
pollLimit = 16 pollLimit = 16
) )
var base32Encoding = base32.StdEncoding.WithPadding(base32.NoPadding) var pool4K = sync.Pool{
New: func() any {
return make([]byte, 4096)
},
}
type packet struct { type packet struct {
p []byte p []byte
addr net.Addr addr net.Addr
} }
type xdnsConnClient struct { type xdnsClient struct {
net.PacketConn dialer *finalmask.Dialer
resolverAddrs []*net.UDPAddr clientID ClientID
resolverTypes []uint16 fragID atomic.Uint32
resolverIdx uint32 domains []*Domain
resolverSend map[string]*atomic.Uint32 extraPoll int32
clientID []byte resolvers []Resolver
domains []Name resolverSends []atomic.Uint32
resolverIndex atomic.Uint32
pollChan chan struct{} readCh chan packet
readQueue chan *packet sendCh chan []byte
writeQueue chan *packet poolCh chan struct{}
closeCh chan struct{}
closed bool wg sync.WaitGroup
mutex sync.Mutex mu sync.Mutex
} }
func NewConnClient(c *Config, raw net.PacketConn) (net.PacketConn, error) { func NewClient(c *Config, dialer *finalmask.Dialer) (net.PacketConn, error) {
if len(c.Domains) == 0 {
return nil, errors.New("empty domains")
}
if len(c.Resolvers) == 0 { if len(c.Resolvers) == 0 {
return nil, errors.New("empty resolvers") return nil, errors.New("empty resolvers")
} }
if c.ExtraPoll < 0 || c.ExtraPoll > 3 {
var domains []Name return nil, errors.New("c.ExtraPoll < 0 || c.ExtraPoll > 3")
var servers []string
var resolverTypes []uint16
for _, rs := range c.Resolvers {
domain, server, resolverType, err := parseResolver(rs)
if err != nil {
return nil, errors.New("invalid resolvers").Base(err)
}
domains = append(domains, domain)
servers = append(servers, server)
resolverTypes = append(resolverTypes, resolverType)
} }
domains := make([]*Domain, 0, len(c.Domains))
var resolverAddrs []*net.UDPAddr for i := range c.Domains {
resolverSend := make(map[string]*atomic.Uint32) types := make([]uint16, 0, len(c.Domains[i].Types))
for _, rs := range servers { for j := range c.Domains[i].Types {
h, p, err := net.SplitHostPort(rs) types = append(types, uint16(c.Domains[i].Types[j]))
}
domain, err := NewDomain(c.Domains[i].Name, int(c.Domains[i].LenLimit), int(c.Domains[i].LabelLimit), types, uint16(c.Domains[i].Edns0))
if err != nil { if err != nil {
return nil, err return nil, err
} }
ip := net.ParseIP(h) domains = append(domains, domain)
if ip == nil { }
return nil, errors.New("invalid ip address") resolvers := make([]Resolver, 0, len(c.Resolvers))
} for i := range c.Resolvers {
port, err := strconv.Atoi(p) resolver, err := NewResolver(c.Resolvers[i], dialer)
if err != nil { if err != nil {
return nil, errors.New("invalid port").Base(err) return nil, err
} }
addr := &net.UDPAddr{IP: ip, Port: port} resolvers = append(resolvers, resolver)
resolverAddrs = append(resolverAddrs, addr)
resolverSend[addr.String()] = &atomic.Uint32{}
} }
client := &xdnsClient{
dialer: dialer,
conn := &xdnsConnClient{ clientID: NewClientID(),
PacketConn: raw, domains: domains,
extraPoll: c.ExtraPoll,
resolverAddrs: resolverAddrs, resolvers: resolvers,
resolverTypes: resolverTypes, resolverSends: make([]atomic.Uint32, len(c.Resolvers)),
resolverIdx: 0,
resolverSend: resolverSend,
clientID: make([]byte, 8), readCh: make(chan packet),
domains: domains, sendCh: make(chan []byte, 16),
poolCh: make(chan struct{}, pollLimit),
pollChan: make(chan struct{}, pollLimit), closeCh: make(chan struct{}),
readQueue: make(chan *packet, 256),
writeQueue: make(chan *packet, 256),
} }
go client.run()
common.Must2(rand.Read(conn.clientID)) return client, nil
go conn.recvLoop()
go conn.sendLoop()
return conn, nil
} }
func (c *xdnsConnClient) recvLoop() { func (c *xdnsClient) closed() bool {
var buf [finalmask.UDPSize]byte select {
case <-c.closeCh:
return true
default:
return false
}
}
for { func (c *xdnsClient) read(buf []byte, addr net.Addr) bool {
if c.closed { msg := dnsmessage.Message{}
if err := msg.Unpack(buf); err != nil {
return false
}
if !msg.Header.Response || msg.Header.Truncated || msg.Header.RCode != dnsmessage.RCodeSuccess || len(msg.Questions) != 1 {
return false
}
var domain *Domain
for i := range c.domains {
if c.domains[i].IsDomain(msg.Questions[0].Name) {
domain = c.domains[i]
break break
} }
}
if domain == nil || !domain.HasType(uint16(msg.Questions[0].Type)) {
return false
}
n, addr, err := c.PacketConn.ReadFrom(buf[:]) edns0 := uint16(0)
for i := range msg.Additionals {
if msg.Additionals[i].Header.Type == dnsmessage.TypeOPT {
edns0 = uint16(msg.Additionals[i].Header.Class)
break
}
}
errors.LogDebug(context.Background(), addr, " edns0 ", edns0, " buf ", len(buf), " ", msg.Questions[0].Type)
resp := NewResp(msg, domain, 0)
p := pool4K.Get().([]byte)
n := resp.Decode(p)
p = p[:n]
b := p
var bs [][]byte
for len(b) > 1 {
last := b[0]&0xC0 == 0xC0
length := int(b[0]&0x3F)<<8 | int(b[1])
b = b[2:]
if length > len(b) {
bs = nil
break
}
packet := make([]byte, length)
copy(packet, b)
bs = append(bs, packet)
if last {
break
}
b = b[length:]
if len(b) < 2 {
bs = nil
}
}
pool4K.Put(p[:cap(p)])
for i := range bs {
select {
case <-c.closeCh:
return true
case c.readCh <- packet{p: bs[i], addr: addr}:
}
}
return len(bs) > 0
}
func (c *xdnsClient) run() {
for i := range len(c.resolvers) {
c.wg.Add(1)
go c.recv(i)
}
c.wg.Add(1)
go c.send()
c.wg.Wait()
close(c.readCh)
close(c.sendCh)
close(c.poolCh)
}
func (c *xdnsClient) recv(i int) {
defer c.wg.Done()
var buf [4096]byte
for {
n, err := c.resolvers[i].Read(buf[:])
if err != nil { if err != nil {
if go_errors.Is(err, net.ErrClosed) { if c.closed() {
break return
} }
continue errors.LogErrorInner(context.Background(), err, "recv err ", i)
return
} }
if c.read(buf[:n], c.resolvers[i].Addr()) {
if addr == nil { c.resolverSends[i].Store(0)
continue
}
send := c.resolverSend[addr.String()]
if send == nil {
continue
}
resp, err := MessageFromWireFormat(buf[:n])
if err != nil {
errors.LogDebug(context.Background(), addr, " xdns from wireformat err ", err)
continue
}
payload := dnsResponsePayload(&resp, c.domains)
r := bytes.NewReader(payload)
anyPacket := false
for {
p, err := nextPacket(r)
if err != nil {
break
}
anyPacket = true
buf := make([]byte, len(p))
copy(buf, p)
select { select {
case c.readQueue <- &packet{ case c.poolCh <- struct{}{}:
p: buf,
addr: addr,
}:
default:
errors.LogDebug(context.Background(), addr, " mask read err queue full")
}
}
if anyPacket {
send.Store(0)
select {
case c.pollChan <- struct{}{}:
default: default:
} }
} }
} }
errors.LogDebug(context.Background(), "xdns closed")
close(c.pollChan)
close(c.readQueue)
c.mutex.Lock()
defer c.mutex.Unlock()
c.closed = true
close(c.writeQueue)
} }
func (c *xdnsConnClient) sendLoop() { func (c *xdnsClient) send() {
pollDelay := initPollDelay defer c.wg.Done()
pollTimer := time.NewTimer(pollDelay)
for {
var p *packet
pollTimerExpired := false
select { var buf [512]byte
case p = <-c.writeQueue: var data [255]byte
default:
select { sendMsg := func(p []byte, domain *Domain, qtype uint16) {
case p = <-c.writeQueue: msg := dnsmessage.Message{
case <-c.pollChan: Header: dnsmessage.Header{
case <-pollTimer.C: RecursionDesired: true,
pollTimerExpired = true },
Questions: []dnsmessage.Question{
{
Name: domain.Encode(p),
Type: dnsmessage.Type(qtype),
Class: dnsmessage.ClassINET,
},
},
}
if domain.edns0 > 0 {
msg.Additionals = []dnsmessage.Resource{
{
Header: dnsmessage.ResourceHeader{
Name: dnsmessage.MustNewName("."),
Type: dnsmessage.TypeOPT,
Class: dnsmessage.Class(domain.edns0),
TTL: 0,
},
Body: &dnsmessage.OPTResource{},
},
} }
} }
pack := common.Must2(msg.AppendPack(buf[:0]))
common.Must2(rand.Read(pack[:2]))
if p != nil { index := c.resolverIndex.Load()
select { cur := c.resolverSends[index].Add(1)
case <-c.pollChan: i := index
default: for {
i++
if i == uint32(len(c.resolvers)) {
i = 0
} }
} else { if i == index {
encoded, _ := encode(nil, c.clientID, c.domains[c.resolverIdx], c.resolverTypes[c.resolverIdx]) break
p = &packet{ }
p: encoded, if cur > c.resolverSends[i].Load() {
break
} }
} }
c.resolverIndex.Store(i)
c.resolvers[index].Send(pack)
}
if pollTimerExpired { send := func(p []byte) {
pollDelay = time.Duration(float64(pollDelay) * pollDelayMultiplier) domain := c.domains[mrand.Intn(len(c.domains))]
if pollDelay > maxPollDelay { qtype := domain.types[mrand.Intn(len(domain.types))]
pollDelay = maxPollDelay
}
} else {
if !pollTimer.Stop() {
<-pollTimer.C
}
pollDelay = initPollDelay
}
pollTimer.Reset(pollDelay)
if c.closed { if len(p) == 0 {
copy(data[:], c.clientID[:])
data[0] |= TypeMap[qtype]
data[8] = 8
common.Must2(rand.Read(data[9:17]))
sendMsg(data[:17], domain, qtype)
return return
} }
cur := c.resolverIdx if len(p) <= domain.cap-12 {
curSend := c.resolverSend[c.resolverAddrs[cur].String()].Add(1) copy(data[:], c.clientID[:])
_, _ = c.PacketConn.WriteTo(p.p, c.resolverAddrs[cur]) data[0] |= TypeMap[qtype]
for { data[8] = 3
c.resolverIdx += 1 common.Must2(rand.Read(data[9:12]))
c.resolverIdx %= uint32(len(c.resolverAddrs)) copy(data[12:], p)
if c.resolverIdx == cur { sendMsg(data[:12+len(p)], domain, qtype)
break return
}
if len(p) <= 255*(domain.cap-15) {
copy(data[:], c.clientID[:])
data[0] |= TypeMap[qtype]
data[8] = 3 | 0xC0
common.Must2(rand.Read(data[9:12]))
fragID := byte(c.fragID.Add(1))
fragN := len(p) / (domain.cap - 15)
if len(p)%(domain.cap-15) > 0 {
fragN++
} }
if c.resolverSend[c.resolverAddrs[c.resolverIdx].String()].Load() < curSend {
break for i := range fragN {
data[12] = fragID
data[13] = byte(i)
data[14] = byte(fragN)
size := min(len(p), domain.cap-15)
copy(data[15:], p[:size])
sendMsg(data[:15+size], domain, qtype)
p = p[size:]
}
return
}
errors.LogError(context.Background(), "err size ", len(p))
}
ticker := time.NewTicker(initPollDelay)
defer ticker.Stop()
delay := initPollDelay
p := []byte(nil)
timeout := false
for {
select {
case <-c.closeCh:
return
default:
select {
case <-c.closeCh:
return
case p = <-c.sendCh:
case <-c.poolCh:
case <-ticker.C:
timeout = true
} }
} }
if len(p) > 0 {
select {
case <-c.poolCh:
default:
}
}
send(p)
for range c.extraPoll {
send(nil)
}
if timeout {
delay *= pollDelayMultiplier
if delay > maxPollDelay {
delay = maxPollDelay
}
timeout = false
} else {
delay = initPollDelay
}
ticker.Reset(delay)
} }
} }
func (c *xdnsConnClient) ReadFrom(p []byte) (n int, addr net.Addr, err error) { func (c *xdnsClient) ReadFrom(p []byte) (n int, addr net.Addr, err error) {
packet, ok := <-c.readQueue packet, ok := <-c.readCh
if !ok { if ok {
return 0, nil, net.ErrClosed return copy(p, packet.p), packet.addr, nil
} }
if len(p) < len(packet.p) { return 0, nil, io.ErrClosedPipe
errors.LogDebug(context.Background(), packet.addr, " mask read err short buffer ", len(p), " ", len(packet.p))
return 0, packet.addr, nil
}
copy(p, packet.p)
return len(packet.p), packet.addr, nil
} }
func (c *xdnsConnClient) WriteTo(p []byte, addr net.Addr) (n int, err error) { func (c *xdnsClient) WriteTo(p []byte, addr net.Addr) (n int, err error) {
c.mutex.Lock() c.mu.Lock()
defer c.mutex.Unlock() defer c.mu.Unlock()
if c.closed() {
if c.closed {
return 0, io.ErrClosedPipe return 0, io.ErrClosedPipe
} }
if len(p) == 0 || len(p) > 4096 {
idx := c.resolverIdx % uint32(len(c.resolverAddrs)) errors.LogError(context.Background(), "err size ", len(p))
encoded, err := encode(p, c.clientID, c.domains[idx], c.resolverTypes[idx]) return 0, errors.New("err size")
if err != nil {
errors.LogDebug(context.Background(), addr, " xdns wireformat err ", err, " ", len(p))
return 0, nil
} }
b := make([]byte, len(p))
copy(b, p)
select { select {
case c.writeQueue <- &packet{ case c.sendCh <- b:
p: encoded,
addr: addr,
}:
return len(p), nil
default: default:
errors.LogDebug(context.Background(), addr, " mask write err queue full")
return 0, nil
} }
return len(p), nil
} }
func (c *xdnsConnClient) Close() error { func (c *xdnsClient) Close() error {
c.closed = true c.mu.Lock()
return c.PacketConn.Close() defer c.mu.Unlock()
} if c.closed() {
func encode(p []byte, clientID []byte, domain Name, qtype uint16) ([]byte, error) {
var decoded []byte
{
if len(p) >= 224 {
return nil, errors.New("too long")
}
var buf bytes.Buffer
buf.Write(clientID[:])
n := numPadding
if len(p) == 0 {
n = numPaddingForPoll
}
buf.WriteByte(byte(224 + n))
_, _ = io.CopyN(&buf, rand.Reader, int64(n))
if len(p) > 0 {
buf.WriteByte(byte(len(p)))
buf.Write(p)
}
decoded = buf.Bytes()
}
encoded := make([]byte, base32Encoding.EncodedLen(len(decoded)))
base32Encoding.Encode(encoded, decoded)
encoded = bytes.ToLower(encoded)
labels := chunks(encoded, 63)
labels = append(labels, domain...)
name, err := NewName(labels)
if err != nil {
return nil, err
}
var id uint16
_ = binary.Read(rand.Reader, binary.BigEndian, &id)
query := &Message{
ID: id,
Flags: 0x0100,
Question: []Question{
{
Name: name,
Type: qtype,
Class: ClassIN,
},
},
Additional: []RR{
{
Name: Name{},
Type: RRTypeOPT,
Class: 4096,
TTL: 0,
Data: []byte{},
},
},
}
buf, err := query.WireFormat()
if err != nil {
return nil, err
}
return buf, nil
}
func chunks(p []byte, n int) [][]byte {
var result [][]byte
for len(p) > 0 {
sz := len(p)
if sz > n {
sz = n
}
result = append(result, p[:sz])
p = p[sz:]
}
return result
}
func nextPacket(r *bytes.Reader) ([]byte, error) {
var n uint16
err := binary.Read(r, binary.BigEndian, &n)
if err != nil {
return nil, err
}
p := make([]byte, n)
_, err = io.ReadFull(r, p)
if err == io.EOF {
err = io.ErrUnexpectedEOF
}
return p, err
}
func dnsResponsePayload(resp *Message, domains []Name) []byte {
if resp.Flags&0x8000 != 0x8000 {
return nil return nil
} }
if resp.Flags&0x000f != RcodeNoError { close(c.closeCh)
return nil for i := range c.resolvers {
c.resolvers[i].Close()
} }
return nil
if len(resp.Answer) == 0 { }
return nil
} func (c *xdnsClient) LocalAddr() net.Addr { return &net.UDPAddr{IP: []byte{0, 0, 0, 0}} }
for _, answer := range resp.Answer { func (c *xdnsClient) SetDeadline(t time.Time) error { return errors.New("not support") }
var ok bool
for _, domain := range domains { func (c *xdnsClient) SetReadDeadline(t time.Time) error { return errors.New("not support") }
_, ok = answer.Name.TrimSuffix(domain)
if ok { func (c *xdnsClient) SetWriteDeadline(t time.Time) error { return errors.New("not support") }
break
} type ClientID [8]byte
}
if !ok { func NewClientID() ClientID {
return nil var id ClientID
} common.Must2(rand.Read(id[:]))
} id[0] &= 0xFC
return id
return decodeResponsePayload(resp.Answer) }
func ClientIDFromRaw(id [8]byte) ClientID {
id[0] &= 0xFC
return id
}
func ClientIDFromAddr(addr *net.UDPAddr) ClientID {
return ClientID(addr.IP[8:])
}
func (id ClientID) Addr() *net.UDPAddr {
var ip [16]byte
ip[0] = 0xFD
copy(ip[8:], id[:])
return &net.UDPAddr{IP: ip[:]}
} }
+2 -2
View File
@@ -6,9 +6,9 @@ import (
) )
func (c *Config) WrapPacketConnClient(conn net.PacketConn, dest *net.Destination, dialer *finalmask.Dialer) (net.PacketConn, error) { func (c *Config) WrapPacketConnClient(conn net.PacketConn, dest *net.Destination, dialer *finalmask.Dialer) (net.PacketConn, error) {
return NewConnClient(c, conn) return NewClient(c, dialer)
} }
func (c *Config) WrapPacketConnServer(conn net.PacketConn, addr net.Addr, lc *finalmask.ListenConfig) (net.PacketConn, error) { func (c *Config) WrapPacketConnServer(conn net.PacketConn, addr net.Addr, lc *finalmask.ListenConfig) (net.PacketConn, error) {
return NewConnServer(c, conn) return NewServer(c, conn)
} }
+211 -19
View File
@@ -7,6 +7,7 @@
package xdns package xdns
import ( import (
serial "github.com/xtls/xray-core/common/serial"
protoreflect "google.golang.org/protobuf/reflect/protoreflect" protoreflect "google.golang.org/protobuf/reflect/protoreflect"
protoimpl "google.golang.org/protobuf/runtime/protoimpl" protoimpl "google.golang.org/protobuf/runtime/protoimpl"
reflect "reflect" reflect "reflect"
@@ -21,17 +22,94 @@ const (
_ = protoimpl.EnforceVersion(protoimpl.MaxVersion - 20) _ = protoimpl.EnforceVersion(protoimpl.MaxVersion - 20)
) )
type DomainProto struct {
state protoimpl.MessageState `protogen:"open.v1"`
Name string `protobuf:"bytes,1,opt,name=name,proto3" json:"name,omitempty"`
LenLimit int32 `protobuf:"varint,2,opt,name=len_limit,json=lenLimit,proto3" json:"len_limit,omitempty"`
LabelLimit int32 `protobuf:"varint,3,opt,name=label_limit,json=labelLimit,proto3" json:"label_limit,omitempty"`
Types []int32 `protobuf:"varint,4,rep,packed,name=types,proto3" json:"types,omitempty"`
Edns0 int32 `protobuf:"varint,5,opt,name=edns0,proto3" json:"edns0,omitempty"`
unknownFields protoimpl.UnknownFields
sizeCache protoimpl.SizeCache
}
func (x *DomainProto) Reset() {
*x = DomainProto{}
mi := &file_transport_internet_finalmask_xdns_config_proto_msgTypes[0]
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
ms.StoreMessageInfo(mi)
}
func (x *DomainProto) String() string {
return protoimpl.X.MessageStringOf(x)
}
func (*DomainProto) ProtoMessage() {}
func (x *DomainProto) ProtoReflect() protoreflect.Message {
mi := &file_transport_internet_finalmask_xdns_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 DomainProto.ProtoReflect.Descriptor instead.
func (*DomainProto) Descriptor() ([]byte, []int) {
return file_transport_internet_finalmask_xdns_config_proto_rawDescGZIP(), []int{0}
}
func (x *DomainProto) GetName() string {
if x != nil {
return x.Name
}
return ""
}
func (x *DomainProto) GetLenLimit() int32 {
if x != nil {
return x.LenLimit
}
return 0
}
func (x *DomainProto) GetLabelLimit() int32 {
if x != nil {
return x.LabelLimit
}
return 0
}
func (x *DomainProto) GetTypes() []int32 {
if x != nil {
return x.Types
}
return nil
}
func (x *DomainProto) GetEdns0() int32 {
if x != nil {
return x.Edns0
}
return 0
}
type Config struct { type Config struct {
state protoimpl.MessageState `protogen:"open.v1"` state protoimpl.MessageState `protogen:"open.v1"`
Domains []string `protobuf:"bytes,1,rep,name=domains,proto3" json:"domains,omitempty"` Domains []*DomainProto `protobuf:"bytes,1,rep,name=domains,proto3" json:"domains,omitempty"`
Resolvers []string `protobuf:"bytes,2,rep,name=resolvers,proto3" json:"resolvers,omitempty"` Resolvers []*serial.TypedMessage `protobuf:"bytes,2,rep,name=resolvers,proto3" json:"resolvers,omitempty"`
ExtraPoll int32 `protobuf:"varint,3,opt,name=extra_poll,json=extraPoll,proto3" json:"extra_poll,omitempty"`
unknownFields protoimpl.UnknownFields unknownFields protoimpl.UnknownFields
sizeCache protoimpl.SizeCache sizeCache protoimpl.SizeCache
} }
func (x *Config) Reset() { func (x *Config) Reset() {
*x = Config{} *x = Config{}
mi := &file_transport_internet_finalmask_xdns_config_proto_msgTypes[0] mi := &file_transport_internet_finalmask_xdns_config_proto_msgTypes[1]
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
ms.StoreMessageInfo(mi) ms.StoreMessageInfo(mi)
} }
@@ -43,7 +121,7 @@ func (x *Config) String() string {
func (*Config) ProtoMessage() {} func (*Config) ProtoMessage() {}
func (x *Config) ProtoReflect() protoreflect.Message { func (x *Config) ProtoReflect() protoreflect.Message {
mi := &file_transport_internet_finalmask_xdns_config_proto_msgTypes[0] mi := &file_transport_internet_finalmask_xdns_config_proto_msgTypes[1]
if x != nil { if x != nil {
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
if ms.LoadMessageInfo() == nil { if ms.LoadMessageInfo() == nil {
@@ -56,31 +134,139 @@ func (x *Config) ProtoReflect() protoreflect.Message {
// Deprecated: Use Config.ProtoReflect.Descriptor instead. // Deprecated: Use Config.ProtoReflect.Descriptor instead.
func (*Config) Descriptor() ([]byte, []int) { func (*Config) Descriptor() ([]byte, []int) {
return file_transport_internet_finalmask_xdns_config_proto_rawDescGZIP(), []int{0} return file_transport_internet_finalmask_xdns_config_proto_rawDescGZIP(), []int{1}
} }
func (x *Config) GetDomains() []string { func (x *Config) GetDomains() []*DomainProto {
if x != nil { if x != nil {
return x.Domains return x.Domains
} }
return nil return nil
} }
func (x *Config) GetResolvers() []string { func (x *Config) GetResolvers() []*serial.TypedMessage {
if x != nil { if x != nil {
return x.Resolvers return x.Resolvers
} }
return nil return nil
} }
func (x *Config) GetExtraPoll() int32 {
if x != nil {
return x.ExtraPoll
}
return 0
}
type TCPResolverProto struct {
state protoimpl.MessageState `protogen:"open.v1"`
Addr string `protobuf:"bytes,1,opt,name=addr,proto3" json:"addr,omitempty"`
unknownFields protoimpl.UnknownFields
sizeCache protoimpl.SizeCache
}
func (x *TCPResolverProto) Reset() {
*x = TCPResolverProto{}
mi := &file_transport_internet_finalmask_xdns_config_proto_msgTypes[2]
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
ms.StoreMessageInfo(mi)
}
func (x *TCPResolverProto) String() string {
return protoimpl.X.MessageStringOf(x)
}
func (*TCPResolverProto) ProtoMessage() {}
func (x *TCPResolverProto) ProtoReflect() protoreflect.Message {
mi := &file_transport_internet_finalmask_xdns_config_proto_msgTypes[2]
if x != nil {
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
if ms.LoadMessageInfo() == nil {
ms.StoreMessageInfo(mi)
}
return ms
}
return mi.MessageOf(x)
}
// Deprecated: Use TCPResolverProto.ProtoReflect.Descriptor instead.
func (*TCPResolverProto) Descriptor() ([]byte, []int) {
return file_transport_internet_finalmask_xdns_config_proto_rawDescGZIP(), []int{2}
}
func (x *TCPResolverProto) GetAddr() string {
if x != nil {
return x.Addr
}
return ""
}
type UDPResolverProto struct {
state protoimpl.MessageState `protogen:"open.v1"`
Addr string `protobuf:"bytes,1,opt,name=addr,proto3" json:"addr,omitempty"`
unknownFields protoimpl.UnknownFields
sizeCache protoimpl.SizeCache
}
func (x *UDPResolverProto) Reset() {
*x = UDPResolverProto{}
mi := &file_transport_internet_finalmask_xdns_config_proto_msgTypes[3]
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
ms.StoreMessageInfo(mi)
}
func (x *UDPResolverProto) String() string {
return protoimpl.X.MessageStringOf(x)
}
func (*UDPResolverProto) ProtoMessage() {}
func (x *UDPResolverProto) ProtoReflect() protoreflect.Message {
mi := &file_transport_internet_finalmask_xdns_config_proto_msgTypes[3]
if x != nil {
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
if ms.LoadMessageInfo() == nil {
ms.StoreMessageInfo(mi)
}
return ms
}
return mi.MessageOf(x)
}
// Deprecated: Use UDPResolverProto.ProtoReflect.Descriptor instead.
func (*UDPResolverProto) Descriptor() ([]byte, []int) {
return file_transport_internet_finalmask_xdns_config_proto_rawDescGZIP(), []int{3}
}
func (x *UDPResolverProto) GetAddr() string {
if x != nil {
return x.Addr
}
return ""
}
var File_transport_internet_finalmask_xdns_config_proto protoreflect.FileDescriptor var File_transport_internet_finalmask_xdns_config_proto protoreflect.FileDescriptor
const file_transport_internet_finalmask_xdns_config_proto_rawDesc = "" + const file_transport_internet_finalmask_xdns_config_proto_rawDesc = "" +
"\n" + "\n" +
".transport/internet/finalmask/xdns/config.proto\x12&xray.transport.internet.finalmask.xdns\"@\n" + ".transport/internet/finalmask/xdns/config.proto\x12&xray.transport.internet.finalmask.xdns\x1a!common/serial/typed_message.proto\"\x8b\x01\n" +
"\x06Config\x12\x18\n" + "\vDomainProto\x12\x12\n" +
"\adomains\x18\x01 \x03(\tR\adomains\x12\x1c\n" + "\x04name\x18\x01 \x01(\tR\x04name\x12\x1b\n" +
"\tresolvers\x18\x02 \x03(\tR\tresolversB\x94\x01\n" + "\tlen_limit\x18\x02 \x01(\x05R\blenLimit\x12\x1f\n" +
"\vlabel_limit\x18\x03 \x01(\x05R\n" +
"labelLimit\x12\x14\n" +
"\x05types\x18\x04 \x03(\x05R\x05types\x12\x14\n" +
"\x05edns0\x18\x05 \x01(\x05R\x05edns0\"\xb6\x01\n" +
"\x06Config\x12M\n" +
"\adomains\x18\x01 \x03(\v23.xray.transport.internet.finalmask.xdns.DomainProtoR\adomains\x12>\n" +
"\tresolvers\x18\x02 \x03(\v2 .xray.common.serial.TypedMessageR\tresolvers\x12\x1d\n" +
"\n" +
"extra_poll\x18\x03 \x01(\x05R\textraPoll\"&\n" +
"\x10TCPResolverProto\x12\x12\n" +
"\x04addr\x18\x01 \x01(\tR\x04addr\"&\n" +
"\x10UDPResolverProto\x12\x12\n" +
"\x04addr\x18\x01 \x01(\tR\x04addrB\x94\x01\n" +
"*com.xray.transport.internet.finalmask.xdnsP\x01Z;github.com/xtls/xray-core/transport/internet/finalmask/xdns\xaa\x02&Xray.Transport.Internet.Finalmask.Xdnsb\x06proto3" "*com.xray.transport.internet.finalmask.xdnsP\x01Z;github.com/xtls/xray-core/transport/internet/finalmask/xdns\xaa\x02&Xray.Transport.Internet.Finalmask.Xdnsb\x06proto3"
var ( var (
@@ -95,16 +281,22 @@ func file_transport_internet_finalmask_xdns_config_proto_rawDescGZIP() []byte {
return file_transport_internet_finalmask_xdns_config_proto_rawDescData return file_transport_internet_finalmask_xdns_config_proto_rawDescData
} }
var file_transport_internet_finalmask_xdns_config_proto_msgTypes = make([]protoimpl.MessageInfo, 1) var file_transport_internet_finalmask_xdns_config_proto_msgTypes = make([]protoimpl.MessageInfo, 4)
var file_transport_internet_finalmask_xdns_config_proto_goTypes = []any{ var file_transport_internet_finalmask_xdns_config_proto_goTypes = []any{
(*Config)(nil), // 0: xray.transport.internet.finalmask.xdns.Config (*DomainProto)(nil), // 0: xray.transport.internet.finalmask.xdns.DomainProto
(*Config)(nil), // 1: xray.transport.internet.finalmask.xdns.Config
(*TCPResolverProto)(nil), // 2: xray.transport.internet.finalmask.xdns.TCPResolverProto
(*UDPResolverProto)(nil), // 3: xray.transport.internet.finalmask.xdns.UDPResolverProto
(*serial.TypedMessage)(nil), // 4: xray.common.serial.TypedMessage
} }
var file_transport_internet_finalmask_xdns_config_proto_depIdxs = []int32{ var file_transport_internet_finalmask_xdns_config_proto_depIdxs = []int32{
0, // [0:0] is the sub-list for method output_type 0, // 0: xray.transport.internet.finalmask.xdns.Config.domains:type_name -> xray.transport.internet.finalmask.xdns.DomainProto
0, // [0:0] is the sub-list for method input_type 4, // 1: xray.transport.internet.finalmask.xdns.Config.resolvers:type_name -> xray.common.serial.TypedMessage
0, // [0:0] is the sub-list for extension type_name 2, // [2:2] is the sub-list for method output_type
0, // [0:0] is the sub-list for extension extendee 2, // [2:2] is the sub-list for method input_type
0, // [0:0] is the sub-list for field type_name 2, // [2:2] is the sub-list for extension type_name
2, // [2:2] is the sub-list for extension extendee
0, // [0:2] is the sub-list for field type_name
} }
func init() { file_transport_internet_finalmask_xdns_config_proto_init() } func init() { file_transport_internet_finalmask_xdns_config_proto_init() }
@@ -118,7 +310,7 @@ func file_transport_internet_finalmask_xdns_config_proto_init() {
GoPackagePath: reflect.TypeOf(x{}).PkgPath(), GoPackagePath: reflect.TypeOf(x{}).PkgPath(),
RawDescriptor: unsafe.Slice(unsafe.StringData(file_transport_internet_finalmask_xdns_config_proto_rawDesc), len(file_transport_internet_finalmask_xdns_config_proto_rawDesc)), RawDescriptor: unsafe.Slice(unsafe.StringData(file_transport_internet_finalmask_xdns_config_proto_rawDesc), len(file_transport_internet_finalmask_xdns_config_proto_rawDesc)),
NumEnums: 0, NumEnums: 0,
NumMessages: 1, NumMessages: 4,
NumExtensions: 0, NumExtensions: 0,
NumServices: 0, NumServices: 0,
}, },
+22 -3
View File
@@ -6,7 +6,26 @@ option go_package = "github.com/xtls/xray-core/transport/internet/finalmask/xdns
option java_package = "com.xray.transport.internet.finalmask.xdns"; option java_package = "com.xray.transport.internet.finalmask.xdns";
option java_multiple_files = true; option java_multiple_files = true;
message Config { import "common/serial/typed_message.proto";
repeated string domains = 1;
repeated string resolvers = 2; message DomainProto {
string name = 1;
int32 len_limit = 2;
int32 label_limit = 3;
repeated int32 types = 4;
int32 edns0 = 5;
}
message Config {
repeated DomainProto domains = 1;
repeated xray.common.serial.TypedMessage resolvers = 2;
int32 extra_poll = 3;
}
message TCPResolverProto {
string addr = 1;
}
message UDPResolverProto {
string addr = 1;
} }
-581
View File
@@ -1,581 +0,0 @@
// Package dns deals with encoding and decoding DNS wire format.
package xdns
import (
"bytes"
"encoding/binary"
"errors"
"fmt"
"io"
"strings"
)
// The maximum number of DNS name compression pointers we are willing to follow.
// Without something like this, infinite loops are possible.
const compressionPointerLimit = 10
var (
// ErrZeroLengthLabel is the error returned for names that contain a
// zero-length label, like "example..com".
ErrZeroLengthLabel = errors.New("name contains a zero-length label")
// ErrLabelTooLong is the error returned for labels that are longer than
// 63 octets.
ErrLabelTooLong = errors.New("name contains a label longer than 63 octets")
// ErrNameTooLong is the error returned for names whose encoded
// representation is longer than 255 octets.
ErrNameTooLong = errors.New("name is longer than 255 octets")
// ErrReservedLabelType is the error returned when reading a label type
// prefix whose two most significant bits are not 00 or 11.
ErrReservedLabelType = errors.New("reserved label type")
// ErrTooManyPointers is the error returned when reading a compressed
// name that has too many compression pointers.
ErrTooManyPointers = errors.New("too many compression pointers")
// ErrTrailingBytes is the error returned when bytes remain in the parse
// buffer after parsing a message.
ErrTrailingBytes = errors.New("trailing bytes after message")
// ErrIntegerOverflow is the error returned when trying to encode an
// integer greater than 65535 into a 16-bit field.
ErrIntegerOverflow = errors.New("integer overflow")
)
const (
// https://tools.ietf.org/html/rfc1035#section-3.2.2
RRTypeA = 1
// https://tools.ietf.org/html/rfc1035#section-3.2.2
RRTypeCNAME = 5
// https://tools.ietf.org/html/rfc1035#section-3.2.2
RRTypeTXT = 16
// https://tools.ietf.org/html/rfc3596#section-2.1
RRTypeAAAA = 28
// https://tools.ietf.org/html/rfc6891#section-6.1.1
RRTypeOPT = 41
// https://tools.ietf.org/html/rfc1035#section-3.2.4
ClassIN = 1
// https://tools.ietf.org/html/rfc1035#section-4.1.1
RcodeNoError = 0 // a.k.a. NOERROR
RcodeFormatError = 1 // a.k.a. FORMERR
RcodeNameError = 3 // a.k.a. NXDOMAIN
RcodeNotImplemented = 4 // a.k.a. NOTIMPL
// https://tools.ietf.org/html/rfc6891#section-9
ExtendedRcodeBadVers = 16 // a.k.a. BADVERS
)
// Name represents a domain name, a sequence of labels each of which is 63
// octets or less in length.
//
// https://tools.ietf.org/html/rfc1035#section-3.1
type Name [][]byte
// NewName returns a Name from a slice of labels, after checking the labels for
// validity. Does not include a zero-length label at the end of the slice.
func NewName(labels [][]byte) (Name, error) {
name := Name(labels)
// https://tools.ietf.org/html/rfc1035#section-2.3.4
// Various objects and parameters in the DNS have size limits.
// labels 63 octets or less
// names 255 octets or less
for _, label := range labels {
if len(label) == 0 {
return nil, ErrZeroLengthLabel
}
if len(label) > 63 {
return nil, ErrLabelTooLong
}
}
// Check the total length.
builder := newMessageBuilder()
builder.WriteName(name)
if len(builder.Bytes()) > 255 {
return nil, ErrNameTooLong
}
return name, nil
}
// ParseName returns a new Name from a string of labels separated by dots, after
// checking the name for validity. A single dot at the end of the string is
// ignored.
func ParseName(s string) (Name, error) {
b := bytes.TrimSuffix([]byte(s), []byte("."))
if len(b) == 0 {
// bytes.Split(b, ".") would return [""] in this case
return NewName([][]byte{})
} else {
return NewName(bytes.Split(b, []byte(".")))
}
}
// String returns a reversible string representation of name. Labels are
// separated by dots, and any bytes in a label that are outside the set
// [0-9A-Za-z-] are replaced with a \xXX hex escape sequence.
func (name Name) String() string {
if len(name) == 0 {
return "."
}
var buf strings.Builder
for i, label := range name {
if i > 0 {
buf.WriteByte('.')
}
for _, b := range label {
if b == '-' ||
('0' <= b && b <= '9') ||
('A' <= b && b <= 'Z') ||
('a' <= b && b <= 'z') {
buf.WriteByte(b)
} else {
fmt.Fprintf(&buf, "\\x%02x", b)
}
}
}
return buf.String()
}
// TrimSuffix returns a Name with the given suffix removed, if it was present.
// The second return value indicates whether the suffix was present. If the
// suffix was not present, the first return value is nil.
func (name Name) TrimSuffix(suffix Name) (Name, bool) {
if len(name) < len(suffix) {
return nil, false
}
split := len(name) - len(suffix)
fore, aft := name[:split], name[split:]
for i := 0; i < len(aft); i++ {
if !bytes.Equal(bytes.ToLower(aft[i]), bytes.ToLower(suffix[i])) {
return nil, false
}
}
return fore, true
}
// Message represents a DNS message.
//
// https://tools.ietf.org/html/rfc1035#section-4.1
type Message struct {
ID uint16
Flags uint16
Question []Question
Answer []RR
Authority []RR
Additional []RR
}
// Opcode extracts the OPCODE part of the Flags field.
//
// https://tools.ietf.org/html/rfc1035#section-4.1.1
func (message *Message) Opcode() uint16 {
return (message.Flags >> 11) & 0xf
}
// Rcode extracts the RCODE part of the Flags field.
//
// https://tools.ietf.org/html/rfc1035#section-4.1.1
func (message *Message) Rcode() uint16 {
return message.Flags & 0x000f
}
// Question represents an entry in the question section of a message.
//
// https://tools.ietf.org/html/rfc1035#section-4.1.2
type Question struct {
Name Name
Type uint16
Class uint16
}
// RR represents a resource record.
//
// https://tools.ietf.org/html/rfc1035#section-4.1.3
type RR struct {
Name Name
Type uint16
Class uint16
TTL uint32
Data []byte
}
// readName parses a DNS name from r. It leaves r positioned just after the
// parsed name.
func readName(r io.ReadSeeker) (Name, error) {
var labels [][]byte
// We limit the number of compression pointers we are willing to follow.
numPointers := 0
// If we followed any compression pointers, we must finally seek to just
// past the first pointer.
var seekTo int64
loop:
for {
var labelType byte
err := binary.Read(r, binary.BigEndian, &labelType)
if err != nil {
return nil, err
}
switch labelType & 0xc0 {
case 0x00:
// This is an ordinary label.
// https://tools.ietf.org/html/rfc1035#section-3.1
length := int(labelType & 0x3f)
if length == 0 {
break loop
}
label := make([]byte, length)
_, err := io.ReadFull(r, label)
if err != nil {
return nil, err
}
labels = append(labels, label)
case 0xc0:
// This is a compression pointer.
// https://tools.ietf.org/html/rfc1035#section-4.1.4
upper := labelType & 0x3f
var lower byte
err := binary.Read(r, binary.BigEndian, &lower)
if err != nil {
return nil, err
}
offset := (uint16(upper) << 8) | uint16(lower)
if numPointers == 0 {
// The first time we encounter a pointer,
// remember our position so we can seek back to
// it when done.
seekTo, err = r.Seek(0, io.SeekCurrent)
if err != nil {
return nil, err
}
}
numPointers++
if numPointers > compressionPointerLimit {
return nil, ErrTooManyPointers
}
// Follow the pointer and continue.
_, err = r.Seek(int64(offset), io.SeekStart)
if err != nil {
return nil, err
}
default:
// "The 10 and 01 combinations are reserved for future
// use."
return nil, ErrReservedLabelType
}
}
// If we followed any pointers, then seek back to just after the first
// one.
if numPointers > 0 {
_, err := r.Seek(seekTo, io.SeekStart)
if err != nil {
return nil, err
}
}
return NewName(labels)
}
// readQuestion parses one entry from the Question section. It leaves r
// positioned just after the parsed entry.
//
// https://tools.ietf.org/html/rfc1035#section-4.1.2
func readQuestion(r io.ReadSeeker) (Question, error) {
var question Question
var err error
question.Name, err = readName(r)
if err != nil {
return question, err
}
for _, ptr := range []*uint16{&question.Type, &question.Class} {
err := binary.Read(r, binary.BigEndian, ptr)
if err != nil {
return question, err
}
}
return question, nil
}
// readRR parses one resource record. It leaves r positioned just after the
// parsed resource record.
//
// https://tools.ietf.org/html/rfc1035#section-4.1.3
func readRR(r io.ReadSeeker) (RR, error) {
var rr RR
var err error
rr.Name, err = readName(r)
if err != nil {
return rr, err
}
for _, ptr := range []*uint16{&rr.Type, &rr.Class} {
err := binary.Read(r, binary.BigEndian, ptr)
if err != nil {
return rr, err
}
}
err = binary.Read(r, binary.BigEndian, &rr.TTL)
if err != nil {
return rr, err
}
var rdLength uint16
err = binary.Read(r, binary.BigEndian, &rdLength)
if err != nil {
return rr, err
}
rr.Data = make([]byte, rdLength)
_, err = io.ReadFull(r, rr.Data)
if err != nil {
return rr, err
}
return rr, nil
}
// readMessage parses a complete DNS message. It leaves r positioned just after
// the parsed message.
func readMessage(r io.ReadSeeker) (Message, error) {
var message Message
// Header section
// https://tools.ietf.org/html/rfc1035#section-4.1.1
var qdCount, anCount, nsCount, arCount uint16
for _, ptr := range []*uint16{
&message.ID, &message.Flags,
&qdCount, &anCount, &nsCount, &arCount,
} {
err := binary.Read(r, binary.BigEndian, ptr)
if err != nil {
return message, err
}
}
// Question section
// https://tools.ietf.org/html/rfc1035#section-4.1.2
for i := 0; i < int(qdCount); i++ {
question, err := readQuestion(r)
if err != nil {
return message, err
}
message.Question = append(message.Question, question)
}
// Answer, Authority, and Additional sections
// https://tools.ietf.org/html/rfc1035#section-4.1.3
for _, rec := range []struct {
ptr *[]RR
count uint16
}{
{&message.Answer, anCount},
{&message.Authority, nsCount},
{&message.Additional, arCount},
} {
for i := 0; i < int(rec.count); i++ {
rr, err := readRR(r)
if err != nil {
return message, err
}
*rec.ptr = append(*rec.ptr, rr)
}
}
return message, nil
}
// MessageFromWireFormat parses a message from buf and returns a Message object.
// It returns ErrTrailingBytes if there are bytes remaining in buf after parsing
// is done.
func MessageFromWireFormat(buf []byte) (Message, error) {
r := bytes.NewReader(buf)
message, err := readMessage(r)
if err == io.EOF {
err = io.ErrUnexpectedEOF
} else if err == nil {
// Check for trailing bytes.
_, err = r.ReadByte()
if err == io.EOF {
err = nil
} else if err == nil {
err = ErrTrailingBytes
}
}
return message, err
}
// messageBuilder manages the state of serializing a DNS message. Its main
// function is to keep track of names already written for the purpose of name
// compression.
type messageBuilder struct {
w bytes.Buffer
nameCache map[string]int
}
// newMessageBuilder creates a new messageBuilder with an empty name cache.
func newMessageBuilder() *messageBuilder {
return &messageBuilder{
nameCache: make(map[string]int),
}
}
// Bytes returns the serialized DNS message as a slice of bytes.
func (builder *messageBuilder) Bytes() []byte {
return builder.w.Bytes()
}
// WriteName appends name to the in-progress messageBuilder, employing
// compression pointers to previously written names if possible.
func (builder *messageBuilder) WriteName(name Name) {
// https://tools.ietf.org/html/rfc1035#section-3.1
for i := range name {
// Has this suffix already been encoded in the message?
if ptr, ok := builder.nameCache[name[i:].String()]; ok && ptr&0x3fff == ptr {
// If so, we can write a compression pointer.
binary.Write(&builder.w, binary.BigEndian, uint16(0xc000|ptr))
return
}
// Not cached; we must encode this label verbatim. Store a cache
// entry pointing to the beginning of it.
builder.nameCache[name[i:].String()] = builder.w.Len()
length := len(name[i])
if length == 0 || length > 63 {
panic(length)
}
builder.w.WriteByte(byte(length))
builder.w.Write(name[i])
}
builder.w.WriteByte(0)
}
// WriteQuestion appends a Question section entry to the in-progress
// messageBuilder.
func (builder *messageBuilder) WriteQuestion(question *Question) {
// https://tools.ietf.org/html/rfc1035#section-4.1.2
builder.WriteName(question.Name)
binary.Write(&builder.w, binary.BigEndian, question.Type)
binary.Write(&builder.w, binary.BigEndian, question.Class)
}
// WriteRR appends a resource record to the in-progress messageBuilder. It
// returns ErrIntegerOverflow if the length of rr.Data does not fit in 16 bits.
func (builder *messageBuilder) WriteRR(rr *RR) error {
// https://tools.ietf.org/html/rfc1035#section-4.1.3
builder.WriteName(rr.Name)
binary.Write(&builder.w, binary.BigEndian, rr.Type)
binary.Write(&builder.w, binary.BigEndian, rr.Class)
binary.Write(&builder.w, binary.BigEndian, rr.TTL)
rdLength := uint16(len(rr.Data))
if int(rdLength) != len(rr.Data) {
return ErrIntegerOverflow
}
binary.Write(&builder.w, binary.BigEndian, rdLength)
builder.w.Write(rr.Data)
return nil
}
// WriteMessage appends a complete DNS message to the in-progress
// messageBuilder. It returns ErrIntegerOverflow if the number of entries in any
// section, or the length of the data in any resource record, does not fit in 16
// bits.
func (builder *messageBuilder) WriteMessage(message *Message) error {
// Header section
// https://tools.ietf.org/html/rfc1035#section-4.1.1
binary.Write(&builder.w, binary.BigEndian, message.ID)
binary.Write(&builder.w, binary.BigEndian, message.Flags)
for _, count := range []int{
len(message.Question),
len(message.Answer),
len(message.Authority),
len(message.Additional),
} {
count16 := uint16(count)
if int(count16) != count {
return ErrIntegerOverflow
}
binary.Write(&builder.w, binary.BigEndian, count16)
}
// Question section
// https://tools.ietf.org/html/rfc1035#section-4.1.2
for _, question := range message.Question {
builder.WriteQuestion(&question)
}
// Answer, Authority, and Additional sections
// https://tools.ietf.org/html/rfc1035#section-4.1.3
for _, rrs := range [][]RR{message.Answer, message.Authority, message.Additional} {
for _, rr := range rrs {
err := builder.WriteRR(&rr)
if err != nil {
return err
}
}
}
return nil
}
// WireFormat encodes a Message as a slice of bytes in DNS wire format. It
// returns ErrIntegerOverflow if the number of entries in any section, or the
// length of the data in any resource record, does not fit in 16 bits.
func (message *Message) WireFormat() ([]byte, error) {
builder := newMessageBuilder()
err := builder.WriteMessage(message)
if err != nil {
return nil, err
}
return builder.Bytes(), nil
}
// DecodeRDataTXT decodes TXT-DATA (as found in the RDATA for a resource record
// with TYPE=TXT) as a raw byte slice, by concatenating all the
// <character-string>s it contains.
//
// https://tools.ietf.org/html/rfc1035#section-3.3.14
func DecodeRDataTXT(p []byte) ([]byte, error) {
var buf bytes.Buffer
for {
if len(p) == 0 {
return nil, io.ErrUnexpectedEOF
}
n := int(p[0])
p = p[1:]
if len(p) < n {
return nil, io.ErrUnexpectedEOF
}
buf.Write(p[:n])
p = p[n:]
if len(p) == 0 {
break
}
}
return buf.Bytes(), nil
}
// EncodeRDataTXT encodes a slice of bytes as TXT-DATA, as appropriate for the
// RDATA of a resource record with TYPE=TXT. No length restriction is enforced
// here; that must be checked at a higher level.
//
// https://tools.ietf.org/html/rfc1035#section-3.3.14
func EncodeRDataTXT(p []byte) []byte {
// https://tools.ietf.org/html/rfc1035#section-3.3
// https://tools.ietf.org/html/rfc1035#section-3.3.14
// TXT data is a sequence of one or more <character-string>s, where
// <character-string> is a length octet followed by that number of
// octets.
var buf bytes.Buffer
for len(p) > 255 {
buf.WriteByte(255)
buf.Write(p[:255])
p = p[255:]
}
// Must write here, even if len(p) == 0, because it's "*one or more*
// <character-string>s".
buf.WriteByte(byte(len(p)))
buf.Write(p)
return buf.Bytes()
}
@@ -1,953 +0,0 @@
package xdns
import (
"bytes"
"fmt"
"io"
"strconv"
"strings"
"testing"
)
func namesEqual(a, b Name) bool {
if len(a) != len(b) {
return false
}
for i := 0; i < len(a); i++ {
if !bytes.Equal(a[i], b[i]) {
return false
}
}
return true
}
func TestName(t *testing.T) {
for _, test := range []struct {
labels [][]byte
err error
s string
}{
{[][]byte{}, nil, "."},
{[][]byte{[]byte("test")}, nil, "test"},
{[][]byte{[]byte("a"), []byte("b"), []byte("c")}, nil, "a.b.c"},
{[][]byte{{}}, ErrZeroLengthLabel, ""},
{[][]byte{[]byte("a"), {}, []byte("c")}, ErrZeroLengthLabel, ""},
// 63 octets.
{
[][]byte{[]byte("0123456789abcdef0123456789ABCDEF0123456789abcdef0123456789ABCDE")},
nil,
"0123456789abcdef0123456789ABCDEF0123456789abcdef0123456789ABCDE",
},
// 64 octets.
{[][]byte{[]byte("0123456789abcdef0123456789ABCDEF0123456789abcdef0123456789ABCDEF")}, ErrLabelTooLong, ""},
// 64+64+64+62 octets.
{
[][]byte{
[]byte("0123456789abcdef0123456789ABCDEF0123456789abcdef0123456789ABCDE"),
[]byte("0123456789abcdef0123456789ABCDEF0123456789abcdef0123456789ABCDE"),
[]byte("0123456789abcdef0123456789ABCDEF0123456789abcdef0123456789ABCDE"),
[]byte("0123456789abcdef0123456789ABCDEF0123456789abcdef0123456789ABC"),
},
nil,
"0123456789abcdef0123456789ABCDEF0123456789abcdef0123456789ABCDE.0123456789abcdef0123456789ABCDEF0123456789abcdef0123456789ABCDE.0123456789abcdef0123456789ABCDEF0123456789abcdef0123456789ABCDE.0123456789abcdef0123456789ABCDEF0123456789abcdef0123456789ABC",
},
// 64+64+64+63 octets.
{[][]byte{
[]byte("0123456789abcdef0123456789ABCDEF0123456789abcdef0123456789ABCDE"),
[]byte("0123456789abcdef0123456789ABCDEF0123456789abcdef0123456789ABCDE"),
[]byte("0123456789abcdef0123456789ABCDEF0123456789abcdef0123456789ABCDE"),
[]byte("0123456789abcdef0123456789ABCDEF0123456789abcdef0123456789ABCD"),
}, ErrNameTooLong, ""},
// 127 one-octet labels.
{
[][]byte{
{'0'},
{'1'},
{'2'},
{'3'},
{'4'},
{'5'},
{'6'},
{'7'},
{'8'},
{'9'},
{'a'},
{'b'},
{'c'},
{'d'},
{'e'},
{'f'},
{'0'},
{'1'},
{'2'},
{'3'},
{'4'},
{'5'},
{'6'},
{'7'},
{'8'},
{'9'},
{'A'},
{'B'},
{'C'},
{'D'},
{'E'},
{'F'},
{'0'},
{'1'},
{'2'},
{'3'},
{'4'},
{'5'},
{'6'},
{'7'},
{'8'},
{'9'},
{'a'},
{'b'},
{'c'},
{'d'},
{'e'},
{'f'},
{'0'},
{'1'},
{'2'},
{'3'},
{'4'},
{'5'},
{'6'},
{'7'},
{'8'},
{'9'},
{'A'},
{'B'},
{'C'},
{'D'},
{'E'},
{'F'},
{'0'},
{'1'},
{'2'},
{'3'},
{'4'},
{'5'},
{'6'},
{'7'},
{'8'},
{'9'},
{'a'},
{'b'},
{'c'},
{'d'},
{'e'},
{'f'},
{'0'},
{'1'},
{'2'},
{'3'},
{'4'},
{'5'},
{'6'},
{'7'},
{'8'},
{'9'},
{'A'},
{'B'},
{'C'},
{'D'},
{'E'},
{'F'},
{'0'},
{'1'},
{'2'},
{'3'},
{'4'},
{'5'},
{'6'},
{'7'},
{'8'},
{'9'},
{'a'},
{'b'},
{'c'},
{'d'},
{'e'},
{'f'},
{'0'},
{'1'},
{'2'},
{'3'},
{'4'},
{'5'},
{'6'},
{'7'},
{'8'},
{'9'},
{'A'},
{'B'},
{'C'},
{'D'},
{'E'},
},
nil,
"0.1.2.3.4.5.6.7.8.9.a.b.c.d.e.f.0.1.2.3.4.5.6.7.8.9.A.B.C.D.E.F.0.1.2.3.4.5.6.7.8.9.a.b.c.d.e.f.0.1.2.3.4.5.6.7.8.9.A.B.C.D.E.F.0.1.2.3.4.5.6.7.8.9.a.b.c.d.e.f.0.1.2.3.4.5.6.7.8.9.A.B.C.D.E.F.0.1.2.3.4.5.6.7.8.9.a.b.c.d.e.f.0.1.2.3.4.5.6.7.8.9.A.B.C.D.E",
},
// 128 one-octet labels.
{[][]byte{
{'0'},
{'1'},
{'2'},
{'3'},
{'4'},
{'5'},
{'6'},
{'7'},
{'8'},
{'9'},
{'a'},
{'b'},
{'c'},
{'d'},
{'e'},
{'f'},
{'0'},
{'1'},
{'2'},
{'3'},
{'4'},
{'5'},
{'6'},
{'7'},
{'8'},
{'9'},
{'A'},
{'B'},
{'C'},
{'D'},
{'E'},
{'F'},
{'0'},
{'1'},
{'2'},
{'3'},
{'4'},
{'5'},
{'6'},
{'7'},
{'8'},
{'9'},
{'a'},
{'b'},
{'c'},
{'d'},
{'e'},
{'f'},
{'0'},
{'1'},
{'2'},
{'3'},
{'4'},
{'5'},
{'6'},
{'7'},
{'8'},
{'9'},
{'A'},
{'B'},
{'C'},
{'D'},
{'E'},
{'F'},
{'0'},
{'1'},
{'2'},
{'3'},
{'4'},
{'5'},
{'6'},
{'7'},
{'8'},
{'9'},
{'a'},
{'b'},
{'c'},
{'d'},
{'e'},
{'f'},
{'0'},
{'1'},
{'2'},
{'3'},
{'4'},
{'5'},
{'6'},
{'7'},
{'8'},
{'9'},
{'A'},
{'B'},
{'C'},
{'D'},
{'E'},
{'F'},
{'0'},
{'1'},
{'2'},
{'3'},
{'4'},
{'5'},
{'6'},
{'7'},
{'8'},
{'9'},
{'a'},
{'b'},
{'c'},
{'d'},
{'e'},
{'f'},
{'0'},
{'1'},
{'2'},
{'3'},
{'4'},
{'5'},
{'6'},
{'7'},
{'8'},
{'9'},
{'A'},
{'B'},
{'C'},
{'D'},
{'E'},
{'F'},
}, ErrNameTooLong, ""},
} {
// Test that NewName returns proper error codes, and otherwise
// returns an equal slice of labels.
name, err := NewName(test.labels)
if err != test.err || (err == nil && !namesEqual(name, test.labels)) {
t.Errorf("%+q returned (%+q, %v), expected (%+q, %v)",
test.labels, name, err, test.labels, test.err)
continue
}
if test.err != nil {
continue
}
// Test that the string version of the name comes out as
// expected.
s := name.String()
if s != test.s {
t.Errorf("%+q became string %+q, expected %+q", test.labels, s, test.s)
continue
}
// Test that parsing from a string back to a Name results in the
// original slice of labels.
name, err = ParseName(s)
if err != nil || !namesEqual(name, test.labels) {
t.Errorf("%+q parsing %+q returned (%+q, %v), expected (%+q, %v)",
test.labels, s, name, err, test.labels, nil)
continue
}
// A trailing dot should be ignored.
if !strings.HasSuffix(s, ".") {
dotName, dotErr := ParseName(s + ".")
if dotErr != err || !namesEqual(dotName, name) {
t.Errorf("%+q parsing %+q returned (%+q, %v), expected (%+q, %v)",
test.labels, s+".", dotName, dotErr, name, err)
continue
}
}
}
}
func TestParseName(t *testing.T) {
for _, test := range []struct {
s string
name Name
err error
}{
// This case can't be tested by TestName above because String
// will never produce "" (it produces "." instead).
{"", [][]byte{}, nil},
} {
name, err := ParseName(test.s)
if err != test.err || (err == nil && !namesEqual(name, test.name)) {
t.Errorf("%+q returned (%+q, %v), expected (%+q, %v)",
test.s, name, err, test.name, test.err)
continue
}
}
}
func unescapeString(s string) ([][]byte, error) {
if s == "." {
return [][]byte{}, nil
}
var result [][]byte
for _, label := range strings.Split(s, ".") {
var buf bytes.Buffer
i := 0
for i < len(label) {
switch label[i] {
case '\\':
if i+3 >= len(label) {
return nil, fmt.Errorf("truncated escape sequence at index %v", i)
}
if label[i+1] != 'x' {
return nil, fmt.Errorf("malformed escape sequence at index %v", i)
}
b, err := strconv.ParseUint(string(label[i+2:i+4]), 16, 8)
if err != nil {
return nil, fmt.Errorf("malformed hex sequence at index %v", i+2)
}
buf.WriteByte(byte(b))
i += 4
default:
buf.WriteByte(label[i])
i++
}
}
result = append(result, buf.Bytes())
}
return result, nil
}
func TestNameString(t *testing.T) {
for _, test := range []struct {
name Name
s string
}{
{[][]byte{}, "."},
{[][]byte{[]byte("\x00"), []byte("a.b"), []byte("c\nd\\")}, "\\x00.a\\x2eb.c\\x0ad\\x5c"},
{[][]byte{
[]byte("\x00\x01\x02\x03\x04\x05\x06\x07\x08\t\n\x0b\x0c\r\x0e\x0f\x10\x11\x12\x13\x14\x15\x16\x17\x18\x19\x1a\x1b\x1c\x1d\x1e\x1f !\"#$%&'()*+,-./0123456789:;<=>"),
[]byte("?@ABCDEFGHIJKLMNOPQRSTUVWXYZ[\\]^_`abcdefghijklmnopqrstuvwxyz{|}"),
[]byte("~\x7f\x80\x81\x82\x83\x84\x85\x86\x87\x88\x89\x8a\x8b\x8c\x8d\x8e\x8f\x90\x91\x92\x93\x94\x95\x96\x97\x98\x99\x9a\x9b\x9c\x9d\x9e\x9f\xa0\xa1\xa2\xa3\xa4\xa5\xa6\xa7\xa8\xa9\xaa\xab\xac\xad\xae\xaf\xb0\xb1\xb2\xb3\xb4\xb5\xb6\xb7\xb8\xb9\xba\xbb\xbc"),
[]byte("\xbd\xbe\xbf\xc0\xc1\xc2\xc3\xc4\xc5\xc6\xc7\xc8\xc9\xca\xcb\xcc\xcd\xce\xcf\xd0\xd1\xd2\xd3\xd4\xd5\xd6\xd7\xd8\xd9\xda\xdb\xdc\xdd\xde\xdf\xe0\xe1\xe2\xe3\xe4\xe5\xe6\xe7\xe8\xe9\xea\xeb\xec\xed\xee\xef\xf0\xf1\xf2\xf3\xf4\xf5\xf6\xf7\xf8\xf9\xfa\xfb"),
[]byte("\xfc\xfd\xfe\xff"),
}, "\\x00\\x01\\x02\\x03\\x04\\x05\\x06\\x07\\x08\\x09\\x0a\\x0b\\x0c\\x0d\\x0e\\x0f\\x10\\x11\\x12\\x13\\x14\\x15\\x16\\x17\\x18\\x19\\x1a\\x1b\\x1c\\x1d\\x1e\\x1f\\x20\\x21\\x22\\x23\\x24\\x25\\x26\\x27\\x28\\x29\\x2a\\x2b\\x2c-\\x2e\\x2f0123456789\\x3a\\x3b\\x3c\\x3d\\x3e.\\x3f\\x40ABCDEFGHIJKLMNOPQRSTUVWXYZ\\x5b\\x5c\\x5d\\x5e\\x5f\\x60abcdefghijklmnopqrstuvwxyz\\x7b\\x7c\\x7d.\\x7e\\x7f\\x80\\x81\\x82\\x83\\x84\\x85\\x86\\x87\\x88\\x89\\x8a\\x8b\\x8c\\x8d\\x8e\\x8f\\x90\\x91\\x92\\x93\\x94\\x95\\x96\\x97\\x98\\x99\\x9a\\x9b\\x9c\\x9d\\x9e\\x9f\\xa0\\xa1\\xa2\\xa3\\xa4\\xa5\\xa6\\xa7\\xa8\\xa9\\xaa\\xab\\xac\\xad\\xae\\xaf\\xb0\\xb1\\xb2\\xb3\\xb4\\xb5\\xb6\\xb7\\xb8\\xb9\\xba\\xbb\\xbc.\\xbd\\xbe\\xbf\\xc0\\xc1\\xc2\\xc3\\xc4\\xc5\\xc6\\xc7\\xc8\\xc9\\xca\\xcb\\xcc\\xcd\\xce\\xcf\\xd0\\xd1\\xd2\\xd3\\xd4\\xd5\\xd6\\xd7\\xd8\\xd9\\xda\\xdb\\xdc\\xdd\\xde\\xdf\\xe0\\xe1\\xe2\\xe3\\xe4\\xe5\\xe6\\xe7\\xe8\\xe9\\xea\\xeb\\xec\\xed\\xee\\xef\\xf0\\xf1\\xf2\\xf3\\xf4\\xf5\\xf6\\xf7\\xf8\\xf9\\xfa\\xfb.\\xfc\\xfd\\xfe\\xff"},
} {
s := test.name.String()
if s != test.s {
t.Errorf("%+q escaped to %+q, expected %+q", test.name, s, test.s)
continue
}
unescaped, err := unescapeString(s)
if err != nil {
t.Errorf("%+q unescaping %+q resulted in error %v", test.name, s, err)
continue
}
if !namesEqual(Name(unescaped), test.name) {
t.Errorf("%+q roundtripped through %+q to %+q", test.name, s, unescaped)
continue
}
}
}
func TestNameTrimSuffix(t *testing.T) {
for _, test := range []struct {
name, suffix string
trimmed string
ok bool
}{
{"", "", ".", true},
{".", ".", ".", true},
{"abc", "", "abc", true},
{"abc", ".", "abc", true},
{"", "abc", ".", false},
{".", "abc", ".", false},
{"example.com", "com", "example", true},
{"example.com", "net", ".", false},
{"example.com", "example.com", ".", true},
{"example.com", "test.com", ".", false},
{"example.com", "xample.com", ".", false},
{"example.com", "example", ".", false},
{"example.com", "COM", "example", true},
{"EXAMPLE.COM", "com", "EXAMPLE", true},
} {
tmp, ok := mustParseName(test.name).TrimSuffix(mustParseName(test.suffix))
trimmed := tmp.String()
if ok != test.ok || trimmed != test.trimmed {
t.Errorf("TrimSuffix %+q %+q returned (%+q, %v), expected (%+q, %v)",
test.name, test.suffix, trimmed, ok, test.trimmed, test.ok)
continue
}
}
}
func TestReadName(t *testing.T) {
// Good tests.
for _, test := range []struct {
start int64
end int64
input string
s string
}{
// Empty name.
{0, 1, "\x00abcd", "."},
// No pointers.
{12, 25, "AAAABBBBCCCC\x07example\x03com\x00", "example.com"},
// Backward pointer.
{25, 31, "AAAABBBBCCCC\x07example\x03com\x00\x03sub\xc0\x0c", "sub.example.com"},
// Forward pointer.
{0, 4, "\x01a\xc0\x04\x03bcd\x00", "a.bcd"},
// Two backwards pointers.
{31, 38, "AAAABBBBCCCC\x07example\x03com\x00\x03sub\xc0\x0c\x04sub2\xc0\x19", "sub2.sub.example.com"},
// Forward then backward pointer.
{25, 31, "AAAABBBBCCCC\x07example\x03com\x00\x03sub\xc0\x1f\x04sub2\xc0\x0c", "sub.sub2.example.com"},
// Overlapping codons.
{0, 4, "\x01a\xc0\x03bcd\x00", "a.bcd"},
// Pointer to empty label.
{0, 10, "\x07example\xc0\x0a\x00", "example"},
{1, 11, "\x00\x07example\xc0\x00", "example"},
// Pointer to pointer to empty label.
{0, 10, "\x07example\xc0\x0a\xc0\x0c\x00", "example"},
{1, 11, "\x00\x07example\xc0\x0c\xc0\x00", "example"},
} {
r := bytes.NewReader([]byte(test.input))
_, err := r.Seek(test.start, io.SeekStart)
if err != nil {
panic(err)
}
name, err := readName(r)
if err != nil {
t.Errorf("%+q returned error %s", test.input, err)
continue
}
s := name.String()
if s != test.s {
t.Errorf("%+q returned %+q, expected %+q", test.input, s, test.s)
continue
}
cur, _ := r.Seek(0, io.SeekCurrent)
if cur != test.end {
t.Errorf("%+q left offset %d, expected %d", test.input, cur, test.end)
continue
}
}
// Bad tests.
for _, test := range []struct {
start int64
input string
err error
}{
{0, "", io.ErrUnexpectedEOF},
// Reserved label type.
{0, "\x80example", ErrReservedLabelType},
// Reserved label type.
{0, "\x40example", ErrReservedLabelType},
// No Terminating empty label.
{0, "\x07example\x03com", io.ErrUnexpectedEOF},
// Pointer past end of buffer.
{0, "\x07example\xc0\xff", io.ErrUnexpectedEOF},
// Pointer to self.
{0, "\x07example\x03com\xc0\x0c", ErrTooManyPointers},
// Pointer to self with intermediate label.
{0, "\x07example\x03com\xc0\x08", ErrTooManyPointers},
// Two pointers that point to each other.
{0, "\xc0\x02\xc0\x00", ErrTooManyPointers},
// Two pointers that point to each other, with intermediate labels.
{0, "\x01a\xc0\x04\x01b\xc0\x00", ErrTooManyPointers},
// EOF while reading label.
{0, "\x0aexample", io.ErrUnexpectedEOF},
// EOF before second byte of pointer.
{0, "\xc0", io.ErrUnexpectedEOF},
{0, "\x07example\xc0", io.ErrUnexpectedEOF},
} {
r := bytes.NewReader([]byte(test.input))
_, err := r.Seek(test.start, io.SeekStart)
if err != nil {
panic(err)
}
name, err := readName(r)
if err == io.EOF {
err = io.ErrUnexpectedEOF
}
if err != test.err {
t.Errorf("%+q returned (%+q, %v), expected %v", test.input, name, err, test.err)
continue
}
}
}
func mustParseName(s string) Name {
name, err := ParseName(s)
if err != nil {
panic(err)
}
return name
}
func questionsEqual(a, b *Question) bool {
if !namesEqual(a.Name, b.Name) {
return false
}
if a.Type != b.Type || a.Class != b.Class {
return false
}
return true
}
func rrsEqual(a, b *RR) bool {
if !namesEqual(a.Name, b.Name) {
return false
}
if a.Type != b.Type || a.Class != b.Class || a.TTL != b.TTL {
return false
}
if !bytes.Equal(a.Data, b.Data) {
return false
}
return true
}
func messagesEqual(a, b *Message) bool {
if a.ID != b.ID || a.Flags != b.Flags {
return false
}
if len(a.Question) != len(b.Question) {
return false
}
for i := 0; i < len(a.Question); i++ {
if !questionsEqual(&a.Question[i], &b.Question[i]) {
return false
}
}
for _, rec := range []struct{ rrA, rrB []RR }{
{a.Answer, b.Answer},
{a.Authority, b.Authority},
{a.Additional, b.Additional},
} {
if len(rec.rrA) != len(rec.rrB) {
return false
}
for i := 0; i < len(rec.rrA); i++ {
if !rrsEqual(&rec.rrA[i], &rec.rrB[i]) {
return false
}
}
}
return true
}
func TestMessageFromWireFormat(t *testing.T) {
for _, test := range []struct {
buf string
expected Message
err error
}{
{
"\x12\x34",
Message{},
io.ErrUnexpectedEOF,
},
{
"\x12\x34\x01\x00\x00\x01\x00\x00\x00\x00\x00\x00\x03www\x07example\x03com\x00\x00\x01\x00\x01",
Message{
ID: 0x1234,
Flags: 0x0100,
Question: []Question{
{
Name: mustParseName("www.example.com"),
Type: 1,
Class: 1,
},
},
Answer: []RR{},
Authority: []RR{},
Additional: []RR{},
},
nil,
},
{
"\x12\x34\x01\x00\x00\x01\x00\x00\x00\x00\x00\x00\x03www\x07example\x03com\x00\x00\x01\x00\x01X",
Message{},
ErrTrailingBytes,
},
{
"\x12\x34\x81\x80\x00\x01\x00\x01\x00\x00\x00\x00\x03www\x07example\x03com\x00\x00\x01\x00\x01\x03www\x07example\x03com\x00\x00\x01\x00\x01\x00\x00\x00\x80\x00\x04\xc0\x00\x02\x01",
Message{
ID: 0x1234,
Flags: 0x8180,
Question: []Question{
{
Name: mustParseName("www.example.com"),
Type: 1,
Class: 1,
},
},
Answer: []RR{
{
Name: mustParseName("www.example.com"),
Type: 1,
Class: 1,
TTL: 128,
Data: []byte{192, 0, 2, 1},
},
},
Authority: []RR{},
Additional: []RR{},
},
nil,
},
} {
message, err := MessageFromWireFormat([]byte(test.buf))
if err != test.err || (err == nil && !messagesEqual(&message, &test.expected)) {
t.Errorf("%+q\nreturned (%+v, %v)\nexpected (%+v, %v)",
test.buf, message, err, test.expected, test.err)
continue
}
}
}
func TestMessageWireFormatRoundTrip(t *testing.T) {
for _, message := range []Message{
{
ID: 0x1234,
Flags: 0x0100,
Question: []Question{
{
Name: mustParseName("www.example.com"),
Type: 1,
Class: 1,
},
{
Name: mustParseName("www2.example.com"),
Type: 2,
Class: 2,
},
},
Answer: []RR{
{
Name: mustParseName("abc"),
Type: 2,
Class: 3,
TTL: 0xffffffff,
Data: []byte{1},
},
{
Name: mustParseName("xyz"),
Type: 2,
Class: 3,
TTL: 255,
Data: []byte{},
},
},
Authority: []RR{
{
Name: mustParseName("."),
Type: 65535,
Class: 65535,
TTL: 0,
Data: []byte("XXXXXXXXXXXXXXXXXXX"),
},
},
Additional: []RR{},
},
} {
buf, err := message.WireFormat()
if err != nil {
t.Errorf("%+v cannot make wire format: %v", message, err)
continue
}
message2, err := MessageFromWireFormat(buf)
if err != nil {
t.Errorf("%+q cannot parse wire format: %v", buf, err)
continue
}
if !messagesEqual(&message, &message2) {
t.Errorf("messages unequal\nbefore: %+v\n after: %+v", message, message2)
continue
}
}
}
func TestDecodeRDataTXT(t *testing.T) {
for _, test := range []struct {
p []byte
decoded []byte
err error
}{
{[]byte{}, nil, io.ErrUnexpectedEOF},
{[]byte("\x00"), []byte{}, nil},
{[]byte("\x01"), nil, io.ErrUnexpectedEOF},
} {
decoded, err := DecodeRDataTXT(test.p)
if err != test.err || (err == nil && !bytes.Equal(decoded, test.decoded)) {
t.Errorf("%+q\nreturned (%+q, %v)\nexpected (%+q, %v)",
test.p, decoded, err, test.decoded, test.err)
continue
}
}
}
func TestEncodeRDataTXT(t *testing.T) {
// Encoding 0 bytes needs to return at least a single length octet of
// zero, not an empty slice.
p := make([]byte, 0)
encoded := EncodeRDataTXT(p)
if len(encoded) < 0 {
t.Errorf("EncodeRDataTXT(%v) returned %v", p, encoded)
}
// 255 bytes should be able to be encoded into 256 bytes.
p = make([]byte, 255)
encoded = EncodeRDataTXT(p)
if len(encoded) > 256 {
t.Errorf("EncodeRDataTXT(%d bytes) returned %d bytes", len(p), len(encoded))
}
fmt.Println(EncodeRDataTXT(nil))
fmt.Println(computeMaxEncodedPayload(maxUDPPayload))
}
func TestRDataTXTRoundTrip(t *testing.T) {
for _, p := range [][]byte{
{},
[]byte("\x00"),
{
0x00, 0x01, 0x02, 0x03, 0x04, 0x05, 0x06, 0x07, 0x08, 0x09, 0x0a, 0x0b, 0x0c, 0x0d, 0x0e, 0x0f,
0x10, 0x11, 0x12, 0x13, 0x14, 0x15, 0x16, 0x17, 0x18, 0x19, 0x1a, 0x1b, 0x1c, 0x1d, 0x1e, 0x1f,
0x20, 0x21, 0x22, 0x23, 0x24, 0x25, 0x26, 0x27, 0x28, 0x29, 0x2a, 0x2b, 0x2c, 0x2d, 0x2e, 0x2f,
0x30, 0x31, 0x32, 0x33, 0x34, 0x35, 0x36, 0x37, 0x38, 0x39, 0x3a, 0x3b, 0x3c, 0x3d, 0x3e, 0x3f,
0x40, 0x41, 0x42, 0x43, 0x44, 0x45, 0x46, 0x47, 0x48, 0x49, 0x4a, 0x4b, 0x4c, 0x4d, 0x4e, 0x4f,
0x50, 0x51, 0x52, 0x53, 0x54, 0x55, 0x56, 0x57, 0x58, 0x59, 0x5a, 0x5b, 0x5c, 0x5d, 0x5e, 0x5f,
0x60, 0x61, 0x62, 0x63, 0x64, 0x65, 0x66, 0x67, 0x68, 0x69, 0x6a, 0x6b, 0x6c, 0x6d, 0x6e, 0x6f,
0x70, 0x71, 0x72, 0x73, 0x74, 0x75, 0x76, 0x77, 0x78, 0x79, 0x7a, 0x7b, 0x7c, 0x7d, 0x7e, 0x7f,
0x80, 0x81, 0x82, 0x83, 0x84, 0x85, 0x86, 0x87, 0x88, 0x89, 0x8a, 0x8b, 0x8c, 0x8d, 0x8e, 0x8f,
0x90, 0x91, 0x92, 0x93, 0x94, 0x95, 0x96, 0x97, 0x98, 0x99, 0x9a, 0x9b, 0x9c, 0x9d, 0x9e, 0x9f,
0xa0, 0xa1, 0xa2, 0xa3, 0xa4, 0xa5, 0xa6, 0xa7, 0xa8, 0xa9, 0xaa, 0xab, 0xac, 0xad, 0xae, 0xaf,
0xb0, 0xb1, 0xb2, 0xb3, 0xb4, 0xb5, 0xb6, 0xb7, 0xb8, 0xb9, 0xba, 0xbb, 0xbc, 0xbd, 0xbe, 0xbf,
0xc0, 0xc1, 0xc2, 0xc3, 0xc4, 0xc5, 0xc6, 0xc7, 0xc8, 0xc9, 0xca, 0xcb, 0xcc, 0xcd, 0xce, 0xcf,
0xd0, 0xd1, 0xd2, 0xd3, 0xd4, 0xd5, 0xd6, 0xd7, 0xd8, 0xd9, 0xda, 0xdb, 0xdc, 0xdd, 0xde, 0xdf,
0xe0, 0xe1, 0xe2, 0xe3, 0xe4, 0xe5, 0xe6, 0xe7, 0xe8, 0xe9, 0xea, 0xeb, 0xec, 0xed, 0xee, 0xef,
0xf0, 0xf1, 0xf2, 0xf3, 0xf4, 0xf5, 0xf6, 0xf7, 0xf8, 0xf9, 0xfa, 0xfb, 0xfc, 0xfd, 0xfe, 0xff,
},
} {
rdata := EncodeRDataTXT(p)
decoded, err := DecodeRDataTXT(rdata)
if err != nil || !bytes.Equal(decoded, p) {
t.Errorf("%+q returned (%+q, %v)", p, decoded, err)
continue
}
}
}
func TestIPAnswerPayloadRoundTrip(t *testing.T) {
for _, rrType := range []uint16{RRTypeA, RRTypeAAAA} {
for _, payload := range [][]byte{
{},
{0x01},
[]byte("hello world"),
bytes.Repeat([]byte{0xab}, payloadChunkSizeForType(rrType)*3+1),
} {
question := Question{
Name: mustParseName("example.com"),
Type: rrType,
Class: ClassIN,
}
answers, err := answersForPayload(question, responseTTL, payload)
if err != nil {
t.Fatalf("answersForPayload(%d) err = %v", rrType, err)
}
if len(answers) > 1 {
answers[0], answers[len(answers)-1] = answers[len(answers)-1], answers[0]
}
decoded := decodeResponsePayload(answers)
if !bytes.Equal(decoded, payload) {
t.Fatalf("rrType=%d decoded %x want %x", rrType, decoded, payload)
}
}
}
}
func TestParseResolver(t *testing.T) {
tests := []struct {
resolver string
rrType uint16
}{
{"example.com+udp://1.1.1.1:53", RRTypeTXT},
{"example.com:txt+udp://1.1.1.1:53", RRTypeTXT},
{"example.com:a+udp://1.1.1.1:53", RRTypeA},
{"example.com:aaaa+udp://1.1.1.1:53", RRTypeAAAA},
}
for _, test := range tests {
domain, server, rrType, err := parseResolver(test.resolver)
if err != nil {
t.Fatalf("parseResolver(%q) err = %v", test.resolver, err)
}
if domain.String() != "example.com" || server != "1.1.1.1:53" || rrType != test.rrType {
t.Fatalf("parseResolver(%q) = (%q, %q, %d)", test.resolver, domain.String(), server, rrType)
}
}
}
func TestParseDomainSpec(t *testing.T) {
tests := []struct {
spec string
def string
rrType uint16
wantErr bool
}{
{"example.com", "", 0, false},
{"example.com", "txt", RRTypeTXT, false},
{"example.com:a", "", RRTypeA, false},
{"example.com:aaaa", "", RRTypeAAAA, false},
{"example.com:doh", "", 0, true},
}
for _, test := range tests {
got, err := parseDomainSpec(test.spec, test.def)
if test.wantErr {
if err == nil {
t.Fatalf("parseDomainSpec(%q, %q) err = nil", test.spec, test.def)
}
continue
}
if err != nil {
t.Fatalf("parseDomainSpec(%q, %q) err = %v", test.spec, test.def, err)
}
if got.name.String() != "example.com" || got.rrType != test.rrType {
t.Fatalf("parseDomainSpec(%q, %q) = (%q, %d)", test.spec, test.def, got.name.String(), got.rrType)
}
}
}
func TestResponseForMethodRestriction(t *testing.T) {
query := &Message{
ID: 1,
Flags: 0x0100,
Question: []Question{{
Name: mustParseName("abc.example.com"),
Type: RRTypeTXT,
Class: ClassIN,
}},
Additional: []RR{{
Name: Name{},
Type: RRTypeOPT,
Class: 4096,
}},
}
resp, _ := responseFor(query, []domainSpec{{name: mustParseName("example.com"), rrType: RRTypeA}})
if resp == nil || resp.Rcode() != RcodeNameError {
t.Fatalf("responseFor method restriction rcode = %v", resp)
}
resp, _ = responseFor(query, []domainSpec{{name: mustParseName("example.com")}})
if resp == nil || resp.Rcode() != RcodeNoError {
t.Fatalf("responseFor unrestricted rcode = %v", resp)
}
}
+215
View File
@@ -0,0 +1,215 @@
package xdns
import (
"encoding/base32"
"errors"
"fmt"
"strings"
"golang.org/x/net/dns/dnsmessage"
"golang.org/x/net/idna"
)
func Lower(c byte) byte {
if c >= 'A' && c <= 'Z' {
return c + ('a' - 'A')
}
return c
}
func ToUpper(b []byte) {
for i, c := range b {
if c >= 'a' && c <= 'z' {
b[i] = c - 'a' + 'A'
}
}
}
func ToLower(b []byte) {
for i, c := range b {
if c >= 'A' && c <= 'Z' {
b[i] = c - 'A' + 'a'
}
}
}
func NewTable() ([256]int, [256]int) {
var t, t_ [256]int
for i := range t {
t[i] = base32Encoding.DecodedLen(i)
}
for i := range t_ {
t_[i] = base32Encoding.EncodedLen(i)
}
return t, t_
}
const (
TypeA uint16 = 1
TypeCNAME uint16 = 5
TypeTXT uint16 = 16
TypeAAAA uint16 = 28
)
var (
base32Encoding = base32.StdEncoding.WithPadding(base32.NoPadding)
table, table_ = NewTable()
TypeMap = map[uint16]byte{
TypeA: 0,
TypeCNAME: 1,
TypeTXT: 2,
TypeAAAA: 3,
}
TypeMap_ = map[byte]uint16{
0: TypeA,
1: TypeCNAME,
2: TypeTXT,
3: TypeAAAA,
}
)
type Domain struct {
name dnsmessage.Name
lenLimit int
labelLimit int
types []uint16
edns0 uint16
cap int
lenMax int
}
func NewDomain(domain string, lenLimit int, labelLimit int, types []uint16, edns0 uint16) (*Domain, error) {
if strings.Contains(domain, "..") {
return nil, errors.New("invalid domain")
}
if lenLimit < 0 || lenLimit > 255 {
return nil, errors.New("lenLimit < 0 || lenLimit > 255")
}
if labelLimit < 0 || labelLimit > 63 {
return nil, errors.New("labelLimit < 0 || labelLimit > 63")
}
if len(types) == 0 {
return nil, errors.New("empty types")
}
for i := range types {
switch types[i] {
case uint16(dnsmessage.TypeA), uint16(dnsmessage.TypeCNAME), uint16(dnsmessage.TypeTXT), uint16(dnsmessage.TypeAAAA):
default:
return nil, errors.New("unknown types")
}
}
if edns0 != 0 && (edns0 < 512 || edns0 > 4096) {
return nil, errors.New("edns0 != 0 && (edns0 < 512 || edns0 > 4096)")
}
ascii, err := idna.ToASCII(domain)
if err != nil {
return nil, err
}
ascii = strings.Trim(ascii, ".")
name, err := dnsmessage.NewName(domain + ".")
if err != nil {
return nil, err
}
if lenLimit < int(name.Length)+1 {
return nil, errors.New("lenLimit < int(name.Length)+1")
}
n := (lenLimit - int(name.Length) - 1) / (labelLimit + 1)
left := (lenLimit - int(name.Length) - 1) % (labelLimit + 1)
total := n * labelLimit
if left > 1 {
total += left - 1
}
cap := table[total]
if cap < 17 {
return nil, errors.New("cap < 17")
}
total = table_[cap]
lenMax := int(name.Length) + 1 + total + total/labelLimit
if total%labelLimit > 0 {
lenMax += 1
}
return &Domain{
name: name,
lenLimit: lenLimit,
labelLimit: labelLimit,
types: types,
edns0: edns0,
cap: cap,
lenMax: lenMax,
}, nil
}
func (d *Domain) Show() string {
return fmt.Sprint(d.name, d.cap)
}
func (d *Domain) IsDomain(name dnsmessage.Name) bool {
if d.name.Length >= name.Length {
return false
}
i := d.name.Length
j := name.Length
for i > 0 {
i--
j--
if Lower(d.name.Data[i]) != Lower(name.Data[j]) {
return false
}
}
return true
}
func (d *Domain) HasType(qtype uint16) bool {
for i := range d.types {
if d.types[i] == qtype {
return true
}
}
return false
}
func (d *Domain) Encode(data []byte) dnsmessage.Name {
var name dnsmessage.Name
var encoded [255]byte
base32Encoding.Encode(encoded[:], data)
ToLower(encoded[:table_[len(data)]])
b1 := name.Data[:0]
b2 := encoded[:table_[len(data)]]
for len(b2) > 0 {
size := min(len(b2), d.labelLimit)
b1 = append(b1, b2[:size]...)
b1 = append(b1, '.')
b2 = b2[size:]
}
b1 = append(b1, d.name.Data[:d.name.Length]...)
if len(b1) > 254 {
panic("len(b1) > 254")
}
name.Length = byte(len(b1))
return name
}
func (d *Domain) Decode(decoded *[255]byte, name dnsmessage.Name) int {
if !d.IsDomain(name) {
return 0
}
var encoded [255]byte
b1 := encoded[:0]
b2 := name.Data[:name.Length-d.name.Length]
for i := range b2 {
if b2[i] != '.' {
b1 = append(b1, b2[i])
}
}
ToUpper(b1)
n, err := base32Encoding.Decode(decoded[:], b1)
if err != nil {
return 0
}
return n
}
+171
View File
@@ -0,0 +1,171 @@
package xdns
import (
"sync"
"time"
)
const (
fragTTL = 8 * time.Second
fragSize = 4096
fragClientIDSize = 16384
fragCount = 4096
)
type FragKey struct {
clientID ClientID
fragID byte
}
type FragEntry struct {
data [][]byte
size int
len int
total byte
deadline time.Time
}
type FragManager struct {
m map[FragKey]*FragEntry
sizem map[ClientID]int
ch chan struct{}
mu sync.Mutex
}
func NewFragManager() *FragManager {
m := &FragManager{
m: make(map[FragKey]*FragEntry),
sizem: make(map[ClientID]int),
ch: make(chan struct{}),
}
go m.gc()
return m
}
func (m *FragManager) closed() bool {
select {
case <-m.ch:
return true
default:
return false
}
}
func (m *FragManager) removeEntey(k FragKey, e *FragEntry) {
m.sizem[k.clientID] -= e.size
delete(m.m, k)
}
func (m *FragManager) tryRemove() {
if len(m.m) < fragCount {
return
}
var key FragKey
var entry *FragEntry
first := true
for k, e := range m.m {
if first || e.deadline.Before(entry.deadline) {
key = k
entry = e
first = false
}
}
m.removeEntey(key, entry)
}
func (m *FragManager) gc() {
ticker := time.NewTicker(fragTTL / 2)
defer ticker.Stop()
for {
select {
case <-m.ch:
return
case now := <-ticker.C:
m.mu.Lock()
for k, e := range m.m {
if now.After(e.deadline) {
m.removeEntey(k, e)
}
}
m.mu.Unlock()
}
}
}
func (m *FragManager) Feed(out []byte, key FragKey, fragIdx, fragN byte, data []byte) int {
m.mu.Lock()
defer m.mu.Unlock()
if m.closed() {
return 0
}
if fragN < 2 {
return 0
}
now := time.Now()
entry := m.m[key]
if entry == nil || now.After(entry.deadline) {
if entry == nil {
m.tryRemove()
} else {
m.removeEntey(key, entry)
}
entry = &FragEntry{
data: make([][]byte, fragN),
total: fragN,
deadline: now.Add(fragTTL),
}
m.m[key] = entry
}
if fragN != entry.total {
return 0
}
if fragIdx >= entry.total {
return 0
}
if entry.data[fragIdx] != nil {
return 0
}
if entry.size+len(data) > fragSize {
return 0
}
if entry.len < int(entry.total)-1 {
if m.sizem[key.clientID]+len(data) > fragClientIDSize {
return 0
}
}
cp := make([]byte, len(data))
copy(cp, data)
entry.data[fragIdx] = cp
entry.size += len(data)
entry.len++
entry.deadline = now.Add(fragTTL)
m.sizem[key.clientID] += len(data)
if entry.len < int(entry.total) {
return 0
}
out = out[:0]
for i := range entry.data {
out = append(out, entry.data[i]...)
}
m.removeEntey(key, entry)
return len(out)
}
func (m *FragManager) Close() {
m.mu.Lock()
defer m.mu.Unlock()
if m.closed() {
return
}
close(m.ch)
for k := range m.m {
delete(m.m, k)
}
}
@@ -1,226 +0,0 @@
package xdns
import "bytes"
const ipRecordHeaderSize = 2
func maxEncodedPayloadForType(rrType uint16) int {
switch rrType {
case RRTypeA:
return maxEncodedPayloadA
case RRTypeAAAA:
return maxEncodedPayloadAAAA
default:
return maxEncodedPayloadTXT
}
}
func rrDataSizeForType(rrType uint16) int {
switch rrType {
case RRTypeA:
return 4
case RRTypeAAAA:
return 16
default:
return 0
}
}
func payloadChunkSizeForType(rrType uint16) int {
size := rrDataSizeForType(rrType)
if size <= ipRecordHeaderSize {
return 0
}
return size - ipRecordHeaderSize
}
func answersForPayload(question Question, ttl uint32, payload []byte) ([]RR, error) {
switch question.Type {
case RRTypeTXT:
return []RR{
{
Name: question.Name,
Type: question.Type,
Class: question.Class,
TTL: ttl,
Data: EncodeRDataTXT(payload),
},
}, nil
case RRTypeA, RRTypeAAAA:
return ipAnswersForPayload(question, ttl, payload)
default:
return nil, ErrIntegerOverflow
}
}
func ipAnswersForPayload(question Question, ttl uint32, payload []byte) ([]RR, error) {
chunkSize := payloadChunkSizeForType(question.Type)
rrDataSize := rrDataSizeForType(question.Type)
if chunkSize == 0 || rrDataSize == 0 {
return nil, ErrIntegerOverflow
}
numRecords := 1
if len(payload) > 0 {
numRecords = (len(payload) + chunkSize - 1) / chunkSize
}
if numRecords > 256 {
return nil, ErrIntegerOverflow
}
answers := make([]RR, 0, numRecords)
for i := 0; i < numRecords; i++ {
offset := i * chunkSize
n := len(payload) - offset
if n < 0 {
n = 0
}
if n > chunkSize {
n = chunkSize
}
data := make([]byte, rrDataSize)
data[0] = byte(i)
data[1] = byte(n)
copy(data[ipRecordHeaderSize:], payload[offset:offset+n])
answers = append(answers, RR{
Name: question.Name,
Type: question.Type,
Class: question.Class,
TTL: ttl,
Data: data,
})
}
return answers, nil
}
func decodeResponsePayload(answers []RR) []byte {
if len(answers) == 0 {
return nil
}
switch answers[0].Type {
case RRTypeTXT:
if len(answers) != 1 {
return nil
}
payload, err := DecodeRDataTXT(answers[0].Data)
if err != nil {
return nil
}
return payload
case RRTypeA, RRTypeAAAA:
return decodeIPAnswerPayload(answers, answers[0].Type)
default:
return nil
}
}
func decodeIPAnswerPayload(answers []RR, rrType uint16) []byte {
chunkSize := payloadChunkSizeForType(rrType)
rrDataSize := rrDataSizeForType(rrType)
if chunkSize == 0 || rrDataSize == 0 || len(answers) > 256 {
return nil
}
parts := make([][]byte, len(answers))
for _, answer := range answers {
if answer.Type != rrType || len(answer.Data) != rrDataSize {
return nil
}
idx := int(answer.Data[0])
n := int(answer.Data[1])
if idx >= len(answers) || n > chunkSize || parts[idx] != nil {
return nil
}
part := make([]byte, n)
copy(part, answer.Data[ipRecordHeaderSize:ipRecordHeaderSize+n])
parts[idx] = part
}
var payload bytes.Buffer
for _, part := range parts {
if part == nil {
return nil
}
payload.Write(part)
}
return payload.Bytes()
}
func computeMaxEncodedPayload(limit int) int {
return computeMaxEncodedPayloadForType(limit, RRTypeTXT)
}
func computeMaxEncodedPayloadForType(limit int, rrType uint16) int {
maxLengthName, err := NewName([][]byte{
[]byte("AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA"),
[]byte("AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA"),
[]byte("AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA"),
[]byte("AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA"),
})
if err != nil {
panic(err)
}
{
n := 0
for _, label := range maxLengthName {
n += len(label) + 1
}
n += 1
if n != 255 {
panic("computeMaxEncodedPayload n != 255")
}
}
queryLimit := uint16(limit)
if int(queryLimit) != limit {
queryLimit = 0xffff
}
query := &Message{
Question: []Question{
{
Name: maxLengthName,
Type: rrType,
Class: ClassIN,
},
},
Additional: []RR{
{
Name: Name{},
Type: RRTypeOPT,
Class: queryLimit,
TTL: 0,
Data: []byte{},
},
},
}
resp, _ := responseFor(query, []domainSpec{{name: Name{[]byte{}}}})
low := 0
high := 32768
if chunkSize := payloadChunkSizeForType(rrType); chunkSize > 0 {
high = 256*chunkSize + 1
}
for low+1 < high {
mid := (low + high) / 2
resp.Answer, err = answersForPayload(query.Question[0], responseTTL, make([]byte, mid))
if err != nil {
panic(err)
}
buf, err := resp.WireFormat()
if err != nil {
panic(err)
}
if len(buf) <= limit {
low = mid
} else {
high = mid
}
}
return low
}
@@ -0,0 +1,31 @@
package xdns
import (
"errors"
"net"
"github.com/xtls/xray-core/common/serial"
"github.com/xtls/xray-core/transport/internet/finalmask"
)
type Resolver interface {
Addr() *net.UDPAddr
Read(p []byte) (int, error)
Send(p []byte)
Close()
}
func NewResolver(proto *serial.TypedMessage, dialer *finalmask.Dialer) (Resolver, error) {
config, err := proto.GetInstance()
if err != nil {
return nil, err
}
switch v := config.(type) {
case *TCPResolverProto:
return NewTCPResolver(v, dialer)
case *UDPResolverProto:
return NewUDPResolver(v, dialer)
default:
return nil, errors.New("unknown proto")
}
}
@@ -0,0 +1,143 @@
package xdns
import (
"encoding/binary"
"errors"
"io"
"sync"
"github.com/xtls/xray-core/common/net"
"github.com/xtls/xray-core/transport/internet/finalmask"
)
type TCPResolver struct {
dest net.Destination
dialer *finalmask.Dialer
conn net.Conn
tcpAddr *net.TCPAddr
udpAddr *net.UDPAddr
readCh chan []byte
closeCh chan struct{}
wg sync.WaitGroup
mu sync.Mutex
}
func NewTCPResolver(config *TCPResolverProto, dialer *finalmask.Dialer) (Resolver, error) {
dest, err := net.ParseDestination("tcp:" + config.Addr)
if err != nil {
return nil, err
}
r := &TCPResolver{
dest: dest,
dialer: dialer,
readCh: make(chan []byte),
closeCh: make(chan struct{}),
}
if err := r.dial(); err != nil {
r.Close()
return nil, err
}
return r, nil
}
func (r *TCPResolver) closed() bool {
select {
case <-r.closeCh:
return true
default:
return false
}
}
func (r *TCPResolver) dial() error {
if r.closed() {
return errors.New("closed")
}
if r.conn != nil {
return nil
}
conn, err := r.dialer.DialTCP(r.dest)
if err != nil {
return err
}
r.conn = conn
r.tcpAddr = conn.RemoteAddr().(*net.TCPAddr)
r.udpAddr = &net.UDPAddr{IP: r.tcpAddr.IP, Port: r.tcpAddr.Port}
r.wg.Add(1)
go r.recv(conn)
return nil
}
func (r *TCPResolver) recv(conn net.Conn) {
defer r.wg.Done()
var buf [4096]byte
for {
_, err := io.ReadFull(conn, buf[:2])
if err != nil {
break
}
n := binary.BigEndian.Uint16(buf[:2])
if n == 0 || n > 4096 {
io.CopyN(io.Discard, conn, int64(n))
continue
}
_, err = io.ReadFull(conn, buf[:n])
if err != nil {
break
}
p := pool4K.Get().([]byte)
copy(p, buf[:n])
select {
case <-r.closeCh:
pool4K.Put(p[:cap(p)])
case r.readCh <- p[:n]:
}
}
r.mu.Lock()
defer r.mu.Unlock()
_ = conn.Close()
r.conn = nil
}
func (r *TCPResolver) Addr() *net.UDPAddr {
return r.udpAddr
}
func (r *TCPResolver) Read(p []byte) (n int, err error) {
packet, ok := <-r.readCh
if ok {
n = copy(p, packet)
pool4K.Put(packet[:cap(packet)])
return n, nil
}
return 0, io.ErrClosedPipe
}
func (r *TCPResolver) Send(p []byte) {
r.mu.Lock()
defer r.mu.Unlock()
if r.dial() != nil {
return
}
_ = binary.Write(r.conn, binary.BigEndian, len(p))
_, _ = r.conn.Write(p)
}
func (r *TCPResolver) Close() {
r.mu.Lock()
defer r.mu.Unlock()
if r.closed() {
return
}
close(r.closeCh)
if r.conn != nil {
_ = r.conn.Close()
}
r.wg.Wait()
close(r.readCh)
}
@@ -0,0 +1,130 @@
package xdns
import (
"errors"
"io"
"sync"
"github.com/xtls/xray-core/common/net"
"github.com/xtls/xray-core/transport/internet/finalmask"
)
type UDPResolver struct {
dest net.Destination
dialer *finalmask.Dialer
conn net.PacketConn
udpAddr *net.UDPAddr
readCh chan []byte
closeCh chan struct{}
wg sync.WaitGroup
mu sync.Mutex
}
func NewUDPResolver(config *UDPResolverProto, dialer *finalmask.Dialer) (Resolver, error) {
dest, err := net.ParseDestination("udp:" + config.Addr)
if err != nil {
return nil, err
}
r := &UDPResolver{
dest: dest,
dialer: dialer,
readCh: make(chan []byte),
closeCh: make(chan struct{}),
}
if err := r.dial(); err != nil {
r.Close()
return nil, err
}
return r, nil
}
func (r *UDPResolver) closed() bool {
select {
case <-r.closeCh:
return true
default:
return false
}
}
func (r *UDPResolver) dial() error {
if r.closed() {
return errors.New("closed")
}
if r.conn != nil {
return nil
}
conn, err := r.dialer.DialUDP(r.dest)
if err != nil {
return err
}
r.conn = conn.(*net.PacketConnWrapper).PacketConn
r.udpAddr = conn.RemoteAddr().(*net.UDPAddr)
r.wg.Add(1)
go r.recv(conn.(*net.PacketConnWrapper).PacketConn)
return nil
}
func (r *UDPResolver) recv(conn net.PacketConn) {
defer r.wg.Done()
var buf [4096]byte
for {
n, _, err := conn.ReadFrom(buf[:])
if err != nil {
break
}
p := pool4K.Get().([]byte)
copy(p, buf[:n])
select {
case <-r.closeCh:
pool4K.Put(p[:cap(p)])
case r.readCh <- p[:n]:
}
}
r.mu.Lock()
defer r.mu.Unlock()
_ = conn.Close()
r.conn = nil
}
func (r *UDPResolver) Addr() *net.UDPAddr {
return r.udpAddr
}
func (r *UDPResolver) Read(p []byte) (n int, err error) {
packet, ok := <-r.readCh
if ok {
n = copy(p, packet)
pool4K.Put(packet[:cap(packet)])
return n, nil
}
return 0, io.ErrClosedPipe
}
func (r *UDPResolver) Send(p []byte) {
r.mu.Lock()
defer r.mu.Unlock()
if err := r.dial(); err != nil {
return
}
_, _ = r.conn.WriteTo(p, r.udpAddr)
}
func (r *UDPResolver) Close() {
r.mu.Lock()
defer r.mu.Unlock()
if r.closed() {
return
}
close(r.closeCh)
if r.conn != nil {
_ = r.conn.Close()
}
r.wg.Wait()
close(r.readCh)
}
+392
View File
@@ -0,0 +1,392 @@
package xdns
import (
"sort"
"sync"
"time"
"github.com/xtls/xray-core/common"
"golang.org/x/net/dns/dnsmessage"
)
const (
sendTTL = 4 * time.Second
)
type Resp struct {
msg dnsmessage.Message
domain *Domain
edns0 uint16
cap int
}
func NewResp(msg dnsmessage.Message, domain *Domain, edns0 uint16) *Resp {
if msg.Header.Response {
return &Resp{
msg: msg,
domain: domain,
}
}
size := min(max(int(edns0), 512), max(int(domain.edns0), 512))
left := size - 12 - int(msg.Questions[0].Name.Length) - 1 - 2 - 2
if edns0 > 0 {
left -= 1 + 2 + 2 + 4 + 2 + 0
}
cap := 0
switch msg.Questions[0].Type {
case dnsmessage.TypeA:
single := 2 + 2 + 2 + 4 + 2 + 4
n := left / single
if n > 255 {
n = 255
}
cap = 4*n - n - 1
case dnsmessage.TypeCNAME:
single := 2 + 2 + 2 + 4 + 2 + domain.lenMax
n := left / single
if n > 255 {
n = 255
}
cap = domain.cap*n - n - 1
case dnsmessage.TypeTXT:
left -= 2 + 2 + 2 + 4 + 2
single := 255
n := left / single
m := left % single
cap = 255*n - n
if m > 1 {
cap += m - 1
}
case dnsmessage.TypeAAAA:
single := 2 + 2 + 2 + 4 + 2 + 16
n := left / single
if n > 255 {
n = 255
}
cap = 16*n - n - 1
}
return &Resp{
msg: msg,
domain: domain,
edns0: edns0,
cap: cap,
}
}
func (r *Resp) Encode(encoded []byte, data []byte) []byte {
msg := r.msg
msg.Header = dnsmessage.Header{
ID: msg.Header.ID,
Response: true,
Authoritative: true,
RCode: dnsmessage.RCodeSuccess,
}
msg.Answers = nil
msg.Authorities = nil
msg.Additionals = nil
switch msg.Questions[0].Type {
case dnsmessage.TypeA:
fragN := 0
if len(data) > 0 {
fragN = 1
}
if (len(data) - (4 - 2)) > 0 {
fragN += (len(data) - (4 - 2)) / (4 - 1)
if (len(data)-(4-2))%(4-1) > 0 {
fragN++
}
}
for i := range fragN {
A := [4]byte{byte(i)}
if i == 0 {
A[1] = byte(fragN)
n := copy(A[2:], data)
data = data[n:]
} else {
n := copy(A[1:], data)
data = data[n:]
}
msg.Answers = append(msg.Answers, dnsmessage.Resource{
Header: dnsmessage.ResourceHeader{
Name: msg.Questions[0].Name,
Type: msg.Questions[0].Type,
Class: dnsmessage.ClassINET,
TTL: 60,
},
Body: &dnsmessage.AResource{A: A},
})
}
case dnsmessage.TypeCNAME:
fragN := 0
if len(data) > 0 {
fragN = 1
}
if (len(data) - (r.domain.cap - 2)) > 0 {
fragN += (len(data) - (r.domain.cap - 2)) / (r.domain.cap - 1)
if (len(data)-(r.domain.cap-2))%(r.domain.cap-1) > 0 {
fragN++
}
}
DATA := make([]byte, r.domain.cap)
for i := range fragN {
DATA[0] = byte(i)
if i == 0 {
DATA[1] = byte(fragN)
n := copy(DATA[2:], data)
data = data[n:]
msg.Answers = append(msg.Answers, dnsmessage.Resource{
Header: dnsmessage.ResourceHeader{
Name: msg.Questions[0].Name,
Type: msg.Questions[0].Type,
Class: dnsmessage.ClassINET,
TTL: 60,
},
Body: &dnsmessage.CNAMEResource{CNAME: r.domain.Encode(DATA[:2+n])},
})
} else {
n := copy(DATA[1:], data)
data = data[n:]
msg.Answers = append(msg.Answers, dnsmessage.Resource{
Header: dnsmessage.ResourceHeader{
Name: msg.Questions[0].Name,
Type: msg.Questions[0].Type,
Class: dnsmessage.ClassINET,
TTL: 60,
},
Body: &dnsmessage.CNAMEResource{CNAME: r.domain.Encode(DATA[:1+n])},
})
}
}
case dnsmessage.TypeTXT:
var txt []string
for len(data) > 0 {
size := min(len(data), 255)
txt = append(txt, string(data[:size]))
data = data[size:]
}
msg.Answers = append(msg.Answers, dnsmessage.Resource{
Header: dnsmessage.ResourceHeader{
Name: msg.Questions[0].Name,
Type: msg.Questions[0].Type,
Class: dnsmessage.ClassINET,
TTL: 60,
},
Body: &dnsmessage.TXTResource{TXT: txt},
})
case dnsmessage.TypeAAAA:
fragN := 0
if len(data) > 0 {
fragN = 1
}
if (len(data) - (16 - 2)) > 0 {
fragN += (len(data) - (16 - 2)) / (16 - 1)
if (len(data)-(16-2))%(16-1) > 0 {
fragN++
}
}
for i := range fragN {
AAAA := [16]byte{byte(i)}
if i == 0 {
AAAA[1] = byte(fragN)
n := copy(AAAA[2:], data)
data = data[n:]
} else {
n := copy(AAAA[1:], data)
data = data[n:]
}
msg.Answers = append(msg.Answers, dnsmessage.Resource{
Header: dnsmessage.ResourceHeader{
Name: msg.Questions[0].Name,
Type: msg.Questions[0].Type,
Class: dnsmessage.ClassINET,
TTL: 60,
},
Body: &dnsmessage.AAAAResource{AAAA: AAAA},
})
}
}
if r.edns0 > 0 {
msg.Additionals = append(msg.Additionals, dnsmessage.Resource{
Header: dnsmessage.ResourceHeader{
Name: dnsmessage.MustNewName("."),
Type: dnsmessage.TypeOPT,
Class: dnsmessage.Class(r.edns0),
TTL: 0,
},
Body: &dnsmessage.OPTResource{},
})
}
return common.Must2(msg.AppendPack(encoded[:0]))
}
func (r *Resp) Decode(decoded []byte) int {
decoded = decoded[:0]
msg := r.msg
if msg.Questions[0].Type == dnsmessage.TypeTXT {
if len(msg.Answers) == 1 && r.domain.IsDomain(msg.Answers[0].Header.Name) && msg.Answers[0].Header.Type == dnsmessage.TypeTXT {
for i := range msg.Answers[0].Body.(*dnsmessage.TXTResource).TXT {
decoded = append(decoded, msg.Answers[0].Body.(*dnsmessage.TXTResource).TXT[i]...)
}
}
return len(decoded)
} else {
var frags [][]byte
for i := range msg.Answers {
if !r.domain.IsDomain(msg.Answers[i].Header.Name) || msg.Answers[i].Header.Type != msg.Questions[0].Type {
continue
}
switch msg.Questions[0].Type {
case dnsmessage.TypeA:
frags = append(frags, msg.Answers[i].Body.(*dnsmessage.AResource).A[:])
case dnsmessage.TypeCNAME:
var decoded [255]byte
n := r.domain.Decode(&decoded, msg.Answers[i].Body.(*dnsmessage.CNAMEResource).CNAME)
if n == 0 {
continue
}
frags = append(frags, decoded[:n])
case dnsmessage.TypeAAAA:
frags = append(frags, msg.Answers[i].Body.(*dnsmessage.AAAAResource).AAAA[:])
}
}
sort.Slice(frags, func(i, j int) bool {
return frags[i][0] < frags[j][0]
})
if len(frags) < 1 || len(frags[0]) < 2 || int(frags[0][1]) > len(frags) {
return 0
}
decoded = append(decoded, frags[0][2:]...)
for i := range frags {
if i > 0 {
if frags[i][0] == frags[i-1][0] {
return 0
}
decoded = append(decoded, frags[i][1:]...)
}
}
return len(decoded)
}
}
type SendInfo struct {
stash chan []byte
ch chan []byte
deadline time.Time
}
type SendManager struct {
m map[ClientID]*SendInfo
ch chan struct{}
mu sync.Mutex
}
func NewSendManager() *SendManager {
m := &SendManager{
m: make(map[ClientID]*SendInfo),
ch: make(chan struct{}),
}
go m.gc()
return m
}
func (m *SendManager) closed() bool {
select {
case <-m.ch:
return true
default:
return false
}
}
func (m *SendManager) gc() {
ticker := time.NewTicker(sendTTL)
defer ticker.Stop()
for {
select {
case <-m.ch:
return
case now := <-ticker.C:
m.mu.Lock()
for key, info := range m.m {
if now.After(info.deadline) {
close(info.stash)
close(info.ch)
delete(m.m, key)
}
}
m.mu.Unlock()
ticker.Reset(sendTTL)
}
}
}
func (m *SendManager) Push(clientID ClientID, p []byte) {
m.mu.Lock()
defer m.mu.Unlock()
info := m.m[clientID]
if info == nil {
info = &SendInfo{
stash: make(chan []byte, 1),
ch: make(chan []byte, 128),
deadline: time.Now().Add(sendTTL),
}
m.m[clientID] = info
}
b := make([]byte, len(p))
copy(b, p)
select {
case info.ch <- b:
default:
}
}
func (m *SendManager) Stash(clientID ClientID, p []byte) {
m.mu.Lock()
defer m.mu.Unlock()
info := m.m[clientID]
if info == nil {
return
}
info.deadline = time.Now().Add(sendTTL)
select {
case info.stash <- p:
default:
}
}
func (m *SendManager) Pop(clientID ClientID) (chan []byte, chan []byte) {
m.mu.Lock()
defer m.mu.Unlock()
info := m.m[clientID]
if info == nil {
info = &SendInfo{
stash: make(chan []byte, 1),
ch: make(chan []byte, 128),
}
m.m[clientID] = info
}
info.deadline = time.Now().Add(sendTTL)
return info.ch, info.stash
}
func (m *SendManager) Close() {
m.mu.Lock()
defer m.mu.Unlock()
if m.closed() {
return
}
close(m.ch)
for key, info := range m.m {
close(info.stash)
close(info.ch)
delete(m.m, key)
}
}
+288 -415
View File
@@ -1,512 +1,385 @@
package xdns package xdns
import ( import (
"bytes"
"context" "context"
"encoding/binary"
go_errors "errors"
"io" "io"
"net"
"sync" "sync"
"time" "time"
"github.com/xtls/xray-core/common"
"github.com/xtls/xray-core/common/errors" "github.com/xtls/xray-core/common/errors"
"github.com/xtls/xray-core/transport/internet/finalmask" "github.com/xtls/xray-core/common/net"
"golang.org/x/net/dns/dnsmessage"
) )
const ( const (
idleTimeout = 10 * time.Second maxResponseDelay = time.Second
responseTTL = 60
maxResponseDelay = 1 * time.Second
) )
var ( type resp struct {
maxUDPPayload = 1280 - 40 - 8 msg dnsmessage.Message
maxEncodedPayloadTXT = computeMaxEncodedPayloadForType(maxUDPPayload, RRTypeTXT) addr net.Addr
maxEncodedPayloadA = computeMaxEncodedPayloadForType(maxUDPPayload, RRTypeA)
maxEncodedPayloadAAAA = computeMaxEncodedPayloadForType(maxUDPPayload, RRTypeAAAA)
)
func clientIDToAddr(clientID [8]byte) *net.UDPAddr {
ip := make(net.IP, 16)
copy(ip, []byte{0xfd, 0x00, 0, 0, 0, 0, 0, 0})
copy(ip[8:], clientID[:])
return &net.UDPAddr{
IP: ip,
}
} }
type record struct { type Rec struct {
Resp *Message resp *Resp
Addr net.Addr clientID ClientID
// ClientID [8]byte addr net.Addr
ClientAddr net.Addr
} }
type queue struct { type xdnsServer struct {
last time.Time
rrType uint16
queue chan []byte
stash chan []byte
}
type xdnsConnServer struct {
net.PacketConn net.PacketConn
domains []domainSpec domains []*Domain
fragManager *FragManager
sendManager *SendManager
ch chan *record readCh chan packet
readQueue chan *packet recCh chan *Rec
writeQueueMap map[string]*queue drCh chan resp
closeCh chan struct{}
closed bool wg sync.WaitGroup
mutex sync.Mutex mu sync.RWMutex
} }
func NewConnServer(c *Config, raw net.PacketConn) (net.PacketConn, error) { func NewServer(c *Config, raw net.PacketConn) (net.PacketConn, error) {
if len(c.Domains) == 0 { if len(c.Domains) == 0 {
return nil, errors.New("empty domains") return nil, errors.New("empty domains")
} }
domains := make([]domainSpec, 0, len(c.Domains)) domains := make([]*Domain, 0, len(c.Domains))
for _, domain := range c.Domains { for i := range c.Domains {
domain, err := parseDomainSpec(domain, "") 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 := NewDomain(c.Domains[i].Name, int(c.Domains[i].LenLimit), int(c.Domains[i].LabelLimit), types, uint16(c.Domains[i].Edns0))
if err != nil { if err != nil {
return nil, err return nil, err
} }
domains = append(domains, domain) domains = append(domains, domain)
} }
server := &xdnsServer{
conn := &xdnsConnServer{
PacketConn: raw, PacketConn: raw,
domains: domains, domains: domains,
fragManager: NewFragManager(),
sendManager: NewSendManager(),
ch: make(chan *record, 500), readCh: make(chan packet),
readQueue: make(chan *packet, 512), recCh: make(chan *Rec, 255),
writeQueueMap: make(map[string]*queue), drCh: make(chan resp),
closeCh: make(chan struct{}),
} }
go server.run()
go conn.clean() return server, nil
go conn.recvLoop()
go conn.sendLoop()
return conn, nil
} }
func (c *xdnsConnServer) clean() { func (c *xdnsServer) closed() bool {
f := func() bool { select {
c.mutex.Lock() case <-c.closeCh:
defer c.mutex.Unlock() return true
default:
if c.closed {
return true
}
now := time.Now()
for key, q := range c.writeQueueMap {
if now.Sub(q.last) >= idleTimeout {
close(q.queue)
close(q.stash)
delete(c.writeQueueMap, key)
}
}
return false return false
} }
for {
time.Sleep(idleTimeout / 2)
if f() {
return
}
}
} }
func (c *xdnsConnServer) ensureQueue(addr net.Addr) *queue { func (c *xdnsServer) decref(msg dnsmessage.Message, addr net.Addr) {
if c.closed {
return nil
}
q, ok := c.writeQueueMap[addr.String()]
if !ok {
q = &queue{
queue: make(chan []byte, 512),
stash: make(chan []byte, 1),
}
c.writeQueueMap[addr.String()] = q
}
q.last = time.Now()
return q
}
func (c *xdnsConnServer) stash(queue *queue, p []byte) {
c.mutex.Lock()
defer c.mutex.Unlock()
if c.closed {
return
}
select { select {
case queue.stash <- p: case c.drCh <- resp{msg: msg, addr: addr}:
default: default:
} }
} }
func (c *xdnsConnServer) recvLoop() { func (c *xdnsServer) read(buf []byte, addr net.Addr) {
var buf [finalmask.UDPSize]byte msg := dnsmessage.Message{}
if err := msg.Unpack(buf); err != nil {
return
}
if msg.Header.Response {
return
}
for { if msg.Header.OpCode != 0 {
if c.closed { msg.Header.Response = true
break msg.Header.RCode = dnsmessage.RCodeNotImplemented
} c.decref(msg, addr)
return
}
n, addr, err := c.PacketConn.ReadFrom(buf[:]) if len(msg.Questions) != 1 {
if err != nil { msg.Header.Response = true
if go_errors.Is(err, net.ErrClosed) { msg.Header.RCode = dnsmessage.RCodeFormatError
break c.decref(msg, addr)
return
}
opt := false
edns0 := uint16(0)
for i := range msg.Additionals {
if msg.Additionals[i].Header.Type == dnsmessage.TypeOPT {
if opt {
msg.Header.RCode = dnsmessage.RCodeFormatError
c.decref(msg, addr)
return
} }
continue opt = true
} edns0 = uint16(msg.Additionals[i].Header.Class)
if ver := (msg.Additionals[i].Header.TTL >> 16) & 0xFF; ver != 0 {
query, err := MessageFromWireFormat(buf[:n]) msg.Header.RCode = dnsmessage.RCodeSuccess
if err != nil { msg.Additionals[i].Header.TTL = 1 << 24
errors.LogDebug(context.Background(), addr, " xdns from wireformat err ", err) c.decref(msg, addr)
continue return
}
resp, payload := responseFor(&query, c.domains)
var clientID [8]byte
n = copy(clientID[:], payload)
payload = payload[n:]
if n == len(clientID) {
r := bytes.NewReader(payload)
for {
p, err := nextPacketServer(r)
if err != nil {
break
}
buf := make([]byte, len(p))
copy(buf, p)
select {
case c.readQueue <- &packet{
p: buf,
addr: clientIDToAddr(clientID),
}:
default:
errors.LogDebug(context.Background(), addr, " ", clientID, " mask read err queue full")
}
}
} else {
if resp != nil && resp.Rcode() == RcodeNoError {
resp.Flags |= RcodeNameError
}
}
if resp != nil {
select {
case c.ch <- &record{resp, addr, clientIDToAddr(clientID)}:
default:
errors.LogDebug(context.Background(), addr, " ", clientID, " mask read err record queue full")
} }
} }
} }
if opt {
if edns0 < 512 {
edns0 = 512
}
if edns0 > 4096 {
edns0 = 4096
}
}
errors.LogDebug(context.Background(), addr, " edns0 ", edns0, " buf ", len(buf), " ", msg.Questions[0].Type)
errors.LogDebug(context.Background(), "xdns closed") var domain *Domain
for i := range c.domains {
if c.domains[i].IsDomain(msg.Questions[0].Name) {
domain = c.domains[i]
break
}
}
if domain == nil {
msg.Header.Response = true
msg.Header.RCode = dnsmessage.RCodeNameError
c.decref(msg, addr)
return
}
if !domain.HasType(uint16(msg.Questions[0].Type)) {
msg.Header.Response = true
msg.Header.Authoritative = true
msg.Header.RCode = dnsmessage.RCodeSuccess
c.decref(msg, addr)
return
}
close(c.ch) var decoded [255]byte
close(c.readQueue) n := domain.Decode(&decoded, msg.Questions[0].Name)
if n < 9 {
msg.Header.Response = true
msg.Header.Authoritative = true
msg.Header.RCode = dnsmessage.RCodeSuccess
c.decref(msg, addr)
return
}
if TypeMap_[decoded[0]&3] != uint16(msg.Questions[0].Type) || (decoded[8]&0x3F != 3 && decoded[8]&0x3F != 8) || (decoded[8]&0x3F == 3 && n < 9+3+1) || (decoded[8]&0x3F == 8 && n != 9+8) {
msg.Header.Response = true
msg.Header.Authoritative = true
msg.Header.RCode = dnsmessage.RCodeSuccess
c.decref(msg, addr)
return
}
clientID := ClientIDFromRaw([8]byte(decoded[:8]))
c.mutex.Lock() r := NewResp(msg, domain, edns0)
defer c.mutex.Unlock() if r == nil {
msg.Header.Response = true
msg.Header.Authoritative = true
msg.Header.RCode = dnsmessage.RCodeSuccess
c.decref(msg, addr)
return
}
select {
case c.recCh <- &Rec{resp: r, clientID: clientID, addr: addr}:
default:
msg.Header.Response = true
msg.Header.Authoritative = true
msg.Header.RCode = dnsmessage.RCodeSuccess
c.decref(msg, addr)
}
c.closed = true if decoded[8]&0x3F == 8 {
for key, q := range c.writeQueueMap { return
close(q.queue) }
close(q.stash) p := pool4K.Get().([]byte)
delete(c.writeQueueMap, key) p = p[:0]
if decoded[8]&0xC0 == 0xC0 {
out := pool4K.Get().([]byte)
n := c.fragManager.Feed(out, FragKey{clientID: clientID, fragID: decoded[12]}, decoded[13], decoded[14], decoded[15:n])
pool4K.Put(p[:cap(p)])
if n > 0 {
p = out[:n]
} else {
pool4K.Put(out[:cap(out)])
return
}
} else {
p = append(p, decoded[12:n]...)
}
select {
case <-c.closeCh:
pool4K.Put(p[:cap(p)])
return
case c.readCh <- packet{p: p, addr: clientID.Addr()}:
return
} }
} }
func (c *xdnsConnServer) sendLoop() { func (c *xdnsServer) run() {
var nextRec *record c.wg.Add(1)
go c.recv()
c.wg.Add(1)
go c.send()
c.wg.Add(1)
go c.dr()
c.wg.Wait()
close(c.readCh)
close(c.recCh)
close(c.drCh)
c.fragManager.Close()
c.sendManager.Close()
}
func (c *xdnsServer) recv() {
defer c.wg.Done()
var buf [512]byte
for {
n, addr, err := c.PacketConn.ReadFrom(buf[:])
if err != nil {
if c.closed() {
return
}
errors.LogErrorInner(context.Background(), err, "recv err")
return
}
c.read(buf[:n], addr)
}
}
func (c *xdnsServer) send() {
defer c.wg.Done()
timer := time.NewTimer(maxResponseDelay)
timer.Stop()
var buf [4096]byte
var data [4096]byte
var nextRec *Rec
for { for {
var err error
rec := nextRec rec := nextRec
nextRec = nil nextRec = nil
if rec == nil { if rec == nil {
var ok bool select {
rec, ok = <-c.ch case rec = <-c.recCh:
if !ok { case <-c.closeCh:
break return
} }
} }
if rec.Resp.Rcode() == RcodeNoError && len(rec.Resp.Question) == 1 { ch, stash := c.sendManager.Pop(rec.clientID)
var payload bytes.Buffer left := rec.resp.cap
limit := maxEncodedPayloadForType(rec.Resp.Question[0].Type) timer.Reset(maxResponseDelay)
timer := time.NewTimer(maxResponseDelay) var ps [][]byte
for {
for { var p []byte
c.mutex.Lock() select {
q := c.ensureQueue(rec.ClientAddr) case p = <-stash:
if q == nil { default:
c.mutex.Unlock()
return
}
q.rrType = rec.Resp.Question[0].Type
c.mutex.Unlock()
var p []byte
select { select {
case p = <-q.stash: case p = <-stash:
case p = <-ch:
default: default:
select { select {
case p = <-q.stash: case p = <-stash:
case p = <-q.queue: case p = <-ch:
default: case <-timer.C:
select { case nextRec = <-c.recCh:
case p = <-q.stash:
case p = <-q.queue:
case <-timer.C:
case nextRec = <-c.ch:
}
} }
} }
}
timer.Reset(0) if len(p) == 0 {
break
if len(p) == 0 { }
timer.Reset(0)
left -= 2 + len(p)
if left < 0 {
if len(ps) == 0 {
errors.LogError(context.Background(), "err size ", len(p))
break break
} }
c.sendManager.Stash(rec.clientID, p)
limit -= 2 + len(p) break
if limit < 0 {
if payload.Len() == 0 {
errors.LogDebug(context.Background(), rec.Addr, " ", rec.ClientAddr, " xdns payload too large for rrtype ", rec.Resp.Question[0].Type, " ", len(p))
continue
}
c.stash(q, p)
break
}
// if len(p) > 65535 {
// panic(len(p))
// }
_ = binary.Write(&payload, binary.BigEndian, uint16(len(p)))
payload.Write(p)
} }
ps = append(ps, p)
}
timer.Stop()
timer.Stop() d := data[:0]
rec.Resp.Answer, err = answersForPayload(rec.Resp.Question[0], responseTTL, payload.Bytes()) for i := range ps {
if err != nil { l := len(ps[i])
errors.LogDebug(context.Background(), rec.Addr, " ", rec.ClientAddr, " xdns encode err ", err) if i == len(ps)-1 {
continue l |= 0xC000
} }
d = append(d, []byte{byte(l >> 8), byte(l)}...)
d = append(d, ps[i]...)
} }
_, _ = c.PacketConn.WriteTo(rec.resp.Encode(buf[:0], d), rec.addr)
}
}
buf, err := rec.Resp.WireFormat() func (c *xdnsServer) dr() {
if err != nil { defer c.wg.Done()
errors.LogDebug(context.Background(), rec.Addr, " ", rec.ClientAddr, " xdns wireformat err ", err)
continue
}
if len(buf) > maxUDPPayload { var buf [512]byte
errors.LogDebug(context.Background(), rec.Addr, " ", rec.ClientAddr, " xdns truncate ", len(buf)) for {
buf = buf[:maxUDPPayload] select {
buf[2] |= 0x02 case <-c.closeCh:
}
if c.closed {
return return
} case r := <-c.drCh:
_, _ = c.PacketConn.WriteTo(common.Must2(r.msg.AppendPack(buf[:0])), r.addr)
_, err = c.PacketConn.WriteTo(buf, rec.Addr)
if go_errors.Is(err, net.ErrClosed) {
c.closed = true
break
} }
} }
} }
func (c *xdnsConnServer) ReadFrom(p []byte) (n int, addr net.Addr, err error) { func (c *xdnsServer) ReadFrom(p []byte) (n int, addr net.Addr, err error) {
packet, ok := <-c.readQueue packet, ok := <-c.readCh
if !ok { if ok {
return 0, nil, net.ErrClosed n = copy(p, packet.p)
pool4K.Put(packet.p[:cap(packet.p)])
return n, packet.addr, nil
} }
if len(p) < len(packet.p) { return 0, nil, io.ErrClosedPipe
errors.LogDebug(context.Background(), packet.addr, " mask read err short buffer ", len(p), " ", len(packet.p))
return 0, packet.addr, nil
}
copy(p, packet.p)
return len(packet.p), packet.addr, nil
} }
func (c *xdnsConnServer) WriteTo(p []byte, addr net.Addr) (n int, err error) { func (c *xdnsServer) WriteTo(p []byte, addr net.Addr) (n int, err error) {
c.mutex.Lock() if c.closed() {
defer c.mutex.Unlock()
q := c.ensureQueue(addr)
if q == nil {
return 0, io.ErrClosedPipe return 0, io.ErrClosedPipe
} }
limit := maxEncodedPayloadForType(q.rrType) if len(p) == 0 || len(p) > 4096 {
if q.rrType == 0 { errors.LogError(context.Background(), "err size ", len(p))
limit = maxEncodedPayloadTXT return 0, errors.New("err size")
}
if len(p)+2 > limit {
errors.LogDebug(context.Background(), addr, " mask write err short write ", len(p), "+2 > ", limit)
return 0, nil
}
buf := make([]byte, len(p))
copy(buf, p)
select {
case q.queue <- buf:
return len(p), nil
default:
// errors.LogDebug(context.Background(), addr, " mask write err queue full")
return 0, nil
} }
c.sendManager.Push(ClientIDFromAddr(addr.(*net.UDPAddr)), p)
return len(p), nil
} }
func (c *xdnsConnServer) Close() error { func (c *xdnsServer) Close() error {
c.closed = true c.mu.Lock()
return c.PacketConn.Close() defer c.mu.Unlock()
if c.closed() {
return nil
}
close(c.closeCh)
_ = c.PacketConn.Close()
return nil
} }
func nextPacketServer(r *bytes.Reader) ([]byte, error) { func (c *xdnsServer) SetDeadline(t time.Time) error { return errors.New("not support") }
eof := func(err error) error {
if err == io.EOF {
err = io.ErrUnexpectedEOF
}
return err
}
for { func (c *xdnsServer) SetReadDeadline(t time.Time) error { return errors.New("not support") }
prefix, err := r.ReadByte()
if err != nil {
return nil, err
}
if prefix >= 224 {
paddingLen := prefix - 224
_, err := io.CopyN(io.Discard, r, int64(paddingLen))
if err != nil {
return nil, eof(err)
}
} else {
p := make([]byte, int(prefix))
_, err = io.ReadFull(r, p)
return p, eof(err)
}
}
}
func responseFor(query *Message, domains []domainSpec) (*Message, []byte) { func (c *xdnsServer) SetWriteDeadline(t time.Time) error { return errors.New("not support") }
resp := &Message{
ID: query.ID,
Flags: 0x8000,
Question: query.Question,
}
if query.Flags&0x8000 != 0 {
return nil, nil
}
payloadSize := 0
for _, rr := range query.Additional {
if rr.Type != RRTypeOPT {
continue
}
if len(resp.Additional) != 0 {
resp.Flags |= RcodeFormatError
return resp, nil
}
resp.Additional = append(resp.Additional, RR{
Name: Name{},
Type: RRTypeOPT,
Class: 4096,
TTL: 0,
Data: []byte{},
})
additional := &resp.Additional[0]
version := (rr.TTL >> 16) & 0xff
if version != 0 {
resp.Flags |= ExtendedRcodeBadVers & 0xf
additional.TTL = (ExtendedRcodeBadVers >> 4) << 24
return resp, nil
}
payloadSize = int(rr.Class)
}
if payloadSize < 512 {
payloadSize = 512
}
if len(query.Question) != 1 {
resp.Flags |= RcodeFormatError
return resp, nil
}
question := query.Question[0]
var (
prefix Name
ok bool
match domainSpec
)
for _, domain := range domains {
prefix, ok = question.Name.TrimSuffix(domain.name)
if ok {
match = domain
break
}
}
if !ok {
resp.Flags |= RcodeNameError
return resp, nil
}
resp.Flags |= 0x0400
if query.Opcode() != 0 {
resp.Flags |= RcodeNotImplemented
return resp, nil
}
switch question.Type {
case RRTypeTXT, RRTypeA, RRTypeAAAA:
default:
resp.Flags |= RcodeNameError
return resp, nil
}
if match.rrType != 0 && question.Type != match.rrType {
resp.Flags |= RcodeNameError
return resp, nil
}
encoded := bytes.ToUpper(bytes.Join(prefix, nil))
payload := make([]byte, base32Encoding.DecodedLen(len(encoded)))
n, err := base32Encoding.Decode(payload, encoded)
if err != nil {
resp.Flags |= RcodeNameError
return resp, nil
}
payload = payload[:n]
if payloadSize < maxUDPPayload {
resp.Flags |= RcodeFormatError
return resp, nil
}
return resp, payload
}
-80
View File
@@ -1,80 +0,0 @@
package xdns
import (
"strings"
"github.com/xtls/xray-core/common/errors"
)
type domainSpec struct {
name Name
rrType uint16
}
func rrTypeFromMethod(method string) (uint16, error) {
switch strings.ToLower(method) {
case "", "txt":
return RRTypeTXT, nil
case "a":
return RRTypeA, nil
case "aaaa":
return RRTypeAAAA, nil
default:
return 0, errors.New("unsupported method")
}
}
func parseDomainSpec(s string, defaultMethod string) (domainSpec, error) {
domainPart := s
method := ""
hasMethod := false
if i := strings.LastIndex(s, ":"); i >= 0 {
domainPart = s[:i]
method = s[i+1:]
hasMethod = true
} else if defaultMethod != "" {
method = defaultMethod
hasMethod = true
}
if domainPart == "" {
return domainSpec{}, errors.New("empty domain")
}
name, err := ParseName(domainPart)
if err != nil {
return domainSpec{}, err
}
rrType := uint16(0)
if hasMethod {
var err error
rrType, err = rrTypeFromMethod(method)
if err != nil {
return domainSpec{}, err
}
}
return domainSpec{
name: name,
rrType: rrType,
}, nil
}
func parseResolver(s string) (Name, string, uint16, error) {
head, server, ok := strings.Cut(s, "+udp://")
if !ok {
return nil, "", 0, errors.New("invalid resolver scheme")
}
if server == "" {
return nil, "", 0, errors.New("empty resolver server")
}
spec, err := parseDomainSpec(head, "txt")
if err != nil {
return nil, "", 0, err
}
return spec.name, server, spec.rrType, nil
}
@@ -0,0 +1,208 @@
package xdns
import (
"bytes"
"crypto/rand"
"fmt"
mrand "math/rand"
"testing"
"github.com/xtls/xray-core/common"
"golang.org/x/net/dns/dnsmessage"
)
func TestXxx(t *testing.T) {
m1 := dnsmessage.Message{
Questions: []dnsmessage.Question{
{
Name: dnsmessage.MustNewName("a.example.com."),
},
},
Answers: []dnsmessage.Resource{
{
Header: dnsmessage.ResourceHeader{
Name: dnsmessage.MustNewName("a.example.com."),
Type: dnsmessage.TypeA,
Class: dnsmessage.ClassINET,
TTL: 60,
Length: 16,
},
Body: &dnsmessage.AResource{A: [4]byte{127, 0, 0, 1}},
},
},
Additionals: []dnsmessage.Resource{
{
Header: dnsmessage.ResourceHeader{
Name: dnsmessage.MustNewName("."),
Type: dnsmessage.TypeOPT,
Class: 255,
TTL: 0,
Length: 16,
},
Body: &dnsmessage.OPTResource{},
},
},
}
p1, e1 := m1.Pack()
if e1 != nil {
t.Fatal(e1)
}
if !bytes.Equal(p1, []byte{
0, 0, 0, 0, 0, 1, 0, 1, 0, 0, 0, 1,
1, 97, 7, 101, 120, 97, 109, 112, 108, 101, 3, 99, 111, 109, 0,
0, 0,
0, 0,
192, 12,
0, 1,
0, 1,
0, 0, 0, 60,
0, 4,
127, 0, 0, 1,
0,
0, 41,
0, 255,
0, 0, 0, 0,
0, 0,
}) {
t.Fatal("!bytes.Equal")
}
domain, _ := NewDomain("a.example.com", 200, 1, []uint16{1}, 0)
fmt.Println(domain.cap, domain.lenMax)
lenMax := domain.lenMax
data := make([]byte, domain.cap)
msg := dnsmessage.Message{}
msg.Unpack(p1)
for range 3 {
msg.Answers = nil
msg.Authorities = nil
msg.Additionals = nil
n := mrand.Intn(255)
for range n {
msg.Answers = append(msg.Answers, dnsmessage.Resource{
Header: dnsmessage.ResourceHeader{
Name: dnsmessage.MustNewName("a.example.com."),
Type: dnsmessage.TypeA,
Class: dnsmessage.ClassINET,
TTL: 60,
},
Body: &dnsmessage.AResource{A: [4]byte{127, 0, 0, 1}},
})
}
if len(common.Must2(msg.Pack())) != 12+15+2+2+n*(2+2+2+4+2+4) {
t.Fatal("fatal a")
}
}
for range 3 {
msg.Answers = nil
msg.Authorities = nil
msg.Additionals = nil
n := mrand.Intn(255)
for range n {
common.Must2(rand.Read(data))
msg.Answers = append(msg.Answers, dnsmessage.Resource{
Header: dnsmessage.ResourceHeader{
Name: dnsmessage.MustNewName("a.example.com."),
Type: dnsmessage.TypeCNAME,
Class: dnsmessage.ClassINET,
TTL: 60,
},
Body: &dnsmessage.CNAMEResource{
CNAME: domain.Encode(data),
},
})
}
if len(common.Must2(msg.Pack())) > 12+15+2+2+n*(2+2+2+4+2+lenMax) {
t.Fatal("fatal cname")
}
}
for range 3 {
msg.Answers = nil
msg.Authorities = nil
msg.Additionals = nil
n := (mrand.Intn(2048) + 1024) % 2048
a := n / 255
b := n % 255
c := 0
var d [255]byte
var s []string
for range a {
s = append(s, string(d[:]))
}
if b > 0 {
c = 1
s = append(s, string(d[:b]))
}
msg.Answers = append(msg.Answers, dnsmessage.Resource{
Header: dnsmessage.ResourceHeader{
Name: dnsmessage.MustNewName("a.example.com."),
Type: dnsmessage.TypeTXT,
Class: dnsmessage.ClassINET,
TTL: 60,
},
Body: &dnsmessage.TXTResource{TXT: s},
})
if len(common.Must2(msg.Pack())) != 12+15+2+2+(2+2+2+4+2+n+n/255+c) {
t.Fatal("fatal txt")
}
}
for range 3 {
msg.Answers = nil
msg.Authorities = nil
msg.Additionals = nil
n := mrand.Intn(255)
for range n {
msg.Answers = append(msg.Answers, dnsmessage.Resource{
Header: dnsmessage.ResourceHeader{
Name: dnsmessage.MustNewName("a.example.com."),
Type: dnsmessage.TypeAAAA,
Class: dnsmessage.ClassINET,
TTL: 60,
},
Body: &dnsmessage.AAAAResource{AAAA: [16]byte{}},
})
}
if len(common.Must2(msg.Pack())) != 12+15+2+2+n*(2+2+2+4+2+16) {
t.Fatal("fatal aaaa")
}
}
}
func TestTXT(t *testing.T) {
txt := [][]byte{{}, {}}
for i := range 255 {
txt[0] = append(txt[0], byte(i))
}
txt[1] = []byte{255}
str := []string{}
for i := range txt {
str = append(str, string(txt[i]))
}
m1 := dnsmessage.Message{
Answers: []dnsmessage.Resource{
{
Header: dnsmessage.ResourceHeader{
Name: dnsmessage.MustNewName("."),
Type: dnsmessage.TypeTXT,
Class: dnsmessage.ClassINET,
TTL: 60,
},
Body: &dnsmessage.TXTResource{
TXT: str,
},
},
},
}
p1 := common.Must2(m1.Pack())
m2 := dnsmessage.Message{}
common.Must(m2.Unpack(p1))
if len(m2.Answers[0].Body.(*dnsmessage.TXTResource).TXT) != len(txt) {
t.Fatal("fatal txt")
}
for i := range txt {
if !bytes.Equal(txt[i], []byte(m2.Answers[0].Body.(*dnsmessage.TXTResource).TXT[i])) {
t.Fatal("fatal txt")
}
}
}
@@ -310,13 +310,6 @@ func (c *xicmpConnClient) Close() error {
_ = c.icmp4.Close() _ = c.icmp4.Close()
_ = c.icmp6.Close() _ = c.icmp6.Close()
c.wg.Wait() c.wg.Wait()
select {
case p := <-c.readCh:
if p.p != nil {
pool.Put(p.p)
}
default:
}
close(c.readCh) close(c.readCh)
return nil return nil
} }
@@ -329,13 +329,6 @@ func (c *xicmpConnServer) Close() error {
_ = c.icmp4.Close() _ = c.icmp4.Close()
_ = c.icmp6.Close() _ = c.icmp6.Close()
c.wg.Wait() c.wg.Wait()
select {
case p := <-c.readCh:
if p.p != nil {
pool.Put(p.p)
}
default:
}
close(c.readCh) close(c.readCh)
return nil return nil
} }
@@ -340,13 +340,6 @@ func (c *xicmpConnServer) Close() error {
_ = c.icmp4.Close() _ = c.icmp4.Close()
_ = c.icmp6.Close() _ = c.icmp6.Close()
c.wg.Wait() c.wg.Wait()
select {
case p := <-c.readCh:
if p.p != nil {
pool.Put(p.p)
}
default:
}
close(c.readCh) close(c.readCh)
return nil return nil
} }
+12
View File
@@ -3,6 +3,8 @@ package httpupgrade
import ( import (
"bufio" "bufio"
"context" "context"
"crypto/rand"
"encoding/base64"
"net/http" "net/http"
"net/url" "net/url"
"strings" "strings"
@@ -97,6 +99,16 @@ func dialhttpUpgrade(ctx context.Context, dest net.Destination, streamSettings *
req.Header.Set("Connection", "Upgrade") req.Header.Set("Connection", "Upgrade")
req.Header.Set("Upgrade", "websocket") req.Header.Set("Upgrade", "websocket")
// make a valid Sec-WebSocket-Key if not present
if len(req.Header.Values("Sec-WebSocket-Key")) == 0 {
var buf [16]byte
rand.Read(buf[:])
req.Header.Set("Sec-WebSocket-Key", base64.StdEncoding.EncodeToString(buf[:]))
}
if len(req.Header.Values("Sec-WebSocket-Version")) == 0 {
req.Header.Set("Sec-WebSocket-Version", "13")
}
err = req.Write(conn) err = req.Write(conn)
if err != nil { if err != nil {
return nil, err return nil, err
+7
View File
@@ -3,7 +3,9 @@ package httpupgrade
import ( import (
"bufio" "bufio"
"context" "context"
"crypto/sha1"
"crypto/tls" "crypto/tls"
"encoding/base64"
"io" "io"
"net/http" "net/http"
"strings" "strings"
@@ -81,6 +83,11 @@ func (s *server) upgrade(conn net.Conn) (stat.Connection, error) {
} }
resp.Header.Set("Connection", "Upgrade") resp.Header.Set("Connection", "Upgrade")
resp.Header.Set("Upgrade", "websocket") resp.Header.Set("Upgrade", "websocket")
// respond a valid Sec-WebSocket-Accept header if received a Sec-WebSocket-Key
if wsKey := req.Header.Get("Sec-WebSocket-Key"); wsKey != "" {
acceptKey := sha1.Sum([]byte(wsKey + "258EAFA5-E914-47DA-95CA-C5AB0DC85B11")) // magic number in RFC 6455
resp.Header.Set("Sec-WebSocket-Accept", base64.StdEncoding.EncodeToString(acceptKey[:]))
}
err = resp.Write(conn) err = resp.Write(conn)
if err != nil { if err != nil {
return nil, err return nil, err