mirror of
https://github.com/XTLS/Xray-core.git
synced 2026-10-01 21:45:44 +00:00
Compare commits
10
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
747b153333 | ||
|
|
94cd83ccb4 | ||
|
|
1c225a8041 | ||
|
|
3a6bdb19ba | ||
|
|
de02da553a | ||
|
|
4ec4fb8aab | ||
|
|
63de6135cb | ||
|
|
edd916b08e | ||
|
|
eb29a4e3de | ||
|
|
2db099b34b |
@@ -82,10 +82,19 @@ func (f *MphDomainMatcherFactory) BuildMatcher(rules []*DomainRule) (DomainMatch
|
||||
}
|
||||
g.Add(m, uint32(i))
|
||||
case *DomainRule_Geosite:
|
||||
err := loadSiteMatchers(v.Geosite, func(m strmatcher.Matcher) { g.Add(m, uint32(i)) })
|
||||
domains, err := loadSiteWithAttrs(v.Geosite.File, v.Geosite.Code, v.Geosite.Attrs)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
for j, d := range domains {
|
||||
domains[j] = nil // peak mem
|
||||
m, err := parseDomain(d)
|
||||
if err != nil {
|
||||
errors.LogError(context.Background(), "ignore invalid geosite entry in ", v.Geosite.File, ":", v.Geosite.Code, " at index ", j, ", ", err)
|
||||
continue
|
||||
}
|
||||
g.Add(m, uint32(i))
|
||||
}
|
||||
default:
|
||||
panic("unknown domain rule type")
|
||||
}
|
||||
@@ -99,12 +108,12 @@ func (f *MphDomainMatcherFactory) BuildMatcher(rules []*DomainRule) (DomainMatch
|
||||
return g, nil
|
||||
}
|
||||
|
||||
type CompactMphDomainMatcherFactory struct {
|
||||
type CompactDomainMatcherFactory struct {
|
||||
sync.Mutex
|
||||
shared *utils.WeakCacheMap[string, strmatcher.MphValueMatcher]
|
||||
shared *utils.WeakCacheMap[string, strmatcher.LinearAnyMatcher]
|
||||
}
|
||||
|
||||
func (f *CompactMphDomainMatcherFactory) getOrCreateFrom(rule *GeoSiteRule) (*strmatcher.MphValueMatcher, error) {
|
||||
func (f *CompactDomainMatcherFactory) getOrCreateFrom(rule *GeoSiteRule) (strmatcher.MatcherSet, error) {
|
||||
key := rule.File + ":" + rule.Code + "@" + rule.Attrs
|
||||
|
||||
f.Lock()
|
||||
@@ -116,23 +125,33 @@ func (f *CompactMphDomainMatcherFactory) getOrCreateFrom(rule *GeoSiteRule) (*st
|
||||
}
|
||||
errors.LogDebug(context.Background(), "geodata geosite matcher cache MISS ", key)
|
||||
|
||||
s := strmatcher.NewMphValueMatcher()
|
||||
if err := loadSiteMatchers(rule, func(m strmatcher.Matcher) { s.Add(m, 0) }); err != nil {
|
||||
s := strmatcher.NewLinearAnyMatcher()
|
||||
domains, err := loadSiteWithAttrs(rule.File, rule.Code, rule.Attrs)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := s.Build(); err != nil {
|
||||
return nil, err
|
||||
for i, d := range domains {
|
||||
domains[i] = nil // peak mem
|
||||
m, err := parseDomain(d)
|
||||
if err != nil {
|
||||
errors.LogError(context.Background(), "ignore invalid geosite entry in ", rule.File, ":", rule.Code, " at index ", i, ", ", err)
|
||||
continue
|
||||
}
|
||||
s.Add(m)
|
||||
}
|
||||
f.shared.Store(key, s)
|
||||
return s, nil
|
||||
return s, err
|
||||
}
|
||||
|
||||
// BuildMatcher implements DomainMatcherFactory.
|
||||
func (f *CompactMphDomainMatcherFactory) BuildMatcher(rules []*DomainRule) (DomainMatcher, error) {
|
||||
func (f *CompactDomainMatcherFactory) BuildMatcher(rules []*DomainRule) (DomainMatcher, error) {
|
||||
if len(rules) == 0 {
|
||||
return nil, errors.New("empty domain rule list")
|
||||
}
|
||||
compact := new(CompactMphDomainMatcher)
|
||||
compact := &CompactDomainMatcher{
|
||||
matchers: make([]strmatcher.MatcherSet, 0, len(rules)),
|
||||
values: make([]uint32, 0, len(rules)),
|
||||
}
|
||||
for i, r := range rules {
|
||||
switch v := r.Value.(type) {
|
||||
case *DomainRule_Custom:
|
||||
@@ -149,7 +168,8 @@ func (f *CompactMphDomainMatcherFactory) BuildMatcher(rules []*DomainRule) (Doma
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
compact.combiner.Add(m, uint32(i))
|
||||
compact.matchers = append(compact.matchers, m)
|
||||
compact.values = append(compact.values, uint32(i))
|
||||
default:
|
||||
panic("unknown domain rule type")
|
||||
}
|
||||
@@ -157,40 +177,37 @@ func (f *CompactMphDomainMatcherFactory) BuildMatcher(rules []*DomainRule) (Doma
|
||||
return compact, nil
|
||||
}
|
||||
|
||||
type CompactMphDomainMatcher struct {
|
||||
type CompactDomainMatcher struct {
|
||||
custom strmatcher.ValueMatcher
|
||||
combiner strmatcher.MphValueMatcherCombiner
|
||||
matchers []strmatcher.MatcherSet
|
||||
values []uint32
|
||||
}
|
||||
|
||||
// Match implements DomainMatcher.
|
||||
func (c *CompactMphDomainMatcher) Match(input string) []uint32 {
|
||||
result := c.combiner.Match(input)
|
||||
func (c *CompactDomainMatcher) Match(input string) []uint32 {
|
||||
var result []uint32
|
||||
if c.custom != nil {
|
||||
result = append(c.custom.Match(input), result...)
|
||||
result = append(result, c.custom.Match(input)...)
|
||||
}
|
||||
for i, m := range c.matchers {
|
||||
if m.MatchAny(input) {
|
||||
result = append(result, c.values[i])
|
||||
}
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
// MatchAny implements DomainMatcher.
|
||||
func (c *CompactMphDomainMatcher) MatchAny(input string) bool {
|
||||
func (c *CompactDomainMatcher) MatchAny(input string) bool {
|
||||
if c.custom != nil && c.custom.MatchAny(input) {
|
||||
return true
|
||||
}
|
||||
return c.combiner.MatchAny(input)
|
||||
}
|
||||
|
||||
// loadSiteMatchers calls add with a matcher for every domain of the geosite rule and logs the invalid ones.
|
||||
func loadSiteMatchers(rule *GeoSiteRule, add func(strmatcher.Matcher)) error {
|
||||
i := 0
|
||||
return loadSite(rule.File, rule.Code, rule.Attrs, func(t Domain_Type, value []byte) {
|
||||
m, err := parseDomain(&Domain{Type: t, Value: string(value)})
|
||||
if err != nil {
|
||||
errors.LogError(context.Background(), "ignore invalid geosite entry in ", rule.File, ":", rule.Code, " at index ", i, ", ", err)
|
||||
} else {
|
||||
add(m)
|
||||
for _, m := range c.matchers {
|
||||
if m.MatchAny(input) {
|
||||
return true
|
||||
}
|
||||
i++
|
||||
})
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func parseDomain(d *Domain) (strmatcher.Matcher, error) {
|
||||
@@ -214,7 +231,7 @@ func parseDomain(d *Domain) (strmatcher.Matcher, error) {
|
||||
func newDomainMatcherFactory() DomainMatcherFactory {
|
||||
switch runtime.GOOS {
|
||||
case "ios", "android":
|
||||
return &CompactMphDomainMatcherFactory{shared: utils.NewWeakCacheMap[string, strmatcher.MphValueMatcher]()}
|
||||
return &CompactDomainMatcherFactory{shared: utils.NewWeakCacheMap[string, strmatcher.LinearAnyMatcher]()}
|
||||
default:
|
||||
return &MphDomainMatcherFactory{shared: utils.NewWeakCacheMap[string, strmatcher.MphValueMatcher]()}
|
||||
}
|
||||
|
||||
@@ -4,7 +4,6 @@ import (
|
||||
"path/filepath"
|
||||
"reflect"
|
||||
"slices"
|
||||
"sync"
|
||||
"testing"
|
||||
|
||||
"github.com/xtls/xray-core/common/geodata/strmatcher"
|
||||
@@ -12,7 +11,7 @@ import (
|
||||
)
|
||||
|
||||
func TestCompactDomainMatcher_PreservesCustomRuleIndices(t *testing.T) {
|
||||
factory := &CompactMphDomainMatcherFactory{shared: utils.NewWeakCacheMap[string, strmatcher.MphValueMatcher]()}
|
||||
factory := &CompactDomainMatcherFactory{shared: utils.NewWeakCacheMap[string, strmatcher.LinearAnyMatcher]()}
|
||||
matcher, err := factory.BuildMatcher([]*DomainRule{
|
||||
{Value: &DomainRule_Custom{Custom: &Domain{Type: Domain_Full, Value: "example.com"}}},
|
||||
{Value: &DomainRule_Custom{Custom: &Domain{Type: Domain_Domain, Value: "example.com"}}},
|
||||
@@ -33,7 +32,7 @@ func TestCompactDomainMatcher_PreservesCustomRuleIndices(t *testing.T) {
|
||||
func TestCompactDomainMatcher_PreservesMixedRuleIndices(t *testing.T) {
|
||||
t.Setenv("xray.location.asset", filepath.Join("..", "..", "resources"))
|
||||
|
||||
factory := &CompactMphDomainMatcherFactory{shared: utils.NewWeakCacheMap[string, strmatcher.MphValueMatcher]()}
|
||||
factory := &CompactDomainMatcherFactory{shared: utils.NewWeakCacheMap[string, strmatcher.LinearAnyMatcher]()}
|
||||
matcher, err := factory.BuildMatcher([]*DomainRule{
|
||||
{Value: &DomainRule_Geosite{Geosite: &GeoSiteRule{File: DefaultGeoSiteDat, Code: "CN"}}},
|
||||
{Value: &DomainRule_Custom{Custom: &Domain{Type: Domain_Full, Value: "163.com"}}},
|
||||
@@ -73,76 +72,3 @@ func TestMphDomainMatcher_MatchReturnsDetachedSlice(t *testing.T) {
|
||||
t.Fatalf("Match() after caller mutation = %v, want %v", gotAgain, []uint32{0, 1})
|
||||
}
|
||||
}
|
||||
|
||||
// DNS sorts every Match result in place, so a matcher must never hand out a
|
||||
// slice it keeps, also when only its keyword or regex part matches.
|
||||
func TestDomainMatcher_MatchResultsCanBeSortedConcurrently(t *testing.T) {
|
||||
t.Setenv("xray.location.asset", filepath.Join("..", "..", "resources"))
|
||||
|
||||
rules := []*DomainRule{
|
||||
{Value: &DomainRule_Custom{Custom: &Domain{Type: Domain_Full, Value: "example.com"}}},
|
||||
{Value: &DomainRule_Custom{Custom: &Domain{Type: Domain_Domain, Value: "example.com"}}},
|
||||
{Value: &DomainRule_Custom{Custom: &Domain{Type: Domain_Substr, Value: "exam"}}},
|
||||
{Value: &DomainRule_Custom{Custom: &Domain{Type: Domain_Regex, Value: `^ex.*\.org$`}}},
|
||||
{Value: &DomainRule_Custom{Custom: &Domain{Type: Domain_Substr, Value: "exam"}}},
|
||||
{Value: &DomainRule_Geosite{Geosite: &GeoSiteRule{File: DefaultGeoSiteDat, Code: "CN"}}},
|
||||
{Value: &DomainRule_Custom{Custom: &Domain{Type: Domain_Full, Value: "only.full.test"}}},
|
||||
}
|
||||
cases := []struct {
|
||||
input string
|
||||
want []uint32
|
||||
}{
|
||||
{"example.com", []uint32{0, 1, 2, 4}},
|
||||
{"www.example.com", []uint32{1, 2, 4}},
|
||||
{"exam.net", []uint32{2, 4}}, // keyword part only
|
||||
{"example.org", []uint32{2, 3, 4}},
|
||||
{"163.com", []uint32{5}},
|
||||
{"www.163.com", []uint32{5}},
|
||||
{"only.full.test", []uint32{6}}, // full part only
|
||||
{"nomatch.test", nil},
|
||||
}
|
||||
factories := map[string]DomainMatcherFactory{
|
||||
"mph": &MphDomainMatcherFactory{shared: utils.NewWeakCacheMap[string, strmatcher.MphValueMatcher]()},
|
||||
"compact": &CompactMphDomainMatcherFactory{shared: utils.NewWeakCacheMap[string, strmatcher.MphValueMatcher]()},
|
||||
}
|
||||
for name, factory := range factories {
|
||||
t.Run(name, func(t *testing.T) {
|
||||
matcher, err := factory.BuildMatcher(rules)
|
||||
if err != nil {
|
||||
t.Fatalf("BuildMatcher() failed: %v", err)
|
||||
}
|
||||
for _, c := range cases {
|
||||
got := matcher.Match(c.input)
|
||||
if sorted := slices.Sorted(slices.Values(got)); !slices.Equal(sorted, c.want) {
|
||||
t.Fatalf("Match(%q) = %v, want %v", c.input, sorted, c.want)
|
||||
}
|
||||
got = got[:cap(got)]
|
||||
for j := range got {
|
||||
got[j] = ^uint32(0)
|
||||
}
|
||||
if again := slices.Sorted(slices.Values(matcher.Match(c.input))); !slices.Equal(again, c.want) {
|
||||
t.Fatalf("Match(%q) after caller mutation = %v, want %v", c.input, again, c.want)
|
||||
}
|
||||
}
|
||||
|
||||
var wg sync.WaitGroup
|
||||
for range 8 {
|
||||
wg.Add(1)
|
||||
go func() {
|
||||
defer wg.Done()
|
||||
for range 500 {
|
||||
for _, c := range cases {
|
||||
got := matcher.Match(c.input)
|
||||
slices.Sort(got)
|
||||
if !slices.Equal(got, c.want) {
|
||||
t.Errorf("Match(%q) = %v, want %v", c.input, got, c.want)
|
||||
return
|
||||
}
|
||||
}
|
||||
}
|
||||
}()
|
||||
}
|
||||
wg.Wait()
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
+62
-213
@@ -5,14 +5,11 @@ import (
|
||||
"bytes"
|
||||
"io"
|
||||
"runtime"
|
||||
"slices"
|
||||
"strings"
|
||||
"unicode/utf8"
|
||||
|
||||
"github.com/xtls/xray-core/common/errors"
|
||||
"github.com/xtls/xray-core/common/platform/filesystem"
|
||||
|
||||
"google.golang.org/protobuf/encoding/protowire"
|
||||
"google.golang.org/protobuf/proto"
|
||||
)
|
||||
|
||||
@@ -55,56 +52,17 @@ func loadIP(file, code string) ([]*CIDR, error) {
|
||||
return geoip.Cidr, nil
|
||||
}
|
||||
|
||||
// loadSite calls fn, in file order, with the type and value of every domain of the geosite code
|
||||
// that has all the "@"-separated attrs. It decodes the entry while reading the file instead of
|
||||
// unmarshalling it into a []*Domain, so value is only valid during fn.
|
||||
func loadSite(file, code, attrs string, fn func(Domain_Type, []byte)) error {
|
||||
runtime.GC() // peak mem
|
||||
r, err := filesystem.OpenAsset(file)
|
||||
func loadSite(file, code string) ([]*Domain, error) {
|
||||
bs, err := loadFile(file, code)
|
||||
if err != nil {
|
||||
return errors.New("failed to open ", file).Base(err)
|
||||
return nil, err
|
||||
}
|
||||
defer r.Close()
|
||||
br := bufio.NewReaderSize(r, 64*1024)
|
||||
n, err := seek(br, []byte(code))
|
||||
if err != nil {
|
||||
return errors.New("failed to load code ", code, " from ", file).Base(err)
|
||||
defer runtime.GC() // peak mem
|
||||
var geosite GeoSite
|
||||
if err := proto.Unmarshal(bs, &geosite); err != nil {
|
||||
return nil, errors.New("error unmarshal Site in ", file, ":", code).Base(err)
|
||||
}
|
||||
loadErr := func(err error) error {
|
||||
if err == io.EOF {
|
||||
err = io.ErrUnexpectedEOF
|
||||
}
|
||||
return errors.New("failed to load code ", code, " from ", file).Base(err)
|
||||
}
|
||||
unmarshalErr := func(err error) error {
|
||||
return errors.New("error unmarshal Site in ", file, ":", code).Base(err)
|
||||
}
|
||||
d := newSiteDecoder(attrs, fn)
|
||||
for n > 0 {
|
||||
w, err := br.Peek(min(n, br.Size()))
|
||||
if err != nil {
|
||||
return loadErr(err)
|
||||
}
|
||||
used, err := d.decode(w, len(w) < n)
|
||||
if err != nil {
|
||||
return unmarshalErr(err)
|
||||
}
|
||||
if used == 0 {
|
||||
break // a field longer than the buffer
|
||||
}
|
||||
br.Discard(used)
|
||||
n -= used
|
||||
}
|
||||
if n > 0 {
|
||||
w := make([]byte, n)
|
||||
if _, err := io.ReadFull(br, w); err != nil {
|
||||
return loadErr(err)
|
||||
}
|
||||
if _, err := d.decode(w, false); err != nil {
|
||||
return unmarshalErr(err)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
return geosite.Domain, nil
|
||||
}
|
||||
|
||||
func decodeVarint(br *bufio.Reader) (uint64, error) {
|
||||
@@ -124,63 +82,68 @@ func decodeVarint(br *bufio.Reader) (uint64, error) {
|
||||
}
|
||||
|
||||
func find(r io.Reader, code []byte, readBody bool) ([]byte, error) {
|
||||
br := bufio.NewReaderSize(r, 64*1024)
|
||||
bodyL, err := seek(br, code)
|
||||
if err != nil || !readBody {
|
||||
return nil, err
|
||||
}
|
||||
out := make([]byte, bodyL)
|
||||
if _, err := io.ReadFull(br, out); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
|
||||
// seek advances br to the body of the entry for code and returns the body length.
|
||||
func seek(br *bufio.Reader, code []byte) (int, error) {
|
||||
codeL := len(code)
|
||||
if codeL == 0 {
|
||||
return 0, errors.New("empty code")
|
||||
return nil, errors.New("empty code")
|
||||
}
|
||||
|
||||
br := bufio.NewReaderSize(r, 64*1024)
|
||||
need := 2 + codeL // TODO: if code too long
|
||||
prefixBuf := make([]byte, need)
|
||||
|
||||
for {
|
||||
if _, err := br.ReadByte(); err != nil {
|
||||
return 0, err
|
||||
return nil, err
|
||||
}
|
||||
|
||||
x, err := decodeVarint(br)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
return nil, err
|
||||
}
|
||||
bodyL := int(x)
|
||||
if bodyL <= 0 {
|
||||
return 0, errors.New("invalid body length: ", bodyL)
|
||||
return nil, errors.New("invalid body length: ", bodyL)
|
||||
}
|
||||
|
||||
// Peek no more than the buffer holds: a code longer than the buffer cannot match a single
|
||||
// length byte anyway, so a short peek only skips it, as base find (io.ReadFull) does.
|
||||
prefix, err := br.Peek(min(bodyL, need, br.Size()))
|
||||
if err != nil {
|
||||
if err == io.EOF && len(prefix) > 0 {
|
||||
err = io.ErrUnexpectedEOF // as io.ReadFull
|
||||
prefixL := bodyL
|
||||
if prefixL > need {
|
||||
prefixL = need
|
||||
}
|
||||
prefix := prefixBuf[:prefixL]
|
||||
if _, err := io.ReadFull(br, prefix); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
match := false
|
||||
if bodyL >= need {
|
||||
if int(prefix[1]) == codeL && bytes.Equal(prefix[2:need], code) {
|
||||
if !readBody {
|
||||
return nil, nil
|
||||
}
|
||||
match = true
|
||||
}
|
||||
return 0, err
|
||||
}
|
||||
if bodyL >= need && len(prefix) >= need && int(prefix[1]) == codeL && bytes.Equal(prefix[2:], code) {
|
||||
return bodyL, nil
|
||||
|
||||
remain := bodyL - prefixL
|
||||
if match {
|
||||
out := make([]byte, bodyL)
|
||||
copy(out, prefix)
|
||||
if remain > 0 {
|
||||
if _, err := io.ReadFull(br, out[prefixL:]); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
if _, err := br.Discard(bodyL); err != nil {
|
||||
return 0, err
|
||||
|
||||
if remain > 0 {
|
||||
if _, err := br.Discard(remain); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// AttributeMatcher, HasAttrMatcher, AllAttrsMatcher and NewAllAttrsMatcher are the exported
|
||||
// attribute helpers that have been part of this package's API since #5814. The streaming loader
|
||||
// above filters attributes itself without building a *Domain, so it does not use them, but they
|
||||
// are kept for external callers. Their behaviour is unchanged.
|
||||
|
||||
type AttributeMatcher interface {
|
||||
Match(*Domain) bool
|
||||
}
|
||||
@@ -222,137 +185,23 @@ func NewAllAttrsMatcher(attrs string) AttributeMatcher {
|
||||
return m
|
||||
}
|
||||
|
||||
var errInvalidUTF8 = errors.New("string field contains invalid UTF-8")
|
||||
func loadSiteWithAttrs(file, code, attrs string) ([]*Domain, error) {
|
||||
domains, err := loadSite(file, code)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
type siteDecoder struct {
|
||||
want []string
|
||||
has []bool
|
||||
fn func(Domain_Type, []byte)
|
||||
}
|
||||
matcher := NewAllAttrsMatcher(attrs)
|
||||
if matcher == nil {
|
||||
return domains, nil
|
||||
}
|
||||
|
||||
func newSiteDecoder(attrs string, fn func(Domain_Type, []byte)) *siteDecoder {
|
||||
d := &siteDecoder{fn: fn}
|
||||
if attrs != "" {
|
||||
d.want = strings.Split(attrs, "@")
|
||||
d.has = make([]bool, len(d.want))
|
||||
filtered := make([]*Domain, 0, len(domains))
|
||||
for _, d := range domains {
|
||||
if matcher.Match(d) {
|
||||
filtered = append(filtered, d)
|
||||
}
|
||||
}
|
||||
return d
|
||||
}
|
||||
|
||||
// decode walks the whole fields at the start of b, a part of an encoded GeoSite (see geodat.proto),
|
||||
// calls fn for every domain that has all attrs and returns how many bytes it used. A field cut off
|
||||
// by the end of b is an error unless more is set. It accepts and rejects what proto.Unmarshal does.
|
||||
func (d *siteDecoder) decode(b []byte, more bool) (int, error) {
|
||||
used := 0
|
||||
for used < len(b) {
|
||||
f, n, err := consumeField(b[used:])
|
||||
if err == io.ErrUnexpectedEOF && more {
|
||||
break
|
||||
}
|
||||
if err != nil {
|
||||
return used, err
|
||||
}
|
||||
used += n
|
||||
if f.typ != protowire.BytesType {
|
||||
continue
|
||||
}
|
||||
switch f.num {
|
||||
case 1: // code
|
||||
if !utf8.Valid(f.v) {
|
||||
return used, errInvalidUTF8
|
||||
}
|
||||
case 2: // domain
|
||||
t, value, err := decodeDomain(f.v, d.want, d.has)
|
||||
if err != nil {
|
||||
return used, err
|
||||
}
|
||||
if !slices.Contains(d.has, false) {
|
||||
d.fn(t, value)
|
||||
}
|
||||
}
|
||||
}
|
||||
return used, nil
|
||||
}
|
||||
|
||||
// decodeDomain decodes an encoded Domain and sets has[i] if one of its attributes has the key want[i].
|
||||
func decodeDomain(b []byte, want []string, has []bool) (t Domain_Type, value []byte, err error) {
|
||||
clear(has)
|
||||
for len(b) > 0 {
|
||||
f, n, err := consumeField(b)
|
||||
if err != nil {
|
||||
return 0, nil, err
|
||||
}
|
||||
b = b[n:]
|
||||
switch {
|
||||
case f.num == 1 && f.typ == protowire.VarintType: // type
|
||||
t = Domain_Type(f.x)
|
||||
case f.num == 2 && f.typ == protowire.BytesType: // value
|
||||
if !utf8.Valid(f.v) {
|
||||
return 0, nil, errInvalidUTF8
|
||||
}
|
||||
value = f.v
|
||||
case f.num == 3 && f.typ == protowire.BytesType: // attribute
|
||||
key, err := decodeAttributeKey(f.v)
|
||||
if err != nil {
|
||||
return 0, nil, err
|
||||
}
|
||||
for i, w := range want {
|
||||
if string(key) == w {
|
||||
has[i] = true
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
return t, value, nil
|
||||
}
|
||||
|
||||
// decodeAttributeKey returns the key of an encoded Domain.Attribute.
|
||||
func decodeAttributeKey(b []byte) ([]byte, error) {
|
||||
var key []byte
|
||||
for len(b) > 0 {
|
||||
f, n, err := consumeField(b)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
b = b[n:]
|
||||
if f.num == 1 && f.typ == protowire.BytesType {
|
||||
if !utf8.Valid(f.v) {
|
||||
return nil, errInvalidUTF8
|
||||
}
|
||||
key = f.v
|
||||
}
|
||||
}
|
||||
return key, nil
|
||||
}
|
||||
|
||||
type protoField struct {
|
||||
num protowire.Number
|
||||
typ protowire.Type
|
||||
v []byte // payload of a length-delimited field
|
||||
x uint64 // value of a varint field
|
||||
}
|
||||
|
||||
// consumeField parses the first field of an encoded message and returns it with its length.
|
||||
func consumeField(b []byte) (protoField, int, error) {
|
||||
num, typ, n := protowire.ConsumeTag(b)
|
||||
if n < 0 {
|
||||
return protoField{}, 0, protowire.ParseError(n)
|
||||
}
|
||||
if num > protowire.MaxValidNumber {
|
||||
return protoField{}, 0, errors.New("invalid field number ", num)
|
||||
}
|
||||
f := protoField{num: num, typ: typ}
|
||||
var m int
|
||||
switch typ {
|
||||
case protowire.BytesType:
|
||||
f.v, m = protowire.ConsumeBytes(b[n:])
|
||||
case protowire.VarintType:
|
||||
f.x, m = protowire.ConsumeVarint(b[n:])
|
||||
default:
|
||||
m = protowire.ConsumeFieldValue(num, typ, b[n:])
|
||||
}
|
||||
if m < 0 {
|
||||
return protoField{}, 0, protowire.ParseError(m)
|
||||
}
|
||||
return f, n + m, nil
|
||||
return filtered, nil
|
||||
}
|
||||
|
||||
@@ -1,283 +0,0 @@
|
||||
package geodata
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"slices"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"google.golang.org/protobuf/encoding/protowire"
|
||||
"google.golang.org/protobuf/proto"
|
||||
)
|
||||
|
||||
type siteEntry struct {
|
||||
Type Domain_Type
|
||||
Value string
|
||||
}
|
||||
|
||||
// unmarshalSite is what loadSite used to do: proto.Unmarshal, then keep the domains that have all attrs.
|
||||
func unmarshalSite(b []byte, attrs string) ([]siteEntry, error) {
|
||||
var site GeoSite
|
||||
if err := proto.Unmarshal(b, &site); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
var entries []siteEntry
|
||||
for _, d := range site.Domain {
|
||||
ok := true
|
||||
for _, key := range strings.Split(attrs, "@") {
|
||||
ok = ok && (attrs == "" || slices.ContainsFunc(d.Attribute, func(a *Domain_Attribute) bool { return a.Key == key }))
|
||||
}
|
||||
if ok {
|
||||
entries = append(entries, siteEntry{d.Type, d.Value})
|
||||
}
|
||||
}
|
||||
return entries, nil
|
||||
}
|
||||
|
||||
func checkDecodeSite(t *testing.T, name string, b []byte, attrs string) {
|
||||
t.Helper()
|
||||
want, wantErr := unmarshalSite(b, attrs)
|
||||
var got []siteEntry
|
||||
_, err := newSiteDecoder(attrs, func(typ Domain_Type, value []byte) {
|
||||
got = append(got, siteEntry{typ, string(value)})
|
||||
}).decode(b, false)
|
||||
if (err == nil) != (wantErr == nil) {
|
||||
t.Fatalf("%s@%s: error %v, proto.Unmarshal: %v", name, attrs, err, wantErr)
|
||||
}
|
||||
if err == nil && !slices.Equal(got, want) {
|
||||
t.Fatalf("%s@%s: got %v, want %v", name, attrs, got, want)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDecodeSiteMatchesUnmarshal(t *testing.T) {
|
||||
bs, err := os.ReadFile(filepath.Join("..", "..", "resources", DefaultGeoSiteDat))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
for len(bs) > 0 {
|
||||
num, typ, n := protowire.ConsumeTag(bs)
|
||||
if n < 0 || num != 1 || typ != protowire.BytesType {
|
||||
t.Fatal("unexpected GeoSiteList field")
|
||||
}
|
||||
entry, m := protowire.ConsumeBytes(bs[n:])
|
||||
if m < 0 {
|
||||
t.Fatal(protowire.ParseError(m))
|
||||
}
|
||||
bs = bs[n+m:]
|
||||
|
||||
var site GeoSite
|
||||
if err := proto.Unmarshal(entry, &site); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
queries := []string{"", "none"}
|
||||
for _, d := range site.Domain {
|
||||
for _, a := range d.Attribute {
|
||||
if !slices.Contains(queries, a.Key) {
|
||||
queries = append(queries, a.Key, a.Key+"@none")
|
||||
}
|
||||
}
|
||||
}
|
||||
for _, attrs := range queries {
|
||||
checkDecodeSite(t, site.Code, entry, attrs)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestDecodeSiteUnusualEncodings(t *testing.T) {
|
||||
field := func(num protowire.Number, v []byte) []byte {
|
||||
return protowire.AppendBytes(protowire.AppendTag(nil, num, protowire.BytesType), v)
|
||||
}
|
||||
typ := func(v Domain_Type) []byte {
|
||||
return protowire.AppendVarint(protowire.AppendTag(nil, 1, protowire.VarintType), uint64(v))
|
||||
}
|
||||
value := func(s string) []byte { return field(2, []byte(s)) }
|
||||
attr := func(keys ...string) []byte {
|
||||
var b []byte
|
||||
for _, k := range keys {
|
||||
b = append(b, field(1, []byte(k))...)
|
||||
}
|
||||
return field(3, b)
|
||||
}
|
||||
domain := func(fields ...[]byte) []byte { return field(2, slices.Concat(fields...)) }
|
||||
unknown := protowire.AppendFixed32(protowire.AppendTag(nil, 9, protowire.Fixed32Type), 1)
|
||||
|
||||
for name, b := range map[string][]byte{
|
||||
"unknown field": domain(typ(Domain_Full), unknown, value("example.com")),
|
||||
"repeated value": domain(value("a.com"), typ(Domain_Full), value("b.com")),
|
||||
"repeated type": domain(typ(Domain_Full), value("a.com"), typ(Domain_Regex)),
|
||||
"repeated key": domain(value("a.com"), attr("cn", "ads")),
|
||||
"type as bytes": domain(field(1, []byte("x")), value("a.com")),
|
||||
"no value": domain(typ(Domain_Domain), attr("cn")),
|
||||
"truncated": domain(typ(Domain_Full), value("example.com"))[:10],
|
||||
"invalid utf8": domain(value("example.\xff")),
|
||||
"invalid key": domain(value("a.com"), attr("\xff")),
|
||||
"bad field": protowire.AppendVarint(protowire.AppendTag(nil, protowire.MaxValidNumber+1, protowire.VarintType), 1),
|
||||
"stray end group": protowire.AppendTag(nil, 5, protowire.EndGroupType),
|
||||
} {
|
||||
for _, attrs := range []string{"", "cn", "ads", "cn@ads"} {
|
||||
checkDecodeSite(t, name, b, attrs)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestLoadSiteReadsInPieces covers what real lists never do: an entry far longer than the read
|
||||
// buffer, with a field longer than the buffer in the middle, and a file cut short.
|
||||
func TestLoadSiteReadsInPieces(t *testing.T) {
|
||||
site := &GeoSite{Code: "BIG"}
|
||||
for i := range 5000 {
|
||||
d := &Domain{Type: Domain_Domain, Value: strings.Repeat("x", i%40) + ".example.com"}
|
||||
if i%3 == 0 {
|
||||
d.Attribute = []*Domain_Attribute{{Key: "cn"}}
|
||||
}
|
||||
if i == 2500 {
|
||||
d = &Domain{Type: Domain_Regex, Value: strings.Repeat("a", 100_000)}
|
||||
}
|
||||
site.Domain = append(site.Domain, d)
|
||||
}
|
||||
list := &GeoSiteList{Entry: []*GeoSite{{Code: "SMALL", Domain: []*Domain{{Type: Domain_Full, Value: "a.com"}}}, site}}
|
||||
bs, err := proto.Marshal(list)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
entry, err := proto.Marshal(site)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
dir := t.TempDir()
|
||||
t.Setenv("xray.location.asset", dir)
|
||||
write := func(b []byte) {
|
||||
if err := os.WriteFile(filepath.Join(dir, "big.dat"), b, 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
for _, attrs := range []string{"", "cn"} {
|
||||
want, _ := unmarshalSite(entry, attrs)
|
||||
var got []siteEntry
|
||||
write(bs)
|
||||
err := loadSite("big.dat", "BIG", attrs, func(typ Domain_Type, value []byte) {
|
||||
got = append(got, siteEntry{typ, string(value)})
|
||||
})
|
||||
if err != nil || !slices.Equal(got, want) {
|
||||
t.Fatalf("attrs %q: %d entries, want %d, error %v", attrs, len(got), len(want), err)
|
||||
}
|
||||
for _, cut := range []int{30_000, len(bs) - 150_000, len(bs) - 1} {
|
||||
write(bs[:cut])
|
||||
if err := loadSite("big.dat", "BIG", attrs, func(Domain_Type, []byte) {}); err == nil {
|
||||
t.Fatalf("file cut at %d of %d: no error", cut, len(bs))
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// oneEntryGeoSiteFile wraps an encoded GeoSite as a one-entry GeoSiteList, the file loadSite reads.
|
||||
func oneEntryGeoSiteFile(entry []byte) []byte {
|
||||
return protowire.AppendBytes(protowire.AppendTag(nil, 1, protowire.BytesType), entry)
|
||||
}
|
||||
|
||||
// TestLoadSiteWindowedMatchesSingleShot checks that the windowed reader in loadSite (its Peek/Discard
|
||||
// loop, the more-break when a field is cut by a window edge, the used==0 fallback for a field longer
|
||||
// than the buffer, and the tail path) reaches exactly the same result as decoding the whole entry at
|
||||
// once, for a category several 64 KiB windows long, valid and then mutated near a window edge and
|
||||
// early in the file: same error-or-not, and the same emitted (type, value) sequence when both accept.
|
||||
func TestLoadSiteWindowedMatchesSingleShot(t *testing.T) {
|
||||
const window = 64 * 1024
|
||||
site := &GeoSite{Code: "BIG"}
|
||||
for i := range 12000 { // ~250 KiB, four windows
|
||||
d := &Domain{Type: Domain_Domain, Value: fmt.Sprintf("host%d.%s.example.com", i, strings.Repeat("y", i%30))}
|
||||
if i%3 == 0 {
|
||||
d.Attribute = []*Domain_Attribute{{Key: "cn"}}
|
||||
}
|
||||
site.Domain = append(site.Domain, d)
|
||||
}
|
||||
// a field longer than the buffer, straddling the third window, to force the used==0 fallback
|
||||
site.Domain = slices.Insert(site.Domain, 8000, &Domain{Type: Domain_Regex, Value: strings.Repeat("a", 90_000)})
|
||||
entry, err := proto.Marshal(site)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
dir := t.TempDir()
|
||||
t.Setenv("xray.location.asset", dir)
|
||||
|
||||
// mutations of the encoded entry: unchanged, a byte flipped at several offsets (early windows and
|
||||
// either side of a window edge), and truncations at the same places.
|
||||
type mut struct {
|
||||
name string
|
||||
make func([]byte) []byte
|
||||
}
|
||||
muts := []mut{{"valid", func(b []byte) []byte { return b }}}
|
||||
for _, off := range []int{3, 40, 4000, window - 1, window, window + 1, 2*window - 2, 2 * window} {
|
||||
if off < len(entry) {
|
||||
off := off
|
||||
muts = append(muts, mut{fmt.Sprintf("flip@%d", off), func(b []byte) []byte {
|
||||
c := slices.Clone(b)
|
||||
c[off] ^= 0xff
|
||||
return c
|
||||
}})
|
||||
muts = append(muts, mut{fmt.Sprintf("cut@%d", off), func(b []byte) []byte { return slices.Clone(b[:off]) }})
|
||||
}
|
||||
}
|
||||
|
||||
for _, attrs := range []string{"", "cn"} {
|
||||
for _, m := range muts {
|
||||
e := m.make(entry)
|
||||
// single-shot reference: decode the whole entry in one call
|
||||
var want []siteEntry
|
||||
_, wantErr := newSiteDecoder(attrs, func(typ Domain_Type, value []byte) {
|
||||
want = append(want, siteEntry{typ, string(value)})
|
||||
}).decode(e, false)
|
||||
// windowed: loadSite reads the file 64 KiB at a time
|
||||
if err := os.WriteFile(filepath.Join(dir, "w.dat"), oneEntryGeoSiteFile(e), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
var got []siteEntry
|
||||
gotErr := loadSite("w.dat", "BIG", attrs, func(typ Domain_Type, value []byte) {
|
||||
got = append(got, siteEntry{typ, string(value)})
|
||||
})
|
||||
if (gotErr == nil) != (wantErr == nil) {
|
||||
t.Fatalf("%s attrs=%q: windowed err %v, single-shot err %v", m.name, attrs, gotErr, wantErr)
|
||||
}
|
||||
if gotErr == nil && !slices.Equal(got, want) {
|
||||
t.Fatalf("%s attrs=%q: windowed got %d entries, single-shot %d", m.name, attrs, len(got), len(want))
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestLoadSiteLongCode covers a geosite entry whose code is longer than the 64 KiB read buffer. seek
|
||||
// must skip it (find compares a single length byte, so it never matches such a code) and still find a
|
||||
// later entry, and looking the long code up must fail cleanly, like a missing code, not panic.
|
||||
func TestLoadSiteLongCode(t *testing.T) {
|
||||
longCode := strings.Repeat("Z", 70000)
|
||||
list := &GeoSiteList{Entry: []*GeoSite{
|
||||
{Code: "FIRST", Domain: []*Domain{{Type: Domain_Full, Value: "first.com"}}},
|
||||
{Code: longCode, Domain: []*Domain{{Type: Domain_Full, Value: "huge.com"}}},
|
||||
{Code: "AFTER", Domain: []*Domain{{Type: Domain_Domain, Value: "after.com"}}},
|
||||
}}
|
||||
bs, err := proto.Marshal(list)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
dir := t.TempDir()
|
||||
t.Setenv("xray.location.asset", dir)
|
||||
if err := os.WriteFile(filepath.Join(dir, "lc.dat"), bs, 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
collect := func(code string) ([]siteEntry, error) {
|
||||
var got []siteEntry
|
||||
err := loadSite("lc.dat", code, "", func(typ Domain_Type, value []byte) {
|
||||
got = append(got, siteEntry{typ, string(value)})
|
||||
})
|
||||
return got, err
|
||||
}
|
||||
if got, err := collect("FIRST"); err != nil || !slices.Equal(got, []siteEntry{{Domain_Full, "first.com"}}) {
|
||||
t.Fatalf("FIRST: %v %v", got, err)
|
||||
}
|
||||
if got, err := collect("AFTER"); err != nil || !slices.Equal(got, []siteEntry{{Domain_Domain, "after.com"}}) {
|
||||
t.Fatalf("AFTER (past the oversized entry): %v %v", got, err)
|
||||
}
|
||||
if _, err := collect(longCode); err == nil {
|
||||
t.Fatal("oversized code: expected a not-found error, got nil")
|
||||
}
|
||||
}
|
||||
@@ -52,9 +52,7 @@ func (g *MphIndexMatcher) Add(matcher Matcher) uint32 {
|
||||
func (g *MphIndexMatcher) Build() error {
|
||||
if g.mph != nil {
|
||||
runtime.GC() // peak mem
|
||||
if err := g.mph.Build(); err != nil {
|
||||
return err
|
||||
}
|
||||
g.mph.Build()
|
||||
}
|
||||
runtime.GC() // peak mem
|
||||
if g.ac != nil {
|
||||
@@ -66,17 +64,23 @@ func (g *MphIndexMatcher) Build() error {
|
||||
|
||||
// Match implements IndexMatcher.Match.
|
||||
func (g *MphIndexMatcher) Match(input string) []uint32 {
|
||||
var result []uint32
|
||||
result := make([][]uint32, 0, 5)
|
||||
if g.mph != nil {
|
||||
result = g.mph.Match(input) // a new slice, returned without another copy
|
||||
if matches := g.mph.Match(input); len(matches) > 0 {
|
||||
result = append(result, matches)
|
||||
}
|
||||
}
|
||||
if g.ac != nil {
|
||||
result = append(result, g.ac.Match(input)...)
|
||||
if matches := g.ac.Match(input); len(matches) > 0 {
|
||||
result = append(result, matches)
|
||||
}
|
||||
}
|
||||
if g.regex != nil {
|
||||
result = append(result, g.regex.Match(input)...)
|
||||
if matches := g.regex.Match(input); len(matches) > 0 {
|
||||
result = append(result, matches)
|
||||
}
|
||||
}
|
||||
return result
|
||||
return CompositeMatches(result)
|
||||
}
|
||||
|
||||
// MatchAny implements IndexMatcher.MatchAny.
|
||||
|
||||
@@ -78,10 +78,6 @@ func TestMphIndexMatcher(t *testing.T) {
|
||||
Input: "example.com",
|
||||
Output: []uint32{10, 4},
|
||||
},
|
||||
{
|
||||
Input: "apis.org",
|
||||
Output: []uint32{2, 6},
|
||||
},
|
||||
}
|
||||
matcherGroup := NewMphIndexMatcher()
|
||||
for _, rule := range rules {
|
||||
@@ -91,13 +87,8 @@ func TestMphIndexMatcher(t *testing.T) {
|
||||
}
|
||||
matcherGroup.Build()
|
||||
for _, test := range cases {
|
||||
m := matcherGroup.Match(test.Input)
|
||||
if !reflect.DeepEqual(m, test.Output) {
|
||||
t.Error("unexpected output: ", m, " for test case ", test)
|
||||
}
|
||||
clear(m) // the caller owns the result, so this must not change the next one
|
||||
if m := matcherGroup.Match(test.Input); !reflect.DeepEqual(m, test.Output) {
|
||||
t.Error("unexpected output after clearing the previous one: ", m, " for test case ", test)
|
||||
t.Error("unexpected output: ", m, " for test case ", test)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,440 +1,231 @@
|
||||
package strmatcher
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"cmp"
|
||||
"encoding/binary"
|
||||
"errors"
|
||||
"math"
|
||||
"slices"
|
||||
"math/bits"
|
||||
"runtime"
|
||||
"sort"
|
||||
"strings"
|
||||
"unsafe"
|
||||
)
|
||||
|
||||
// Flags of a level1 slot, stored above the record offset.
|
||||
const (
|
||||
mphDomain = 1 << 31 // matches the pattern and its subdomains
|
||||
mphFull = 1 << 30 // matches the pattern only
|
||||
mphParent = 1 << 29 // matches subdomains only, from a pattern with a leading dot
|
||||
mphOffMask = mphParent - 1
|
||||
)
|
||||
// PrimeRK is the prime base used in Rabin-Karp algorithm.
|
||||
const PrimeRK = 16777619
|
||||
|
||||
// Kinds of an added pattern, indexes of mphKinds.
|
||||
const (
|
||||
mphKindFull = iota
|
||||
mphKindParent
|
||||
mphKindDomain
|
||||
)
|
||||
|
||||
// mphKinds are the slot flags in the order Match reports their values.
|
||||
var mphKinds = [...]uint32{mphFull, mphParent, mphDomain}
|
||||
|
||||
// mphMultipliers are odd multipliers for the suffix hash. Build moves to the next one if two patterns collide.
|
||||
var mphMultipliers = [...]uint64{0x9e3779b97f4a7c15, 0xc2b2ae3d27d4eb4f, 0x165667b19e3779f9, 0x27d4eb2f165667c5}
|
||||
|
||||
var (
|
||||
errMphCollision = errors.New("strmatcher: suffix hash collision in MphMatcherGroup")
|
||||
errMphBuilt = errors.New("strmatcher: MphMatcherGroup is already built")
|
||||
)
|
||||
|
||||
type mphEntry struct {
|
||||
off uint32 // pattern start in buf
|
||||
value uint32
|
||||
n uint32 // pattern length
|
||||
kind uint8
|
||||
// RollingHash calculates the rolling murmurHash of given string based on a provided suffix hash.
|
||||
func RollingHash(hash uint32, input string) uint32 {
|
||||
for i := len(input) - 1; i >= 0; i-- {
|
||||
hash = hash*PrimeRK + uint32(input[i])
|
||||
}
|
||||
return hash
|
||||
}
|
||||
|
||||
// MphMatcherGroup is an implementation of MatcherGroup for Full and Domain matchers.
|
||||
// Each distinct pattern is stored once as a record in arena: its length (255 means a uvarint length follows),
|
||||
// its bytes and, if the group holds more than one distinct value, its values. A minimal perfect hash table
|
||||
// built with hash, displace and compress (http://cmph.sourceforge.net/papers/esa09.pdf) maps a pattern to its
|
||||
// record. Patterns are hashed from the right, so one pass over the input hashes all its parent domains.
|
||||
type MphMatcherGroup struct {
|
||||
arena string
|
||||
level0 []uint16 // bucket -> seed
|
||||
level1 []uint32 // slot -> flags | record offset
|
||||
fp []uint8 // slot -> low byte of its pattern's hash, rejects most misses without reading arena
|
||||
n0, n1 uint32
|
||||
mul uint64 // multiplier of the suffix hash
|
||||
single uint32 // the only value if !multi
|
||||
multi bool
|
||||
// MemHash is the hash function used by go map, it utilizes available hardware instructions(behaves
|
||||
// as aeshash if aes instruction is available).
|
||||
// With different seed, each MemHash<seed> performs as distinct hash functions.
|
||||
func MemHash(seed uint32, input string) uint32 {
|
||||
return uint32(strhash(unsafe.Pointer(&input), uintptr(seed))) // nosemgrep
|
||||
}
|
||||
|
||||
buf []byte // build only, patterns in Add order
|
||||
entries []mphEntry
|
||||
const (
|
||||
mphMatchTypeCount = 2 // Full and Domain
|
||||
)
|
||||
|
||||
type mphRuleInfo struct {
|
||||
rollingHash uint32
|
||||
matchers [mphMatchTypeCount][]uint32
|
||||
}
|
||||
|
||||
// MphMatcherGroup is an implementation of MatcherGroup.
|
||||
// It implements Rabin-Karp algorithm and minimal perfect hash table for Full and Domain matcher.
|
||||
type MphMatcherGroup struct {
|
||||
patterns string // All rule patterns concatenated
|
||||
patternOffs []uint32 // RuleIdx -> patterns[patternOffs[i]:patternOffs[i+1]], index 0 reserved for failed lookup
|
||||
values []uint32 // All registered matcher values concatenated
|
||||
valueOffs []uint32 // RuleIdx -> values[valueOffs[i]:valueOffs[i+1]] (Full Matcher takes precedence)
|
||||
level0 []uint32 // RollingHash & Mask -> seed for Memhash
|
||||
level0Mask uint32 // Mask restricting RollingHash to 0 ~ len(level0)
|
||||
level1 []uint32 // Memhash<seed> & Mask -> stored index for rules
|
||||
level1Mask uint32 // Mask for restricting Memhash<seed> to 0 ~ len(level1)
|
||||
rules []string // RuleIdx -> pattern string, only used for building
|
||||
ruleInfos *map[string]mphRuleInfo
|
||||
}
|
||||
|
||||
func NewMphMatcherGroup() *MphMatcherGroup {
|
||||
return new(MphMatcherGroup)
|
||||
return &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.
|
||||
func (g *MphMatcherGroup) AddFullMatcher(matcher FullMatcher, value uint32) {
|
||||
g.add(matcher.Pattern(), mphKindFull, value)
|
||||
pattern := strings.ToLower(matcher.Pattern())
|
||||
g.addPattern(0, "", pattern, matcher.Type(), value)
|
||||
}
|
||||
|
||||
// AddDomainMatcher implements MatcherGroupForDomain.
|
||||
func (g *MphMatcherGroup) AddDomainMatcher(matcher DomainMatcher, value uint32) {
|
||||
g.add(matcher.Pattern(), mphKindDomain, value)
|
||||
pattern := strings.ToLower(matcher.Pattern())
|
||||
hash := g.addPattern(0, "", pattern, matcher.Type(), value) // For full domain match
|
||||
g.addPattern(hash, pattern, ".", matcher.Type(), value) // For partial domain match
|
||||
}
|
||||
|
||||
func (g *MphMatcherGroup) add(pattern string, kind uint8, value uint32) {
|
||||
if g.arena != "" {
|
||||
panic(errMphBuilt)
|
||||
}
|
||||
pattern = strings.ToLower(pattern)
|
||||
off := uint32(len(g.buf))
|
||||
g.buf = append(g.buf, pattern...)
|
||||
g.entries = append(g.entries, mphEntry{off: off, value: value, n: uint32(len(pattern)), kind: kind})
|
||||
if len(pattern) > 0 && pattern[0] == '.' {
|
||||
// ".x" has always matched "*.x" as well, so it also gets a parent-only record for "x"
|
||||
g.entries = append(g.entries, mphEntry{off: off + 1, value: value, n: uint32(len(pattern) - 1), kind: mphKindParent})
|
||||
func (g *MphMatcherGroup) addPattern(suffixHash uint32, suffixPattern string, pattern string, matcherType Type, value uint32) uint32 {
|
||||
fullPattern := pattern + suffixPattern
|
||||
info, found := (*g.ruleInfos)[fullPattern]
|
||||
if !found {
|
||||
info = mphRuleInfo{rollingHash: RollingHash(suffixHash, pattern)}
|
||||
g.rules = append(g.rules, fullPattern)
|
||||
}
|
||||
info.matchers[matcherType] = append(info.matchers[matcherType], value)
|
||||
(*g.ruleInfos)[fullPattern] = info
|
||||
return info.rollingHash
|
||||
}
|
||||
|
||||
func (g *MphMatcherGroup) key(i uint32) []byte {
|
||||
e := &g.entries[i]
|
||||
return g.buf[e.off : e.off+e.n]
|
||||
}
|
||||
|
||||
// Build builds the hash table. It must be called once, after the last Add.
|
||||
// Build builds a minimal perfect hash table for insert rules.
|
||||
// Algorithm used: Hash, displace, and compress. See http://cmph.sourceforge.net/papers/esa09.pdf
|
||||
func (g *MphMatcherGroup) Build() error {
|
||||
if g.arena != "" {
|
||||
return errMphBuilt
|
||||
ruleCount := len(*g.ruleInfos)
|
||||
g.level0 = make([]uint32, nextPow2(ruleCount/4))
|
||||
g.level0Mask = uint32(len(g.level0) - 1)
|
||||
g.level1 = make([]uint32, nextPow2(ruleCount))
|
||||
g.level1Mask = uint32(len(g.level1) - 1)
|
||||
|
||||
// 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])
|
||||
}
|
||||
if uint64(len(g.buf)) > math.MaxUint32 {
|
||||
g.patterns = strings.Join(g.rules, "")
|
||||
if uint64(len(g.patterns)) > math.MaxUint32 || uint64(valueCount) > math.MaxUint32 {
|
||||
return errors.New("too many rules for MphMatcherGroup")
|
||||
}
|
||||
recs := g.writeRecords()
|
||||
if len(g.arena) > mphOffMask {
|
||||
return errors.New("too many rules for MphMatcherGroup")
|
||||
}
|
||||
hashes := make([]uint64, len(recs))
|
||||
for _, mul := range mphMultipliers {
|
||||
for i, rec := range recs {
|
||||
hashes[i] = mphMix(mphHash(mul, g.recKey(rec)))
|
||||
}
|
||||
g.mul = mul
|
||||
if err := g.place(recs, hashes); err != errMphCollision {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return errMphCollision
|
||||
}
|
||||
g.patternOffs = make([]uint32, len(g.rules)+1)
|
||||
g.values = make([]uint32, 0, valueCount)
|
||||
g.valueOffs = make([]uint32, len(g.rules)+1)
|
||||
|
||||
// writeRecords writes one record per distinct pattern to arena and returns flags | offset of each.
|
||||
func (g *MphMatcherGroup) writeRecords() []uint32 {
|
||||
g.multi = false
|
||||
if len(g.entries) > 0 {
|
||||
g.single = g.entries[0].value
|
||||
for _, e := range g.entries {
|
||||
if e.value != g.single {
|
||||
g.multi = true
|
||||
break
|
||||
}
|
||||
}
|
||||
// Create buckets based on all rule's rolling hash
|
||||
buckets := make([][]uint32, len(g.level0))
|
||||
for ruleIdx := 1; ruleIdx < len(g.rules); ruleIdx++ { // Traverse rules starting from index 1 (0 reserved for failed lookup)
|
||||
ruleInfo := (*g.ruleInfos)[g.rules[ruleIdx]]
|
||||
bucketIdx := ruleInfo.rollingHash & g.level0Mask
|
||||
buckets[bucketIdx] = append(buckets[bucketIdx], uint32(ruleIdx))
|
||||
g.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))
|
||||
}
|
||||
// Equal patterns become neighbours in Add order, so their values keep their priority
|
||||
order := make([]uint32, len(g.entries))
|
||||
for i := range order {
|
||||
order[i] = uint32(i)
|
||||
}
|
||||
slices.SortFunc(order, func(a, b uint32) int {
|
||||
return cmp.Or(bytes.Compare(g.key(a), g.key(b)), cmp.Compare(a, b))
|
||||
})
|
||||
g.rules = nil
|
||||
g.ruleInfos = nil // Set ruleInfos nil to release memory
|
||||
runtime.GC() // peak mem
|
||||
|
||||
size := len(g.buf) + len(g.entries) + 2
|
||||
if g.multi {
|
||||
size += 3 * len(g.entries)
|
||||
// Sort buckets in descending order with respect to each bucket's size
|
||||
bucketIdxs := make([]int, len(buckets))
|
||||
for bucketIdx := range buckets {
|
||||
bucketIdxs[bucketIdx] = bucketIdx
|
||||
}
|
||||
arena := make([]byte, 0, size)
|
||||
recs := make([]uint32, 0, len(order))
|
||||
var vals [len(mphKinds)][]uint32
|
||||
for i := 0; i < len(order); {
|
||||
k := g.key(order[i])
|
||||
for t := range vals {
|
||||
vals[t] = vals[t][:0]
|
||||
}
|
||||
for ; i < len(order) && bytes.Equal(g.key(order[i]), k); i++ {
|
||||
e := &g.entries[order[i]]
|
||||
if !slices.Contains(vals[e.kind], e.value) {
|
||||
vals[e.kind] = append(vals[e.kind], e.value)
|
||||
}
|
||||
}
|
||||
rec := uint32(len(arena))
|
||||
if len(k) < 255 {
|
||||
arena = append(arena, byte(len(k)))
|
||||
} else {
|
||||
arena = binary.AppendUvarint(append(arena, 255), uint64(len(k)))
|
||||
}
|
||||
arena = append(arena, k...)
|
||||
for t, v := range vals {
|
||||
if len(v) == 0 {
|
||||
continue
|
||||
}
|
||||
rec |= mphKinds[t]
|
||||
if g.multi {
|
||||
arena = binary.AppendUvarint(arena, uint64(len(v)))
|
||||
for _, x := range v {
|
||||
arena = binary.AppendUvarint(arena, uint64(x))
|
||||
sort.Slice(bucketIdxs, func(i, j int) bool { return len(buckets[bucketIdxs[i]]) > len(buckets[bucketIdxs[j]]) })
|
||||
|
||||
// Exercise Hash, Displace, and Compress algorithm to construct minimal perfect hash table
|
||||
occupied := make([]bool, len(g.level1)) // Whether a second-level hash has been already used
|
||||
hashedBucket := make([]uint32, 0, 4) // Second-level hashes for each rule in a specific bucket
|
||||
for _, bucketIdx := range bucketIdxs {
|
||||
bucket := buckets[bucketIdx]
|
||||
hashedBucket = hashedBucket[:0]
|
||||
seed := uint32(0)
|
||||
for len(hashedBucket) != len(bucket) {
|
||||
for _, ruleIdx := range bucket {
|
||||
memHash := MemHash(seed, g.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)
|
||||
}
|
||||
}
|
||||
recs = append(recs, rec)
|
||||
}
|
||||
// Lookups may point one byte past a pattern, and an empty group needs a record at offset 0 for empty slots
|
||||
arena = append(arena, 0)
|
||||
if len(recs) == 0 {
|
||||
arena = append(arena, 0)
|
||||
}
|
||||
g.buf, g.entries = nil, nil
|
||||
if cap(arena)-len(arena) > len(arena)/32 {
|
||||
arena = slices.Clone(arena)
|
||||
}
|
||||
g.arena = unsafe.String(unsafe.SliceData(arena), len(arena)) // arena is not written after this
|
||||
return recs
|
||||
}
|
||||
|
||||
// place fills level0, level1 and fp: records are bucketed by hash, and each bucket, largest first, gets
|
||||
// the first seed that puts all its records in free slots.
|
||||
func (g *MphMatcherGroup) place(recs []uint32, hashes []uint64) error {
|
||||
r := len(recs)
|
||||
n0, n1 := max(1, r/3), max(1, r+r/99)
|
||||
g.n0, g.n1 = uint32(n0), uint32(n1)
|
||||
g.level0 = make([]uint16, n0)
|
||||
g.level1 = make([]uint32, n1)
|
||||
g.fp = make([]uint8, n1)
|
||||
|
||||
start := make([]uint32, n0+1)
|
||||
for _, h := range hashes {
|
||||
start[g.bucket(h)+1]++
|
||||
}
|
||||
for b := range n0 {
|
||||
start[b+1] += start[b]
|
||||
}
|
||||
members := make([]uint32, r)
|
||||
fill := slices.Clone(start[:n0])
|
||||
for i, h := range hashes {
|
||||
b := g.bucket(h)
|
||||
members[fill[b]] = uint32(i)
|
||||
fill[b]++
|
||||
}
|
||||
fill = nil
|
||||
buckets := make([]uint32, n0)
|
||||
for b := range buckets {
|
||||
buckets[b] = uint32(b)
|
||||
}
|
||||
slices.SortStableFunc(buckets, func(a, b uint32) int {
|
||||
return cmp.Compare(start[b+1]-start[b], start[a+1]-start[a])
|
||||
})
|
||||
|
||||
occupied := make([]uint64, (n1+63)/64)
|
||||
var slots []uint32
|
||||
next:
|
||||
for _, b := range buckets {
|
||||
m := members[start[b]:start[b+1]]
|
||||
if len(m) == 0 {
|
||||
break
|
||||
}
|
||||
for i := range m {
|
||||
for j := range i {
|
||||
if hashes[m[i]] == hashes[m[j]] {
|
||||
return errMphCollision // no seed can separate them
|
||||
}
|
||||
}
|
||||
}
|
||||
search:
|
||||
for seed := range math.MaxUint16 + 1 {
|
||||
slots = slots[:0]
|
||||
for _, ri := range m {
|
||||
s := g.slot(hashes[ri], uint16(seed))
|
||||
if occupied[s/64]&(1<<(s%64)) != 0 || slices.Contains(slots, s) {
|
||||
continue search
|
||||
}
|
||||
slots = append(slots, s)
|
||||
}
|
||||
for k, ri := range m {
|
||||
s := slots[k]
|
||||
occupied[s/64] |= 1 << (s % 64)
|
||||
g.level1[s] = recs[ri]
|
||||
g.fp[s] = uint8(hashes[ri])
|
||||
}
|
||||
g.level0[b] = uint16(seed)
|
||||
continue next
|
||||
}
|
||||
return errors.New("strmatcher: no seed found for a bucket in MphMatcherGroup")
|
||||
g.level0[bucketIdx] = seed // Displacement value for this bucket
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// mphHash is the suffix hash of s, taken from the right: the hash of s[i:] is the state after reading s[i].
|
||||
func mphHash(mul uint64, s string) uint64 {
|
||||
h := uint64(0)
|
||||
for i := len(s) - 1; i >= 0; i-- {
|
||||
h = h*mul + uint64(s[i])
|
||||
}
|
||||
return h
|
||||
func (g *MphMatcherGroup) pattern(ruleIdx uint32) string {
|
||||
return g.patterns[g.patternOffs[ruleIdx]:g.patternOffs[ruleIdx+1]]
|
||||
}
|
||||
|
||||
// mphMix spreads the weak low bits of a suffix hash.
|
||||
func mphMix(h uint64) uint64 {
|
||||
h ^= h >> 32
|
||||
h *= 0xd6e8feb86659fd93
|
||||
return h ^ h>>32
|
||||
// valuesOf caps the capacity, so appending to a Match result can't overwrite the next rule's values.
|
||||
func (g *MphMatcherGroup) valuesOf(ruleIdx uint32) []uint32 {
|
||||
start, end := g.valueOffs[ruleIdx], g.valueOffs[ruleIdx+1]
|
||||
return g.values[start:end:end]
|
||||
}
|
||||
|
||||
func (g *MphMatcherGroup) bucket(f uint64) uint32 {
|
||||
return uint32(((f >> 32) * uint64(g.n0)) >> 32)
|
||||
}
|
||||
|
||||
func (g *MphMatcherGroup) slot(f uint64, seed uint16) uint32 {
|
||||
x := ((f ^ uint64(seed)*0x9e3779b97f4a7c15) * 0xc4ceb9fe1a85ec53) >> 32
|
||||
return uint32((x * uint64(g.n1)) >> 32)
|
||||
}
|
||||
|
||||
func (g *MphMatcherGroup) uvarint(p uint32) (x, next uint32) {
|
||||
for shift := 0; ; shift += 7 {
|
||||
c := g.arena[p]
|
||||
p++
|
||||
x |= uint32(c&0x7f) << shift
|
||||
if c < 0x80 {
|
||||
return x, p
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// recSpan returns where the pattern of the record at off starts and how long it is.
|
||||
func (g *MphMatcherGroup) recSpan(off uint32) (p, n uint32) {
|
||||
n, p = uint32(g.arena[off]), off+1
|
||||
if n == 255 {
|
||||
n, p = g.uvarint(p)
|
||||
}
|
||||
return p, n
|
||||
}
|
||||
|
||||
func (g *MphMatcherGroup) recKey(rec uint32) string {
|
||||
p, n := g.recSpan(rec & mphOffMask)
|
||||
return g.arena[p : p+n]
|
||||
}
|
||||
|
||||
// lookup returns the level1 entry of s, or 0 if s is not a pattern. h is the suffix hash of s.
|
||||
func (g *MphMatcherGroup) lookup(h uint64, s string) uint32 {
|
||||
f := mphMix(h)
|
||||
// bucket < n0 == len(level0) and slot < n1 == len(level1) == len(fp), skip the bounds checks
|
||||
seed := *(*uint16)(unsafe.Add(unsafe.Pointer(unsafe.SliceData(g.level0)), uintptr(g.bucket(f))*2))
|
||||
slot := uintptr(g.slot(f, seed))
|
||||
if *(*uint8)(unsafe.Add(unsafe.Pointer(unsafe.SliceData(g.fp)), slot)) != uint8(f) {
|
||||
return 0
|
||||
}
|
||||
e := *(*uint32)(unsafe.Add(unsafe.Pointer(unsafe.SliceData(g.level1)), slot*4))
|
||||
if len(s) < 255 {
|
||||
// A record whose length byte is len(s) has len(s) pattern bytes after it
|
||||
p := unsafe.Add(unsafe.Pointer(unsafe.StringData(g.arena)), e&mphOffMask)
|
||||
if int(*(*byte)(p)) == len(s) && unsafe.String((*byte)(unsafe.Add(p, 1)), len(s)) == s {
|
||||
return e
|
||||
}
|
||||
return 0
|
||||
}
|
||||
if g.recKey(e) == s {
|
||||
return e
|
||||
// Lookup searches for input in minimal perfect hash table and returns its index. 0 indicates not found.
|
||||
func (g *MphMatcherGroup) Lookup(rollingHash uint32, input string) uint32 {
|
||||
i0 := rollingHash & g.level0Mask
|
||||
seed := g.level0[i0]
|
||||
i1 := MemHash(seed, input) & g.level1Mask
|
||||
n := g.level1[i1]
|
||||
// Build only puts valid rule indices in level1, so n+1 < len(patternOffs) and the span is inside patterns.
|
||||
// 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 {
|
||||
return n
|
||||
}
|
||||
return 0
|
||||
}
|
||||
|
||||
// appendValues appends the values of record e for the flags in want, in mphKinds order.
|
||||
func (g *MphMatcherGroup) appendValues(dst []uint32, e, want uint32) []uint32 {
|
||||
if !g.multi {
|
||||
for _, flag := range mphKinds {
|
||||
if e&want&flag != 0 {
|
||||
dst = append(dst, g.single)
|
||||
}
|
||||
}
|
||||
return dst
|
||||
}
|
||||
if e&want == 0 {
|
||||
return dst
|
||||
}
|
||||
p, n := g.recSpan(e & mphOffMask)
|
||||
p += n
|
||||
for _, flag := range mphKinds {
|
||||
if e&flag == 0 {
|
||||
continue
|
||||
}
|
||||
var count, v uint32
|
||||
for count, p = g.uvarint(p); count > 0; count-- {
|
||||
v, p = g.uvarint(p)
|
||||
if want&flag != 0 {
|
||||
dst = append(dst, v)
|
||||
}
|
||||
}
|
||||
}
|
||||
return dst
|
||||
}
|
||||
|
||||
// Match implements MatcherGroup.Match. Values of an exact match come first (Full, then Domain), then those of
|
||||
// the parent domains, nearest first.
|
||||
// Match implements MatcherGroup.Match.
|
||||
func (g *MphMatcherGroup) Match(input string) []uint32 {
|
||||
var stack [8]uint32
|
||||
parents := stack[:0] // TLD side first
|
||||
h, mul := uint64(0), g.mul
|
||||
matches := make([][]uint32, 0, 5)
|
||||
hash := uint32(0)
|
||||
for i := len(input) - 1; i >= 0; i-- {
|
||||
hash = hash*PrimeRK + uint32(input[i])
|
||||
if input[i] == '.' {
|
||||
if e := g.lookup(h, input[i+1:]); e&(mphDomain|mphParent) != 0 {
|
||||
parents = append(parents, e)
|
||||
if mphIdx := g.Lookup(hash, input[i:]); mphIdx != 0 {
|
||||
matches = append(matches, g.valuesOf(mphIdx))
|
||||
}
|
||||
}
|
||||
h = h*mul + uint64(input[i])
|
||||
}
|
||||
exact := g.lookup(h, input)
|
||||
if exact&(mphFull|mphDomain) == 0 && len(parents) == 0 {
|
||||
return nil
|
||||
if mphIdx := g.Lookup(hash, input); mphIdx != 0 {
|
||||
matches = append(matches, g.valuesOf(mphIdx))
|
||||
}
|
||||
result := g.appendValues(make([]uint32, 0, len(parents)+1), exact, mphFull|mphDomain)
|
||||
for k := len(parents) - 1; k >= 0; k-- {
|
||||
result = g.appendValues(result, parents[k], mphParent|mphDomain)
|
||||
}
|
||||
return result
|
||||
return CompositeMatchesReverse(matches)
|
||||
}
|
||||
|
||||
// MatchAny implements MatcherGroup.MatchAny.
|
||||
func (g *MphMatcherGroup) MatchAny(input string) bool {
|
||||
h, mul := uint64(0), g.mul
|
||||
for i := len(input) - 1; i >= 0; i-- {
|
||||
if input[i] == '.' && g.lookup(h, input[i+1:])&(mphDomain|mphParent) != 0 {
|
||||
return true
|
||||
}
|
||||
h = h*mul + uint64(input[i])
|
||||
}
|
||||
return g.lookup(h, input)&(mphFull|mphDomain) != 0
|
||||
}
|
||||
|
||||
// mphSuffix is the suffix hash of input[off:], a parent domain of the input.
|
||||
type mphSuffix struct {
|
||||
h uint64
|
||||
off int
|
||||
}
|
||||
|
||||
// mphSuffixes appends the suffix hashes of the parent domains of input to dst, TLD side first, and returns them
|
||||
// with the hash of input itself: what MatchAny computes, computed once for several groups.
|
||||
func mphSuffixes(dst []mphSuffix, mul uint64, input string) ([]mphSuffix, uint64) {
|
||||
h := uint64(0)
|
||||
hash := uint32(0)
|
||||
for i := len(input) - 1; i >= 0; i-- {
|
||||
hash = hash*PrimeRK + uint32(input[i])
|
||||
if input[i] == '.' {
|
||||
dst = append(dst, mphSuffix{h, i + 1})
|
||||
if g.Lookup(hash, input[i:]) != 0 {
|
||||
return true
|
||||
}
|
||||
}
|
||||
h = h*mul + uint64(input[i])
|
||||
}
|
||||
return dst, h
|
||||
return g.Lookup(hash, input) != 0
|
||||
}
|
||||
|
||||
// matchAnyHashed is MatchAny with parents and h from mphSuffixes(_, mul, input).
|
||||
func (g *MphMatcherGroup) matchAnyHashed(input string, parents []mphSuffix, h, mul uint64) bool {
|
||||
if g.mul != mul {
|
||||
return g.MatchAny(input) // built with a later multiplier after a collision
|
||||
func nextPow2(v int) int {
|
||||
if v <= 1 {
|
||||
return 1
|
||||
}
|
||||
for _, p := range parents {
|
||||
if g.lookup(p.h, input[p.off:])&(mphDomain|mphParent) != 0 {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return g.lookup(h, input)&(mphFull|mphDomain) != 0
|
||||
const MaxUInt = ^uint(0)
|
||||
n := (MaxUInt >> bits.LeadingZeros(uint(v))) + 1
|
||||
return int(n)
|
||||
}
|
||||
|
||||
//go:noescape
|
||||
//go:linkname strhash runtime.strhash
|
||||
func strhash(p unsafe.Pointer, h uintptr) uintptr
|
||||
|
||||
@@ -1,108 +0,0 @@
|
||||
package strmatcher
|
||||
|
||||
import (
|
||||
"slices"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestMphMatcherGroupHashCollision(t *testing.T) {
|
||||
saved := mphMultipliers
|
||||
defer func() { mphMultipliers = saved }()
|
||||
|
||||
mphMultipliers[0] = 1 // anagrams collide
|
||||
g := NewMphMatcherGroup()
|
||||
g.AddFullMatcher(FullMatcher("ab.com"), 1)
|
||||
g.AddDomainMatcher(DomainMatcher("ba.com"), 2)
|
||||
g.AddDomainMatcher(DomainMatcher("com"), 3)
|
||||
if err := g.Build(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if g.mul != saved[1] {
|
||||
t.Errorf("multiplier %#x, want the second one %#x", g.mul, saved[1])
|
||||
}
|
||||
for input, want := range map[string][]uint32{"ab.com": {1, 3}, "x.ba.com": {2, 3}, "x.ab.com": {3}, "ba.com": {2, 3}} {
|
||||
if m := g.Match(input); !slices.Equal(m, want) {
|
||||
t.Errorf("Match(%q) = %v, want %v", input, m, want)
|
||||
}
|
||||
}
|
||||
|
||||
// Thue-Morse strings of 2048 bytes and their complements collide for every odd multiplier
|
||||
mphMultipliers = saved
|
||||
a, b := make([]byte, 2048), make([]byte, 2048)
|
||||
for i := range a {
|
||||
a[i], b[i] = "ab"[bitsOnes(i)%2], "ba"[bitsOnes(i)%2]
|
||||
}
|
||||
g = NewMphMatcherGroup()
|
||||
g.AddFullMatcher(FullMatcher(a), 1)
|
||||
g.AddFullMatcher(FullMatcher(b), 1)
|
||||
if err := g.Build(); err != errMphCollision {
|
||||
t.Errorf("Build() = %v, want %v", err, errMphCollision)
|
||||
}
|
||||
}
|
||||
|
||||
func bitsOnes(i int) int {
|
||||
n := 0
|
||||
for ; i > 0; i &= i - 1 {
|
||||
n++
|
||||
}
|
||||
return n
|
||||
}
|
||||
|
||||
func TestMphValueMatcherCombiner(t *testing.T) {
|
||||
build := func(matchers ...Matcher) *MphValueMatcher {
|
||||
m := NewMphValueMatcher()
|
||||
for _, x := range matchers {
|
||||
m.Add(x, 0)
|
||||
}
|
||||
if err := m.Build(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
return m
|
||||
}
|
||||
regex, err := Regex.New(`^a\d+\.net$`)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
saved := mphMultipliers
|
||||
t.Cleanup(func() { mphMultipliers = saved })
|
||||
mphMultipliers[0] = 1 // anagrams collide, so this one falls back to its own hash pass
|
||||
collided := build(FullMatcher("ab.com"), DomainMatcher("ba.com"))
|
||||
mphMultipliers = saved
|
||||
if collided.mph.mul == mphMultipliers[0] {
|
||||
t.Fatal("collided matcher uses the first multiplier")
|
||||
}
|
||||
matchers := []*MphValueMatcher{
|
||||
build(DomainMatcher("example.com"), FullMatcher("full.org"), DomainMatcher(".dot.io")),
|
||||
collided,
|
||||
build(regex, SubstrMatcher("keyword")),
|
||||
build(),
|
||||
build(DomainMatcher("com"), DomainMatcher("a.b.c.d.e.f.g.h.i.j.k.l.m.n.o.p.q.r.s")),
|
||||
}
|
||||
var s MphValueMatcherCombiner
|
||||
for i, m := range matchers {
|
||||
s.Add(m, uint32(10+i))
|
||||
}
|
||||
inputs := []string{
|
||||
"", ".", "..", "com", "example.com", "www.example.com", "xexample.com", "example.com.", "full.org", "x.full.org",
|
||||
"dot.io", "x.dot.io", ".dot.io", "ab.com", "x.ab.com", "ba.com", "x.ba.com", "a12.net", "a12.net.x", "my-keyword.org",
|
||||
"a.b.c.d.e.f.g.h.i.j.k.l.m.n.o.p.q.r.s", "0.a.b.c.d.e.f.g.h.i.j.k.l.m.n.o.p.q.r.s", "b.c.d.e.f.g.h.i.j.k.l.m.n.o.p.q.r.s",
|
||||
"x.y.z.1.2.3.4.5.6.7.8.9.10.11.12.13.14.15.16.17.ab.com", "x.y.z.1.2.3.4.5.6.7.8.9.10.11.12.13.14.15.16.17.org",
|
||||
}
|
||||
for _, input := range inputs {
|
||||
var want []uint32
|
||||
for i, m := range matchers {
|
||||
if m.MatchAny(input) {
|
||||
want = append(want, uint32(10+i))
|
||||
}
|
||||
}
|
||||
if got := s.Match(input); !slices.Equal(got, want) {
|
||||
t.Errorf("Match(%q) = %v, want %v", input, got, want)
|
||||
}
|
||||
if got := s.MatchAny(input); got != (len(want) > 0) {
|
||||
t.Errorf("MatchAny(%q) = %v", input, got)
|
||||
}
|
||||
}
|
||||
if n := testing.AllocsPerRun(100, func() { s.MatchAny("www.a.b.c.example.org") }); n != 0 {
|
||||
t.Errorf("MatchAny allocates %v times", n)
|
||||
}
|
||||
}
|
||||
@@ -4,7 +4,6 @@ import (
|
||||
"math/rand"
|
||||
"reflect"
|
||||
"slices"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/xtls/xray-core/common"
|
||||
@@ -305,7 +304,7 @@ func TestMphMatcherGroupRandom(t *testing.T) {
|
||||
domain["."+p] = append(domain["."+p], value)
|
||||
}
|
||||
}
|
||||
common.Must(g.Build())
|
||||
g.Build()
|
||||
for _, input := range inputs {
|
||||
keys := []string{input} // Whole input first, then "." suffixes from longest to shortest
|
||||
for i := range len(input) {
|
||||
@@ -317,10 +316,7 @@ func TestMphMatcherGroupRandom(t *testing.T) {
|
||||
for _, k := range keys {
|
||||
want = append(append(want, full[k]...), domain[k]...)
|
||||
}
|
||||
// Compared as sets: Match reports a value once per matching pattern, and orders them differently
|
||||
// from want for patterns and inputs with a leading dot
|
||||
m := g.Match(input)
|
||||
if !slices.Equal(sortedSet(m), sortedSet(want)) {
|
||||
if m := g.Match(input); !slices.Equal(m, want) {
|
||||
t.Fatalf("seed %d: Match(%q) = %v, want %v", seed, input, m, want)
|
||||
}
|
||||
if m := g.MatchAny(input); m != (len(want) > 0) {
|
||||
@@ -342,79 +338,3 @@ func TestMphMatcherGroupAppend(t *testing.T) {
|
||||
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)
|
||||
}
|
||||
|
||||
@@ -2,12 +2,10 @@ package strmatcher
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"math/bits"
|
||||
"regexp"
|
||||
"regexp/syntax"
|
||||
"slices"
|
||||
"strings"
|
||||
"unicode"
|
||||
"unicode/utf8"
|
||||
|
||||
"golang.org/x/net/idna"
|
||||
@@ -77,9 +75,7 @@ func (m SubstrMatcher) Match(s string) bool {
|
||||
// RegexMatcher is an implementation of Matcher.
|
||||
type RegexMatcher struct {
|
||||
pattern *regexp.Regexp
|
||||
literals []string // every match contains all of them, longest first
|
||||
tail []byteSet // tail[i] holds the bytes a matching input can have i bytes before its end
|
||||
rest *byteSet // the bytes it can have further before, nil if any
|
||||
literals []string // every match contains all of them, longest first
|
||||
}
|
||||
|
||||
func newRegexMatcher(pattern string) (Matcher, error) {
|
||||
@@ -91,239 +87,10 @@ func newRegexMatcher(pattern string) (Matcher, error) {
|
||||
if re, err := syntax.Parse(pattern, syntax.Perl); err == nil { // same flags as regexp.Compile
|
||||
m.literals = requiredLiterals(re, nil)
|
||||
slices.SortStableFunc(m.literals, func(a, b string) int { return len(b) - len(a) })
|
||||
m.tail, m.rest = tailGuard(re)
|
||||
}
|
||||
return m, nil
|
||||
}
|
||||
|
||||
// byteSet is a set of bytes. The bytes >= 0x80 share one bit with 0x7f.
|
||||
type byteSet [4]uint32
|
||||
|
||||
func (s *byteSet) add(c byte) { c = min(c, 0x7f); s[c>>5] |= 1 << (c & 31) }
|
||||
func (s *byteSet) has(c byte) bool { c = min(c, 0x7f); return s[c>>5]&(1<<(c&31)) != 0 }
|
||||
func (s *byteSet) or(t *byteSet) {
|
||||
for i := range s {
|
||||
s[i] |= t[i]
|
||||
}
|
||||
}
|
||||
|
||||
var allBytes = byteSet{^uint32(0), ^uint32(0), ^uint32(0), ^uint32(0)}
|
||||
|
||||
// tailLen is how many positions before the end of the input tailGuard tells apart.
|
||||
const tailLen = 8
|
||||
|
||||
// tailBudget caps how many repetition steps tailGuard walks. Only nested repeats can make the
|
||||
// walk explode, so only they are charged: a flat pattern, however long, is walked once and keeps
|
||||
// its guard.
|
||||
const tailBudget = 100000
|
||||
|
||||
// tailWalk is a set of positions in the input, counted in bytes before its end.
|
||||
type tailWalk struct {
|
||||
at uint32 // bit i: exactly i bytes before the end, for i < tailLen
|
||||
far bool // tailLen or more bytes before the end
|
||||
free bool // not tied to the end of the input yet
|
||||
}
|
||||
|
||||
func (w tailWalk) union(v tailWalk) tailWalk {
|
||||
return tailWalk{w.at | v.at, w.far || v.far, w.free || v.free}
|
||||
}
|
||||
|
||||
type tailBuilder struct {
|
||||
tail [tailLen]byteSet
|
||||
rest byteSet
|
||||
void bool
|
||||
work int
|
||||
}
|
||||
|
||||
// tailGuard walks re backwards from the end of the input and collects the bytes an input
|
||||
// matching re can have at each position before its end. It returns nil, nil when a branch
|
||||
// of re does not end with $ or when nested repeats push the walk past tailBudget.
|
||||
func tailGuard(re *syntax.Regexp) ([]byteSet, *byteSet) {
|
||||
var b tailBuilder
|
||||
w := b.walk(re, tailWalk{free: true})
|
||||
b.stop(w)
|
||||
if b.void {
|
||||
return nil, nil
|
||||
}
|
||||
if w.at != 0 { // a match can start here, so any bytes can come before
|
||||
for i := bits.TrailingZeros32(w.at); i < tailLen; i++ {
|
||||
b.tail[i] = allBytes
|
||||
}
|
||||
}
|
||||
if w.at != 0 || w.far {
|
||||
b.rest = allBytes
|
||||
}
|
||||
n := tailLen
|
||||
for n > 0 && b.tail[n-1] == b.rest {
|
||||
n--
|
||||
}
|
||||
var tail []byteSet
|
||||
if n > 0 {
|
||||
tail = slices.Clone(b.tail[:n])
|
||||
}
|
||||
if b.rest != allBytes {
|
||||
rest := b.rest
|
||||
return tail, &rest
|
||||
}
|
||||
return tail, nil
|
||||
}
|
||||
|
||||
// stop ends the paths of w. One that never met $ lets its match be followed by anything.
|
||||
func (b *tailBuilder) stop(w tailWalk) {
|
||||
if w.free {
|
||||
b.void = true
|
||||
}
|
||||
}
|
||||
|
||||
func (b *tailBuilder) walk(re *syntax.Regexp, w tailWalk) tailWalk {
|
||||
if w == (tailWalk{}) || b.void {
|
||||
return w
|
||||
}
|
||||
switch re.Op {
|
||||
case syntax.OpNoMatch:
|
||||
return tailWalk{}
|
||||
case syntax.OpLiteral:
|
||||
for i := len(re.Rune) - 1; i >= 0; i-- {
|
||||
var set byteSet
|
||||
set.add(byte(min(re.Rune[i], utf8.RuneSelf)))
|
||||
if re.Flags&syntax.FoldCase != 0 {
|
||||
for f := unicode.SimpleFold(re.Rune[i]); f != re.Rune[i]; f = unicode.SimpleFold(f) {
|
||||
set.add(byte(min(f, utf8.RuneSelf)))
|
||||
}
|
||||
}
|
||||
w = b.step(w, &set)
|
||||
}
|
||||
return w
|
||||
case syntax.OpCharClass:
|
||||
var set byteSet
|
||||
for i := 0; i+1 < len(re.Rune); i += 2 {
|
||||
for r := min(re.Rune[i], utf8.RuneSelf); r <= min(re.Rune[i+1], utf8.RuneSelf); r++ {
|
||||
set.add(byte(r))
|
||||
}
|
||||
}
|
||||
return b.step(w, &set)
|
||||
case syntax.OpAnyChar, syntax.OpAnyCharNotNL: // a domain has no \n to reject
|
||||
return b.step(w, &allBytes)
|
||||
case syntax.OpBeginText: // nothing comes before
|
||||
b.stop(w)
|
||||
return tailWalk{}
|
||||
case syntax.OpEndText:
|
||||
out := tailWalk{at: w.at & 1}
|
||||
if w.free {
|
||||
out.at = 1
|
||||
}
|
||||
return out
|
||||
case syntax.OpCapture:
|
||||
return b.walk(re.Sub[0], w)
|
||||
case syntax.OpConcat:
|
||||
for i := len(re.Sub) - 1; i >= 0; i-- {
|
||||
w = b.walk(re.Sub[i], w)
|
||||
}
|
||||
return w
|
||||
case syntax.OpAlternate:
|
||||
var out tailWalk
|
||||
for _, sub := range re.Sub {
|
||||
out = out.union(b.walk(sub, w))
|
||||
}
|
||||
return out
|
||||
case syntax.OpQuest:
|
||||
return b.repeat(re.Sub[0], w, 1)
|
||||
case syntax.OpStar:
|
||||
return b.repeat(re.Sub[0], w, -1)
|
||||
case syntax.OpPlus:
|
||||
return b.repeat(re.Sub[0], b.walk(re.Sub[0], w), -1)
|
||||
case syntax.OpRepeat:
|
||||
for i := 0; i < re.Min; i++ {
|
||||
if b.charge() {
|
||||
return w
|
||||
}
|
||||
w = b.walk(re.Sub[0], w)
|
||||
}
|
||||
if re.Max < 0 {
|
||||
return b.repeat(re.Sub[0], w, -1)
|
||||
}
|
||||
return b.repeat(re.Sub[0], w, re.Max-re.Min)
|
||||
}
|
||||
return w // empty match, line and word boundaries: no constraint
|
||||
}
|
||||
|
||||
// charge counts one repetition step and reports whether the walk has run out of budget. Only
|
||||
// repeats re-walk their body, so charging them alone bounds the blow-up of nested repeats while
|
||||
// leaving a single linear pass, of any length, free.
|
||||
func (b *tailBuilder) charge() bool {
|
||||
b.work++
|
||||
if b.work > tailBudget {
|
||||
b.void = true
|
||||
}
|
||||
return b.void
|
||||
}
|
||||
|
||||
// repeat walks back over up to n more repetitions of re, any number if n < 0.
|
||||
func (b *tailBuilder) repeat(re *syntax.Regexp, w tailWalk, n int) tailWalk {
|
||||
for ; n != 0; n-- {
|
||||
if b.charge() {
|
||||
return w
|
||||
}
|
||||
next := w.union(b.walk(re, w))
|
||||
if next == w {
|
||||
break
|
||||
}
|
||||
w = next
|
||||
}
|
||||
return w
|
||||
}
|
||||
|
||||
// step walks back over one character whose last byte is in set. A character that can be
|
||||
// non-ASCII can take up to 4 bytes, all >= 0x80; regexp matches an invalid byte as U+FFFD.
|
||||
func (b *tailBuilder) step(w tailWalk, set *byteSet) tailWalk {
|
||||
out := tailWalk{far: w.far, free: w.free}
|
||||
if w.far {
|
||||
b.rest.or(set)
|
||||
}
|
||||
width := 1
|
||||
if set.has(0x80) {
|
||||
width = utf8.UTFMax
|
||||
}
|
||||
for i := 0; i < tailLen; i++ {
|
||||
if w.at&(1<<i) == 0 {
|
||||
continue
|
||||
}
|
||||
b.tail[i].or(set)
|
||||
for n := 1; n <= width; n++ {
|
||||
if j := i + n; j < tailLen {
|
||||
out.at |= 1 << j
|
||||
if n < width {
|
||||
b.tail[j].add(0x80)
|
||||
}
|
||||
} else {
|
||||
out.far = true
|
||||
if n < width {
|
||||
b.rest.add(0x80)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
// mayMatch reports whether s passes the tail guard.
|
||||
func (m *RegexMatcher) mayMatch(s string) bool {
|
||||
n := len(s)
|
||||
if m.rest == nil {
|
||||
n = min(n, len(m.tail))
|
||||
}
|
||||
for i := 0; i < n; i++ {
|
||||
set := m.rest
|
||||
if i < len(m.tail) {
|
||||
set = &m.tail[i]
|
||||
}
|
||||
if !set.has(s[len(s)-1-i]) {
|
||||
return false
|
||||
}
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
// requiredLiterals appends to dst the case-sensitive strings that every match of re contains.
|
||||
func requiredLiterals(re *syntax.Regexp, dst []string) []string {
|
||||
switch re.Op {
|
||||
@@ -359,9 +126,6 @@ func (m *RegexMatcher) String() string {
|
||||
}
|
||||
|
||||
func (m *RegexMatcher) Match(s string) bool {
|
||||
if !m.mayMatch(s) {
|
||||
return false
|
||||
}
|
||||
for _, l := range m.literals {
|
||||
if !strings.Contains(s, l) {
|
||||
return false
|
||||
|
||||
@@ -1,16 +1,9 @@
|
||||
package strmatcher
|
||||
|
||||
import (
|
||||
"hash/fnv"
|
||||
"math/rand/v2"
|
||||
"regexp"
|
||||
"regexp/syntax"
|
||||
"slices"
|
||||
"strconv"
|
||||
"strings"
|
||||
"testing"
|
||||
"unicode"
|
||||
"unicode/utf8"
|
||||
)
|
||||
|
||||
var regexLiteralCases = []struct {
|
||||
@@ -44,147 +37,6 @@ 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) {
|
||||
inputs := []string{
|
||||
"", "x", "yy", "abd", "ccc", "ABC", "abCDef", "abcdef", "abababcc", "a.b", "a\xffb", "a\uFFFDb",
|
||||
@@ -195,39 +47,14 @@ func FuzzRegexMatcher(f *testing.F) {
|
||||
f.Add(test.pattern, s)
|
||||
}
|
||||
}
|
||||
for _, test := range regexTailCases {
|
||||
for _, s := range append(test.match, test.reject...) {
|
||||
f.Add(test.pattern, s)
|
||||
}
|
||||
}
|
||||
f.Fuzz(func(t *testing.T, pattern, s string) {
|
||||
re, err := regexp.Compile(pattern)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
m, _ := newRegexMatcher(pattern)
|
||||
check := func(s string) {
|
||||
if got, want := m.Match(s), re.MatchString(s); got != want {
|
||||
t.Errorf("pattern %q, input %q: got %v, want %v", pattern, s, got, want)
|
||||
}
|
||||
}
|
||||
check(s)
|
||||
// random inputs seldom match, so also try strings built from the pattern
|
||||
parsed, _ := syntax.Parse(pattern, syntax.Perl)
|
||||
h := fnv.New64a()
|
||||
h.Write([]byte(s))
|
||||
rnd := rand.New(rand.NewPCG(h.Sum64(), 1))
|
||||
for range 8 {
|
||||
var sb strings.Builder
|
||||
budget := 256
|
||||
sampleMatch(&sb, parsed, rnd, &budget)
|
||||
sample := sb.String()
|
||||
check(sample)
|
||||
check(s + sample)
|
||||
if len(sample) > 0 && len(s) > 0 {
|
||||
i := rnd.IntN(len(sample))
|
||||
check(sample[:i] + s[:1] + sample[i+1:])
|
||||
}
|
||||
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)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
@@ -46,9 +46,7 @@ func (g *MphValueMatcher) Add(matcher Matcher, value uint32) {
|
||||
func (g *MphValueMatcher) Build() error {
|
||||
if g.mph != nil {
|
||||
runtime.GC() // peak mem
|
||||
if err := g.mph.Build(); err != nil {
|
||||
return err
|
||||
}
|
||||
g.mph.Build()
|
||||
}
|
||||
runtime.GC() // peak mem
|
||||
if g.ac != nil {
|
||||
@@ -60,17 +58,23 @@ func (g *MphValueMatcher) Build() error {
|
||||
|
||||
// Match implements ValueMatcher.Match.
|
||||
func (g *MphValueMatcher) Match(input string) []uint32 {
|
||||
var result []uint32
|
||||
result := make([][]uint32, 0, 5)
|
||||
if g.mph != nil {
|
||||
result = g.mph.Match(input) // a new slice, returned without another copy
|
||||
if matches := g.mph.Match(input); len(matches) > 0 {
|
||||
result = append(result, matches)
|
||||
}
|
||||
}
|
||||
if g.ac != nil {
|
||||
result = append(result, g.ac.Match(input)...)
|
||||
if matches := g.ac.Match(input); len(matches) > 0 {
|
||||
result = append(result, matches)
|
||||
}
|
||||
}
|
||||
if g.regex != nil {
|
||||
result = append(result, g.regex.Match(input)...)
|
||||
if matches := g.regex.Match(input); len(matches) > 0 {
|
||||
result = append(result, matches)
|
||||
}
|
||||
}
|
||||
return result
|
||||
return CompositeMatches(result)
|
||||
}
|
||||
|
||||
// MatchAny implements ValueMatcher.MatchAny.
|
||||
@@ -83,62 +87,3 @@ func (g *MphValueMatcher) MatchAny(input string) bool {
|
||||
}
|
||||
return g.regex != nil && g.regex.MatchAny(input)
|
||||
}
|
||||
|
||||
func (g *MphValueMatcher) matchAnyHashed(input string, parents []mphSuffix, h, mul uint64) bool {
|
||||
if g.mph != nil && g.mph.matchAnyHashed(input, parents, h, mul) {
|
||||
return true
|
||||
}
|
||||
if g.ac != nil && g.ac.MatchAny(input) {
|
||||
return true
|
||||
}
|
||||
return g.regex != nil && g.regex.MatchAny(input)
|
||||
}
|
||||
|
||||
// MphValueMatcherCombiner combines several built MphValueMatchers, each bound to one value, and matches an input
|
||||
// against them as their MatchAny would, hashing the input once for all of them.
|
||||
type MphValueMatcherCombiner struct {
|
||||
matchers []*MphValueMatcher
|
||||
values []uint32
|
||||
}
|
||||
|
||||
// Add adds a built matcher that stands for value.
|
||||
func (s *MphValueMatcherCombiner) Add(m *MphValueMatcher, value uint32) {
|
||||
s.matchers = append(s.matchers, m)
|
||||
s.values = append(s.values, value)
|
||||
}
|
||||
|
||||
// Match returns the values of the matchers that match input, in Add order.
|
||||
func (s *MphValueMatcherCombiner) Match(input string) []uint32 {
|
||||
if len(s.matchers) == 0 {
|
||||
return nil
|
||||
}
|
||||
var stack [16]mphSuffix
|
||||
mul := mphMultipliers[0]
|
||||
parents, h := mphSuffixes(stack[:0], mul, input)
|
||||
var result []uint32
|
||||
for i, m := range s.matchers {
|
||||
if m.matchAnyHashed(input, parents, h, mul) {
|
||||
result = append(result, s.values[i])
|
||||
}
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
// MatchAny returns true as soon as one matcher matches input.
|
||||
func (s *MphValueMatcherCombiner) MatchAny(input string) bool {
|
||||
switch len(s.matchers) {
|
||||
case 0:
|
||||
return false
|
||||
case 1:
|
||||
return s.matchers[0].MatchAny(input) // nothing to share, and it stops at the first matching suffix
|
||||
}
|
||||
var stack [16]mphSuffix
|
||||
mul := mphMultipliers[0]
|
||||
parents, h := mphSuffixes(stack[:0], mul, input)
|
||||
for _, m := range s.matchers {
|
||||
if m.matchAnyHashed(input, parents, h, mul) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
@@ -1,20 +0,0 @@
|
||||
package net
|
||||
|
||||
// PacketConnWrapper wraps a PacketConn into a Conn with a fixed destination address.
|
||||
type PacketConnWrapper struct {
|
||||
PacketConn
|
||||
Dest Addr
|
||||
}
|
||||
|
||||
func (c *PacketConnWrapper) Read(p []byte) (int, error) {
|
||||
n, _, err := c.PacketConn.ReadFrom(p)
|
||||
return n, err
|
||||
}
|
||||
|
||||
func (c *PacketConnWrapper) Write(p []byte) (int, error) {
|
||||
return c.PacketConn.WriteTo(p, c.Dest)
|
||||
}
|
||||
|
||||
func (c *PacketConnWrapper) RemoteAddr() Addr {
|
||||
return c.Dest
|
||||
}
|
||||
+1
-1
@@ -20,7 +20,7 @@ import (
|
||||
var (
|
||||
Version_x byte = 26
|
||||
Version_y byte = 9
|
||||
Version_z byte = 30
|
||||
Version_z byte = 9
|
||||
)
|
||||
|
||||
var (
|
||||
|
||||
@@ -1,7 +1,6 @@
|
||||
package conf
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/x509"
|
||||
"encoding/base64"
|
||||
"encoding/hex"
|
||||
@@ -15,7 +14,6 @@ import (
|
||||
googleuuid "github.com/google/uuid"
|
||||
"github.com/xtls/xray-core/common/errors"
|
||||
"github.com/xtls/xray-core/common/net"
|
||||
"github.com/xtls/xray-core/common/serial"
|
||||
"github.com/xtls/xray-core/transport/internet/finalmask/fragment"
|
||||
"github.com/xtls/xray-core/transport/internet/finalmask/header/custom"
|
||||
"github.com/xtls/xray-core/transport/internet/finalmask/mkcp/aes128gcm"
|
||||
@@ -83,7 +81,7 @@ var (
|
||||
"noise": func() interface{} { return new(NoiseMask) },
|
||||
"salamander": func() interface{} { return new(Salamander) },
|
||||
"sudoku": func() interface{} { return new(Sudoku) },
|
||||
"xdns": func() interface{} { return new(XDNS) },
|
||||
"xdns": func() interface{} { return new(Xdns) },
|
||||
"xicmp": func() interface{} { return new(Xicmp) },
|
||||
"realm": func() interface{} { return new(Realm) },
|
||||
"udphop": func() interface{} { return new(UDPHop) },
|
||||
@@ -310,27 +308,14 @@ type NoiseMask struct {
|
||||
}
|
||||
|
||||
func (c *NoiseMask) Build() (proto.Message, error) {
|
||||
noiseSlice := make([]*noise.Item, 0, len(c.Noise))
|
||||
for _, item := range c.Noise {
|
||||
if len(item.Packet) > 0 && item.Rand.To > 0 {
|
||||
return nil, errors.New("len(item.Packet) > 0 && item.Rand.To > 0")
|
||||
}
|
||||
if strings.ToLower(item.Type) == "exp" {
|
||||
var exp string
|
||||
if err := json.Unmarshal(item.Packet, &exp); err != nil {
|
||||
return nil, errors.New(`"packet" of noise "type": "exp" must be a string`).Base(err)
|
||||
}
|
||||
segments, err := parseNoiseExp(exp)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
noiseSlice = append(noiseSlice, &noise.Item{
|
||||
Segments: segments,
|
||||
DelayMin: int64(item.Delay.From),
|
||||
DelayMax: int64(item.Delay.To),
|
||||
})
|
||||
continue
|
||||
}
|
||||
}
|
||||
|
||||
noiseSlice := make([]*noise.Item, 0, len(c.Noise))
|
||||
for _, item := range c.Noise {
|
||||
if item.RandRange == nil {
|
||||
item.RandRange = &Int32Range{From: 0, To: 255}
|
||||
}
|
||||
@@ -359,88 +344,6 @@ func (c *NoiseMask) Build() (proto.Message, error) {
|
||||
}, nil
|
||||
}
|
||||
|
||||
var noiseExpPattern = regexp.MustCompile(`<\s*([a-z]+)(?:\s+([^>]*?))?\s*>`)
|
||||
|
||||
func parseNoiseExp(exp string) ([]*noise.Segment, error) {
|
||||
var segments []*noise.Segment
|
||||
matches := noiseExpPattern.FindAllStringSubmatchIndex(exp, -1)
|
||||
last := 0
|
||||
for _, m := range matches {
|
||||
if strings.TrimSpace(exp[last:m[0]]) != "" {
|
||||
return nil, errors.New("invalid noise exp near ", exp[last:m[0]])
|
||||
}
|
||||
last = m[1]
|
||||
key := exp[m[2]:m[3]]
|
||||
arg := ""
|
||||
if m[4] >= 0 {
|
||||
arg = exp[m[4]:m[5]]
|
||||
}
|
||||
segment, err := buildNoiseSegment(key, arg)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
segments = append(segments, segment)
|
||||
}
|
||||
if strings.TrimSpace(exp[last:]) != "" {
|
||||
return nil, errors.New("invalid noise exp near ", exp[last:])
|
||||
}
|
||||
if len(segments) == 0 {
|
||||
return nil, errors.New("empty noise exp: ", exp)
|
||||
}
|
||||
return segments, nil
|
||||
}
|
||||
|
||||
func buildNoiseSegment(key, arg string) (*noise.Segment, error) {
|
||||
sizeSegment := func(kind noise.Segment_Kind) (*noise.Segment, error) {
|
||||
if arg == "" {
|
||||
return nil, errors.New("<", key, "> in noise exp needs a size")
|
||||
}
|
||||
lo, hi, err := ParseRangeString(arg)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if lo < 0 || hi < lo || hi > 65535 {
|
||||
return nil, errors.New("invalid size in noise exp: ", arg)
|
||||
}
|
||||
return &noise.Segment{Kind: kind, MinSize: int64(lo), MaxSize: int64(hi)}, nil
|
||||
}
|
||||
switch key {
|
||||
case "b":
|
||||
hexStr := strings.TrimPrefix(strings.TrimPrefix(strings.Join(strings.Fields(arg), ""), "0x"), "0X")
|
||||
if len(hexStr) == 0 {
|
||||
return nil, errors.New("empty bytes in noise exp")
|
||||
}
|
||||
raw, err := hex.DecodeString(hexStr)
|
||||
if err != nil {
|
||||
return nil, errors.New("invalid hex in noise exp: ", arg).Base(err)
|
||||
}
|
||||
return &noise.Segment{Kind: noise.Segment_BYTES, Bytes: raw}, nil
|
||||
case "r":
|
||||
return sizeSegment(noise.Segment_RANDOM)
|
||||
case "rc":
|
||||
return sizeSegment(noise.Segment_RANDOM_ASCII)
|
||||
case "rd":
|
||||
return sizeSegment(noise.Segment_RANDOM_DIGIT)
|
||||
case "t":
|
||||
if arg != "" {
|
||||
return nil, errors.New("<t> in noise exp takes no argument")
|
||||
}
|
||||
return &noise.Segment{Kind: noise.Segment_TIMESTAMP}, nil
|
||||
case "c":
|
||||
if arg != "" {
|
||||
return nil, errors.New("<c> in noise exp takes no argument")
|
||||
}
|
||||
return &noise.Segment{Kind: noise.Segment_COUNTER}, nil
|
||||
case "n":
|
||||
if arg != "" {
|
||||
return nil, errors.New("<n> in noise exp takes no argument")
|
||||
}
|
||||
return &noise.Segment{Kind: noise.Segment_NONCE}, nil
|
||||
default:
|
||||
return nil, errors.New("unknown <", key, "> in noise exp")
|
||||
}
|
||||
}
|
||||
|
||||
type UDPItem struct {
|
||||
Rand int32 `json:"rand"`
|
||||
RandRange *Int32Range `json:"randRange"`
|
||||
@@ -791,88 +694,32 @@ func (c *Sudoku) Build() (proto.Message, error) {
|
||||
}, nil
|
||||
}
|
||||
|
||||
type XDNSDomain struct {
|
||||
Name string `json:"name"`
|
||||
LenLimit int32 `json:"lenLimit"`
|
||||
LabelLimit int32 `json:"labelLimit"`
|
||||
Types []int32 `json:"types"`
|
||||
Edns0 int32 `json:"edns0"`
|
||||
type Xdns struct {
|
||||
Domain json.RawMessage `json:"domain"`
|
||||
|
||||
Domains []string `json:"domains"`
|
||||
Resolvers []string `json:"resolvers"`
|
||||
}
|
||||
|
||||
type XDNSResolverTCP struct {
|
||||
Addr string `json:"addr"`
|
||||
}
|
||||
|
||||
func (c *XDNSResolverTCP) Build() (proto.Message, error) {
|
||||
return &xdns.TCPResolverProto{Addr: c.Addr}, nil
|
||||
}
|
||||
|
||||
type XDNSResolverUDP struct {
|
||||
Addr string `json:"addr"`
|
||||
}
|
||||
|
||||
func (c *XDNSResolverUDP) Build() (proto.Message, error) {
|
||||
return &xdns.UDPResolverProto{Addr: c.Addr}, nil
|
||||
}
|
||||
|
||||
var xdnsLoader = NewJSONConfigLoader(ConfigCreatorCache{
|
||||
"tcp": func() interface{} { return new(XDNSResolverTCP) },
|
||||
"udp": func() interface{} { return new(XDNSResolverUDP) },
|
||||
}, "type", "settings")
|
||||
|
||||
type XDNSResolver struct {
|
||||
Type string `json:"type"`
|
||||
Settings json.RawMessage `json:"settings"`
|
||||
}
|
||||
|
||||
type XDNS struct {
|
||||
Domains []XDNSDomain `json:"domains"`
|
||||
Resolvers []XDNSResolver `json:"resolvers"`
|
||||
ExtraPoll int32 `json:"extraPoll"`
|
||||
}
|
||||
|
||||
func (c *XDNS) Build() (proto.Message, error) {
|
||||
var domains []*xdns.DomainProto
|
||||
var resolvers []*serial.TypedMessage
|
||||
for i := range c.Domains {
|
||||
if c.Domains[i].LenLimit == 0 {
|
||||
c.Domains[i].LenLimit = 255
|
||||
}
|
||||
if c.Domains[i].LabelLimit == 0 {
|
||||
c.Domains[i].LabelLimit = 63
|
||||
}
|
||||
types := make([]uint16, 0, len(c.Domains[i].Types))
|
||||
for j := range c.Domains[i].Types {
|
||||
types = append(types, uint16(c.Domains[i].Types[j]))
|
||||
}
|
||||
domain, err := xdns.NewDomain(c.Domains[i].Name, int(c.Domains[i].LenLimit), int(c.Domains[i].LabelLimit), types, uint16(c.Domains[i].Edns0))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
errors.LogInfo(context.Background(), domain.Show())
|
||||
domains = append(domains, &xdns.DomainProto{
|
||||
Name: c.Domains[i].Name,
|
||||
LenLimit: c.Domains[i].LenLimit,
|
||||
LabelLimit: c.Domains[i].LabelLimit,
|
||||
Types: c.Domains[i].Types,
|
||||
Edns0: c.Domains[i].Edns0,
|
||||
})
|
||||
func (c *Xdns) Build() (proto.Message, error) {
|
||||
if c.Domain != nil {
|
||||
return nil, errors.PrintRemovedFeatureError("domain", "domains(server) & resolvers(client)")
|
||||
}
|
||||
for i := range c.Resolvers {
|
||||
config, err := xdnsLoader.LoadWithID(c.Resolvers[i].Settings, c.Resolvers[i].Type)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
pm, err := config.(interface{ Build() (proto.Message, error) }).Build()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
resolvers = append(resolvers, serial.ToTypedMessage(pm))
|
||||
|
||||
if len(c.Domains) == 0 && len(c.Resolvers) == 0 {
|
||||
return nil, errors.New("empty domains & empty resolvers")
|
||||
}
|
||||
if c.ExtraPoll < 0 || c.ExtraPoll > 3 {
|
||||
return nil, errors.New("c.ExtraPoll < 0 || c.ExtraPoll > 3")
|
||||
|
||||
for _, r := range c.Resolvers {
|
||||
if !strings.Contains(r, "+udp://") {
|
||||
return nil, errors.New("invalid resolver ", r)
|
||||
}
|
||||
}
|
||||
return &xdns.Config{Domains: domains, Resolvers: resolvers, ExtraPoll: c.ExtraPoll}, nil
|
||||
|
||||
return &xdns.Config{
|
||||
Domains: c.Domains,
|
||||
Resolvers: c.Resolvers,
|
||||
}, nil
|
||||
}
|
||||
|
||||
type XMC struct {
|
||||
|
||||
@@ -1,136 +0,0 @@
|
||||
package conf
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"testing"
|
||||
|
||||
"github.com/xtls/xray-core/transport/internet/finalmask/noise"
|
||||
)
|
||||
|
||||
func expPacket(exp string) json.RawMessage {
|
||||
b, _ := json.Marshal(exp)
|
||||
return b
|
||||
}
|
||||
|
||||
func buildNoiseExp(exp string) (*noise.Config, error) {
|
||||
msg, err := (&NoiseMask{Noise: []NoiseItem{{Type: "exp", Packet: expPacket(exp)}}}).Build()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return msg.(*noise.Config), nil
|
||||
}
|
||||
|
||||
func TestNoiseExp(t *testing.T) {
|
||||
cfg, err := buildNoiseExp("<b 0d0a0d0a><t><r 24><rc 20-40><rd 8><c><n>")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
segments := cfg.Items[0].Segments
|
||||
if len(segments) != 7 {
|
||||
t.Fatalf("got %d segments, want 7", len(segments))
|
||||
}
|
||||
want := []struct {
|
||||
kind noise.Segment_Kind
|
||||
bytes []byte
|
||||
min, max int64
|
||||
}{
|
||||
{noise.Segment_BYTES, []byte{0x0d, 0x0a, 0x0d, 0x0a}, 0, 0},
|
||||
{noise.Segment_TIMESTAMP, nil, 0, 0},
|
||||
{noise.Segment_RANDOM, nil, 24, 24},
|
||||
{noise.Segment_RANDOM_ASCII, nil, 20, 40},
|
||||
{noise.Segment_RANDOM_DIGIT, nil, 8, 8},
|
||||
{noise.Segment_COUNTER, nil, 0, 0},
|
||||
{noise.Segment_NONCE, nil, 0, 0},
|
||||
}
|
||||
for i, w := range want {
|
||||
s := segments[i]
|
||||
if s.Kind != w.kind || s.MinSize != w.min || s.MaxSize != w.max || string(s.Bytes) != string(w.bytes) {
|
||||
t.Errorf("segment %d = %+v, want %+v", i, s, w)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestNoiseExpStripsHexPrefix(t *testing.T) {
|
||||
cfg, err := buildNoiseExp("<b 0x16030100>")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if got := cfg.Items[0].Segments[0].Bytes; string(got) != string([]byte{0x16, 0x03, 0x01, 0x00}) {
|
||||
t.Errorf("got %x", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestNoiseExpWhitespace(t *testing.T) {
|
||||
if _, err := buildNoiseExp(" <b 00> <t> "); err != nil {
|
||||
t.Errorf("surrounding whitespace should be allowed: %v", err)
|
||||
}
|
||||
cfg, err := buildNoiseExp("<b 0d 0a 0d 0a>")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if got := cfg.Items[0].Segments[0].Bytes; string(got) != "\r\n\r\n" {
|
||||
t.Errorf("got %x", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestNoiseExpRejects(t *testing.T) {
|
||||
for _, exp := range []string{
|
||||
"<x 1>",
|
||||
"<b>",
|
||||
"<b zz>",
|
||||
"<b 0d0>",
|
||||
"<r>",
|
||||
"<r -1>",
|
||||
"<r 40-20>",
|
||||
"<r 70000>",
|
||||
"<t 5>",
|
||||
"<n 5>",
|
||||
"garbage<t>",
|
||||
"<t> tail",
|
||||
"<t><b>",
|
||||
} {
|
||||
if _, err := buildNoiseExp(exp); err == nil {
|
||||
t.Errorf("expected an error for %q", exp)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestNoiseExpConflicts(t *testing.T) {
|
||||
if _, err := (&NoiseMask{Noise: []NoiseItem{{Type: "exp", Packet: expPacket("<t>"), Rand: Int32Range{From: 10, To: 20}}}}).Build(); err == nil {
|
||||
t.Error("exp with rand should be rejected")
|
||||
}
|
||||
for _, packet := range []string{``, `[1, 2]`, `5`} {
|
||||
if _, err := (&NoiseMask{Noise: []NoiseItem{{Type: "exp", Packet: json.RawMessage(packet)}}}).Build(); err == nil {
|
||||
t.Errorf("expected an error for packet %q", packet)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestNoiseExpFromJSON(t *testing.T) {
|
||||
var mask NoiseMask
|
||||
if err := json.Unmarshal([]byte(`{"noise": [
|
||||
{"type": "exp", "packet": "<b 504f5354><rd 10-20>", "delay": "1-3"},
|
||||
{"type": "EXP", "packet": "<t>"},
|
||||
{"type": "str", "packet": "<t>"},
|
||||
{"rand": "10-20"}
|
||||
]}`), &mask); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
msg, err := mask.Build()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
items := msg.(*noise.Config).Items
|
||||
if len(items[0].Segments) != 2 || items[0].DelayMin != 1 || items[0].DelayMax != 3 {
|
||||
t.Errorf("item 0 = %+v", items[0])
|
||||
}
|
||||
if len(items[1].Segments) != 1 || items[1].Segments[0].Kind != noise.Segment_TIMESTAMP {
|
||||
t.Errorf("item 1 = %+v", items[1])
|
||||
}
|
||||
if len(items[2].Segments) != 0 || string(items[2].Packet) != "<t>" {
|
||||
t.Errorf("item 2 = %+v", items[2])
|
||||
}
|
||||
if len(items[3].Segments) != 0 || items[3].RandMin != 10 || items[3].RandMax != 20 {
|
||||
t.Errorf("item 3 = %+v", items[3])
|
||||
}
|
||||
}
|
||||
@@ -12,7 +12,6 @@ import (
|
||||
"github.com/xtls/xray-core/core"
|
||||
"github.com/xtls/xray-core/infra/conf"
|
||||
"github.com/xtls/xray-core/infra/conf/serial"
|
||||
"github.com/xtls/xray-core/proxy/hysteria"
|
||||
"github.com/xtls/xray-core/proxy/masque"
|
||||
"github.com/xtls/xray-core/proxy/shadowsocks"
|
||||
"github.com/xtls/xray-core/proxy/shadowsocks_2022"
|
||||
@@ -92,8 +91,6 @@ func extractInboundUsers(inb *core.InboundHandlerConfig) []*protocol.User {
|
||||
return ty.Users
|
||||
case *masque.ServerConfig:
|
||||
return ty.Users
|
||||
case *hysteria.ServerConfig:
|
||||
return ty.Users
|
||||
default:
|
||||
fmt.Println("unsupported inbound type")
|
||||
}
|
||||
|
||||
@@ -467,7 +467,7 @@ func NewPacketReader(conn net.Conn, h *Handler, defaultRule *FinalRule, UDPOverr
|
||||
if statConn != nil {
|
||||
counter = statConn.ReadCounter
|
||||
}
|
||||
if c, ok := iConn.(*net.PacketConnWrapper); ok {
|
||||
if c, ok := iConn.(*internet.PacketConnWrapper); ok {
|
||||
isOverridden := false
|
||||
if UDPOverride.Address != nil || UDPOverride.Port != 0 {
|
||||
isOverridden = true
|
||||
@@ -487,7 +487,7 @@ func NewPacketReader(conn net.Conn, h *Handler, defaultRule *FinalRule, UDPOverr
|
||||
}
|
||||
|
||||
type PacketReader struct {
|
||||
*net.PacketConnWrapper
|
||||
*internet.PacketConnWrapper
|
||||
stats.Counter
|
||||
Handler *Handler
|
||||
DefaultRule *FinalRule
|
||||
@@ -542,7 +542,7 @@ func NewPacketWriter(conn net.Conn, h *Handler, defaultRule *FinalRule, UDPOverr
|
||||
if statConn != nil {
|
||||
counter = statConn.WriteCounter
|
||||
}
|
||||
if c, ok := iConn.(*net.PacketConnWrapper); ok {
|
||||
if c, ok := iConn.(*internet.PacketConnWrapper); ok {
|
||||
// If DialDest is a domain, it will be resolved in dialer
|
||||
// check this behavior and add it to map
|
||||
resolvedUDPAddr := utils.NewTypedSyncMap[string, net.Address]()
|
||||
@@ -563,7 +563,7 @@ func NewPacketWriter(conn net.Conn, h *Handler, defaultRule *FinalRule, UDPOverr
|
||||
}
|
||||
|
||||
type PacketWriter struct {
|
||||
*net.PacketConnWrapper
|
||||
*internet.PacketConnWrapper
|
||||
stats.Counter
|
||||
*Handler
|
||||
DefaultRule *FinalRule
|
||||
|
||||
@@ -151,7 +151,7 @@ func (c *Client) Process(ctx context.Context, link *transport.Link, dialer inter
|
||||
}
|
||||
defer conn.Close()
|
||||
uc := &wireguard.UDPConnClient{
|
||||
PacketConn: conn.(*net.PacketConnWrapper).PacketConn,
|
||||
PacketConn: conn.(*internet.PacketConnWrapper).PacketConn,
|
||||
Dest: conn.RemoteAddr().(*net.UDPAddr),
|
||||
}
|
||||
reader = uc
|
||||
|
||||
@@ -277,7 +277,6 @@ func (w *VisionReader) ReadMultiBuffer() (buf.MultiBuffer, error) {
|
||||
w.ob.CanSpliceCopy = 1
|
||||
}
|
||||
}
|
||||
SuppressOuterCloseNotify(w.conn)
|
||||
readerConn, readCounter, _ := UnwrapRawConn(w.conn)
|
||||
w.directReadCounter = readCounter
|
||||
w.Reader = buf.NewReader(readerConn)
|
||||
@@ -341,7 +340,6 @@ func (w *VisionWriter) WriteMultiBuffer(mb buf.MultiBuffer) error {
|
||||
// w.ob.CanSpliceCopy = 1
|
||||
// }
|
||||
}
|
||||
SuppressOuterCloseNotify(w.conn)
|
||||
rawConn, _, writerCounter := UnwrapRawConn(w.conn)
|
||||
w.Writer = buf.NewWriter(rawConn)
|
||||
w.directWriteCounter = writerCounter
|
||||
@@ -671,19 +669,6 @@ func XtlsFilterTls(buffer buf.MultiBuffer, trafficState *TrafficState, ctx conte
|
||||
}
|
||||
}
|
||||
|
||||
type CloseNotifySuppressor interface {
|
||||
SuppressCloseNotify()
|
||||
}
|
||||
|
||||
// Close our local TLS conn instance might send a incorrect close_notify alert
|
||||
// if the XTLS direct copy mode is enabled and cause TLS BAD_RECORD_MAC on users' browser
|
||||
// Close the underlying connection directly to avoid this issue.
|
||||
func SuppressOuterCloseNotify(conn net.Conn) {
|
||||
if suppressor, ok := stat.TryUnwrapStatsConn(conn).(CloseNotifySuppressor); ok {
|
||||
suppressor.SuppressCloseNotify()
|
||||
}
|
||||
}
|
||||
|
||||
// UnwrapRawConn support unwrap encryption, stats, mask wrappers, tls, utls, reality, proxyproto, uds-wrapper conn and get raw tcp/uds conn from it
|
||||
func UnwrapRawConn(conn net.Conn) (net.Conn, stats.Counter, stats.Counter) {
|
||||
var readCounter, writerCounter stats.Counter
|
||||
|
||||
@@ -2,6 +2,7 @@ package shadowsocks_2022
|
||||
|
||||
import (
|
||||
"context"
|
||||
"io"
|
||||
"time"
|
||||
|
||||
"github.com/xtls/xray-core/common"
|
||||
@@ -12,6 +13,9 @@ import (
|
||||
"github.com/xtls/xray-core/common/net"
|
||||
"github.com/xtls/xray-core/common/protocol"
|
||||
"github.com/xtls/xray-core/common/session"
|
||||
"github.com/xtls/xray-core/common/signal"
|
||||
"github.com/xtls/xray-core/common/task"
|
||||
"github.com/xtls/xray-core/common/utils"
|
||||
"github.com/xtls/xray-core/core"
|
||||
"github.com/xtls/xray-core/features/policy"
|
||||
"github.com/xtls/xray-core/features/routing"
|
||||
@@ -97,29 +101,35 @@ func (i *Inbound) processTCP(ctx context.Context, conn net.Conn, dispatcher rout
|
||||
return errors.New("unable to set read deadline").Base(err)
|
||||
}
|
||||
|
||||
// 1. Single read call for Salt + Fixed-length header chunk per SIP022 §3.1.4
|
||||
headerLen := i.method.KeySaltLength + RequestHeaderFixedChunkLength + AEADTagSize
|
||||
headerBuf := make([]byte, headerLen)
|
||||
n, err := conn.Read(headerBuf)
|
||||
if err != nil || n < headerLen {
|
||||
ResetTCPConn(conn)
|
||||
return errors.New("failed to read complete handshake header")
|
||||
}
|
||||
|
||||
var salt [32]byte
|
||||
copy(salt[:i.method.KeySaltLength], headerBuf[:i.method.KeySaltLength])
|
||||
saltSlice := salt[:i.method.KeySaltLength]
|
||||
fixedChunk := headerBuf[i.method.KeySaltLength:]
|
||||
|
||||
reader, reqHeader, err := InitServerStream(conn, i.method, i.psk, saltSlice, salt, fixedChunk, i.saltFilter)
|
||||
if err != nil {
|
||||
ResetTCPConn(conn)
|
||||
if _, err := io.ReadFull(conn, saltSlice); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
if !i.saltFilter.Check(salt) {
|
||||
return ErrSaltNotUnique
|
||||
}
|
||||
|
||||
sessionKey := DeriveSessionSubKey(i.psk, saltSlice, i.method.KeySaltLength)
|
||||
aead, err := i.method.NewAEAD(sessionKey)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
reader := NewStreamReader(conn, aead)
|
||||
|
||||
reqHeader, err := ReadClientRequestHeader(conn, reader)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
conn.SetReadDeadline(time.Time{})
|
||||
dest := reqHeader.Destination
|
||||
|
||||
writer := NewServerStreamWriter(conn, i.method, i.psk, saltSlice)
|
||||
writer, err := WriteTCPResponse(conn, i.method, i.psk, saltSlice, nil)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
ctx = log.ContextWithAccessMessage(ctx, &log.AccessMessage{
|
||||
From: conn.RemoteAddr(),
|
||||
@@ -136,17 +146,42 @@ func (i *Inbound) processTCP(ctx context.Context, conn net.Conn, dispatcher rout
|
||||
}
|
||||
|
||||
if len(reqHeader.EarlyData) > 0 {
|
||||
mb := buf.MergeBytes(nil, reqHeader.EarlyData)
|
||||
if err := link.Writer.WriteMultiBuffer(mb); err != nil {
|
||||
earlyBuf := buf.New()
|
||||
earlyBuf.Write(reqHeader.EarlyData)
|
||||
if err := link.Writer.WriteMultiBuffer(buf.MultiBuffer{earlyBuf}); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
|
||||
return TransportTCP(ctx, i.policyManager.ForLevel(uint32(i.user.Level)), reader, writer, link)
|
||||
sessionPolicy = i.policyManager.ForLevel(uint32(i.user.Level))
|
||||
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 {
|
||||
reader := buf.NewPacketReader(conn)
|
||||
udpConns := utils.NewTypedSyncMap[uint64, *udpConnEntry]()
|
||||
defer func() {
|
||||
udpConns.Range(func(key uint64, entry *udpConnEntry) bool {
|
||||
entry.timer.SetTimeout(0)
|
||||
return true
|
||||
})
|
||||
}()
|
||||
|
||||
reader := buf.NewReader(conn)
|
||||
for {
|
||||
mb, err := reader.ReadMultiBuffer()
|
||||
if err != nil {
|
||||
@@ -156,30 +191,75 @@ func (i *Inbound) processUDP(ctx context.Context, conn stat.Connection, dispatch
|
||||
|
||||
for _, b := range mb {
|
||||
decoded, err := i.udpCodec.DecodePacket(b.Bytes())
|
||||
b.Release()
|
||||
if err != nil || decoded.HeaderType != HeaderTypeClient {
|
||||
continue
|
||||
}
|
||||
|
||||
sessionItem := i.udpCodec.GetSession(decoded.SessionID)
|
||||
if sessionItem.User == nil {
|
||||
sessionItem.Lock()
|
||||
if sessionItem.User == nil {
|
||||
sessionItem.User = i.user
|
||||
}
|
||||
sessionItem.Unlock()
|
||||
}
|
||||
link, err := sessionItem.EnsureLink(ctx, conn, decoded.Destination, dispatcher, i.policyManager, func(dest net.Destination, payload []byte) ([]byte, error) {
|
||||
return i.udpCodec.EncodeServerPacket(decoded.SessionID, dest, payload)
|
||||
})
|
||||
if err != nil {
|
||||
b.Release()
|
||||
continue
|
||||
}
|
||||
|
||||
entry, ok := udpConns.Load(decoded.SessionID)
|
||||
if !ok {
|
||||
sessCtx, cancel := context.WithCancel(ctx)
|
||||
sessCtx = log.ContextWithAccessMessage(sessCtx, &log.AccessMessage{
|
||||
From: conn.RemoteAddr(),
|
||||
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)
|
||||
}
|
||||
}
|
||||
|
||||
entry.timer.Update()
|
||||
payloadBuf := buf.New()
|
||||
payloadBuf.Write(decoded.Payload)
|
||||
payloadBuf.UDP = &decoded.Destination
|
||||
_ = link.Writer.WriteMultiBuffer(buf.MultiBuffer{payloadBuf})
|
||||
b.Release()
|
||||
_ = entry.link.Writer.WriteMultiBuffer(buf.MultiBuffer{payloadBuf})
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -4,6 +4,7 @@ import (
|
||||
"context"
|
||||
"crypto/cipher"
|
||||
"encoding/binary"
|
||||
"io"
|
||||
"strconv"
|
||||
"strings"
|
||||
"sync"
|
||||
@@ -18,6 +19,8 @@ import (
|
||||
"github.com/xtls/xray-core/common/net"
|
||||
"github.com/xtls/xray-core/common/protocol"
|
||||
"github.com/xtls/xray-core/common/session"
|
||||
"github.com/xtls/xray-core/common/signal"
|
||||
"github.com/xtls/xray-core/common/task"
|
||||
"github.com/xtls/xray-core/common/utils"
|
||||
"github.com/xtls/xray-core/common/uuid"
|
||||
"github.com/xtls/xray-core/core"
|
||||
@@ -204,46 +207,64 @@ func (i *MultiUserInbound) processTCP(ctx context.Context, conn net.Conn, dispat
|
||||
return errors.New("unable to set read deadline").Base(err)
|
||||
}
|
||||
|
||||
// 1. Single read call for Salt + EIH + Fixed-length header chunk per SIP022 §3.1.4
|
||||
headerLen := i.method.KeySaltLength + AESBlockSize + RequestHeaderFixedChunkLength + AEADTagSize
|
||||
headerBuf := make([]byte, headerLen)
|
||||
n, err := conn.Read(headerBuf)
|
||||
if err != nil || n < headerLen {
|
||||
ResetTCPConn(conn)
|
||||
return errors.New("failed to read complete handshake header")
|
||||
}
|
||||
|
||||
// 1. Read Request Salt (16 or 32 bytes)
|
||||
var salt [32]byte
|
||||
copy(salt[:i.method.KeySaltLength], headerBuf[:i.method.KeySaltLength])
|
||||
saltSlice := salt[:i.method.KeySaltLength]
|
||||
eih := headerBuf[i.method.KeySaltLength : i.method.KeySaltLength+AESBlockSize]
|
||||
fixedChunk := headerBuf[i.method.KeySaltLength+AESBlockSize:]
|
||||
|
||||
decryptedHash, err := DecryptEIH(i.method, i.masterPSK, saltSlice, eih)
|
||||
if err != nil {
|
||||
ResetTCPConn(conn)
|
||||
if _, err := io.ReadFull(conn, saltSlice); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
if !i.saltFilter.Check(salt) {
|
||||
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 {
|
||||
return err
|
||||
}
|
||||
|
||||
var decryptedHash [AESBlockSize]byte
|
||||
block.Decrypt(decryptedHash[:], eih[:])
|
||||
|
||||
// Lookup user
|
||||
user, ok := i.usersByHash.Load(decryptedHash)
|
||||
if !ok {
|
||||
ResetTCPConn(conn)
|
||||
if !ok || user == nil {
|
||||
return ErrInvalidRequest
|
||||
}
|
||||
userPSK := user.Account.(*MemoryAccount).Key
|
||||
|
||||
reader, reqHeader, err := InitServerStream(conn, i.method, userPSK, saltSlice, salt, fixedChunk, i.saltFilter)
|
||||
// 3. Derive Session Subkey using matched user's PSK
|
||||
sessionKey := DeriveSessionSubKey(userPSK, saltSlice, i.method.KeySaltLength)
|
||||
aead, err := i.method.NewAEAD(sessionKey)
|
||||
if err != nil {
|
||||
ResetTCPConn(conn)
|
||||
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
|
||||
|
||||
writer := NewServerStreamWriter(conn, i.method, userPSK, saltSlice)
|
||||
// 6. Send Server Response Handshake
|
||||
writer, err := WriteTCPResponse(conn, i.method, userPSK, saltSlice, nil)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// Dispatch Connection to Xray routing with matched User
|
||||
// 7. Dispatch Connection to Xray routing with matched User
|
||||
inbound := session.InboundFromContext(ctx)
|
||||
inbound.User = user
|
||||
|
||||
@@ -262,17 +283,42 @@ func (i *MultiUserInbound) processTCP(ctx context.Context, conn net.Conn, dispat
|
||||
}
|
||||
|
||||
if len(reqHeader.EarlyData) > 0 {
|
||||
mb := buf.MergeBytes(nil, reqHeader.EarlyData)
|
||||
if err := link.Writer.WriteMultiBuffer(mb); err != nil {
|
||||
earlyBuf := buf.New()
|
||||
earlyBuf.Write(reqHeader.EarlyData)
|
||||
if err := link.Writer.WriteMultiBuffer(buf.MultiBuffer{earlyBuf}); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
|
||||
return TransportTCP(ctx, i.policyManager.ForLevel(user.Level), reader, writer, link)
|
||||
sessionPolicy = i.policyManager.ForLevel(user.Level)
|
||||
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 {
|
||||
reader := buf.NewPacketReader(conn)
|
||||
udpConns := utils.NewTypedSyncMap[uint64, *udpConnEntry]()
|
||||
defer func() {
|
||||
udpConns.Range(func(key uint64, entry *udpConnEntry) bool {
|
||||
entry.timer.SetTimeout(0)
|
||||
return true
|
||||
})
|
||||
}()
|
||||
|
||||
reader := buf.NewReader(conn)
|
||||
for {
|
||||
mb, err := reader.ReadMultiBuffer()
|
||||
if err != nil {
|
||||
@@ -296,61 +342,168 @@ func (i *MultiUserInbound) processUDP(ctx context.Context, conn stat.Connection,
|
||||
sessionID := binary.BigEndian.Uint64(rawHeader[:8])
|
||||
packetID := binary.BigEndian.Uint64(rawHeader[8:16])
|
||||
|
||||
// Replay protection & session lookup
|
||||
sessionItem := i.udpSessions.GetOrCreate(sessionID)
|
||||
|
||||
if !sessionItem.CheckPacketID(packetID) {
|
||||
sessionItem.Lock()
|
||||
if !sessionItem.Window.Check(packetID) {
|
||||
sessionItem.Unlock()
|
||||
b.Release()
|
||||
continue
|
||||
}
|
||||
|
||||
var userPSK []byte
|
||||
var currentUser *protocol.MemoryUser
|
||||
sessionItem.Lock()
|
||||
currentUser = sessionItem.User
|
||||
userPSK = sessionItem.UserPSK
|
||||
sessionItem.Unlock()
|
||||
|
||||
if currentUser == nil {
|
||||
if sessionItem.User != nil {
|
||||
currentUser = sessionItem.User
|
||||
userPSK = sessionItem.UserPSK
|
||||
sessionItem.Unlock()
|
||||
} else {
|
||||
sessionItem.Unlock()
|
||||
// Decrypt EIH
|
||||
decryptedHash := DecryptUDPEIH(i.udpMasterCipher, rawHeader[:], packetBytes[16:32])
|
||||
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
|
||||
idBlock.Decrypt(decryptedHash[:], packetBytes[16:32])
|
||||
|
||||
user, ok := i.usersByHash.Load(decryptedHash)
|
||||
if !ok {
|
||||
if !ok || user == nil {
|
||||
b.Release()
|
||||
continue
|
||||
}
|
||||
currentUser = user
|
||||
userPSK = user.Account.(*MemoryAccount).Key
|
||||
|
||||
sessionItem.Lock()
|
||||
sessionItem.User = user
|
||||
sessionItem.UserPSK = userPSK
|
||||
sessionItem.Unlock()
|
||||
}
|
||||
|
||||
decoded, err := sessionItem.DecryptAESPayload(i.method, userPSK, sessionID, packetID, rawHeader[:], packetBytes[32:])
|
||||
// Decrypt Body (with AEAD caching per session)
|
||||
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()
|
||||
if err != nil {
|
||||
if err != nil || len(bodyPlain) < 1+8+2 {
|
||||
continue
|
||||
}
|
||||
|
||||
sessionItem.Lock()
|
||||
if sessionItem.User == nil {
|
||||
sessionItem.User = currentUser
|
||||
sessionItem.UserPSK = userPSK
|
||||
}
|
||||
sessionItem.Window.Add(packetID)
|
||||
sessionItem.Unlock()
|
||||
|
||||
link, err := sessionItem.EnsureLink(ctx, conn, decoded.Destination, dispatcher, i.policyManager, func(replyDest net.Destination, payload []byte) ([]byte, error) {
|
||||
return i.encodeServerUDPPacket(sessionID, userPSK, replyDest, payload)
|
||||
})
|
||||
if 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 {
|
||||
continue
|
||||
}
|
||||
|
||||
payload := bodyPlain[offset+addrLen:]
|
||||
|
||||
entry, ok := udpConns.Load(sessionID)
|
||||
if !ok {
|
||||
sessCtx, cancel := context.WithCancel(ctx)
|
||||
inbound := session.InboundFromContext(sessCtx)
|
||||
inbound.User = currentUser
|
||||
|
||||
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.Write(decoded.Payload)
|
||||
pBuf.UDP = &decoded.Destination
|
||||
_ = link.Writer.WriteMultiBuffer(buf.MultiBuffer{pBuf})
|
||||
pBuf.Write(payload)
|
||||
_ = entry.link.Writer.WriteMultiBuffer(buf.MultiBuffer{pBuf})
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (i *MultiUserInbound) encodeServerUDPPacket(clientSessionID uint64, userPSK []byte, dest net.Destination, payload []byte) ([]byte, error) {
|
||||
return i.udpSessions.EncodeServerPacket(i.method, userPSK, clientSessionID, dest, payload)
|
||||
sessionItem := i.udpSessions.GetOrCreate(clientSessionID)
|
||||
if err := sessionItem.EnsureServerState(i.method, i.udpMasterCipher, nil, userPSK); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return sessionItem.EncodeServerPacket(i.method, clientSessionID, dest, payload)
|
||||
}
|
||||
|
||||
@@ -4,6 +4,7 @@ import (
|
||||
"context"
|
||||
"crypto/cipher"
|
||||
"encoding/binary"
|
||||
"io"
|
||||
"strconv"
|
||||
"time"
|
||||
|
||||
@@ -14,6 +15,9 @@ import (
|
||||
"github.com/xtls/xray-core/common/net"
|
||||
"github.com/xtls/xray-core/common/protocol"
|
||||
"github.com/xtls/xray-core/common/session"
|
||||
"github.com/xtls/xray-core/common/signal"
|
||||
"github.com/xtls/xray-core/common/task"
|
||||
"github.com/xtls/xray-core/common/utils"
|
||||
"github.com/xtls/xray-core/common/uuid"
|
||||
"github.com/xtls/xray-core/core"
|
||||
"github.com/xtls/xray-core/features/policy"
|
||||
@@ -31,17 +35,18 @@ type relayDest struct {
|
||||
destination net.Destination
|
||||
email string
|
||||
level uint32
|
||||
key []byte
|
||||
blockCipher cipher.Block
|
||||
}
|
||||
|
||||
type RelayInbound struct {
|
||||
networks []net.Network
|
||||
method *CipherMethod
|
||||
relayPSK []byte
|
||||
relayBlock cipher.Block
|
||||
destinations map[[AESBlockSize]byte]*relayDest
|
||||
udpSessions *UDPSessionManager
|
||||
policyManager policy.Manager
|
||||
networks []net.Network
|
||||
method *CipherMethod
|
||||
relayPSK []byte
|
||||
relayBlock cipher.Block
|
||||
destinations map[[AESBlockSize]byte]*relayDest
|
||||
rawDestinations []*RelayDestination
|
||||
policyManager policy.Manager
|
||||
}
|
||||
|
||||
func NewRelayServer(ctx context.Context, config *RelayServerConfig) (*RelayInbound, error) {
|
||||
@@ -73,13 +78,13 @@ func NewRelayServer(ctx context.Context, config *RelayServerConfig) (*RelayInbou
|
||||
|
||||
v := core.MustFromContext(ctx)
|
||||
i := &RelayInbound{
|
||||
networks: networks,
|
||||
method: method,
|
||||
relayPSK: relayPSK,
|
||||
relayBlock: relayBlock,
|
||||
destinations: make(map[[AESBlockSize]byte]*relayDest),
|
||||
udpSessions: NewUDPSessionManager(500 * time.Second),
|
||||
policyManager: v.GetFeature(policy.ManagerType()).(policy.Manager),
|
||||
networks: networks,
|
||||
method: method,
|
||||
relayPSK: relayPSK,
|
||||
relayBlock: relayBlock,
|
||||
destinations: make(map[[AESBlockSize]byte]*relayDest),
|
||||
rawDestinations: config.Destinations,
|
||||
policyManager: v.GetFeature(policy.ManagerType()).(policy.Manager),
|
||||
}
|
||||
|
||||
for idx, d := range config.Destinations {
|
||||
@@ -103,6 +108,7 @@ func NewRelayServer(ctx context.Context, config *RelayServerConfig) (*RelayInbou
|
||||
destination: net.TCPDestination(d.Address.AsAddress(), net.Port(d.Port)),
|
||||
email: d.Email,
|
||||
level: uint32(d.Level),
|
||||
key: destKey,
|
||||
blockCipher: destBlock,
|
||||
}
|
||||
}
|
||||
@@ -133,36 +139,28 @@ func (i *RelayInbound) processTCP(ctx context.Context, conn net.Conn, dispatcher
|
||||
return errors.New("unable to set read deadline").Base(err)
|
||||
}
|
||||
|
||||
// Read initial handshake in a single read call per SIP022 §3.1.3 & §3.1.4
|
||||
// Read Salt + Outer EIH
|
||||
needed := i.method.KeySaltLength + AESBlockSize
|
||||
requestHeader := buf.New()
|
||||
n, err := requestHeader.ReadFrom(conn)
|
||||
if err != nil {
|
||||
requestHeader.Release()
|
||||
ResetTCPConn(conn)
|
||||
var headerBuf [48]byte
|
||||
headerSlice := headerBuf[:needed]
|
||||
if _, err := io.ReadFull(conn, headerSlice); err != nil {
|
||||
return err
|
||||
}
|
||||
if int(n) < needed {
|
||||
requestHeader.Release()
|
||||
ResetTCPConn(conn)
|
||||
return ErrInvalidRequest
|
||||
}
|
||||
|
||||
headerSlice := requestHeader.Bytes()
|
||||
salt := headerSlice[:i.method.KeySaltLength]
|
||||
eih := headerSlice[i.method.KeySaltLength:needed]
|
||||
eih := headerSlice[i.method.KeySaltLength:]
|
||||
|
||||
decryptedHash, err := DecryptEIH(i.method, i.relayPSK, salt, eih)
|
||||
identitySubkey := DeriveIdentitySubKey(i.relayPSK, salt, i.method.KeySaltLength)
|
||||
block, err := i.method.NewBlock(identitySubkey)
|
||||
if err != nil {
|
||||
requestHeader.Release()
|
||||
ResetTCPConn(conn)
|
||||
return err
|
||||
}
|
||||
|
||||
var decryptedHash [AESBlockSize]byte
|
||||
block.Decrypt(decryptedHash[:], eih)
|
||||
|
||||
targetDest, ok := i.destinations[decryptedHash]
|
||||
if !ok {
|
||||
requestHeader.Release()
|
||||
ResetTCPConn(conn)
|
||||
return ErrInvalidRequest
|
||||
}
|
||||
conn.SetReadDeadline(time.Time{})
|
||||
@@ -184,26 +182,45 @@ func (i *RelayInbound) processTCP(ctx context.Context, conn net.Conn, dispatcher
|
||||
|
||||
link, err := dispatcher.Dispatch(ctx, targetDest.destination)
|
||||
if err != nil {
|
||||
requestHeader.Release()
|
||||
return err
|
||||
}
|
||||
|
||||
// Unwrap outer EIH: send client salt and remaining handshake bytes to next hop
|
||||
// in a single write call, satisfying downstream server's single-read handshake expectation (SIP022 §3.1.3).
|
||||
var saltCopy [32]byte
|
||||
copy(saltCopy[:i.method.KeySaltLength], salt)
|
||||
copy(requestHeader.Bytes()[AESBlockSize:AESBlockSize+i.method.KeySaltLength], saltCopy[:i.method.KeySaltLength])
|
||||
requestHeader.Advance(AESBlockSize)
|
||||
|
||||
if err := link.Writer.WriteMultiBuffer(buf.MultiBuffer{requestHeader}); err != nil {
|
||||
// Unwrap outer EIH: send client salt to next hop, stripping this hop's EIH
|
||||
saltBuf := buf.New()
|
||||
saltBuf.Write(salt)
|
||||
if err := link.Writer.WriteMultiBuffer(buf.MultiBuffer{saltBuf}); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
return TransportTCP(ctx, i.policyManager.ForLevel(targetDest.level), buf.NewReader(conn), buf.NewWriter(conn), link)
|
||||
sessionPolicy = i.policyManager.ForLevel(targetDest.level)
|
||||
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 {
|
||||
reader := buf.NewPacketReader(conn)
|
||||
udpConns := utils.NewTypedSyncMap[uint64, *udpConnEntry]()
|
||||
defer func() {
|
||||
udpConns.Range(func(key uint64, entry *udpConnEntry) bool {
|
||||
entry.timer.SetTimeout(0)
|
||||
return true
|
||||
})
|
||||
}()
|
||||
|
||||
reader := buf.NewReader(conn)
|
||||
for {
|
||||
mb, err := reader.ReadMultiBuffer()
|
||||
if err != nil {
|
||||
@@ -221,7 +238,11 @@ func (i *RelayInbound) processUDP(ctx context.Context, conn stat.Connection, dis
|
||||
var packetHeader [AESBlockSize]byte
|
||||
i.relayBlock.Decrypt(packetHeader[:], data[:AESBlockSize])
|
||||
|
||||
eiHeader := DecryptUDPEIH(i.relayBlock, packetHeader[:], data[AESBlockSize:2*AESBlockSize])
|
||||
var eiHeader [AESBlockSize]byte
|
||||
i.relayBlock.Decrypt(eiHeader[:], data[AESBlockSize:2*AESBlockSize])
|
||||
for idx := 0; idx < AESBlockSize; idx++ {
|
||||
eiHeader[idx] ^= packetHeader[idx]
|
||||
}
|
||||
|
||||
targetDest, ok := i.destinations[eiHeader]
|
||||
if !ok {
|
||||
@@ -242,24 +263,68 @@ func (i *RelayInbound) processUDP(ctx context.Context, conn stat.Connection, dis
|
||||
dest := targetDest.destination
|
||||
dest.Network = net.Network_UDP
|
||||
|
||||
sessionItem := i.udpSessions.GetOrCreate(sessionID)
|
||||
if sessionItem.User == nil {
|
||||
sessionItem.Lock()
|
||||
if sessionItem.User == nil {
|
||||
sessionItem.User = &protocol.MemoryUser{
|
||||
Email: targetDest.email,
|
||||
Level: targetDest.level,
|
||||
}
|
||||
entry, ok := udpConns.Load(sessionID)
|
||||
if !ok {
|
||||
sessCtx, cancel := context.WithCancel(ctx)
|
||||
inbound := session.InboundFromContext(sessCtx)
|
||||
inbound.User = &protocol.MemoryUser{
|
||||
Email: targetDest.email,
|
||||
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
|
||||
}
|
||||
|
||||
_ = link.Writer.WriteMultiBuffer(buf.MultiBuffer{b})
|
||||
entry.timer.Update()
|
||||
_ = entry.link.Writer.WriteMultiBuffer(buf.MultiBuffer{b})
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -61,14 +61,3 @@ func DeriveUserPSKHash(userPSK []byte) [AESBlockSize]byte {
|
||||
copy(out[:], h[:AESBlockSize])
|
||||
return out
|
||||
}
|
||||
|
||||
func DecryptEIH(method *CipherMethod, key, salt, eih []byte) ([AESBlockSize]byte, error) {
|
||||
identitySubkey := DeriveIdentitySubKey(key, salt, method.KeySaltLength)
|
||||
block, err := method.NewBlock(identitySubkey)
|
||||
if err != nil {
|
||||
return [AESBlockSize]byte{}, err
|
||||
}
|
||||
var decryptedHash [AESBlockSize]byte
|
||||
block.Decrypt(decryptedHash[:], eih)
|
||||
return decryptedHash, nil
|
||||
}
|
||||
|
||||
@@ -4,6 +4,7 @@ import (
|
||||
"context"
|
||||
"crypto/rand"
|
||||
"io"
|
||||
"time"
|
||||
|
||||
"github.com/xtls/xray-core/common"
|
||||
"github.com/xtls/xray-core/common/buf"
|
||||
@@ -45,12 +46,8 @@ func NewClient(ctx context.Context, config *ClientConfig) (*Outbound, error) {
|
||||
return nil, errors.New("invalid key: ", config.Key).Base(err)
|
||||
}
|
||||
|
||||
if method.IsChaCha && len(pskList) > 1 {
|
||||
return nil, errors.New("multi-key is not supported for chacha20-poly1305")
|
||||
}
|
||||
|
||||
finalPSK := pskList[len(pskList)-1]
|
||||
udpCodec, err := NewUDPPacketCodec(method, pskList)
|
||||
udpCodec, err := NewUDPPacketCodec(method, finalPSK)
|
||||
if err != nil {
|
||||
return nil, errors.New("failed to create udp packet codec").Base(err)
|
||||
}
|
||||
@@ -129,30 +126,18 @@ func (o *Outbound) Process(ctx context.Context, link *transport.Link, dialer int
|
||||
|
||||
requestDone := func() error {
|
||||
defer timer.SetTimeout(sessionPolicy.Timeouts.DownlinkOnly)
|
||||
|
||||
var initialPayload []byte
|
||||
var firstBuf *buf.Buffer
|
||||
var remainingMB buf.MultiBuffer
|
||||
if timeoutReader, ok := link.Reader.(buf.TimeoutReader); ok {
|
||||
if mb, err := timeoutReader.ReadMultiBufferTimeout(0); err == nil && !mb.IsEmpty() {
|
||||
remainingMB, firstBuf = buf.SplitFirst(mb)
|
||||
initialPayload = firstBuf.Bytes()
|
||||
}
|
||||
}
|
||||
|
||||
bodyWriter, err := WriteTCPRequest(conn, o.method, o.pskList, destination, clientSaltSlice, initialPayload)
|
||||
if firstBuf != nil {
|
||||
firstBuf.Release()
|
||||
}
|
||||
bufferedWriter := buf.NewBufferedWriter(buf.NewWriter(conn))
|
||||
bodyWriter, err := WriteTCPRequest(bufferedWriter, o.method, o.pskList, destination, clientSaltSlice, nil)
|
||||
if err != nil {
|
||||
buf.ReleaseMulti(remainingMB)
|
||||
return errors.New("failed to write request").Base(err)
|
||||
}
|
||||
|
||||
if !remainingMB.IsEmpty() {
|
||||
if err := bodyWriter.WriteMultiBuffer(remainingMB); err != nil {
|
||||
return err
|
||||
}
|
||||
if err = buf.CopyOnceTimeout(link.Reader, bodyWriter, time.Millisecond*100); err != nil && err != buf.ErrNotTimeoutReader && err != buf.ErrReadTimeout {
|
||||
return errors.New("failed to write A request payload").Base(err)
|
||||
}
|
||||
|
||||
if err := bufferedWriter.SetBuffered(false); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
return buf.Copy(link.Reader, bodyWriter, buf.UpdateActivity(timer))
|
||||
@@ -178,18 +163,13 @@ func (o *Outbound) Process(ctx context.Context, link *transport.Link, dialer int
|
||||
}
|
||||
|
||||
if network == net.Network_UDP {
|
||||
session, err := o.udpCodec.NewClientSession()
|
||||
if err != nil {
|
||||
return errors.New("failed to create client udp session").Base(err)
|
||||
}
|
||||
|
||||
requestDone := func() error {
|
||||
defer timer.SetTimeout(sessionPolicy.Timeouts.DownlinkOnly)
|
||||
|
||||
writer := &UDPWriter{
|
||||
Writer: conn,
|
||||
Destination: destination,
|
||||
Session: session,
|
||||
Codec: o.udpCodec,
|
||||
}
|
||||
|
||||
if err := buf.Copy(link.Reader, writer, buf.UpdateActivity(timer)); err != nil {
|
||||
@@ -202,8 +182,8 @@ func (o *Outbound) Process(ctx context.Context, link *transport.Link, dialer int
|
||||
defer timer.SetTimeout(sessionPolicy.Timeouts.UplinkOnly)
|
||||
|
||||
reader := &UDPReader{
|
||||
Reader: conn,
|
||||
Session: session,
|
||||
Reader: conn,
|
||||
Codec: o.udpCodec,
|
||||
}
|
||||
|
||||
if err := buf.Copy(reader, link.Writer, buf.UpdateActivity(timer)); err != nil {
|
||||
|
||||
+187
-441
@@ -16,13 +16,14 @@ import (
|
||||
)
|
||||
|
||||
type UDPCodec struct {
|
||||
method *CipherMethod
|
||||
pskList [][]byte
|
||||
psk []byte
|
||||
blockCipher cipher.Block
|
||||
blockCiphers []cipher.Block
|
||||
chachaCipher cipher.AEAD
|
||||
sessions *UDPSessionManager
|
||||
method *CipherMethod
|
||||
psk []byte
|
||||
blockCipher cipher.Block
|
||||
chachaCipher cipher.AEAD
|
||||
clientBodyCipher cipher.AEAD
|
||||
clientSessionID uint64
|
||||
nextPacketID atomic.Uint64
|
||||
sessions *UDPSessionManager
|
||||
}
|
||||
|
||||
type (
|
||||
@@ -47,23 +48,22 @@ func newUDPCodec(method *CipherMethod, psk []byte) (*UDPCodec, error) {
|
||||
return c, nil
|
||||
}
|
||||
|
||||
func NewUDPPacketCodec(method *CipherMethod, pskList [][]byte) (*UDPCodec, error) {
|
||||
if method.IsChaCha && len(pskList) > 1 {
|
||||
return nil, errors.New("multi-key is not supported for chacha20-poly1305")
|
||||
}
|
||||
finalPSK := pskList[len(pskList)-1]
|
||||
c, err := newUDPCodec(method, finalPSK)
|
||||
func NewUDPPacketCodec(method *CipherMethod, psk []byte) (*UDPCodec, error) {
|
||||
c, err := newUDPCodec(method, psk)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
c.pskList = pskList
|
||||
if len(pskList) > 1 {
|
||||
c.blockCiphers = make([]cipher.Block, len(pskList))
|
||||
for i, psk := range pskList {
|
||||
c.blockCiphers[i], err = method.NewBlock(psk)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
var sessID [8]byte
|
||||
if _, err := io.ReadFull(rand.Reader, sessID[:]); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
c.clientSessionID = binary.BigEndian.Uint64(sessID[:])
|
||||
|
||||
if !method.IsChaCha {
|
||||
clientBodyKey := DeriveSessionSubKey(psk, sessID[:], method.KeySaltLength)
|
||||
c.clientBodyCipher, err = method.NewAEAD(clientBodyKey)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
return c, nil
|
||||
@@ -78,37 +78,108 @@ func NewUDPServerCodec(method *CipherMethod, psk []byte, sessionTimeout time.Dur
|
||||
return c, nil
|
||||
}
|
||||
|
||||
func (c *UDPCodec) Sessions() *UDPSessionManager {
|
||||
return c.sessions
|
||||
}
|
||||
func (c *UDPCodec) EncodeClientPacket(dest net.Destination, payload []byte) (*buf.Buffer, error) {
|
||||
packetID := c.nextPacketID.Add(1)
|
||||
sessID := c.clientSessionID
|
||||
|
||||
func (c *UDPCodec) GetSession(sessionID uint64) *ServerUDPSession {
|
||||
if c.sessions == nil {
|
||||
return nil
|
||||
// Padding determination (e.g. DNS port 53 disguise)
|
||||
var paddingLen int
|
||||
if dest.Port == 53 && len(payload) < MaxPaddingLength {
|
||||
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 {
|
||||
SessionID uint64
|
||||
PacketID uint64
|
||||
HeaderType byte
|
||||
Timestamp uint64
|
||||
ClientSessionID uint64
|
||||
Destination net.Destination
|
||||
Payload []byte
|
||||
SessionID uint64
|
||||
PacketID uint64
|
||||
HeaderType byte
|
||||
Timestamp uint64
|
||||
Destination net.Destination
|
||||
Payload []byte
|
||||
}
|
||||
|
||||
func DecryptUDPEIH(block cipher.Block, rawHeader, eih []byte) [AESBlockSize]byte {
|
||||
var decryptedHash [AESBlockSize]byte
|
||||
block.Decrypt(decryptedHash[:], eih)
|
||||
for k := 0; k < AESBlockSize; k++ {
|
||||
decryptedHash[k] ^= rawHeader[k]
|
||||
}
|
||||
return decryptedHash
|
||||
}
|
||||
|
||||
func ParseAddressPort(data []byte) (net.Destination, int, error) {
|
||||
func parseAddressPort(data []byte) (net.Destination, int, error) {
|
||||
if len(data) < 1 {
|
||||
return net.Destination{}, 0, ErrPacketTooShort
|
||||
}
|
||||
@@ -149,9 +220,6 @@ func parsePlainUDPPacket(sessionID, packetID uint64, bodyPlain []byte) (DecodedU
|
||||
}
|
||||
|
||||
headerType := bodyPlain[0]
|
||||
if headerType != HeaderTypeClient && headerType != HeaderTypeServer {
|
||||
return DecodedUDPPacket{}, ErrBadHeaderType
|
||||
}
|
||||
epoch := binary.BigEndian.Uint64(bodyPlain[1:9])
|
||||
diff := int(math.Abs(float64(time.Now().Unix() - int64(epoch))))
|
||||
if diff > 30 {
|
||||
@@ -159,13 +227,11 @@ func parsePlainUDPPacket(sessionID, packetID uint64, bodyPlain []byte) (DecodedU
|
||||
}
|
||||
|
||||
offset := 9
|
||||
var clientSessionID uint64
|
||||
if headerType == HeaderTypeServer {
|
||||
if len(bodyPlain) < offset+8+2 {
|
||||
return DecodedUDPPacket{}, ErrPacketTooShort
|
||||
}
|
||||
clientSessionID = binary.BigEndian.Uint64(bodyPlain[offset : offset+8])
|
||||
offset += 8
|
||||
offset += 8 // skip clientSessionID
|
||||
}
|
||||
|
||||
paddingLen := int(binary.BigEndian.Uint16(bodyPlain[offset : offset+2]))
|
||||
@@ -176,20 +242,19 @@ func parsePlainUDPPacket(sessionID, packetID uint64, bodyPlain []byte) (DecodedU
|
||||
}
|
||||
offset += paddingLen
|
||||
|
||||
dest, addrLen, err := ParseAddressPort(bodyPlain[offset:])
|
||||
dest, addrLen, err := parseAddressPort(bodyPlain[offset:])
|
||||
if err != nil {
|
||||
return DecodedUDPPacket{}, err
|
||||
}
|
||||
payload := bodyPlain[offset+addrLen:]
|
||||
|
||||
return DecodedUDPPacket{
|
||||
SessionID: sessionID,
|
||||
PacketID: packetID,
|
||||
HeaderType: headerType,
|
||||
Timestamp: epoch,
|
||||
ClientSessionID: clientSessionID,
|
||||
Destination: dest,
|
||||
Payload: payload,
|
||||
SessionID: sessionID,
|
||||
PacketID: packetID,
|
||||
HeaderType: headerType,
|
||||
Timestamp: epoch,
|
||||
Destination: dest,
|
||||
Payload: payload,
|
||||
}, nil
|
||||
}
|
||||
|
||||
@@ -204,7 +269,7 @@ func (c *UDPCodec) DecodePacket(data []byte) (DecodedUDPPacket, error) {
|
||||
}
|
||||
nonce := data[:PacketNonceSize]
|
||||
ciphertext := data[PacketNonceSize:]
|
||||
plain, err := c.chachaCipher.Open(nil, nonce, ciphertext, nil)
|
||||
plain, err := c.chachaCipher.Open(ciphertext[:0], nonce, ciphertext, nil)
|
||||
if err != nil {
|
||||
return DecodedUDPPacket{}, errors.New("failed to decrypt chacha udp packet").Base(err)
|
||||
}
|
||||
@@ -215,22 +280,17 @@ func (c *UDPCodec) DecodePacket(data []byte) (DecodedUDPPacket, error) {
|
||||
sessionID := binary.BigEndian.Uint64(plain[:8])
|
||||
packetID := binary.BigEndian.Uint64(plain[8:16])
|
||||
|
||||
sessionItem := c.sessions.GetOrCreate(sessionID)
|
||||
if !sessionItem.CheckPacketID(packetID) {
|
||||
return DecodedUDPPacket{}, ErrPacketIdNotUnique
|
||||
if c.sessions != nil {
|
||||
sessionItem := c.sessions.GetOrCreate(sessionID)
|
||||
sessionItem.Lock()
|
||||
if !sessionItem.Window.CheckAndAdd(packetID) {
|
||||
sessionItem.Unlock()
|
||||
return DecodedUDPPacket{}, ErrPacketIdNotUnique
|
||||
}
|
||||
sessionItem.Unlock()
|
||||
}
|
||||
|
||||
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
|
||||
return parsePlainUDPPacket(sessionID, packetID, plain[16:])
|
||||
}
|
||||
|
||||
// AES mode
|
||||
@@ -239,52 +299,54 @@ func (c *UDPCodec) DecodePacket(data []byte) (DecodedUDPPacket, error) {
|
||||
sessionID := binary.BigEndian.Uint64(rawHeader[:8])
|
||||
packetID := binary.BigEndian.Uint64(rawHeader[8:16])
|
||||
|
||||
sessionItem := c.sessions.GetOrCreate(sessionID)
|
||||
if !sessionItem.CheckPacketID(packetID) {
|
||||
return DecodedUDPPacket{}, ErrPacketIdNotUnique
|
||||
}
|
||||
var bodyAead cipher.AEAD
|
||||
var sessionItem *ServerUDPSession
|
||||
|
||||
return sessionItem.DecryptAESPayload(c.method, c.psk, sessionID, packetID, rawHeader[:], data[16:])
|
||||
}
|
||||
if c.sessions != nil {
|
||||
sessionItem = c.sessions.GetOrCreate(sessionID)
|
||||
sessionItem.Lock()
|
||||
if !sessionItem.Window.Check(packetID) {
|
||||
sessionItem.Unlock()
|
||||
return DecodedUDPPacket{}, ErrPacketIdNotUnique
|
||||
}
|
||||
sessionItem.Unlock()
|
||||
|
||||
func (s *ServerUDPSession) DecryptAESPayload(method *CipherMethod, psk []byte, sessionID, packetID uint64, rawHeader, bodyCipher []byte) (DecodedUDPPacket, error) {
|
||||
bodyAead := s.clientBodyCipher
|
||||
isNewCipher := false
|
||||
if bodyAead == nil {
|
||||
bodyKey := DeriveSessionSubKey(psk, rawHeader[:8], method.KeySaltLength)
|
||||
bodyAead = sessionItem.GetRemoteCipher()
|
||||
if bodyAead == nil {
|
||||
bodyKey := DeriveSessionSubKey(c.psk, rawHeader[:8], c.method.KeySaltLength)
|
||||
var err error
|
||||
bodyAead, err = c.method.NewAEAD(bodyKey)
|
||||
if err != nil {
|
||||
return DecodedUDPPacket{}, err
|
||||
}
|
||||
sessionItem.SetRemoteCipher(bodyAead)
|
||||
}
|
||||
} else {
|
||||
bodyKey := DeriveSessionSubKey(c.psk, rawHeader[:8], c.method.KeySaltLength)
|
||||
var err error
|
||||
bodyAead, err = method.NewAEAD(bodyKey)
|
||||
bodyAead, err = c.method.NewAEAD(bodyKey)
|
||||
if err != nil {
|
||||
return DecodedUDPPacket{}, err
|
||||
}
|
||||
isNewCipher = true
|
||||
}
|
||||
|
||||
bodyNonce := rawHeader[4:16]
|
||||
bodyPlain, err := bodyAead.Open(nil, bodyNonce, bodyCipher, nil)
|
||||
bodyCipher := data[16:]
|
||||
bodyPlain, err := bodyAead.Open(bodyCipher[:0], 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 sessionItem != nil {
|
||||
sessionItem.Lock()
|
||||
sessionItem.Window.Add(packetID)
|
||||
sessionItem.Unlock()
|
||||
}
|
||||
|
||||
if decoded.HeaderType != HeaderTypeClient {
|
||||
return DecodedUDPPacket{}, ErrBadHeaderType
|
||||
}
|
||||
|
||||
s.AddPacketID(packetID)
|
||||
|
||||
if isNewCipher {
|
||||
s.clientBodyCipher = bodyAead
|
||||
}
|
||||
|
||||
return decoded, nil
|
||||
return parsePlainUDPPacket(sessionID, packetID, bodyPlain)
|
||||
}
|
||||
|
||||
func (s *ServerUDPSession) EnsureServerState(method *CipherMethod, psk []byte) error {
|
||||
func (s *ServerUDPSession) EnsureServerState(method *CipherMethod, headerBlock cipher.Block, chachaCipher cipher.AEAD, psk []byte) error {
|
||||
s.Lock()
|
||||
defer s.Unlock()
|
||||
if s.ServerSessionID != 0 {
|
||||
@@ -301,29 +363,23 @@ func (s *ServerUDPSession) EnsureServerState(method *CipherMethod, psk []byte) e
|
||||
}
|
||||
}
|
||||
if method.IsChaCha {
|
||||
var err error
|
||||
s.serverChaCha, err = method.NewUDPCipher(psk)
|
||||
return err
|
||||
}
|
||||
|
||||
var err error
|
||||
s.serverHeaderBlock, err = method.NewBlock(psk)
|
||||
if err != nil {
|
||||
s.ServerSessionID = 0
|
||||
return err
|
||||
}
|
||||
bodyKey := DeriveSessionSubKey(psk, sidBuf[:], method.KeySaltLength)
|
||||
s.serverBodyCipher, err = method.NewAEAD(bodyKey)
|
||||
if err != nil {
|
||||
s.ServerSessionID = 0
|
||||
return err
|
||||
s.ServerChaCha = chachaCipher
|
||||
} else {
|
||||
s.ServerBlockCipher = headerBlock
|
||||
bodyKey := DeriveSessionSubKey(psk, sidBuf[:], method.KeySaltLength)
|
||||
bodyAead, err := method.NewAEAD(bodyKey)
|
||||
if err != nil {
|
||||
s.ServerSessionID = 0
|
||||
return err
|
||||
}
|
||||
s.ServerCipher = bodyAead
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *ServerUDPSession) EncodeServerPacket(method *CipherMethod, clientSessionID uint64, dest net.Destination, payload []byte) ([]byte, error) {
|
||||
serverSessionID := s.ServerSessionID
|
||||
serverPacketID := s.ServerPacketID.Add(1) - 1
|
||||
serverPacketID := s.ServerPacketID.Add(1)
|
||||
|
||||
if method.IsChaCha {
|
||||
var nonce [PacketNonceSize]byte
|
||||
@@ -348,7 +404,7 @@ func (s *ServerUDPSession) EncodeServerPacket(method *CipherMethod, clientSessio
|
||||
}
|
||||
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))
|
||||
copy(res[:PacketNonceSize], nonce[:])
|
||||
copy(res[PacketNonceSize:], sealed)
|
||||
@@ -361,7 +417,7 @@ func (s *ServerUDPSession) EncodeServerPacket(method *CipherMethod, clientSessio
|
||||
binary.BigEndian.PutUint64(rawHeader[8:16], serverPacketID)
|
||||
|
||||
var encryptedHeader [16]byte
|
||||
s.serverHeaderBlock.Encrypt(encryptedHeader[:], rawHeader[:])
|
||||
s.ServerBlockCipher.Encrypt(encryptedHeader[:], rawHeader[:])
|
||||
|
||||
bodyBuf := buf.New()
|
||||
defer bodyBuf.Release()
|
||||
@@ -379,7 +435,7 @@ func (s *ServerUDPSession) EncodeServerPacket(method *CipherMethod, clientSessio
|
||||
bodyBuf.Write(payload)
|
||||
|
||||
bodyNonce := rawHeader[4:16]
|
||||
sealedBody := s.serverBodyCipher.Seal(nil, bodyNonce, bodyBuf.Bytes(), nil)
|
||||
sealedBody := s.ServerCipher.Seal(nil, bodyNonce, bodyBuf.Bytes(), nil)
|
||||
|
||||
res := make([]byte, 16+len(sealedBody))
|
||||
copy(res[:16], encryptedHeader[:])
|
||||
@@ -388,327 +444,17 @@ func (s *ServerUDPSession) EncodeServerPacket(method *CipherMethod, clientSessio
|
||||
}
|
||||
|
||||
func (c *UDPCodec) EncodeServerPacket(clientSessionID uint64, dest net.Destination, payload []byte) ([]byte, error) {
|
||||
return c.sessions.EncodeServerPacket(c.method, c.psk, clientSessionID, dest, payload)
|
||||
}
|
||||
|
||||
type serverSessionState struct {
|
||||
sessionID uint64
|
||||
window *SlidingWindow
|
||||
cipher cipher.AEAD
|
||||
lastSeen atomic.Int64
|
||||
}
|
||||
|
||||
func (st *serverSessionState) check(packetID uint64) bool {
|
||||
if st.window == nil {
|
||||
st.window = new(SlidingWindow)
|
||||
}
|
||||
return st.window.Check(packetID)
|
||||
}
|
||||
|
||||
func (st *serverSessionState) add(packetID uint64) {
|
||||
if st.window == nil {
|
||||
st.window = new(SlidingWindow)
|
||||
}
|
||||
st.window.Add(packetID)
|
||||
}
|
||||
|
||||
type ClientUDPSession struct {
|
||||
codec *UDPCodec
|
||||
clientSessionID uint64
|
||||
nextPacketID atomic.Uint64
|
||||
clientBodyCipher cipher.AEAD
|
||||
current atomic.Pointer[serverSessionState]
|
||||
old atomic.Pointer[serverSessionState]
|
||||
}
|
||||
|
||||
func (c *UDPCodec) NewClientSession() (*ClientUDPSession, error) {
|
||||
var sessID [8]byte
|
||||
if _, err := io.ReadFull(rand.Reader, sessID[:]); err != nil {
|
||||
sessionItem := c.sessions.GetOrCreate(clientSessionID)
|
||||
if err := sessionItem.EnsureServerState(c.method, c.blockCipher, c.chachaCipher, c.psk); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
clientSessionID := binary.BigEndian.Uint64(sessID[:])
|
||||
|
||||
var clientBodyCipher cipher.AEAD
|
||||
var err error
|
||||
if !c.method.IsChaCha {
|
||||
finalPSK := c.psk
|
||||
clientBodyKey := DeriveSessionSubKey(finalPSK, sessID[:], c.method.KeySaltLength)
|
||||
clientBodyCipher, err = c.method.NewAEAD(clientBodyKey)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
|
||||
return &ClientUDPSession{
|
||||
codec: c,
|
||||
clientSessionID: clientSessionID,
|
||||
clientBodyCipher: clientBodyCipher,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (s *ClientUDPSession) getServerSession(sessionID uint64, now int64) (*serverSessionState, error) {
|
||||
cur := s.current.Load()
|
||||
if cur != nil && cur.sessionID == sessionID {
|
||||
return cur, nil
|
||||
}
|
||||
|
||||
old := s.old.Load()
|
||||
if old != nil && old.sessionID == sessionID {
|
||||
if now-old.lastSeen.Load() > 60 {
|
||||
s.old.CompareAndSwap(old, nil)
|
||||
return nil, errors.New("old server session expired")
|
||||
}
|
||||
return old, nil
|
||||
}
|
||||
|
||||
// New server session:
|
||||
// Spec §3.2.4: reject newer server sessions when the last packet received from the old session is less than 1 minute old.
|
||||
if old != nil && now-old.lastSeen.Load() < 60 {
|
||||
return nil, errors.New("newer server session rejected: old session is less than 1 minute old")
|
||||
}
|
||||
|
||||
var bodyAead cipher.AEAD
|
||||
if !s.codec.method.IsChaCha {
|
||||
var sessBytes [8]byte
|
||||
binary.BigEndian.PutUint64(sessBytes[:], sessionID)
|
||||
bodyKey := DeriveSessionSubKey(s.codec.psk, sessBytes[:], s.codec.method.KeySaltLength)
|
||||
var err error
|
||||
bodyAead, err = s.codec.method.NewAEAD(bodyKey)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
|
||||
newState := &serverSessionState{
|
||||
sessionID: sessionID,
|
||||
cipher: bodyAead,
|
||||
}
|
||||
newState.lastSeen.Store(now)
|
||||
|
||||
if cur == nil {
|
||||
s.current.CompareAndSwap(nil, newState)
|
||||
return s.current.Load(), nil
|
||||
}
|
||||
|
||||
s.old.Store(cur)
|
||||
s.current.Store(newState)
|
||||
return newState, nil
|
||||
}
|
||||
|
||||
func (s *ClientUDPSession) ClientSessionID() uint64 {
|
||||
return s.clientSessionID
|
||||
}
|
||||
|
||||
func (s *ClientUDPSession) EncodePacket(dest net.Destination, payload []byte) (*buf.Buffer, error) {
|
||||
packetID := s.nextPacketID.Add(1) - 1
|
||||
sessID := s.clientSessionID
|
||||
|
||||
var paddingLen int
|
||||
if dest.Port == 53 && len(payload) < MaxPaddingLength {
|
||||
paddingLen = mrand.IntN(MaxPaddingLength) + 1
|
||||
}
|
||||
|
||||
addrPortLen := AddrPortLength(dest)
|
||||
|
||||
if s.codec.method.IsChaCha {
|
||||
totalLen := PacketNonceSize + 27 + paddingLen + addrPortLen + len(payload) + AEADTagSize
|
||||
if totalLen > buf.Size {
|
||||
return nil, ErrPacketTooLarge
|
||||
}
|
||||
|
||||
outBuf := buf.New()
|
||||
|
||||
var nonce [PacketNonceSize]byte
|
||||
if _, err := io.ReadFull(rand.Reader, nonce[:]); err != nil {
|
||||
outBuf.Release()
|
||||
return nil, err
|
||||
}
|
||||
outBuf.Write(nonce[:])
|
||||
|
||||
var hdr [16 + 1 + 8 + 2]byte
|
||||
binary.BigEndian.PutUint64(hdr[0:8], sessID)
|
||||
binary.BigEndian.PutUint64(hdr[8:16], packetID)
|
||||
hdr[16] = HeaderTypeClient
|
||||
binary.BigEndian.PutUint64(hdr[17:25], uint64(time.Now().Unix()))
|
||||
binary.BigEndian.PutUint16(hdr[25:27], uint16(paddingLen))
|
||||
outBuf.Write(hdr[:])
|
||||
if paddingLen > 0 {
|
||||
outBuf.Write(zeroPadding[:paddingLen])
|
||||
}
|
||||
|
||||
if err := WriteAddressPort(outBuf, dest); err != nil {
|
||||
outBuf.Release()
|
||||
return nil, err
|
||||
}
|
||||
outBuf.Write(payload)
|
||||
|
||||
plainBytes := outBuf.Bytes()[PacketNonceSize:]
|
||||
outBuf.Extend(int32(s.codec.chachaCipher.Overhead()))
|
||||
s.codec.chachaCipher.Seal(plainBytes[:0], nonce[:], plainBytes, nil)
|
||||
return outBuf, nil
|
||||
}
|
||||
|
||||
// AES mode
|
||||
var sessBytes [8]byte
|
||||
binary.BigEndian.PutUint64(sessBytes[:], sessID)
|
||||
|
||||
var rawHeader [16]byte
|
||||
copy(rawHeader[:8], sessBytes[:])
|
||||
binary.BigEndian.PutUint64(rawHeader[8:16], packetID)
|
||||
|
||||
eihCount := 0
|
||||
if len(s.codec.pskList) > 1 {
|
||||
eihCount = len(s.codec.pskList) - 1
|
||||
}
|
||||
|
||||
totalLen := 16 + eihCount*16 + 11 + paddingLen + addrPortLen + len(payload) + AEADTagSize
|
||||
if totalLen > buf.Size {
|
||||
return nil, ErrPacketTooLarge
|
||||
}
|
||||
|
||||
outBuf := buf.New()
|
||||
|
||||
if len(s.codec.pskList) > 1 {
|
||||
var encryptedHeader [16]byte
|
||||
s.codec.blockCiphers[0].Encrypt(encryptedHeader[:], rawHeader[:])
|
||||
outBuf.Write(encryptedHeader[:])
|
||||
|
||||
for i := 0; i < len(s.codec.pskList)-1; i++ {
|
||||
nextPSK := s.codec.pskList[i+1]
|
||||
pskHash := DeriveUserPSKHash(nextPSK)
|
||||
var eihPlain [16]byte
|
||||
for k := 0; k < 16; k++ {
|
||||
eihPlain[k] = pskHash[k] ^ rawHeader[k]
|
||||
}
|
||||
var encryptedEIH [16]byte
|
||||
s.codec.blockCiphers[i].Encrypt(encryptedEIH[:], eihPlain[:])
|
||||
outBuf.Write(encryptedEIH[:])
|
||||
}
|
||||
} else {
|
||||
var encryptedHeader [16]byte
|
||||
s.codec.blockCipher.Encrypt(encryptedHeader[:], rawHeader[:])
|
||||
outBuf.Write(encryptedHeader[:])
|
||||
}
|
||||
|
||||
bodyAead := s.clientBodyCipher
|
||||
|
||||
var hdr [1 + 8 + 2]byte
|
||||
hdr[0] = HeaderTypeClient
|
||||
binary.BigEndian.PutUint64(hdr[1:9], uint64(time.Now().Unix()))
|
||||
binary.BigEndian.PutUint16(hdr[9:11], uint16(paddingLen))
|
||||
outBuf.Write(hdr[:])
|
||||
if paddingLen > 0 {
|
||||
outBuf.Write(zeroPadding[:paddingLen])
|
||||
}
|
||||
|
||||
if err := WriteAddressPort(outBuf, dest); err != nil {
|
||||
outBuf.Release()
|
||||
return nil, err
|
||||
}
|
||||
outBuf.Write(payload)
|
||||
|
||||
headerOffset := 16 + eihCount*16
|
||||
plainBytes := outBuf.Bytes()[headerOffset:]
|
||||
bodyNonce := rawHeader[4:16]
|
||||
outBuf.Extend(int32(bodyAead.Overhead()))
|
||||
bodyAead.Seal(plainBytes[:0], bodyNonce, plainBytes, nil)
|
||||
return outBuf, nil
|
||||
}
|
||||
|
||||
func (s *ClientUDPSession) DecodePacket(data []byte) (DecodedUDPPacket, error) {
|
||||
if len(data) < PacketMinimalHeaderSize {
|
||||
return DecodedUDPPacket{}, ErrPacketTooShort
|
||||
}
|
||||
|
||||
if s.codec.method.IsChaCha {
|
||||
if len(data) < PacketNonceSize+AEADTagSize {
|
||||
return DecodedUDPPacket{}, ErrPacketTooShort
|
||||
}
|
||||
nonce := data[:PacketNonceSize]
|
||||
ciphertext := data[PacketNonceSize:]
|
||||
plain, err := s.codec.chachaCipher.Open(nil, nonce, ciphertext, nil)
|
||||
if err != nil {
|
||||
return DecodedUDPPacket{}, errors.New("failed to decrypt chacha udp packet").Base(err)
|
||||
}
|
||||
if len(plain) < 16+1+8+2 {
|
||||
return DecodedUDPPacket{}, ErrPacketTooShort
|
||||
}
|
||||
|
||||
sessionID := binary.BigEndian.Uint64(plain[:8])
|
||||
packetID := binary.BigEndian.Uint64(plain[8:16])
|
||||
|
||||
now := time.Now().Unix()
|
||||
st, err := s.getServerSession(sessionID, now)
|
||||
if err != nil {
|
||||
return DecodedUDPPacket{}, err
|
||||
}
|
||||
if !st.check(packetID) {
|
||||
return DecodedUDPPacket{}, ErrPacketIdNotUnique
|
||||
}
|
||||
|
||||
decoded, err := parsePlainUDPPacket(sessionID, packetID, plain[16:])
|
||||
if err != nil {
|
||||
return DecodedUDPPacket{}, err
|
||||
}
|
||||
|
||||
if decoded.HeaderType != HeaderTypeServer {
|
||||
return DecodedUDPPacket{}, ErrBadHeaderType
|
||||
}
|
||||
if decoded.ClientSessionID != s.clientSessionID {
|
||||
return DecodedUDPPacket{}, errors.New("client session ID mismatch")
|
||||
}
|
||||
|
||||
st.add(packetID)
|
||||
st.lastSeen.Store(now)
|
||||
|
||||
return decoded, nil
|
||||
}
|
||||
|
||||
// AES mode
|
||||
var rawHeader [16]byte
|
||||
s.codec.blockCipher.Decrypt(rawHeader[:], data[:16])
|
||||
sessionID := binary.BigEndian.Uint64(rawHeader[:8])
|
||||
packetID := binary.BigEndian.Uint64(rawHeader[8:16])
|
||||
|
||||
now := time.Now().Unix()
|
||||
st, err := s.getServerSession(sessionID, now)
|
||||
if err != nil {
|
||||
return DecodedUDPPacket{}, err
|
||||
}
|
||||
if !st.check(packetID) {
|
||||
return DecodedUDPPacket{}, ErrPacketIdNotUnique
|
||||
}
|
||||
bodyAead := st.cipher
|
||||
|
||||
bodyNonce := rawHeader[4:16]
|
||||
bodyCipher := data[16:]
|
||||
bodyPlain, err := bodyAead.Open(nil, bodyNonce, bodyCipher, nil)
|
||||
if err != nil {
|
||||
return DecodedUDPPacket{}, errors.New("failed to decrypt aes udp body").Base(err)
|
||||
}
|
||||
|
||||
decoded, err := parsePlainUDPPacket(sessionID, packetID, bodyPlain)
|
||||
if err != nil {
|
||||
return DecodedUDPPacket{}, err
|
||||
}
|
||||
|
||||
if decoded.HeaderType != HeaderTypeServer {
|
||||
return DecodedUDPPacket{}, ErrBadHeaderType
|
||||
}
|
||||
if decoded.ClientSessionID != s.clientSessionID {
|
||||
return DecodedUDPPacket{}, errors.New("client session ID mismatch")
|
||||
}
|
||||
|
||||
st.add(packetID)
|
||||
st.lastSeen.Store(now)
|
||||
|
||||
return decoded, nil
|
||||
return sessionItem.EncodeServerPacket(c.method, clientSessionID, dest, payload)
|
||||
}
|
||||
|
||||
type UDPWriter struct {
|
||||
Writer io.Writer
|
||||
Destination net.Destination
|
||||
Session *ClientUDPSession
|
||||
Codec *UDPPacketCodec
|
||||
}
|
||||
|
||||
func (w *UDPWriter) WriteMultiBuffer(mb buf.MultiBuffer) error {
|
||||
@@ -722,7 +468,7 @@ func (w *UDPWriter) WriteMultiBuffer(mb buf.MultiBuffer) error {
|
||||
if b.UDP != nil {
|
||||
dest = *b.UDP
|
||||
}
|
||||
pktBuf, err := w.Session.EncodePacket(dest, b.Bytes())
|
||||
pktBuf, err := w.Codec.EncodeClientPacket(dest, b.Bytes())
|
||||
b.Release()
|
||||
if err != nil {
|
||||
buf.ReleaseMulti(mb)
|
||||
@@ -739,8 +485,8 @@ func (w *UDPWriter) WriteMultiBuffer(mb buf.MultiBuffer) error {
|
||||
}
|
||||
|
||||
type UDPReader struct {
|
||||
Reader io.Reader
|
||||
Session *ClientUDPSession
|
||||
Reader io.Reader
|
||||
Codec *UDPPacketCodec
|
||||
}
|
||||
|
||||
func (r *UDPReader) ReadMultiBuffer() (buf.MultiBuffer, error) {
|
||||
@@ -752,7 +498,7 @@ func (r *UDPReader) ReadMultiBuffer() (buf.MultiBuffer, error) {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
decoded, err := r.Session.DecodePacket(buffer.Bytes())
|
||||
decoded, err := r.Codec.DecodePacket(buffer.Bytes())
|
||||
if err != nil {
|
||||
buffer.Release()
|
||||
continue
|
||||
|
||||
@@ -2,11 +2,9 @@ package shadowsocks_2022_test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/rand"
|
||||
"encoding/base64"
|
||||
"encoding/binary"
|
||||
"errors"
|
||||
"io"
|
||||
gonet "net"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
@@ -271,106 +269,3 @@ func (c *dummyStatConn) WriteMultiBuffer(mb buf.MultiBuffer) error {
|
||||
}
|
||||
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))
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -6,11 +6,8 @@ import (
|
||||
"sync/atomic"
|
||||
"time"
|
||||
|
||||
"github.com/xtls/xray-core/common/net"
|
||||
"github.com/xtls/xray-core/common/protocol"
|
||||
"github.com/xtls/xray-core/common/signal"
|
||||
"github.com/xtls/xray-core/common/utils"
|
||||
"github.com/xtls/xray-core/transport"
|
||||
)
|
||||
|
||||
const (
|
||||
@@ -77,42 +74,30 @@ func (f *SlidingWindow) CheckAndAdd(counter uint64) bool {
|
||||
|
||||
type ServerUDPSession struct {
|
||||
sync.Mutex
|
||||
SessionID uint64
|
||||
Window *SlidingWindow
|
||||
User *protocol.MemoryUser
|
||||
UserPSK []byte
|
||||
LastActive atomic.Int64 // Unix timestamp in seconds
|
||||
|
||||
clientBodyCipher cipher.AEAD
|
||||
SessionID uint64
|
||||
RemoteCipher atomic.Pointer[cipher.AEAD]
|
||||
Window SlidingWindow
|
||||
User *protocol.MemoryUser
|
||||
UserPSK []byte
|
||||
LastActive atomic.Int64 // Unix timestamp in seconds
|
||||
|
||||
ServerSessionID uint64
|
||||
ServerPacketID atomic.Uint64
|
||||
serverBodyCipher cipher.AEAD
|
||||
serverHeaderBlock cipher.Block
|
||||
serverChaCha cipher.AEAD
|
||||
|
||||
manager *UDPSessionManager
|
||||
link atomic.Pointer[transport.Link]
|
||||
timer *signal.ActivityTimer
|
||||
currentConn atomic.Value // stores stat.Connection
|
||||
ServerCipher cipher.AEAD
|
||||
ServerBlockCipher cipher.Block
|
||||
ServerChaCha cipher.AEAD
|
||||
}
|
||||
|
||||
func (s *ServerUDPSession) CheckPacketID(packetID uint64) bool {
|
||||
s.Lock()
|
||||
defer s.Unlock()
|
||||
if s.Window == nil {
|
||||
s.Window = new(SlidingWindow)
|
||||
func (s *ServerUDPSession) GetRemoteCipher() cipher.AEAD {
|
||||
ptr := s.RemoteCipher.Load()
|
||||
if ptr == nil {
|
||||
return nil
|
||||
}
|
||||
return s.Window.Check(packetID)
|
||||
return *ptr
|
||||
}
|
||||
|
||||
func (s *ServerUDPSession) AddPacketID(packetID uint64) {
|
||||
s.Lock()
|
||||
defer s.Unlock()
|
||||
if s.Window == nil {
|
||||
s.Window = new(SlidingWindow)
|
||||
}
|
||||
s.Window.Add(packetID)
|
||||
func (s *ServerUDPSession) SetRemoteCipher(c cipher.AEAD) {
|
||||
s.RemoteCipher.Store(&c)
|
||||
}
|
||||
|
||||
type UDPSessionManager struct {
|
||||
@@ -137,7 +122,6 @@ func (m *UDPSessionManager) GetOrCreate(sessionID uint64) *ServerUDPSession {
|
||||
|
||||
s := &ServerUDPSession{
|
||||
SessionID: sessionID,
|
||||
manager: m,
|
||||
}
|
||||
s.LastActive.Store(now)
|
||||
|
||||
@@ -164,7 +148,6 @@ func (m *UDPSessionManager) cleanup(now int64) {
|
||||
m.sessions.Range(func(k uint64, v *ServerUDPSession) bool {
|
||||
if now-v.LastActive.Load() > timeoutSec {
|
||||
m.sessions.Delete(k)
|
||||
v.Close()
|
||||
}
|
||||
return true
|
||||
})
|
||||
@@ -173,11 +156,3 @@ func (m *UDPSessionManager) cleanup(now int64) {
|
||||
func (m *UDPSessionManager) Delete(sessionID uint64) {
|
||||
m.sessions.Delete(sessionID)
|
||||
}
|
||||
|
||||
func (m *UDPSessionManager) EncodeServerPacket(method *CipherMethod, psk []byte, clientSessionID uint64, dest net.Destination, payload []byte) ([]byte, error) {
|
||||
sessionItem := m.GetOrCreate(clientSessionID)
|
||||
if err := sessionItem.EnsureServerState(method, psk); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return sessionItem.EncodeServerPacket(method, clientSessionID, dest, payload)
|
||||
}
|
||||
|
||||
@@ -2,161 +2,18 @@ package shadowsocks_2022
|
||||
|
||||
import (
|
||||
"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/log"
|
||||
"github.com/xtls/xray-core/common/net"
|
||||
"github.com/xtls/xray-core/common/session"
|
||||
"github.com/xtls/xray-core/common/signal"
|
||||
"github.com/xtls/xray-core/features/policy"
|
||||
"github.com/xtls/xray-core/features/routing"
|
||||
"github.com/xtls/xray-core/proxy"
|
||||
"github.com/xtls/xray-core/transport"
|
||||
"github.com/xtls/xray-core/transport/internet/stat"
|
||||
)
|
||||
|
||||
func (s *ServerUDPSession) UpdateConn(conn stat.Connection) {
|
||||
if s.currentConn.Load() == nil {
|
||||
s.currentConn.Store(conn)
|
||||
}
|
||||
if s.timer != nil {
|
||||
s.timer.Update()
|
||||
}
|
||||
}
|
||||
|
||||
func (s *ServerUDPSession) WriteToClient(b []byte) error {
|
||||
connVal := s.currentConn.Load()
|
||||
if connVal == nil {
|
||||
return errors.New("client connection closed")
|
||||
}
|
||||
conn, ok := connVal.(stat.Connection)
|
||||
if !ok || conn == nil {
|
||||
return errors.New("client connection closed")
|
||||
}
|
||||
_, err := conn.Write(b)
|
||||
return err
|
||||
}
|
||||
|
||||
func (s *ServerUDPSession) Close() {
|
||||
if s.timer != nil {
|
||||
s.timer.SetTimeout(0)
|
||||
}
|
||||
if link := s.link.Load(); link != nil {
|
||||
common.Interrupt(link.Reader)
|
||||
common.Interrupt(link.Writer)
|
||||
}
|
||||
}
|
||||
|
||||
func (s *ServerUDPSession) EnsureLink(
|
||||
ctx context.Context,
|
||||
conn stat.Connection,
|
||||
dest net.Destination,
|
||||
dispatcher routing.Dispatcher,
|
||||
policyManager policy.Manager,
|
||||
responseEncoder func(dest net.Destination, payload []byte) ([]byte, error),
|
||||
) (*transport.Link, error) {
|
||||
s.UpdateConn(conn)
|
||||
|
||||
if link := s.link.Load(); link != nil {
|
||||
return link, nil
|
||||
}
|
||||
|
||||
s.Lock()
|
||||
defer s.Unlock()
|
||||
|
||||
if link := s.link.Load(); link != nil {
|
||||
return link, nil
|
||||
}
|
||||
|
||||
sessCtx, cancel := context.WithCancel(ctx)
|
||||
inbound := session.InboundFromContext(sessCtx)
|
||||
if inbound != nil && s.User != nil {
|
||||
inbound.User = s.User
|
||||
}
|
||||
var email string
|
||||
var level uint32
|
||||
if s.User != nil {
|
||||
email = s.User.Email
|
||||
level = s.User.Level
|
||||
}
|
||||
sessCtx = log.ContextWithAccessMessage(sessCtx, &log.AccessMessage{
|
||||
From: conn.RemoteAddr(),
|
||||
To: dest,
|
||||
Status: log.AccessAccepted,
|
||||
Email: email,
|
||||
})
|
||||
|
||||
link, err := dispatcher.Dispatch(sessCtx, dest)
|
||||
if err != nil {
|
||||
cancel()
|
||||
return nil, err
|
||||
}
|
||||
|
||||
s.link.Store(link)
|
||||
sessionPolicy := policyManager.ForLevel(level)
|
||||
s.timer = signal.CancelAfterInactivity(sessCtx, func() {
|
||||
if s.manager != nil {
|
||||
s.manager.Delete(s.SessionID)
|
||||
}
|
||||
s.Close()
|
||||
cancel()
|
||||
}, sessionPolicy.Timeouts.ConnectionIdle)
|
||||
|
||||
go handleUDPResponse(s, link, dest, responseEncoder)
|
||||
return link, nil
|
||||
}
|
||||
|
||||
// ResetTCPConn sets SO_LINGER to 0 per SIP022 §3.1.4 to consistently send RST on close
|
||||
// when handshake or header validation fails.
|
||||
func ResetTCPConn(conn net.Conn) {
|
||||
rawConn, _, _ := proxy.UnwrapRawConn(conn)
|
||||
if tcpConn, ok := rawConn.(*net.TCPConn); ok {
|
||||
_ = tcpConn.SetLinger(0)
|
||||
}
|
||||
}
|
||||
|
||||
func handleUDPResponse(s *ServerUDPSession, link *transport.Link, fallbackDest net.Destination, encode func(dest net.Destination, payload []byte) ([]byte, error)) {
|
||||
defer func() {
|
||||
if s.timer != nil {
|
||||
s.timer.SetTimeout(0)
|
||||
}
|
||||
}()
|
||||
for {
|
||||
resMb, err := link.Reader.ReadMultiBuffer()
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
if s.timer != nil {
|
||||
s.timer.Update()
|
||||
}
|
||||
for i, rb := range resMb {
|
||||
b := rb.Bytes()
|
||||
if encode != nil {
|
||||
replyDest := fallbackDest
|
||||
if rb.UDP != nil {
|
||||
replyDest = *rb.UDP
|
||||
}
|
||||
encPacket, err := encode(replyDest, b)
|
||||
rb.Release()
|
||||
if err != nil {
|
||||
continue
|
||||
}
|
||||
if err := s.WriteToClient(encPacket); err != nil {
|
||||
buf.ReleaseMulti(resMb[i+1:])
|
||||
return
|
||||
}
|
||||
} else {
|
||||
err := s.WriteToClient(b)
|
||||
rb.Release()
|
||||
if err != nil {
|
||||
buf.ReleaseMulti(resMb[i+1:])
|
||||
return
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
type udpConnEntry struct {
|
||||
sync.Mutex
|
||||
link *transport.Link
|
||||
timer *signal.ActivityTimer
|
||||
cancel context.CancelFunc
|
||||
}
|
||||
|
||||
const (
|
||||
|
||||
@@ -182,48 +182,57 @@ func TestTCPStream(t *testing.T) {
|
||||
common.Must(err)
|
||||
IncreaseNonce(reader.Nonce())
|
||||
|
||||
dest, addrLen, err := ParseAddressPort(plainVar)
|
||||
vBuf := buf.New()
|
||||
vBuf.Write(plainVar)
|
||||
receivedDest, err = ReadAddressPort(vBuf)
|
||||
common.Must(err)
|
||||
receivedDest = net.TCPDestination(dest.Address, dest.Port)
|
||||
plainVar = plainVar[addrLen:]
|
||||
padLen := int(binary.BigEndian.Uint16(plainVar[:2]))
|
||||
receivedPayload = plainVar[2+padLen:]
|
||||
|
||||
// Server sends response stream with receivedPayload as first payload
|
||||
writer := NewServerStreamWriter(serverConn, method, rawKey, salt)
|
||||
pBuf := buf.New()
|
||||
pBuf.Write(receivedPayload)
|
||||
_ = writer.WriteMultiBuffer(buf.MultiBuffer{pBuf})
|
||||
// Skip padding
|
||||
var padBytes [2]byte
|
||||
_, _ = vBuf.Read(padBytes[:])
|
||||
padLen := int(padBytes[0])<<8 | int(padBytes[1])
|
||||
vBuf.Advance(int32(padLen))
|
||||
|
||||
// Read and echo additional stream data
|
||||
receivedPayload = make([]byte, vBuf.Len())
|
||||
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()
|
||||
common.Must(err)
|
||||
_ = writer.WriteMultiBuffer(mb)
|
||||
_ = writer.Close()
|
||||
}()
|
||||
|
||||
// Client goroutine
|
||||
go func() {
|
||||
defer wg.Done()
|
||||
clientSalt := make([]byte, method.KeySaltLength)
|
||||
common.Must2(io.ReadFull(rand.Reader, clientSalt))
|
||||
writer, err := WriteTCPRequest(clientConn, method, [][]byte{rawKey}, dest, clientSalt, testPayload)
|
||||
clientSalt, writer, err := ClientHandshake(clientConn, method, [][]byte{rawKey}, dest, testPayload)
|
||||
common.Must(err)
|
||||
|
||||
reader, err := ReadTCPResponse(clientConn, method, rawKey, clientSalt)
|
||||
reader, _, err := ClientVerifyServerResponse(clientConn, method, rawKey, clientSalt)
|
||||
common.Must(err)
|
||||
|
||||
// The first ReadMultiBuffer drains initialPayload from reader cache
|
||||
mbInit, err := reader.ReadMultiBuffer()
|
||||
common.Must(err)
|
||||
if !bytes.Equal(mbInit[0].Bytes(), testPayload) {
|
||||
t.Errorf("drained initial payload mismatch: got %s, want %s", mbInit[0].Bytes(), testPayload)
|
||||
}
|
||||
buf.ReleaseMulti(mbInit)
|
||||
|
||||
// Send additional stream data
|
||||
streamData := []byte("stream chunk test")
|
||||
_ = writer.WriteMultiBuffer(buf.MultiBuffer{buf.FromBytes(streamData)})
|
||||
_ = writer.WriteChunk(streamData)
|
||||
|
||||
mb, err := reader.ReadMultiBuffer()
|
||||
common.Must(err)
|
||||
@@ -263,14 +272,12 @@ func TestUDPCodec(t *testing.T) {
|
||||
psk := make([]byte, method.KeySaltLength)
|
||||
_, _ = rand.Read(psk)
|
||||
|
||||
clientCodec, err := NewUDPPacketCodec(method, [][]byte{psk})
|
||||
clientCodec, err := NewUDPPacketCodec(method, psk)
|
||||
common.Must(err)
|
||||
serverCodec, err := NewUDPServerCodec(method, psk, time.Minute)
|
||||
common.Must(err)
|
||||
|
||||
session, err := clientCodec.NewClientSession()
|
||||
common.Must(err)
|
||||
pktBuf, err := session.EncodePacket(dest, payload)
|
||||
pktBuf, err := clientCodec.EncodeClientPacket(dest, payload)
|
||||
common.Must(err)
|
||||
defer pktBuf.Release()
|
||||
|
||||
@@ -353,146 +360,3 @@ func TestMultiUserManager(t *testing.T) {
|
||||
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")
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
+115
-235
@@ -1,25 +1,18 @@
|
||||
package shadowsocks_2022
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/cipher"
|
||||
"crypto/rand"
|
||||
"encoding/binary"
|
||||
"io"
|
||||
"math"
|
||||
mrand "math/rand/v2"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/xtls/xray-core/common/antireplay"
|
||||
"github.com/xtls/xray-core/common/buf"
|
||||
"github.com/xtls/xray-core/common/errors"
|
||||
"github.com/xtls/xray-core/common/net"
|
||||
"github.com/xtls/xray-core/common/protocol"
|
||||
"github.com/xtls/xray-core/common/signal"
|
||||
"github.com/xtls/xray-core/common/task"
|
||||
"github.com/xtls/xray-core/features/policy"
|
||||
"github.com/xtls/xray-core/transport"
|
||||
)
|
||||
|
||||
var addrParser = protocol.NewAddressParser(
|
||||
@@ -45,6 +38,15 @@ func WriteAddressPort(w io.Writer, dest net.Destination) error {
|
||||
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
|
||||
func AddrPortLength(dest net.Destination) int {
|
||||
switch dest.Address.Family() {
|
||||
@@ -117,16 +119,8 @@ func (w *StreamWriter) Write(p []byte) (int, error) {
|
||||
func (w *StreamWriter) WriteMultiBuffer(mb buf.MultiBuffer) error {
|
||||
defer buf.ReleaseMulti(mb)
|
||||
for _, b := range mb {
|
||||
p := b.Bytes()
|
||||
for len(p) > 0 {
|
||||
chunkSize := len(p)
|
||||
if chunkSize > MaxPacketSize {
|
||||
chunkSize = MaxPacketSize
|
||||
}
|
||||
if err := w.WriteChunk(p[:chunkSize]); err != nil {
|
||||
return err
|
||||
}
|
||||
p = p[chunkSize:]
|
||||
if err := w.WriteChunk(b.Bytes()); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
@@ -174,7 +168,7 @@ func (r *StreamReader) Read(p []byte) (int, error) {
|
||||
IncreaseNonce(r.nonce[:])
|
||||
|
||||
payloadLen := int(binary.BigEndian.Uint16(decryptedLen))
|
||||
if payloadLen == 0 || payloadLen > MaxPacketSize {
|
||||
if payloadLen == 0 {
|
||||
return 0, ErrInvalidRequest
|
||||
}
|
||||
|
||||
@@ -200,10 +194,11 @@ func (r *StreamReader) Read(p []byte) (int, error) {
|
||||
|
||||
func (r *StreamReader) ReadMultiBuffer() (buf.MultiBuffer, error) {
|
||||
if r.cached > 0 {
|
||||
mb := buf.MergeBytes(nil, r.buffer[r.offset:r.offset+r.cached])
|
||||
b := buf.New()
|
||||
b.Write(r.buffer[r.offset : r.offset+r.cached])
|
||||
r.cached = 0
|
||||
r.offset = 0
|
||||
return mb, nil
|
||||
return buf.MultiBuffer{b}, nil
|
||||
}
|
||||
|
||||
if _, err := io.ReadFull(r.reader, r.lenBuf[:]); err != nil {
|
||||
@@ -217,7 +212,7 @@ func (r *StreamReader) ReadMultiBuffer() (buf.MultiBuffer, error) {
|
||||
IncreaseNonce(r.nonce[:])
|
||||
|
||||
payloadLen := int(binary.BigEndian.Uint16(decryptedLen))
|
||||
if payloadLen == 0 || payloadLen > MaxPacketSize {
|
||||
if payloadLen == 0 {
|
||||
return nil, ErrInvalidRequest
|
||||
}
|
||||
|
||||
@@ -232,8 +227,9 @@ func (r *StreamReader) ReadMultiBuffer() (buf.MultiBuffer, error) {
|
||||
}
|
||||
IncreaseNonce(r.nonce[:])
|
||||
|
||||
mb := buf.MergeBytes(nil, decryptedPayload)
|
||||
return mb, nil
|
||||
b := buf.New()
|
||||
b.Write(decryptedPayload)
|
||||
return buf.MultiBuffer{b}, nil
|
||||
}
|
||||
|
||||
type ClientRequestHeader struct {
|
||||
@@ -241,8 +237,13 @@ type ClientRequestHeader struct {
|
||||
EarlyData []byte
|
||||
}
|
||||
|
||||
func ReadClientRequestHeaderWithFixed(reader *StreamReader, fixedChunk []byte) (*ClientRequestHeader, error) {
|
||||
plainFixed, err := reader.cipher.Open(fixedChunk[:0], reader.Nonce(), fixedChunk, nil)
|
||||
func ReadClientRequestHeader(conn io.Reader, reader *StreamReader) (*ClientRequestHeader, error) {
|
||||
var fixedBuf [RequestHeaderFixedChunkLength + AEADTagSize]byte
|
||||
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 {
|
||||
return nil, errors.New("failed to decrypt client request header").Base(err)
|
||||
}
|
||||
@@ -271,7 +272,7 @@ func ReadClientRequestHeaderWithFixed(reader *StreamReader, fixedChunk []byte) (
|
||||
} else {
|
||||
varChunkCipher = make([]byte, needed)
|
||||
}
|
||||
if _, err := io.ReadFull(reader.reader, varChunkCipher); err != nil {
|
||||
if _, err := io.ReadFull(conn, varChunkCipher); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
@@ -281,34 +282,31 @@ func ReadClientRequestHeaderWithFixed(reader *StreamReader, fixedChunk []byte) (
|
||||
}
|
||||
IncreaseNonce(reader.Nonce())
|
||||
|
||||
dest, addrLen, err := ParseAddressPort(plainVar)
|
||||
b := buf.New()
|
||||
b.Write(plainVar)
|
||||
defer b.Release()
|
||||
|
||||
dest, err := ReadAddressPort(b)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
dest.Network = net.Network_TCP
|
||||
|
||||
offset := addrLen
|
||||
if len(plainVar) < offset+2 {
|
||||
return nil, ErrPacketTooShort
|
||||
var padLenBytes [2]byte
|
||||
if _, err := b.Read(padLenBytes[:]); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
paddingLen := int(binary.BigEndian.Uint16(plainVar[offset : offset+2]))
|
||||
offset += 2
|
||||
|
||||
if len(plainVar) < offset+paddingLen {
|
||||
paddingLen := int(binary.BigEndian.Uint16(padLenBytes[:]))
|
||||
if int(b.Len()) < paddingLen {
|
||||
return nil, ErrNoPadding
|
||||
}
|
||||
offset += paddingLen
|
||||
|
||||
var earlyData []byte
|
||||
var payloadLen int
|
||||
if len(plainVar) > offset {
|
||||
earlyData = plainVar[offset:]
|
||||
payloadLen = len(earlyData)
|
||||
if paddingLen > 0 {
|
||||
b.Advance(int32(paddingLen))
|
||||
}
|
||||
|
||||
// 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")
|
||||
var earlyData []byte
|
||||
if b.Len() > 0 {
|
||||
earlyData = make([]byte, b.Len())
|
||||
copy(earlyData, b.Bytes())
|
||||
}
|
||||
|
||||
return &ClientRequestHeader{
|
||||
@@ -317,6 +315,34 @@ func ReadClientRequestHeaderWithFixed(reader *StreamReader, fixedChunk []byte) (
|
||||
}, 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.
|
||||
func WriteTCPRequest(w io.Writer, method *CipherMethod, pskList [][]byte, dest net.Destination, clientSalt []byte, payload []byte) (buf.Writer, error) {
|
||||
finalPSK := pskList[len(pskList)-1]
|
||||
@@ -328,16 +354,7 @@ func WriteTCPRequest(w io.Writer, method *CipherMethod, pskList [][]byte, dest n
|
||||
|
||||
writer := NewStreamWriter(w, aead)
|
||||
|
||||
payloadLen := len(payload)
|
||||
var paddingLen int
|
||||
if payloadLen < MaxPaddingLength {
|
||||
paddingLen = mrand.IntN(MaxPaddingLength) + 1
|
||||
}
|
||||
addrPortLen := AddrPortLength(dest)
|
||||
varHeaderLen := addrPortLen + 2 + paddingLen + payloadLen
|
||||
|
||||
totalHandshakeLen := int32(method.KeySaltLength + len(pskList)*AESBlockSize + RequestHeaderFixedChunkLength + AEADTagSize + varHeaderLen + AEADTagSize)
|
||||
handshakeBuf := buf.NewWithSize(totalHandshakeLen)
|
||||
handshakeBuf := buf.New()
|
||||
defer handshakeBuf.Release()
|
||||
|
||||
handshakeBuf.Write(clientSalt)
|
||||
@@ -355,6 +372,14 @@ func WriteTCPRequest(w io.Writer, method *CipherMethod, pskList [][]byte, dest n
|
||||
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
|
||||
fixedHeaderPlaintext[0] = HeaderTypeClient
|
||||
binary.BigEndian.PutUint64(fixedHeaderPlaintext[1:9], uint64(time.Now().Unix()))
|
||||
@@ -364,7 +389,7 @@ func WriteTCPRequest(w io.Writer, method *CipherMethod, pskList [][]byte, dest n
|
||||
IncreaseNonce(writer.nonce[:])
|
||||
handshakeBuf.Write(fixedChunk)
|
||||
|
||||
varHeaderBuf := buf.NewWithSize(int32(varHeaderLen))
|
||||
varHeaderBuf := buf.New()
|
||||
defer varHeaderBuf.Release()
|
||||
|
||||
if err := WriteAddressPort(varHeaderBuf, dest); err != nil {
|
||||
@@ -396,21 +421,12 @@ 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.
|
||||
func ReadTCPResponse(r io.Reader, method *CipherMethod, psk []byte, clientSalt []byte) (buf.Reader, error) {
|
||||
fixedPlainLen := 1 + 8 + method.KeySaltLength + 2
|
||||
chunkCipherLen := fixedPlainLen + AEADTagSize
|
||||
headerLen := method.KeySaltLength + chunkCipherLen
|
||||
|
||||
// Single read call for Salt + Fixed-length response header chunk per SIP022 §3.1.4
|
||||
var headerBuf [128]byte
|
||||
headerSlice := headerBuf[:headerLen]
|
||||
n, err := r.Read(headerSlice)
|
||||
if err != nil || n < headerLen {
|
||||
return nil, errors.New("failed to read complete server response header")
|
||||
var serverSalt [32]byte
|
||||
serverSaltSlice := serverSalt[:method.KeySaltLength]
|
||||
if _, err := io.ReadFull(r, serverSaltSlice); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
serverSaltSlice := headerSlice[:method.KeySaltLength]
|
||||
chunkSlice := headerSlice[method.KeySaltLength:headerLen]
|
||||
|
||||
sessionKey := DeriveSessionSubKey(psk, serverSaltSlice, method.KeySaltLength)
|
||||
aead, err := method.NewAEAD(sessionKey)
|
||||
if err != nil {
|
||||
@@ -419,6 +435,14 @@ func ReadTCPResponse(r io.Reader, method *CipherMethod, psk []byte, clientSalt [
|
||||
|
||||
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)
|
||||
if err != nil {
|
||||
return nil, errors.New("failed to decrypt server response header").Base(err)
|
||||
@@ -460,190 +484,46 @@ func ReadTCPResponse(r io.Reader, method *CipherMethod, psk []byte, clientSalt [
|
||||
return reader, nil
|
||||
}
|
||||
|
||||
// ServerStreamWriter lazily sends the response header along with the first payload chunk per SIP022 §3.1.2 & §3.1.4.
|
||||
type ServerStreamWriter struct {
|
||||
mu sync.Mutex
|
||||
w io.Writer
|
||||
method *CipherMethod
|
||||
psk []byte
|
||||
clientSalt []byte
|
||||
streamWriter *StreamWriter
|
||||
}
|
||||
|
||||
func NewServerStreamWriter(w io.Writer, method *CipherMethod, psk []byte, clientSalt []byte) *ServerStreamWriter {
|
||||
return &ServerStreamWriter{
|
||||
w: w,
|
||||
method: method,
|
||||
psk: psk,
|
||||
clientSalt: clientSalt,
|
||||
}
|
||||
}
|
||||
|
||||
func (s *ServerStreamWriter) sendHeaderWithFirstPayload(payload []byte) (*StreamWriter, error) {
|
||||
// WriteTCPResponse writes the server handshake response and returns a body writer for server stream.
|
||||
func WriteTCPResponse(w io.Writer, method *CipherMethod, psk []byte, clientSalt []byte, initialPayload []byte) (buf.Writer, error) {
|
||||
var serverSalt [32]byte
|
||||
serverSaltSlice := serverSalt[:s.method.KeySaltLength]
|
||||
serverSaltSlice := serverSalt[:method.KeySaltLength]
|
||||
if _, err := io.ReadFull(rand.Reader, serverSaltSlice); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
respKey := DeriveSessionSubKey(s.psk, serverSaltSlice, s.method.KeySaltLength)
|
||||
respAead, err := s.method.NewAEAD(respKey)
|
||||
respKey := DeriveSessionSubKey(psk, serverSaltSlice, method.KeySaltLength)
|
||||
respAead, err := method.NewAEAD(respKey)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
sw := NewStreamWriter(s.w, respAead)
|
||||
writer := NewStreamWriter(w, respAead)
|
||||
|
||||
totalHeaderLen := int32(s.method.KeySaltLength + 1 + 8 + s.method.KeySaltLength + 2 + AEADTagSize + len(payload) + AEADTagSize)
|
||||
outBuf := buf.NewWithSize(totalHeaderLen)
|
||||
defer outBuf.Release()
|
||||
respBuf := buf.New()
|
||||
defer respBuf.Release()
|
||||
|
||||
outBuf.Write(serverSaltSlice)
|
||||
respBuf.Write(serverSaltSlice)
|
||||
|
||||
var fixedRespPlain [1 + 8 + 32 + 2]byte
|
||||
fixedRespSlice := fixedRespPlain[:1+8+s.method.KeySaltLength+2]
|
||||
fixedRespSlice := fixedRespPlain[:1+8+method.KeySaltLength+2]
|
||||
fixedRespSlice[0] = HeaderTypeServer
|
||||
binary.BigEndian.PutUint64(fixedRespSlice[1:9], uint64(time.Now().Unix()))
|
||||
copy(fixedRespSlice[9:9+s.method.KeySaltLength], s.clientSalt)
|
||||
binary.BigEndian.PutUint16(fixedRespSlice[9+s.method.KeySaltLength:11+s.method.KeySaltLength], uint16(len(payload)))
|
||||
copy(fixedRespSlice[9:9+method.KeySaltLength], clientSalt)
|
||||
binary.BigEndian.PutUint16(fixedRespSlice[9+method.KeySaltLength:11+method.KeySaltLength], uint16(len(initialPayload)))
|
||||
|
||||
fixedRespChunk := sw.cipher.Seal(nil, sw.nonce[:], fixedRespSlice, nil)
|
||||
IncreaseNonce(sw.nonce[:])
|
||||
outBuf.Write(fixedRespChunk)
|
||||
fixedRespChunk := writer.cipher.Seal(nil, writer.nonce[:], fixedRespSlice, nil)
|
||||
IncreaseNonce(writer.nonce[:])
|
||||
respBuf.Write(fixedRespChunk)
|
||||
|
||||
if len(payload) > 0 {
|
||||
payloadChunk := sw.cipher.Seal(nil, sw.nonce[:], payload, nil)
|
||||
IncreaseNonce(sw.nonce[:])
|
||||
outBuf.Write(payloadChunk)
|
||||
if len(initialPayload) > 0 {
|
||||
initialChunk := writer.cipher.Seal(nil, writer.nonce[:], initialPayload, nil)
|
||||
IncreaseNonce(writer.nonce[:])
|
||||
respBuf.Write(initialChunk)
|
||||
}
|
||||
|
||||
if _, err := s.w.Write(outBuf.Bytes()); err != nil {
|
||||
if _, err := w.Write(respBuf.Bytes()); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return sw, nil
|
||||
}
|
||||
|
||||
func (s *ServerStreamWriter) WriteMultiBuffer(mb buf.MultiBuffer) error {
|
||||
if mb.IsEmpty() {
|
||||
return nil
|
||||
}
|
||||
|
||||
if s.streamWriter == nil {
|
||||
s.mu.Lock()
|
||||
if s.streamWriter == nil {
|
||||
firstBuf := mb[0]
|
||||
firstBytes := firstBuf.Bytes()
|
||||
chunkSize := len(firstBytes)
|
||||
if chunkSize > MaxPacketSize {
|
||||
chunkSize = MaxPacketSize
|
||||
}
|
||||
firstPayload := firstBytes[:chunkSize]
|
||||
sw, err := s.sendHeaderWithFirstPayload(firstPayload)
|
||||
if err != nil {
|
||||
s.mu.Unlock()
|
||||
buf.ReleaseMulti(mb)
|
||||
return err
|
||||
}
|
||||
s.streamWriter = sw
|
||||
|
||||
firstBuf.Advance(int32(chunkSize))
|
||||
if firstBuf.IsEmpty() {
|
||||
firstBuf.Release()
|
||||
mb = mb[1:]
|
||||
}
|
||||
}
|
||||
s.mu.Unlock()
|
||||
if len(mb) == 0 {
|
||||
return nil
|
||||
}
|
||||
}
|
||||
|
||||
return s.streamWriter.WriteMultiBuffer(mb)
|
||||
}
|
||||
|
||||
func (s *ServerStreamWriter) Write(p []byte) (int, error) {
|
||||
n := len(p)
|
||||
if s.streamWriter == nil {
|
||||
s.mu.Lock()
|
||||
if s.streamWriter == nil {
|
||||
chunkSize := len(p)
|
||||
if chunkSize > MaxPacketSize {
|
||||
chunkSize = MaxPacketSize
|
||||
}
|
||||
firstPayload := p[:chunkSize]
|
||||
sw, err := s.sendHeaderWithFirstPayload(firstPayload)
|
||||
if err != nil {
|
||||
s.mu.Unlock()
|
||||
return 0, err
|
||||
}
|
||||
s.streamWriter = sw
|
||||
p = p[chunkSize:]
|
||||
}
|
||||
s.mu.Unlock()
|
||||
if len(p) == 0 {
|
||||
return n, nil
|
||||
}
|
||||
}
|
||||
|
||||
_, err := s.streamWriter.Write(p)
|
||||
return n, err
|
||||
}
|
||||
|
||||
func (s *ServerStreamWriter) Close() error {
|
||||
if s.streamWriter == nil {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
if s.streamWriter == nil {
|
||||
sw, err := s.sendHeaderWithFirstPayload(nil)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
s.streamWriter = sw
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// InitServerStream decrypts the client request header, verifies the timestamp and replay filter,
|
||||
// and returns a StreamReader for subsequent stream chunks.
|
||||
func InitServerStream(conn net.Conn, method *CipherMethod, psk, saltSlice []byte, salt [32]byte, fixedChunk []byte, saltFilter *antireplay.ReplayFilter[[32]byte]) (*StreamReader, *ClientRequestHeader, error) {
|
||||
sessionKey := DeriveSessionSubKey(psk, saltSlice, method.KeySaltLength)
|
||||
aead, err := method.NewAEAD(sessionKey)
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
|
||||
reader := NewStreamReader(conn, aead)
|
||||
|
||||
reqHeader, err := ReadClientRequestHeaderWithFixed(reader, fixedChunk)
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
_ = conn.SetReadDeadline(time.Time{})
|
||||
|
||||
if !saltFilter.Check(salt) {
|
||||
return nil, nil, ErrSaltNotUnique
|
||||
}
|
||||
return reader, reqHeader, nil
|
||||
}
|
||||
|
||||
func TransportTCP(ctx context.Context, sessionPolicy policy.Session, reader buf.Reader, writer buf.Writer, link *transport.Link) error {
|
||||
ctx, cancel := context.WithCancel(ctx)
|
||||
timer := signal.CancelAfterInactivity(ctx, cancel, sessionPolicy.Timeouts.ConnectionIdle)
|
||||
ctx = policy.ContextWithBufferPolicy(ctx, sessionPolicy.Buffer)
|
||||
|
||||
requestDone := func() error {
|
||||
defer timer.SetTimeout(sessionPolicy.Timeouts.DownlinkOnly)
|
||||
return buf.Copy(reader, link.Writer, buf.UpdateActivity(timer))
|
||||
}
|
||||
|
||||
responseDone := func() error {
|
||||
defer timer.SetTimeout(sessionPolicy.Timeouts.UplinkOnly)
|
||||
if c, ok := writer.(io.Closer); ok {
|
||||
defer c.Close()
|
||||
}
|
||||
return buf.Copy(link.Reader, writer, buf.UpdateActivity(timer))
|
||||
}
|
||||
|
||||
responseDoneAndCloseWriter := task.OnSuccess(responseDone, task.Close(link.Writer))
|
||||
return task.Run(ctx, requestDone, responseDoneAndCloseWriter)
|
||||
|
||||
return writer, nil
|
||||
}
|
||||
|
||||
+2
-12
@@ -52,12 +52,9 @@ func (b *bind) Open(port uint16) (fns []conn.ReceiveFunc, actualPort uint16, err
|
||||
case <-ch:
|
||||
default:
|
||||
errors.LogErrorInner(context.Background(), err, "unexpected closed")
|
||||
b.mu.Lock()
|
||||
downFunc := b.downFunc
|
||||
b.mu.Unlock()
|
||||
if downFunc != nil {
|
||||
if b.downFunc != nil {
|
||||
go func() {
|
||||
common.Must(downFunc())
|
||||
common.Must(b.downFunc())
|
||||
}()
|
||||
}
|
||||
}
|
||||
@@ -79,13 +76,6 @@ func (b *bind) Open(port uint16) (fns []conn.ReceiveFunc, actualPort uint16, err
|
||||
}, uint16(c.LocalAddr().(*net.UDPAddr).Port), nil
|
||||
}
|
||||
|
||||
// setDownFunc sets downFunc after the device is created, since the device may already be using the bind.
|
||||
func (b *bind) setDownFunc(f func() error) {
|
||||
b.mu.Lock()
|
||||
defer b.mu.Unlock()
|
||||
b.downFunc = f
|
||||
}
|
||||
|
||||
func (b *bind) Close() error {
|
||||
b.mu.Lock()
|
||||
defer b.mu.Unlock()
|
||||
|
||||
@@ -27,6 +27,7 @@ import (
|
||||
"github.com/xtls/xray-core/features/stats"
|
||||
"github.com/xtls/xray-core/transport"
|
||||
"github.com/xtls/xray-core/transport/internet"
|
||||
"github.com/xtls/xray-core/transport/internet/finalmask"
|
||||
"golang.zx2c4.com/wireguard/device"
|
||||
)
|
||||
|
||||
@@ -199,7 +200,7 @@ func (h *Handler) Process(ctx context.Context, link *transport.Link, dialer inte
|
||||
}
|
||||
defer conn.Close()
|
||||
c := &UDPConnClient{
|
||||
PacketConn: conn.(*net.PacketConnWrapper).PacketConn,
|
||||
PacketConn: conn.(*internet.PacketConnWrapper).PacketConn,
|
||||
Dest: conn.RemoteAddr().(*net.UDPAddr),
|
||||
}
|
||||
reader = c
|
||||
@@ -263,14 +264,14 @@ func (h *Handler) init(ctx context.Context) error {
|
||||
if err != nil {
|
||||
return nil, errors.New("failed to dial to dest").Base(err)
|
||||
}
|
||||
pktConn = conn.(*net.PacketConnWrapper).PacketConn
|
||||
pktConn = conn.(*finalmask.PacketConnWrapper).PacketConn
|
||||
} else {
|
||||
conn, err := internet.DialSystem(ctx, dest, h.streamSettings.SocketSettings)
|
||||
if err != nil {
|
||||
return nil, errors.New("failed to dial to dest").Base(err)
|
||||
}
|
||||
switch c := conn.(type) {
|
||||
case *net.PacketConnWrapper:
|
||||
case *internet.PacketConnWrapper:
|
||||
pktConn = c.PacketConn
|
||||
case *cnc.Connection:
|
||||
pktConn = &internet.FakePacketConn{Conn: c}
|
||||
@@ -287,13 +288,7 @@ func (h *Handler) init(ctx context.Context) error {
|
||||
}
|
||||
return pktConn, nil
|
||||
}
|
||||
// device.NewDevice may use the bind right away (Up -> BindUpdate -> Open),
|
||||
// so everything it reads must be set before creating the device.
|
||||
bind := &bind{
|
||||
resolveFunc: resolveFunc,
|
||||
listenFunc: listenFunc,
|
||||
reserved: h.conf.Reserved,
|
||||
}
|
||||
bind := &bind{}
|
||||
logger := &device.Logger{
|
||||
Verbosef: func(format string, args ...any) {
|
||||
log.Record(&log.GeneralMessage{
|
||||
@@ -309,7 +304,10 @@ func (h *Handler) init(ctx context.Context) error {
|
||||
},
|
||||
}
|
||||
dev := device.NewDevice(h.tun, bind, logger)
|
||||
bind.setDownFunc(dev.Down)
|
||||
bind.resolveFunc = resolveFunc
|
||||
bind.listenFunc = listenFunc
|
||||
bind.downFunc = dev.Down
|
||||
bind.reserved = h.conf.Reserved
|
||||
var cfg strings.Builder
|
||||
cfg.WriteString("private_key=" + h.conf.SecretKey + "\n")
|
||||
for _, peer := range h.conf.Peers {
|
||||
|
||||
@@ -21,7 +21,7 @@ import (
|
||||
"syscall"
|
||||
"time"
|
||||
|
||||
xnet "github.com/xtls/xray-core/common/net"
|
||||
"github.com/xtls/xray-core/transport/internet"
|
||||
"golang.zx2c4.com/wireguard/tun"
|
||||
|
||||
"golang.org/x/net/dns/dnsmessage"
|
||||
@@ -220,7 +220,7 @@ func (tun *netTun) DialUDPAddrPort(laddr, raddr netip.AddrPort) (net.Conn, error
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &xnet.PacketConnWrapper{
|
||||
return &internet.PacketConnWrapper{
|
||||
PacketConn: conn,
|
||||
Dest: net.UDPAddrFromAddrPort(raddr),
|
||||
}, nil
|
||||
|
||||
@@ -113,7 +113,7 @@ func NewServer(ctx context.Context, conf *DeviceConfig) (*Server, error) {
|
||||
users.Store(user.Account.(*MemoryAccount).Pub, user)
|
||||
}
|
||||
|
||||
s := &Server{
|
||||
return &Server{
|
||||
conf: conf,
|
||||
ctx: core.ToBackgroundDetachedContext(ctx),
|
||||
policyManager: p,
|
||||
@@ -131,10 +131,7 @@ func NewServer(ctx context.Context, conf *DeviceConfig) (*Server, error) {
|
||||
|
||||
pub: pub,
|
||||
users: users,
|
||||
}
|
||||
// Install the stack's protocol handlers before the device can deliver packets to it (Start -> dev.Up).
|
||||
CreateForwarder(stack, s.HandleConnection)
|
||||
return s, nil
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (s *Server) AddUser(ctx context.Context, user *protocol.MemoryUser) error {
|
||||
@@ -323,6 +320,7 @@ func (s *Server) Start() error {
|
||||
return err
|
||||
}
|
||||
s.dev = dev
|
||||
CreateForwarder(s.stack, s.HandleConnection)
|
||||
return nil
|
||||
}
|
||||
|
||||
|
||||
@@ -16,7 +16,7 @@ import (
|
||||
|
||||
"github.com/vishvananda/netlink"
|
||||
"github.com/xtls/xray-core/common/errors"
|
||||
xnet "github.com/xtls/xray-core/common/net"
|
||||
"github.com/xtls/xray-core/transport/internet"
|
||||
"golang.zx2c4.com/wireguard/tun"
|
||||
)
|
||||
|
||||
@@ -263,7 +263,7 @@ func (tun *kernelTun) DialUDPAddrPort(laddr, raddr netip.AddrPort) (net.Conn, er
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &xnet.PacketConnWrapper{
|
||||
return &internet.PacketConnWrapper{
|
||||
PacketConn: conn,
|
||||
Dest: net.UDPAddrFromAddrPort(raddr),
|
||||
}, nil
|
||||
|
||||
@@ -82,7 +82,7 @@ func (fm *FinalMask) DialTCP(ctx context.Context, dest net.Destination) (net.Con
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &net.PacketConnWrapper{PacketConn: conn, Dest: addr}, err
|
||||
return &PacketConnWrapper{PacketConn: conn, udpAddr: addr}, err
|
||||
},
|
||||
}
|
||||
for i := range fm.tcpMasks {
|
||||
@@ -144,7 +144,7 @@ func (fm *FinalMask) DialUDP(ctx context.Context, dest net.Destination) (net.Con
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &net.PacketConnWrapper{PacketConn: conn, Dest: addr}, nil
|
||||
return &PacketConnWrapper{PacketConn: conn, udpAddr: addr}, nil
|
||||
}
|
||||
for i := range fm.udpMasks {
|
||||
if i > 0 {
|
||||
@@ -171,7 +171,7 @@ func (fm *FinalMask) DialUDP(ctx context.Context, dest net.Destination) (net.Con
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &net.PacketConnWrapper{PacketConn: conn, Dest: addr}, err
|
||||
return &PacketConnWrapper{PacketConn: conn, udpAddr: addr}, err
|
||||
},
|
||||
}
|
||||
var sizes []int
|
||||
@@ -208,7 +208,7 @@ func (fm *FinalMask) DialUDP(ctx context.Context, dest net.Destination) (net.Con
|
||||
if addr == nil {
|
||||
addr = &net.UDPAddr{IP: []byte{0, 0, 0, 0}}
|
||||
}
|
||||
return &net.PacketConnWrapper{PacketConn: conn, Dest: addr}, nil
|
||||
return &PacketConnWrapper{PacketConn: conn, udpAddr: addr}, nil
|
||||
}
|
||||
|
||||
func (fm *FinalMask) ListenPacket(ctx context.Context, addr net.Addr) (net.PacketConn, error) {
|
||||
@@ -272,6 +272,24 @@ const (
|
||||
UDPSize = 4096
|
||||
)
|
||||
|
||||
type PacketConnWrapper struct {
|
||||
net.PacketConn
|
||||
udpAddr net.Addr
|
||||
}
|
||||
|
||||
func (c *PacketConnWrapper) RemoteAddr() net.Addr {
|
||||
return c.udpAddr
|
||||
}
|
||||
|
||||
func (c *PacketConnWrapper) Read(b []byte) (n int, err error) {
|
||||
n, _, err = c.PacketConn.ReadFrom(b)
|
||||
return
|
||||
}
|
||||
|
||||
func (c *PacketConnWrapper) Write(b []byte) (n int, err error) {
|
||||
return c.PacketConn.WriteTo(b, c.udpAddr)
|
||||
}
|
||||
|
||||
type headerManagerConn struct {
|
||||
net.PacketConn
|
||||
|
||||
|
||||
@@ -21,135 +21,6 @@ const (
|
||||
_ = 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 {
|
||||
state protoimpl.MessageState `protogen:"open.v1"`
|
||||
RandMin int64 `protobuf:"varint,1,opt,name=rand_min,json=randMin,proto3" json:"rand_min,omitempty"`
|
||||
@@ -159,14 +30,13 @@ type Item struct {
|
||||
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"`
|
||||
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
|
||||
sizeCache protoimpl.SizeCache
|
||||
}
|
||||
|
||||
func (x *Item) Reset() {
|
||||
*x = Item{}
|
||||
mi := &file_transport_internet_finalmask_noise_config_proto_msgTypes[1]
|
||||
mi := &file_transport_internet_finalmask_noise_config_proto_msgTypes[0]
|
||||
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
|
||||
ms.StoreMessageInfo(mi)
|
||||
}
|
||||
@@ -178,7 +48,7 @@ func (x *Item) String() string {
|
||||
func (*Item) ProtoMessage() {}
|
||||
|
||||
func (x *Item) ProtoReflect() protoreflect.Message {
|
||||
mi := &file_transport_internet_finalmask_noise_config_proto_msgTypes[1]
|
||||
mi := &file_transport_internet_finalmask_noise_config_proto_msgTypes[0]
|
||||
if x != nil {
|
||||
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
|
||||
if ms.LoadMessageInfo() == nil {
|
||||
@@ -191,7 +61,7 @@ func (x *Item) ProtoReflect() protoreflect.Message {
|
||||
|
||||
// Deprecated: Use Item.ProtoReflect.Descriptor instead.
|
||||
func (*Item) Descriptor() ([]byte, []int) {
|
||||
return file_transport_internet_finalmask_noise_config_proto_rawDescGZIP(), []int{1}
|
||||
return file_transport_internet_finalmask_noise_config_proto_rawDescGZIP(), []int{0}
|
||||
}
|
||||
|
||||
func (x *Item) GetRandMin() int64 {
|
||||
@@ -243,13 +113,6 @@ func (x *Item) GetDelayMax() int64 {
|
||||
return 0
|
||||
}
|
||||
|
||||
func (x *Item) GetSegments() []*Segment {
|
||||
if x != nil {
|
||||
return x.Segments
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
type Config struct {
|
||||
state protoimpl.MessageState `protogen:"open.v1"`
|
||||
ResetMin int64 `protobuf:"varint,1,opt,name=reset_min,json=resetMin,proto3" json:"reset_min,omitempty"`
|
||||
@@ -261,7 +124,7 @@ type Config struct {
|
||||
|
||||
func (x *Config) Reset() {
|
||||
*x = Config{}
|
||||
mi := &file_transport_internet_finalmask_noise_config_proto_msgTypes[2]
|
||||
mi := &file_transport_internet_finalmask_noise_config_proto_msgTypes[1]
|
||||
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
|
||||
ms.StoreMessageInfo(mi)
|
||||
}
|
||||
@@ -273,7 +136,7 @@ func (x *Config) String() string {
|
||||
func (*Config) ProtoMessage() {}
|
||||
|
||||
func (x *Config) ProtoReflect() protoreflect.Message {
|
||||
mi := &file_transport_internet_finalmask_noise_config_proto_msgTypes[2]
|
||||
mi := &file_transport_internet_finalmask_noise_config_proto_msgTypes[1]
|
||||
if x != nil {
|
||||
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
|
||||
if ms.LoadMessageInfo() == nil {
|
||||
@@ -286,7 +149,7 @@ func (x *Config) ProtoReflect() protoreflect.Message {
|
||||
|
||||
// Deprecated: Use Config.ProtoReflect.Descriptor instead.
|
||||
func (*Config) Descriptor() ([]byte, []int) {
|
||||
return file_transport_internet_finalmask_noise_config_proto_rawDescGZIP(), []int{2}
|
||||
return file_transport_internet_finalmask_noise_config_proto_rawDescGZIP(), []int{1}
|
||||
}
|
||||
|
||||
func (x *Config) GetResetMin() int64 {
|
||||
@@ -314,21 +177,7 @@ var File_transport_internet_finalmask_noise_config_proto protoreflect.FileDescri
|
||||
|
||||
const file_transport_internet_finalmask_noise_config_proto_rawDesc = "" +
|
||||
"\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" +
|
||||
"/transport/internet/finalmask/noise/config.proto\x12'xray.transport.internet.finalmask.noise\"\xda\x01\n" +
|
||||
"\x04Item\x12\x19\n" +
|
||||
"\brand_min\x18\x01 \x01(\x03R\arandMin\x12\x19\n" +
|
||||
"\brand_max\x18\x02 \x01(\x03R\arandMax\x12$\n" +
|
||||
@@ -336,8 +185,7 @@ const file_transport_internet_finalmask_noise_config_proto_rawDesc = "" +
|
||||
"\x0erand_range_max\x18\x04 \x01(\x05R\frandRangeMax\x12\x16\n" +
|
||||
"\x06packet\x18\x05 \x01(\fR\x06packet\x12\x1b\n" +
|
||||
"\tdelay_min\x18\x06 \x01(\x03R\bdelayMin\x12\x1b\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" +
|
||||
"\tdelay_max\x18\a \x01(\x03R\bdelayMax\"\x87\x01\n" +
|
||||
"\x06Config\x12\x1b\n" +
|
||||
"\treset_min\x18\x01 \x01(\x03R\bresetMin\x12\x1b\n" +
|
||||
"\treset_max\x18\x02 \x01(\x03R\bresetMax\x12C\n" +
|
||||
@@ -356,23 +204,18 @@ func file_transport_internet_finalmask_noise_config_proto_rawDescGZIP() []byte {
|
||||
return file_transport_internet_finalmask_noise_config_proto_rawDescData
|
||||
}
|
||||
|
||||
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_msgTypes = make([]protoimpl.MessageInfo, 2)
|
||||
var file_transport_internet_finalmask_noise_config_proto_goTypes = []any{
|
||||
(Segment_Kind)(0), // 0: xray.transport.internet.finalmask.noise.Segment.Kind
|
||||
(*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
|
||||
(*Item)(nil), // 0: xray.transport.internet.finalmask.noise.Item
|
||||
(*Config)(nil), // 1: xray.transport.internet.finalmask.noise.Config
|
||||
}
|
||||
var file_transport_internet_finalmask_noise_config_proto_depIdxs = []int32{
|
||||
0, // 0: xray.transport.internet.finalmask.noise.Segment.kind:type_name -> xray.transport.internet.finalmask.noise.Segment.Kind
|
||||
1, // 1: xray.transport.internet.finalmask.noise.Item.segments:type_name -> xray.transport.internet.finalmask.noise.Segment
|
||||
2, // 2: xray.transport.internet.finalmask.noise.Config.items:type_name -> xray.transport.internet.finalmask.noise.Item
|
||||
3, // [3:3] is the sub-list for method output_type
|
||||
3, // [3:3] is the sub-list for method input_type
|
||||
3, // [3:3] is the sub-list for extension type_name
|
||||
3, // [3:3] is the sub-list for extension extendee
|
||||
0, // [0:3] is the sub-list for field type_name
|
||||
0, // 0: xray.transport.internet.finalmask.noise.Config.items:type_name -> xray.transport.internet.finalmask.noise.Item
|
||||
1, // [1:1] is the sub-list for method output_type
|
||||
1, // [1:1] is the sub-list for method input_type
|
||||
1, // [1:1] is the sub-list for extension type_name
|
||||
1, // [1:1] is the sub-list for extension extendee
|
||||
0, // [0:1] is the sub-list for field type_name
|
||||
}
|
||||
|
||||
func init() { file_transport_internet_finalmask_noise_config_proto_init() }
|
||||
@@ -385,14 +228,13 @@ func file_transport_internet_finalmask_noise_config_proto_init() {
|
||||
File: protoimpl.DescBuilder{
|
||||
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)),
|
||||
NumEnums: 1,
|
||||
NumMessages: 3,
|
||||
NumEnums: 0,
|
||||
NumMessages: 2,
|
||||
NumExtensions: 0,
|
||||
NumServices: 0,
|
||||
},
|
||||
GoTypes: file_transport_internet_finalmask_noise_config_proto_goTypes,
|
||||
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,
|
||||
}.Build()
|
||||
File_transport_internet_finalmask_noise_config_proto = out.File
|
||||
|
||||
@@ -6,22 +6,6 @@ option go_package = "github.com/xtls/xray-core/transport/internet/finalmask/nois
|
||||
option java_package = "com.xray.transport.internet.finalmask.noise";
|
||||
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 {
|
||||
int64 rand_min = 1;
|
||||
int64 rand_max = 2;
|
||||
@@ -30,7 +14,6 @@ message Item {
|
||||
bytes packet = 5;
|
||||
int64 delay_min = 6;
|
||||
int64 delay_max = 7;
|
||||
repeated Segment segments = 8;
|
||||
}
|
||||
|
||||
message Config {
|
||||
|
||||
@@ -1,25 +1,18 @@
|
||||
package noise
|
||||
|
||||
import (
|
||||
"crypto/rand"
|
||||
"encoding/binary"
|
||||
"net"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"time"
|
||||
|
||||
"github.com/xtls/xray-core/common"
|
||||
"github.com/xtls/xray-core/common/crypto"
|
||||
)
|
||||
|
||||
const asciiLetters = "abcdefghijklmnopqrstuvwxyzABCDEFGHIJKLMNOPQRSTUVWXYZ"
|
||||
|
||||
type noiseConn struct {
|
||||
net.PacketConn
|
||||
config *Config
|
||||
m map[string]time.Time
|
||||
mu sync.Mutex
|
||||
counter atomic.Uint32
|
||||
config *Config
|
||||
m map[string]time.Time
|
||||
mu sync.Mutex
|
||||
}
|
||||
|
||||
func NewConnClient(c *Config, raw net.PacketConn) (net.PacketConn, error) {
|
||||
@@ -34,62 +27,6 @@ func NewConnServer(c *Config, raw net.PacketConn) (net.PacketConn, error) {
|
||||
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) {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
@@ -98,7 +35,13 @@ func (c *noiseConn) WriteTo(p []byte, addr net.Addr) (n int, err error) {
|
||||
|
||||
if t.IsZero() || (c.config.ResetMax > 0 && time.Now().After(t)) {
|
||||
for _, item := range c.config.Items {
|
||||
c.PacketConn.WriteTo(c.buildPacket(item), addr)
|
||||
if item.RandMax > 0 {
|
||||
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)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,137 +0,0 @@
|
||||
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)
|
||||
}
|
||||
@@ -380,7 +380,7 @@ func TestPacketConnReadWrite(t *testing.T) {
|
||||
t.Fatal(err)
|
||||
}
|
||||
t.Cleanup(func() { clientConn.Close() })
|
||||
client := clientConn.(*net.PacketConnWrapper).PacketConn
|
||||
client := clientConn.(*finalmask.PacketConnWrapper).PacketConn
|
||||
|
||||
_ = client.SetDeadline(time.Now().Add(time.Second))
|
||||
_ = server.SetDeadline(time.Now().Add(time.Second))
|
||||
|
||||
@@ -73,7 +73,7 @@ func NewUDPHopConn(c *Config, dest *net.Destination, dialer *finalmask.Dialer) (
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
cur := conn.(*net.PacketConnWrapper).PacketConn
|
||||
cur := conn.(*finalmask.PacketConnWrapper).PacketConn
|
||||
addr := conn.RemoteAddr().(*net.UDPAddr)
|
||||
client := &udpHopConn{
|
||||
dialer: dialer,
|
||||
@@ -150,7 +150,7 @@ func (c *udpHopConn) hop() {
|
||||
_ = c.pre.Close()
|
||||
}
|
||||
c.pre = c.cur
|
||||
c.cur = conn.(*net.PacketConnWrapper).PacketConn
|
||||
c.cur = conn.(*finalmask.PacketConnWrapper).PacketConn
|
||||
c.wg.Add(1)
|
||||
go c.recv(c.cur)
|
||||
}
|
||||
@@ -223,6 +223,13 @@ func (c *udpHopConn) Close() error {
|
||||
}
|
||||
_ = c.cur.Close()
|
||||
c.wg.Wait()
|
||||
select {
|
||||
case packet := <-c.readCh:
|
||||
if packet.p != nil {
|
||||
pool.Put(packet.p[:cap(packet.p)])
|
||||
}
|
||||
default:
|
||||
}
|
||||
close(c.readCh)
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -1,441 +1,417 @@
|
||||
package xdns
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"crypto/rand"
|
||||
"encoding/base32"
|
||||
"encoding/binary"
|
||||
go_errors "errors"
|
||||
"io"
|
||||
mrand "math/rand"
|
||||
"net"
|
||||
"strconv"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"time"
|
||||
|
||||
"github.com/xtls/xray-core/common"
|
||||
"github.com/xtls/xray-core/common/errors"
|
||||
"github.com/xtls/xray-core/common/net"
|
||||
"github.com/xtls/xray-core/transport/internet/finalmask"
|
||||
"golang.org/x/net/dns/dnsmessage"
|
||||
)
|
||||
|
||||
const (
|
||||
numPadding = 3
|
||||
numPaddingForPoll = 8
|
||||
initPollDelay = 500 * time.Millisecond
|
||||
maxPollDelay = 10 * time.Second
|
||||
pollDelayMultiplier = 2.0
|
||||
pollLimit = 16
|
||||
)
|
||||
|
||||
var pool4K = sync.Pool{
|
||||
New: func() any {
|
||||
return make([]byte, 4096)
|
||||
},
|
||||
}
|
||||
var base32Encoding = base32.StdEncoding.WithPadding(base32.NoPadding)
|
||||
|
||||
type packet struct {
|
||||
p []byte
|
||||
addr net.Addr
|
||||
}
|
||||
|
||||
type xdnsClient struct {
|
||||
dialer *finalmask.Dialer
|
||||
type xdnsConnClient struct {
|
||||
net.PacketConn
|
||||
|
||||
clientID ClientID
|
||||
fragID atomic.Uint32
|
||||
domains []*Domain
|
||||
extraPoll int32
|
||||
resolverAddrs []*net.UDPAddr
|
||||
resolverTypes []uint16
|
||||
resolverIdx uint32
|
||||
resolverSend map[string]*atomic.Uint32
|
||||
|
||||
resolvers []Resolver
|
||||
resolverSends []atomic.Uint32
|
||||
resolverIndex atomic.Uint32
|
||||
clientID []byte
|
||||
domains []Name
|
||||
|
||||
readCh chan packet
|
||||
sendCh chan []byte
|
||||
poolCh chan struct{}
|
||||
closeCh chan struct{}
|
||||
wg sync.WaitGroup
|
||||
mu sync.Mutex
|
||||
pollChan chan struct{}
|
||||
readQueue chan *packet
|
||||
writeQueue chan *packet
|
||||
|
||||
closed bool
|
||||
mutex sync.Mutex
|
||||
}
|
||||
|
||||
func NewClient(c *Config, dialer *finalmask.Dialer) (net.PacketConn, error) {
|
||||
if len(c.Domains) == 0 {
|
||||
return nil, errors.New("empty domains")
|
||||
}
|
||||
func NewConnClient(c *Config, raw net.PacketConn) (net.PacketConn, error) {
|
||||
if len(c.Resolvers) == 0 {
|
||||
return nil, errors.New("empty resolvers")
|
||||
}
|
||||
if c.ExtraPoll < 0 || c.ExtraPoll > 3 {
|
||||
return nil, errors.New("c.ExtraPoll < 0 || c.ExtraPoll > 3")
|
||||
}
|
||||
domains := make([]*Domain, 0, len(c.Domains))
|
||||
for i := range c.Domains {
|
||||
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))
|
||||
|
||||
var domains []Name
|
||||
var servers []string
|
||||
var resolverTypes []uint16
|
||||
for _, rs := range c.Resolvers {
|
||||
domain, server, resolverType, err := parseResolver(rs)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
return nil, errors.New("invalid resolvers").Base(err)
|
||||
}
|
||||
domains = append(domains, domain)
|
||||
servers = append(servers, server)
|
||||
resolverTypes = append(resolverTypes, resolverType)
|
||||
}
|
||||
resolvers := make([]Resolver, 0, len(c.Resolvers))
|
||||
for i := range c.Resolvers {
|
||||
resolver, err := NewResolver(c.Resolvers[i], dialer)
|
||||
|
||||
var resolverAddrs []*net.UDPAddr
|
||||
resolverSend := make(map[string]*atomic.Uint32)
|
||||
for _, rs := range servers {
|
||||
h, p, err := net.SplitHostPort(rs)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
resolvers = append(resolvers, resolver)
|
||||
}
|
||||
client := &xdnsClient{
|
||||
dialer: dialer,
|
||||
|
||||
clientID: NewClientID(),
|
||||
domains: domains,
|
||||
extraPoll: c.ExtraPoll,
|
||||
|
||||
resolvers: resolvers,
|
||||
resolverSends: make([]atomic.Uint32, len(c.Resolvers)),
|
||||
|
||||
readCh: make(chan packet),
|
||||
sendCh: make(chan []byte, 16),
|
||||
poolCh: make(chan struct{}, pollLimit),
|
||||
closeCh: make(chan struct{}),
|
||||
}
|
||||
go client.run()
|
||||
return client, nil
|
||||
}
|
||||
|
||||
func (c *xdnsClient) closed() bool {
|
||||
select {
|
||||
case <-c.closeCh:
|
||||
return true
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
func (c *xdnsClient) read(buf []byte, addr net.Addr) bool {
|
||||
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
|
||||
ip := net.ParseIP(h)
|
||||
if ip == nil {
|
||||
return nil, errors.New("invalid ip address")
|
||||
}
|
||||
}
|
||||
if domain == nil || !domain.HasType(uint16(msg.Questions[0].Type)) {
|
||||
return false
|
||||
}
|
||||
|
||||
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[:])
|
||||
port, err := strconv.Atoi(p)
|
||||
if err != nil {
|
||||
if c.closed() {
|
||||
return
|
||||
}
|
||||
errors.LogErrorInner(context.Background(), err, "recv err ", i)
|
||||
return
|
||||
return nil, errors.New("invalid port").Base(err)
|
||||
}
|
||||
if c.read(buf[:n], c.resolvers[i].Addr()) {
|
||||
c.resolverSends[i].Store(0)
|
||||
addr := &net.UDPAddr{IP: ip, Port: port}
|
||||
resolverAddrs = append(resolverAddrs, addr)
|
||||
resolverSend[addr.String()] = &atomic.Uint32{}
|
||||
}
|
||||
|
||||
conn := &xdnsConnClient{
|
||||
PacketConn: raw,
|
||||
|
||||
resolverAddrs: resolverAddrs,
|
||||
resolverTypes: resolverTypes,
|
||||
resolverIdx: 0,
|
||||
resolverSend: resolverSend,
|
||||
|
||||
clientID: make([]byte, 8),
|
||||
domains: domains,
|
||||
|
||||
pollChan: make(chan struct{}, pollLimit),
|
||||
readQueue: make(chan *packet, 256),
|
||||
writeQueue: make(chan *packet, 256),
|
||||
}
|
||||
|
||||
common.Must2(rand.Read(conn.clientID))
|
||||
|
||||
go conn.recvLoop()
|
||||
go conn.sendLoop()
|
||||
|
||||
return conn, nil
|
||||
}
|
||||
|
||||
func (c *xdnsConnClient) recvLoop() {
|
||||
var buf [finalmask.UDPSize]byte
|
||||
|
||||
for {
|
||||
if c.closed {
|
||||
break
|
||||
}
|
||||
|
||||
n, addr, err := c.PacketConn.ReadFrom(buf[:])
|
||||
if err != nil {
|
||||
if go_errors.Is(err, net.ErrClosed) {
|
||||
break
|
||||
}
|
||||
continue
|
||||
}
|
||||
|
||||
if addr == nil {
|
||||
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 {
|
||||
case c.poolCh <- struct{}{}:
|
||||
case c.readQueue <- &packet{
|
||||
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:
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
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 *xdnsClient) send() {
|
||||
defer c.wg.Done()
|
||||
|
||||
var buf [512]byte
|
||||
var data [255]byte
|
||||
|
||||
sendMsg := func(p []byte, domain *Domain, qtype uint16) {
|
||||
msg := dnsmessage.Message{
|
||||
Header: dnsmessage.Header{
|
||||
RecursionDesired: true,
|
||||
},
|
||||
Questions: []dnsmessage.Question{
|
||||
{
|
||||
Name: domain.Encode(p),
|
||||
Type: dnsmessage.Type(qtype),
|
||||
Class: dnsmessage.ClassINET,
|
||||
},
|
||||
},
|
||||
}
|
||||
if domain.edns0 > 0 {
|
||||
msg.Additionals = []dnsmessage.Resource{
|
||||
{
|
||||
Header: dnsmessage.ResourceHeader{
|
||||
Name: dnsmessage.MustNewName("."),
|
||||
Type: dnsmessage.TypeOPT,
|
||||
Class: dnsmessage.Class(domain.edns0),
|
||||
TTL: 0,
|
||||
},
|
||||
Body: &dnsmessage.OPTResource{},
|
||||
},
|
||||
}
|
||||
}
|
||||
pack := common.Must2(msg.AppendPack(buf[:0]))
|
||||
common.Must2(rand.Read(pack[:2]))
|
||||
|
||||
index := c.resolverIndex.Load()
|
||||
cur := c.resolverSends[index].Add(1)
|
||||
i := index
|
||||
for {
|
||||
i++
|
||||
if i == uint32(len(c.resolvers)) {
|
||||
i = 0
|
||||
}
|
||||
if i == index {
|
||||
break
|
||||
}
|
||||
if cur > c.resolverSends[i].Load() {
|
||||
break
|
||||
}
|
||||
}
|
||||
c.resolverIndex.Store(i)
|
||||
c.resolvers[index].Send(pack)
|
||||
}
|
||||
|
||||
send := func(p []byte) {
|
||||
domain := c.domains[mrand.Intn(len(c.domains))]
|
||||
qtype := domain.types[mrand.Intn(len(domain.types))]
|
||||
|
||||
if len(p) == 0 {
|
||||
copy(data[:], c.clientID[:])
|
||||
data[0] |= TypeMap[qtype]
|
||||
data[8] = 8
|
||||
common.Must2(rand.Read(data[9:17]))
|
||||
sendMsg(data[:17], domain, qtype)
|
||||
return
|
||||
}
|
||||
|
||||
if len(p) <= domain.cap-12 {
|
||||
copy(data[:], c.clientID[:])
|
||||
data[0] |= TypeMap[qtype]
|
||||
data[8] = 3
|
||||
common.Must2(rand.Read(data[9:12]))
|
||||
copy(data[12:], p)
|
||||
sendMsg(data[:12+len(p)], domain, qtype)
|
||||
return
|
||||
}
|
||||
|
||||
if len(p) <= 255*(domain.cap-15) {
|
||||
copy(data[:], c.clientID[:])
|
||||
data[0] |= TypeMap[qtype]
|
||||
data[8] = 3 | 0xC0
|
||||
common.Must2(rand.Read(data[9:12]))
|
||||
|
||||
fragID := byte(c.fragID.Add(1))
|
||||
fragN := len(p) / (domain.cap - 15)
|
||||
if len(p)%(domain.cap-15) > 0 {
|
||||
fragN++
|
||||
}
|
||||
|
||||
for i := range fragN {
|
||||
data[12] = fragID
|
||||
data[13] = byte(i)
|
||||
data[14] = byte(fragN)
|
||||
size := min(len(p), domain.cap-15)
|
||||
copy(data[15:], p[:size])
|
||||
sendMsg(data[:15+size], domain, qtype)
|
||||
p = p[size:]
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
errors.LogError(context.Background(), "err size ", len(p))
|
||||
}
|
||||
|
||||
ticker := time.NewTicker(initPollDelay)
|
||||
defer ticker.Stop()
|
||||
delay := initPollDelay
|
||||
p := []byte(nil)
|
||||
timeout := false
|
||||
func (c *xdnsConnClient) sendLoop() {
|
||||
pollDelay := initPollDelay
|
||||
pollTimer := time.NewTimer(pollDelay)
|
||||
for {
|
||||
var p *packet
|
||||
pollTimerExpired := false
|
||||
|
||||
select {
|
||||
case <-c.closeCh:
|
||||
return
|
||||
case p = <-c.writeQueue:
|
||||
default:
|
||||
select {
|
||||
case <-c.closeCh:
|
||||
return
|
||||
case p = <-c.sendCh:
|
||||
case <-c.poolCh:
|
||||
case <-ticker.C:
|
||||
timeout = true
|
||||
case p = <-c.writeQueue:
|
||||
case <-c.pollChan:
|
||||
case <-pollTimer.C:
|
||||
pollTimerExpired = true
|
||||
}
|
||||
}
|
||||
|
||||
if len(p) > 0 {
|
||||
if p != nil {
|
||||
select {
|
||||
case <-c.poolCh:
|
||||
case <-c.pollChan:
|
||||
default:
|
||||
}
|
||||
}
|
||||
|
||||
send(p)
|
||||
for range c.extraPoll {
|
||||
send(nil)
|
||||
}
|
||||
|
||||
if timeout {
|
||||
delay *= pollDelayMultiplier
|
||||
if delay > maxPollDelay {
|
||||
delay = maxPollDelay
|
||||
}
|
||||
timeout = false
|
||||
} else {
|
||||
delay = initPollDelay
|
||||
encoded, _ := encode(nil, c.clientID, c.domains[c.resolverIdx], c.resolverTypes[c.resolverIdx])
|
||||
p = &packet{
|
||||
p: encoded,
|
||||
}
|
||||
}
|
||||
|
||||
if pollTimerExpired {
|
||||
pollDelay = time.Duration(float64(pollDelay) * pollDelayMultiplier)
|
||||
if pollDelay > maxPollDelay {
|
||||
pollDelay = maxPollDelay
|
||||
}
|
||||
} else {
|
||||
if !pollTimer.Stop() {
|
||||
<-pollTimer.C
|
||||
}
|
||||
pollDelay = initPollDelay
|
||||
}
|
||||
pollTimer.Reset(pollDelay)
|
||||
|
||||
if c.closed {
|
||||
return
|
||||
}
|
||||
|
||||
cur := c.resolverIdx
|
||||
curSend := c.resolverSend[c.resolverAddrs[cur].String()].Add(1)
|
||||
_, _ = c.PacketConn.WriteTo(p.p, c.resolverAddrs[cur])
|
||||
for {
|
||||
c.resolverIdx += 1
|
||||
c.resolverIdx %= uint32(len(c.resolverAddrs))
|
||||
if c.resolverIdx == cur {
|
||||
break
|
||||
}
|
||||
if c.resolverSend[c.resolverAddrs[c.resolverIdx].String()].Load() < curSend {
|
||||
break
|
||||
}
|
||||
}
|
||||
ticker.Reset(delay)
|
||||
}
|
||||
}
|
||||
|
||||
func (c *xdnsClient) ReadFrom(p []byte) (n int, addr net.Addr, err error) {
|
||||
packet, ok := <-c.readCh
|
||||
if ok {
|
||||
return copy(p, packet.p), packet.addr, nil
|
||||
func (c *xdnsConnClient) ReadFrom(p []byte) (n int, addr net.Addr, err error) {
|
||||
packet, ok := <-c.readQueue
|
||||
if !ok {
|
||||
return 0, nil, net.ErrClosed
|
||||
}
|
||||
return 0, nil, io.ErrClosedPipe
|
||||
if len(p) < len(packet.p) {
|
||||
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 *xdnsClient) WriteTo(p []byte, addr net.Addr) (n int, err error) {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
if c.closed() {
|
||||
func (c *xdnsConnClient) WriteTo(p []byte, addr net.Addr) (n int, err error) {
|
||||
c.mutex.Lock()
|
||||
defer c.mutex.Unlock()
|
||||
|
||||
if c.closed {
|
||||
return 0, io.ErrClosedPipe
|
||||
}
|
||||
if len(p) == 0 || len(p) > 4096 {
|
||||
errors.LogError(context.Background(), "err size ", len(p))
|
||||
return 0, errors.New("err size")
|
||||
|
||||
idx := c.resolverIdx % uint32(len(c.resolverAddrs))
|
||||
encoded, err := encode(p, c.clientID, c.domains[idx], c.resolverTypes[idx])
|
||||
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 {
|
||||
case c.sendCh <- b:
|
||||
case c.writeQueue <- &packet{
|
||||
p: encoded,
|
||||
addr: addr,
|
||||
}:
|
||||
return len(p), nil
|
||||
default:
|
||||
errors.LogDebug(context.Background(), addr, " mask write err queue full")
|
||||
return 0, nil
|
||||
}
|
||||
return len(p), nil
|
||||
}
|
||||
|
||||
func (c *xdnsClient) Close() error {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
if c.closed() {
|
||||
func (c *xdnsConnClient) Close() error {
|
||||
c.closed = true
|
||||
return c.PacketConn.Close()
|
||||
}
|
||||
|
||||
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
|
||||
}
|
||||
close(c.closeCh)
|
||||
for i := range c.resolvers {
|
||||
c.resolvers[i].Close()
|
||||
if resp.Flags&0x000f != RcodeNoError {
|
||||
return nil
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (c *xdnsClient) LocalAddr() net.Addr { return &net.UDPAddr{IP: []byte{0, 0, 0, 0}} }
|
||||
|
||||
func (c *xdnsClient) SetDeadline(t time.Time) error { return errors.New("not support") }
|
||||
|
||||
func (c *xdnsClient) SetReadDeadline(t time.Time) error { return errors.New("not support") }
|
||||
|
||||
func (c *xdnsClient) SetWriteDeadline(t time.Time) error { return errors.New("not support") }
|
||||
|
||||
type ClientID [8]byte
|
||||
|
||||
func NewClientID() ClientID {
|
||||
var id ClientID
|
||||
common.Must2(rand.Read(id[:]))
|
||||
id[0] &= 0xFC
|
||||
return id
|
||||
}
|
||||
|
||||
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[:]}
|
||||
|
||||
if len(resp.Answer) == 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
for _, answer := range resp.Answer {
|
||||
var ok bool
|
||||
for _, domain := range domains {
|
||||
_, ok = answer.Name.TrimSuffix(domain)
|
||||
if ok {
|
||||
break
|
||||
}
|
||||
}
|
||||
if !ok {
|
||||
return nil
|
||||
}
|
||||
}
|
||||
|
||||
return decodeResponsePayload(resp.Answer)
|
||||
}
|
||||
|
||||
@@ -6,9 +6,9 @@ import (
|
||||
)
|
||||
|
||||
func (c *Config) WrapPacketConnClient(conn net.PacketConn, dest *net.Destination, dialer *finalmask.Dialer) (net.PacketConn, error) {
|
||||
return NewClient(c, dialer)
|
||||
return NewConnClient(c, conn)
|
||||
}
|
||||
|
||||
func (c *Config) WrapPacketConnServer(conn net.PacketConn, addr net.Addr, lc *finalmask.ListenConfig) (net.PacketConn, error) {
|
||||
return NewServer(c, conn)
|
||||
return NewConnServer(c, conn)
|
||||
}
|
||||
|
||||
@@ -7,7 +7,6 @@
|
||||
package xdns
|
||||
|
||||
import (
|
||||
serial "github.com/xtls/xray-core/common/serial"
|
||||
protoreflect "google.golang.org/protobuf/reflect/protoreflect"
|
||||
protoimpl "google.golang.org/protobuf/runtime/protoimpl"
|
||||
reflect "reflect"
|
||||
@@ -22,94 +21,17 @@ const (
|
||||
_ = 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 {
|
||||
state protoimpl.MessageState `protogen:"open.v1"`
|
||||
Domains []*DomainProto `protobuf:"bytes,1,rep,name=domains,proto3" json:"domains,omitempty"`
|
||||
Resolvers []*serial.TypedMessage `protobuf:"bytes,2,rep,name=resolvers,proto3" json:"resolvers,omitempty"`
|
||||
ExtraPoll int32 `protobuf:"varint,3,opt,name=extra_poll,json=extraPoll,proto3" json:"extra_poll,omitempty"`
|
||||
Domains []string `protobuf:"bytes,1,rep,name=domains,proto3" json:"domains,omitempty"`
|
||||
Resolvers []string `protobuf:"bytes,2,rep,name=resolvers,proto3" json:"resolvers,omitempty"`
|
||||
unknownFields protoimpl.UnknownFields
|
||||
sizeCache protoimpl.SizeCache
|
||||
}
|
||||
|
||||
func (x *Config) Reset() {
|
||||
*x = Config{}
|
||||
mi := &file_transport_internet_finalmask_xdns_config_proto_msgTypes[1]
|
||||
mi := &file_transport_internet_finalmask_xdns_config_proto_msgTypes[0]
|
||||
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
|
||||
ms.StoreMessageInfo(mi)
|
||||
}
|
||||
@@ -121,7 +43,7 @@ func (x *Config) String() string {
|
||||
func (*Config) ProtoMessage() {}
|
||||
|
||||
func (x *Config) ProtoReflect() protoreflect.Message {
|
||||
mi := &file_transport_internet_finalmask_xdns_config_proto_msgTypes[1]
|
||||
mi := &file_transport_internet_finalmask_xdns_config_proto_msgTypes[0]
|
||||
if x != nil {
|
||||
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
|
||||
if ms.LoadMessageInfo() == nil {
|
||||
@@ -134,139 +56,31 @@ func (x *Config) ProtoReflect() protoreflect.Message {
|
||||
|
||||
// Deprecated: Use Config.ProtoReflect.Descriptor instead.
|
||||
func (*Config) Descriptor() ([]byte, []int) {
|
||||
return file_transport_internet_finalmask_xdns_config_proto_rawDescGZIP(), []int{1}
|
||||
return file_transport_internet_finalmask_xdns_config_proto_rawDescGZIP(), []int{0}
|
||||
}
|
||||
|
||||
func (x *Config) GetDomains() []*DomainProto {
|
||||
func (x *Config) GetDomains() []string {
|
||||
if x != nil {
|
||||
return x.Domains
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (x *Config) GetResolvers() []*serial.TypedMessage {
|
||||
func (x *Config) GetResolvers() []string {
|
||||
if x != nil {
|
||||
return x.Resolvers
|
||||
}
|
||||
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
|
||||
|
||||
const file_transport_internet_finalmask_xdns_config_proto_rawDesc = "" +
|
||||
"\n" +
|
||||
".transport/internet/finalmask/xdns/config.proto\x12&xray.transport.internet.finalmask.xdns\x1a!common/serial/typed_message.proto\"\x8b\x01\n" +
|
||||
"\vDomainProto\x12\x12\n" +
|
||||
"\x04name\x18\x01 \x01(\tR\x04name\x12\x1b\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" +
|
||||
".transport/internet/finalmask/xdns/config.proto\x12&xray.transport.internet.finalmask.xdns\"@\n" +
|
||||
"\x06Config\x12\x18\n" +
|
||||
"\adomains\x18\x01 \x03(\tR\adomains\x12\x1c\n" +
|
||||
"\tresolvers\x18\x02 \x03(\tR\tresolversB\x94\x01\n" +
|
||||
"*com.xray.transport.internet.finalmask.xdnsP\x01Z;github.com/xtls/xray-core/transport/internet/finalmask/xdns\xaa\x02&Xray.Transport.Internet.Finalmask.Xdnsb\x06proto3"
|
||||
|
||||
var (
|
||||
@@ -281,22 +95,16 @@ func file_transport_internet_finalmask_xdns_config_proto_rawDescGZIP() []byte {
|
||||
return file_transport_internet_finalmask_xdns_config_proto_rawDescData
|
||||
}
|
||||
|
||||
var file_transport_internet_finalmask_xdns_config_proto_msgTypes = make([]protoimpl.MessageInfo, 4)
|
||||
var file_transport_internet_finalmask_xdns_config_proto_msgTypes = make([]protoimpl.MessageInfo, 1)
|
||||
var file_transport_internet_finalmask_xdns_config_proto_goTypes = []any{
|
||||
(*DomainProto)(nil), // 0: xray.transport.internet.finalmask.xdns.DomainProto
|
||||
(*Config)(nil), // 1: xray.transport.internet.finalmask.xdns.Config
|
||||
(*TCPResolverProto)(nil), // 2: xray.transport.internet.finalmask.xdns.TCPResolverProto
|
||||
(*UDPResolverProto)(nil), // 3: xray.transport.internet.finalmask.xdns.UDPResolverProto
|
||||
(*serial.TypedMessage)(nil), // 4: xray.common.serial.TypedMessage
|
||||
(*Config)(nil), // 0: xray.transport.internet.finalmask.xdns.Config
|
||||
}
|
||||
var file_transport_internet_finalmask_xdns_config_proto_depIdxs = []int32{
|
||||
0, // 0: xray.transport.internet.finalmask.xdns.Config.domains:type_name -> xray.transport.internet.finalmask.xdns.DomainProto
|
||||
4, // 1: xray.transport.internet.finalmask.xdns.Config.resolvers:type_name -> xray.common.serial.TypedMessage
|
||||
2, // [2:2] is the sub-list for method output_type
|
||||
2, // [2:2] is the sub-list for method input_type
|
||||
2, // [2:2] is the sub-list for extension type_name
|
||||
2, // [2:2] is the sub-list for extension extendee
|
||||
0, // [0:2] is the sub-list for field type_name
|
||||
0, // [0:0] is the sub-list for method output_type
|
||||
0, // [0:0] is the sub-list for method input_type
|
||||
0, // [0:0] is the sub-list for extension type_name
|
||||
0, // [0:0] is the sub-list for extension extendee
|
||||
0, // [0:0] is the sub-list for field type_name
|
||||
}
|
||||
|
||||
func init() { file_transport_internet_finalmask_xdns_config_proto_init() }
|
||||
@@ -310,7 +118,7 @@ func file_transport_internet_finalmask_xdns_config_proto_init() {
|
||||
GoPackagePath: reflect.TypeOf(x{}).PkgPath(),
|
||||
RawDescriptor: unsafe.Slice(unsafe.StringData(file_transport_internet_finalmask_xdns_config_proto_rawDesc), len(file_transport_internet_finalmask_xdns_config_proto_rawDesc)),
|
||||
NumEnums: 0,
|
||||
NumMessages: 4,
|
||||
NumMessages: 1,
|
||||
NumExtensions: 0,
|
||||
NumServices: 0,
|
||||
},
|
||||
|
||||
@@ -6,26 +6,7 @@ option go_package = "github.com/xtls/xray-core/transport/internet/finalmask/xdns
|
||||
option java_package = "com.xray.transport.internet.finalmask.xdns";
|
||||
option java_multiple_files = true;
|
||||
|
||||
import "common/serial/typed_message.proto";
|
||||
|
||||
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;
|
||||
repeated string domains = 1;
|
||||
repeated string resolvers = 2;
|
||||
}
|
||||
@@ -0,0 +1,581 @@
|
||||
// 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()
|
||||
}
|
||||
@@ -0,0 +1,953 @@
|
||||
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)
|
||||
}
|
||||
}
|
||||
@@ -1,215 +0,0 @@
|
||||
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
|
||||
}
|
||||
@@ -1,171 +0,0 @@
|
||||
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)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,226 @@
|
||||
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
|
||||
}
|
||||
@@ -1,31 +0,0 @@
|
||||
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")
|
||||
}
|
||||
}
|
||||
@@ -1,143 +0,0 @@
|
||||
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)
|
||||
}
|
||||
@@ -1,130 +0,0 @@
|
||||
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)
|
||||
}
|
||||
@@ -1,392 +0,0 @@
|
||||
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)
|
||||
}
|
||||
}
|
||||
@@ -1,385 +1,512 @@
|
||||
package xdns
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/binary"
|
||||
go_errors "errors"
|
||||
"io"
|
||||
"net"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/xtls/xray-core/common"
|
||||
"github.com/xtls/xray-core/common/errors"
|
||||
"github.com/xtls/xray-core/common/net"
|
||||
"golang.org/x/net/dns/dnsmessage"
|
||||
"github.com/xtls/xray-core/transport/internet/finalmask"
|
||||
)
|
||||
|
||||
const (
|
||||
maxResponseDelay = time.Second
|
||||
idleTimeout = 10 * time.Second
|
||||
responseTTL = 60
|
||||
maxResponseDelay = 1 * time.Second
|
||||
)
|
||||
|
||||
type resp struct {
|
||||
msg dnsmessage.Message
|
||||
addr net.Addr
|
||||
var (
|
||||
maxUDPPayload = 1280 - 40 - 8
|
||||
maxEncodedPayloadTXT = computeMaxEncodedPayloadForType(maxUDPPayload, RRTypeTXT)
|
||||
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 Rec struct {
|
||||
resp *Resp
|
||||
clientID ClientID
|
||||
addr net.Addr
|
||||
type record struct {
|
||||
Resp *Message
|
||||
Addr net.Addr
|
||||
// ClientID [8]byte
|
||||
ClientAddr net.Addr
|
||||
}
|
||||
|
||||
type xdnsServer struct {
|
||||
type queue struct {
|
||||
last time.Time
|
||||
rrType uint16
|
||||
queue chan []byte
|
||||
stash chan []byte
|
||||
}
|
||||
|
||||
type xdnsConnServer struct {
|
||||
net.PacketConn
|
||||
|
||||
domains []*Domain
|
||||
fragManager *FragManager
|
||||
sendManager *SendManager
|
||||
domains []domainSpec
|
||||
|
||||
readCh chan packet
|
||||
recCh chan *Rec
|
||||
drCh chan resp
|
||||
closeCh chan struct{}
|
||||
wg sync.WaitGroup
|
||||
mu sync.RWMutex
|
||||
ch chan *record
|
||||
readQueue chan *packet
|
||||
writeQueueMap map[string]*queue
|
||||
|
||||
closed bool
|
||||
mutex sync.Mutex
|
||||
}
|
||||
|
||||
func NewServer(c *Config, raw net.PacketConn) (net.PacketConn, error) {
|
||||
func NewConnServer(c *Config, raw net.PacketConn) (net.PacketConn, error) {
|
||||
if len(c.Domains) == 0 {
|
||||
return nil, errors.New("empty domains")
|
||||
}
|
||||
domains := make([]*Domain, 0, len(c.Domains))
|
||||
for i := range c.Domains {
|
||||
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))
|
||||
domains := make([]domainSpec, 0, len(c.Domains))
|
||||
for _, domain := range c.Domains {
|
||||
domain, err := parseDomainSpec(domain, "")
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
domains = append(domains, domain)
|
||||
}
|
||||
server := &xdnsServer{
|
||||
|
||||
conn := &xdnsConnServer{
|
||||
PacketConn: raw,
|
||||
|
||||
domains: domains,
|
||||
fragManager: NewFragManager(),
|
||||
sendManager: NewSendManager(),
|
||||
domains: domains,
|
||||
|
||||
readCh: make(chan packet),
|
||||
recCh: make(chan *Rec, 255),
|
||||
drCh: make(chan resp),
|
||||
closeCh: make(chan struct{}),
|
||||
ch: make(chan *record, 500),
|
||||
readQueue: make(chan *packet, 512),
|
||||
writeQueueMap: make(map[string]*queue),
|
||||
}
|
||||
go server.run()
|
||||
return server, nil
|
||||
|
||||
go conn.clean()
|
||||
go conn.recvLoop()
|
||||
go conn.sendLoop()
|
||||
|
||||
return conn, nil
|
||||
}
|
||||
|
||||
func (c *xdnsServer) closed() bool {
|
||||
select {
|
||||
case <-c.closeCh:
|
||||
return true
|
||||
default:
|
||||
func (c *xdnsConnServer) clean() {
|
||||
f := func() bool {
|
||||
c.mutex.Lock()
|
||||
defer c.mutex.Unlock()
|
||||
|
||||
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
|
||||
}
|
||||
|
||||
for {
|
||||
time.Sleep(idleTimeout / 2)
|
||||
if f() {
|
||||
return
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (c *xdnsServer) decref(msg dnsmessage.Message, addr net.Addr) {
|
||||
func (c *xdnsConnServer) ensureQueue(addr net.Addr) *queue {
|
||||
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 {
|
||||
case c.drCh <- resp{msg: msg, addr: addr}:
|
||||
case queue.stash <- p:
|
||||
default:
|
||||
}
|
||||
}
|
||||
|
||||
func (c *xdnsServer) read(buf []byte, addr net.Addr) {
|
||||
msg := dnsmessage.Message{}
|
||||
if err := msg.Unpack(buf); err != nil {
|
||||
return
|
||||
}
|
||||
if msg.Header.Response {
|
||||
return
|
||||
}
|
||||
func (c *xdnsConnServer) recvLoop() {
|
||||
var buf [finalmask.UDPSize]byte
|
||||
|
||||
if msg.Header.OpCode != 0 {
|
||||
msg.Header.Response = true
|
||||
msg.Header.RCode = dnsmessage.RCodeNotImplemented
|
||||
c.decref(msg, addr)
|
||||
return
|
||||
}
|
||||
|
||||
if len(msg.Questions) != 1 {
|
||||
msg.Header.Response = true
|
||||
msg.Header.RCode = dnsmessage.RCodeFormatError
|
||||
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
|
||||
}
|
||||
opt = true
|
||||
edns0 = uint16(msg.Additionals[i].Header.Class)
|
||||
if ver := (msg.Additionals[i].Header.TTL >> 16) & 0xFF; ver != 0 {
|
||||
msg.Header.RCode = dnsmessage.RCodeSuccess
|
||||
msg.Additionals[i].Header.TTL = 1 << 24
|
||||
c.decref(msg, addr)
|
||||
return
|
||||
}
|
||||
}
|
||||
}
|
||||
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)
|
||||
|
||||
var domain *Domain
|
||||
for i := range c.domains {
|
||||
if c.domains[i].IsDomain(msg.Questions[0].Name) {
|
||||
domain = c.domains[i]
|
||||
for {
|
||||
if c.closed {
|
||||
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
|
||||
}
|
||||
|
||||
var decoded [255]byte
|
||||
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]))
|
||||
|
||||
r := NewResp(msg, domain, edns0)
|
||||
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)
|
||||
}
|
||||
|
||||
if decoded[8]&0x3F == 8 {
|
||||
return
|
||||
}
|
||||
p := pool4K.Get().([]byte)
|
||||
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 *xdnsServer) run() {
|
||||
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
|
||||
if go_errors.Is(err, net.ErrClosed) {
|
||||
break
|
||||
}
|
||||
errors.LogErrorInner(context.Background(), err, "recv err")
|
||||
return
|
||||
continue
|
||||
}
|
||||
c.read(buf[:n], addr)
|
||||
|
||||
query, err := MessageFromWireFormat(buf[:n])
|
||||
if err != nil {
|
||||
errors.LogDebug(context.Background(), addr, " xdns from wireformat err ", err)
|
||||
continue
|
||||
}
|
||||
|
||||
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")
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
errors.LogDebug(context.Background(), "xdns closed")
|
||||
|
||||
close(c.ch)
|
||||
close(c.readQueue)
|
||||
|
||||
c.mutex.Lock()
|
||||
defer c.mutex.Unlock()
|
||||
|
||||
c.closed = true
|
||||
for key, q := range c.writeQueueMap {
|
||||
close(q.queue)
|
||||
close(q.stash)
|
||||
delete(c.writeQueueMap, key)
|
||||
}
|
||||
}
|
||||
|
||||
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
|
||||
func (c *xdnsConnServer) sendLoop() {
|
||||
var nextRec *record
|
||||
for {
|
||||
var err error
|
||||
rec := nextRec
|
||||
nextRec = nil
|
||||
|
||||
if rec == nil {
|
||||
select {
|
||||
case rec = <-c.recCh:
|
||||
case <-c.closeCh:
|
||||
return
|
||||
var ok bool
|
||||
rec, ok = <-c.ch
|
||||
if !ok {
|
||||
break
|
||||
}
|
||||
}
|
||||
|
||||
ch, stash := c.sendManager.Pop(rec.clientID)
|
||||
left := rec.resp.cap
|
||||
timer.Reset(maxResponseDelay)
|
||||
var ps [][]byte
|
||||
for {
|
||||
var p []byte
|
||||
select {
|
||||
case p = <-stash:
|
||||
default:
|
||||
if rec.Resp.Rcode() == RcodeNoError && len(rec.Resp.Question) == 1 {
|
||||
var payload bytes.Buffer
|
||||
limit := maxEncodedPayloadForType(rec.Resp.Question[0].Type)
|
||||
timer := time.NewTimer(maxResponseDelay)
|
||||
|
||||
for {
|
||||
c.mutex.Lock()
|
||||
q := c.ensureQueue(rec.ClientAddr)
|
||||
if q == nil {
|
||||
c.mutex.Unlock()
|
||||
return
|
||||
}
|
||||
q.rrType = rec.Resp.Question[0].Type
|
||||
c.mutex.Unlock()
|
||||
|
||||
var p []byte
|
||||
|
||||
select {
|
||||
case p = <-stash:
|
||||
case p = <-ch:
|
||||
case p = <-q.stash:
|
||||
default:
|
||||
select {
|
||||
case p = <-stash:
|
||||
case p = <-ch:
|
||||
case <-timer.C:
|
||||
case nextRec = <-c.recCh:
|
||||
case p = <-q.stash:
|
||||
case p = <-q.queue:
|
||||
default:
|
||||
select {
|
||||
case p = <-q.stash:
|
||||
case p = <-q.queue:
|
||||
case <-timer.C:
|
||||
case nextRec = <-c.ch:
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
if len(p) == 0 {
|
||||
break
|
||||
}
|
||||
timer.Reset(0)
|
||||
left -= 2 + len(p)
|
||||
if left < 0 {
|
||||
if len(ps) == 0 {
|
||||
errors.LogError(context.Background(), "err size ", len(p))
|
||||
|
||||
timer.Reset(0)
|
||||
|
||||
if len(p) == 0 {
|
||||
break
|
||||
}
|
||||
c.sendManager.Stash(rec.clientID, p)
|
||||
break
|
||||
|
||||
limit -= 2 + len(p)
|
||||
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()
|
||||
|
||||
d := data[:0]
|
||||
for i := range ps {
|
||||
l := len(ps[i])
|
||||
if i == len(ps)-1 {
|
||||
l |= 0xC000
|
||||
timer.Stop()
|
||||
rec.Resp.Answer, err = answersForPayload(rec.Resp.Question[0], responseTTL, payload.Bytes())
|
||||
if err != nil {
|
||||
errors.LogDebug(context.Background(), rec.Addr, " ", rec.ClientAddr, " xdns encode err ", err)
|
||||
continue
|
||||
}
|
||||
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)
|
||||
}
|
||||
}
|
||||
|
||||
func (c *xdnsServer) dr() {
|
||||
defer c.wg.Done()
|
||||
buf, err := rec.Resp.WireFormat()
|
||||
if err != nil {
|
||||
errors.LogDebug(context.Background(), rec.Addr, " ", rec.ClientAddr, " xdns wireformat err ", err)
|
||||
continue
|
||||
}
|
||||
|
||||
var buf [512]byte
|
||||
for {
|
||||
select {
|
||||
case <-c.closeCh:
|
||||
if len(buf) > maxUDPPayload {
|
||||
errors.LogDebug(context.Background(), rec.Addr, " ", rec.ClientAddr, " xdns truncate ", len(buf))
|
||||
buf = buf[:maxUDPPayload]
|
||||
buf[2] |= 0x02
|
||||
}
|
||||
|
||||
if c.closed {
|
||||
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 *xdnsServer) ReadFrom(p []byte) (n int, addr net.Addr, err error) {
|
||||
packet, ok := <-c.readCh
|
||||
if ok {
|
||||
n = copy(p, packet.p)
|
||||
pool4K.Put(packet.p[:cap(packet.p)])
|
||||
return n, packet.addr, nil
|
||||
func (c *xdnsConnServer) ReadFrom(p []byte) (n int, addr net.Addr, err error) {
|
||||
packet, ok := <-c.readQueue
|
||||
if !ok {
|
||||
return 0, nil, net.ErrClosed
|
||||
}
|
||||
return 0, nil, io.ErrClosedPipe
|
||||
if len(p) < len(packet.p) {
|
||||
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 *xdnsServer) WriteTo(p []byte, addr net.Addr) (n int, err error) {
|
||||
if c.closed() {
|
||||
func (c *xdnsConnServer) WriteTo(p []byte, addr net.Addr) (n int, err error) {
|
||||
c.mutex.Lock()
|
||||
defer c.mutex.Unlock()
|
||||
|
||||
q := c.ensureQueue(addr)
|
||||
if q == nil {
|
||||
return 0, io.ErrClosedPipe
|
||||
}
|
||||
if len(p) == 0 || len(p) > 4096 {
|
||||
errors.LogError(context.Background(), "err size ", len(p))
|
||||
return 0, errors.New("err size")
|
||||
limit := maxEncodedPayloadForType(q.rrType)
|
||||
if q.rrType == 0 {
|
||||
limit = maxEncodedPayloadTXT
|
||||
}
|
||||
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 *xdnsServer) Close() error {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
if c.closed() {
|
||||
return nil
|
||||
}
|
||||
close(c.closeCh)
|
||||
_ = c.PacketConn.Close()
|
||||
return nil
|
||||
func (c *xdnsConnServer) Close() error {
|
||||
c.closed = true
|
||||
return c.PacketConn.Close()
|
||||
}
|
||||
|
||||
func (c *xdnsServer) SetDeadline(t time.Time) error { return errors.New("not support") }
|
||||
func nextPacketServer(r *bytes.Reader) ([]byte, error) {
|
||||
eof := func(err error) error {
|
||||
if err == io.EOF {
|
||||
err = io.ErrUnexpectedEOF
|
||||
}
|
||||
return err
|
||||
}
|
||||
|
||||
func (c *xdnsServer) SetReadDeadline(t time.Time) error { return errors.New("not support") }
|
||||
for {
|
||||
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 (c *xdnsServer) SetWriteDeadline(t time.Time) error { return errors.New("not support") }
|
||||
func responseFor(query *Message, domains []domainSpec) (*Message, []byte) {
|
||||
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
|
||||
}
|
||||
|
||||
@@ -0,0 +1,80 @@
|
||||
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
|
||||
}
|
||||
@@ -1,208 +0,0 @@
|
||||
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,6 +310,13 @@ func (c *xicmpConnClient) Close() error {
|
||||
_ = c.icmp4.Close()
|
||||
_ = c.icmp6.Close()
|
||||
c.wg.Wait()
|
||||
select {
|
||||
case p := <-c.readCh:
|
||||
if p.p != nil {
|
||||
pool.Put(p.p)
|
||||
}
|
||||
default:
|
||||
}
|
||||
close(c.readCh)
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -329,6 +329,13 @@ func (c *xicmpConnServer) Close() error {
|
||||
_ = c.icmp4.Close()
|
||||
_ = c.icmp6.Close()
|
||||
c.wg.Wait()
|
||||
select {
|
||||
case p := <-c.readCh:
|
||||
if p.p != nil {
|
||||
pool.Put(p.p)
|
||||
}
|
||||
default:
|
||||
}
|
||||
close(c.readCh)
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -340,6 +340,13 @@ func (c *xicmpConnServer) Close() error {
|
||||
_ = c.icmp4.Close()
|
||||
_ = c.icmp6.Close()
|
||||
c.wg.Wait()
|
||||
select {
|
||||
case p := <-c.readCh:
|
||||
if p.p != nil {
|
||||
pool.Put(p.p)
|
||||
}
|
||||
default:
|
||||
}
|
||||
close(c.readCh)
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -3,8 +3,6 @@ package httpupgrade
|
||||
import (
|
||||
"bufio"
|
||||
"context"
|
||||
"crypto/rand"
|
||||
"encoding/base64"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"strings"
|
||||
@@ -99,16 +97,6 @@ func dialhttpUpgrade(ctx context.Context, dest net.Destination, streamSettings *
|
||||
req.Header.Set("Connection", "Upgrade")
|
||||
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)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
|
||||
@@ -3,9 +3,7 @@ package httpupgrade
|
||||
import (
|
||||
"bufio"
|
||||
"context"
|
||||
"crypto/sha1"
|
||||
"crypto/tls"
|
||||
"encoding/base64"
|
||||
"io"
|
||||
"net/http"
|
||||
"strings"
|
||||
@@ -83,11 +81,6 @@ func (s *server) upgrade(conn net.Conn) (stat.Connection, error) {
|
||||
}
|
||||
resp.Header.Set("Connection", "Upgrade")
|
||||
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)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
|
||||
@@ -119,7 +119,7 @@ func (c *client) dial(ctx context.Context) error {
|
||||
if err != nil {
|
||||
return errors.New("failed to dial to dest").Base(err)
|
||||
}
|
||||
pktConn = conn.(*net.PacketConnWrapper).PacketConn
|
||||
pktConn = conn.(*finalmask.PacketConnWrapper).PacketConn
|
||||
udpAddr = conn.RemoteAddr()
|
||||
} else {
|
||||
conn, err := internet.DialSystem(ctx, c.dest, c.socketConfig)
|
||||
@@ -127,7 +127,7 @@ func (c *client) dial(ctx context.Context) error {
|
||||
return errors.New("failed to dial to dest").Base(err)
|
||||
}
|
||||
switch c := conn.(type) {
|
||||
case *net.PacketConnWrapper:
|
||||
case *internet.PacketConnWrapper:
|
||||
pktConn = c.PacketConn
|
||||
udpAddr = c.RemoteAddr()
|
||||
case *cnc.Connection:
|
||||
|
||||
@@ -19,6 +19,7 @@ import (
|
||||
"github.com/xtls/xray-core/common/net/cnc"
|
||||
"github.com/xtls/xray-core/common/utils"
|
||||
"github.com/xtls/xray-core/transport/internet"
|
||||
"github.com/xtls/xray-core/transport/internet/finalmask"
|
||||
"github.com/xtls/xray-core/transport/internet/hysteria/congestion"
|
||||
"github.com/xtls/xray-core/transport/internet/hysteria/congestion/bbr"
|
||||
"github.com/xtls/xray-core/transport/internet/masque/connectip"
|
||||
@@ -79,7 +80,7 @@ func Dial(ctx context.Context, dest net.Destination, streamSettings *internet.Me
|
||||
if err != nil {
|
||||
return nil, errors.New("failed to dial to dest").Base(err)
|
||||
}
|
||||
pktConn = conn.(*net.PacketConnWrapper).PacketConn
|
||||
pktConn = conn.(*finalmask.PacketConnWrapper).PacketConn
|
||||
udpAddr = conn.RemoteAddr()
|
||||
} else {
|
||||
conn, err := internet.DialSystem(ctx, dest, streamSettings.SocketSettings)
|
||||
@@ -87,7 +88,7 @@ func Dial(ctx context.Context, dest net.Destination, streamSettings *internet.Me
|
||||
return nil, errors.New("failed to dial to dest").Base(err)
|
||||
}
|
||||
switch c := conn.(type) {
|
||||
case *net.PacketConnWrapper:
|
||||
case *internet.PacketConnWrapper:
|
||||
pktConn = c.PacketConn
|
||||
udpAddr = c.RemoteAddr()
|
||||
case *cnc.Connection:
|
||||
|
||||
@@ -54,10 +54,11 @@ func ToMemoryStreamConfig(s *StreamConfig) (*MemoryStreamConfig, error) {
|
||||
mss.SecurityType = s.SecurityType
|
||||
mss.SecuritySettings = ess
|
||||
}
|
||||
if s != nil && (len(s.Tcpmasks) != 0 || len(s.Udpmasks) != 0) {
|
||||
var tcpMasks []finalmask.TCPMask
|
||||
var udpMasks []finalmask.UDPMask
|
||||
|
||||
var tcpMasks []finalmask.TCPMask
|
||||
var udpMasks []finalmask.UDPMask
|
||||
|
||||
if s != nil {
|
||||
for i := range s.Tcpmasks {
|
||||
instance := common.Must2(s.Tcpmasks[i].GetInstance())
|
||||
tcpMasks = append(tcpMasks, instance.(finalmask.TCPMask))
|
||||
@@ -66,38 +67,38 @@ func ToMemoryStreamConfig(s *StreamConfig) (*MemoryStreamConfig, error) {
|
||||
instance := common.Must2(s.Udpmasks[i].GetInstance())
|
||||
udpMasks = append(udpMasks, instance.(finalmask.UDPMask))
|
||||
}
|
||||
|
||||
dialTCP := func(ctx context.Context, dest net.Destination) (net.Conn, error) {
|
||||
return DialSystem(ctx, dest, mss.SocketSettings)
|
||||
}
|
||||
listen := func(ctx context.Context, addr net.Addr) (net.Listener, error) {
|
||||
return ListenSystem(ctx, addr, mss.SocketSettings)
|
||||
}
|
||||
dialUDP := func(ctx context.Context, dest net.Destination) (net.PacketConn, net.Addr, error) {
|
||||
conn, err := DialSystem(ctx, dest, mss.SocketSettings)
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
var newConn net.PacketConn
|
||||
var udpAddr net.Addr
|
||||
switch c := conn.(type) {
|
||||
case *net.PacketConnWrapper:
|
||||
newConn = c.PacketConn
|
||||
udpAddr = conn.RemoteAddr()
|
||||
case *cnc.Connection:
|
||||
newConn = &FakePacketConn{Conn: c}
|
||||
udpAddr = &net.UDPAddr{IP: []byte{0, 0, 0, 0}, Port: 0}
|
||||
default:
|
||||
panic(reflect.TypeOf(c))
|
||||
}
|
||||
return newConn, udpAddr, nil
|
||||
}
|
||||
listenPacket := func(ctx context.Context, addr net.Addr) (net.PacketConn, error) {
|
||||
return ListenSystemPacket(ctx, addr, mss.SocketSettings)
|
||||
}
|
||||
mss.FinalMask = finalmask.NewFinalMask(tcpMasks, udpMasks, dialTCP, listen, dialUDP, listenPacket)
|
||||
}
|
||||
|
||||
dialTCP := func(ctx context.Context, dest net.Destination) (net.Conn, error) {
|
||||
return DialSystem(ctx, dest, mss.SocketSettings)
|
||||
}
|
||||
listen := func(ctx context.Context, addr net.Addr) (net.Listener, error) {
|
||||
return ListenSystem(ctx, addr, mss.SocketSettings)
|
||||
}
|
||||
dialUDP := func(ctx context.Context, dest net.Destination) (net.PacketConn, net.Addr, error) {
|
||||
conn, err := DialSystem(ctx, dest, mss.SocketSettings)
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
var newConn net.PacketConn
|
||||
var udpAddr net.Addr
|
||||
switch c := conn.(type) {
|
||||
case *PacketConnWrapper:
|
||||
newConn = c.PacketConn
|
||||
udpAddr = conn.RemoteAddr()
|
||||
case *cnc.Connection:
|
||||
newConn = &FakePacketConn{Conn: c}
|
||||
udpAddr = &net.UDPAddr{IP: []byte{0, 0, 0, 0}, Port: 0}
|
||||
default:
|
||||
panic(reflect.TypeOf(c))
|
||||
}
|
||||
return newConn, udpAddr, nil
|
||||
}
|
||||
listenPacket := func(ctx context.Context, addr net.Addr) (net.PacketConn, error) {
|
||||
return ListenSystemPacket(ctx, addr, mss.SocketSettings)
|
||||
}
|
||||
mss.FinalMask = finalmask.NewFinalMask(tcpMasks, udpMasks, dialTCP, listen, dialUDP, listenPacket)
|
||||
|
||||
if s != nil && s.QuicParams != nil {
|
||||
mss.QuicParams = s.QuicParams
|
||||
}
|
||||
|
||||
@@ -18,7 +18,6 @@ import (
|
||||
"regexp"
|
||||
"strings"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"time"
|
||||
"unsafe"
|
||||
|
||||
@@ -37,18 +36,6 @@ import (
|
||||
|
||||
type Conn struct {
|
||||
*reality.Conn
|
||||
suppressCloseNotify atomic.Bool
|
||||
}
|
||||
|
||||
func (c *Conn) SuppressCloseNotify() {
|
||||
c.suppressCloseNotify.Store(true)
|
||||
}
|
||||
|
||||
func (c *Conn) Close() error {
|
||||
if c.suppressCloseNotify.Load() {
|
||||
return c.Conn.NetConn().Close()
|
||||
}
|
||||
return c.Conn.Close()
|
||||
}
|
||||
|
||||
func (c *Conn) HandshakeAddress() net.Address {
|
||||
@@ -69,22 +56,10 @@ func Server(c net.Conn, config *reality.Config) (net.Conn, error) {
|
||||
|
||||
type UConn struct {
|
||||
*utls.UConn
|
||||
Config *Config
|
||||
ServerName string
|
||||
AuthKey []byte
|
||||
Verified bool
|
||||
suppressCloseNotify atomic.Bool
|
||||
}
|
||||
|
||||
func (c *UConn) SuppressCloseNotify() {
|
||||
c.suppressCloseNotify.Store(true)
|
||||
}
|
||||
|
||||
func (c *UConn) Close() error {
|
||||
if c.suppressCloseNotify.Load() {
|
||||
return c.NetConn().Close()
|
||||
}
|
||||
return c.UConn.Close()
|
||||
Config *Config
|
||||
ServerName string
|
||||
AuthKey []byte
|
||||
Verified bool
|
||||
}
|
||||
|
||||
func (c *UConn) HandshakeAddress() net.Address {
|
||||
|
||||
@@ -25,6 +25,7 @@ import (
|
||||
"github.com/xtls/xray-core/common/signal/done"
|
||||
"github.com/xtls/xray-core/transport/internet"
|
||||
"github.com/xtls/xray-core/transport/internet/browser_dialer"
|
||||
"github.com/xtls/xray-core/transport/internet/finalmask"
|
||||
"github.com/xtls/xray-core/transport/internet/hysteria/congestion"
|
||||
"github.com/xtls/xray-core/transport/internet/hysteria/congestion/bbr"
|
||||
"github.com/xtls/xray-core/transport/internet/reality"
|
||||
@@ -199,7 +200,7 @@ func createHTTPClient(dest net.Destination, streamSettings *internet.MemoryStrea
|
||||
if err != nil {
|
||||
return nil, errors.New("failed to dial to dest").Base(err)
|
||||
}
|
||||
pktConn = conn.(*net.PacketConnWrapper).PacketConn
|
||||
pktConn = conn.(*finalmask.PacketConnWrapper).PacketConn
|
||||
udpAddr = conn.RemoteAddr()
|
||||
} else {
|
||||
conn, err := internet.DialSystem(ctx, dest, streamSettings.SocketSettings)
|
||||
@@ -207,7 +208,7 @@ func createHTTPClient(dest net.Destination, streamSettings *internet.MemoryStrea
|
||||
return nil, errors.New("failed to dial to dest").Base(err)
|
||||
}
|
||||
switch c := conn.(type) {
|
||||
case *net.PacketConnWrapper:
|
||||
case *internet.PacketConnWrapper:
|
||||
pktConn = c.PacketConn
|
||||
udpAddr = c.RemoteAddr()
|
||||
case *cnc.Connection:
|
||||
|
||||
@@ -86,7 +86,7 @@ func (d *DefaultSystemDialer) Dial(ctx context.Context, src net.Address, dest ne
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &net.PacketConnWrapper{
|
||||
return &PacketConnWrapper{
|
||||
PacketConn: packetConn,
|
||||
Dest: destAddr,
|
||||
}, nil
|
||||
@@ -148,6 +148,24 @@ func (d *DefaultSystemDialer) DestIpAddress() net.IP {
|
||||
return nil
|
||||
}
|
||||
|
||||
type PacketConnWrapper struct {
|
||||
net.PacketConn
|
||||
Dest net.Addr
|
||||
}
|
||||
|
||||
func (c *PacketConnWrapper) Read(p []byte) (int, error) {
|
||||
n, _, err := c.PacketConn.ReadFrom(p)
|
||||
return n, err
|
||||
}
|
||||
|
||||
func (c *PacketConnWrapper) Write(p []byte) (int, error) {
|
||||
return c.PacketConn.WriteTo(p, c.Dest)
|
||||
}
|
||||
|
||||
func (c *PacketConnWrapper) RemoteAddr() net.Addr {
|
||||
return c.Dest
|
||||
}
|
||||
|
||||
type SystemDialerAdapter interface {
|
||||
Dial(network string, address string) (net.Conn, error)
|
||||
}
|
||||
|
||||
@@ -6,7 +6,6 @@ import (
|
||||
"crypto/tls"
|
||||
"math/big"
|
||||
"slices"
|
||||
"sync/atomic"
|
||||
"time"
|
||||
|
||||
utls "github.com/refraction-networking/utls"
|
||||
@@ -30,19 +29,11 @@ var (
|
||||
|
||||
type Conn struct {
|
||||
*tls.Conn
|
||||
suppressCloseNotify atomic.Bool
|
||||
}
|
||||
|
||||
const tlsCloseTimeout = 250 * time.Millisecond
|
||||
|
||||
func (c *Conn) SuppressCloseNotify() {
|
||||
c.suppressCloseNotify.Store(true)
|
||||
}
|
||||
|
||||
func (c *Conn) Close() error {
|
||||
if c.suppressCloseNotify.Load() {
|
||||
return c.Conn.NetConn().Close()
|
||||
}
|
||||
timer := time.AfterFunc(tlsCloseTimeout, func() {
|
||||
c.Conn.NetConn().Close()
|
||||
})
|
||||
@@ -83,19 +74,11 @@ func Server(c net.Conn, config *tls.Config) net.Conn {
|
||||
|
||||
type UConn struct {
|
||||
*utls.UConn
|
||||
suppressCloseNotify atomic.Bool
|
||||
}
|
||||
|
||||
var _ Interface = (*UConn)(nil)
|
||||
|
||||
func (c *UConn) SuppressCloseNotify() {
|
||||
c.suppressCloseNotify.Store(true)
|
||||
}
|
||||
|
||||
func (c *UConn) Close() error {
|
||||
if c.suppressCloseNotify.Load() {
|
||||
return c.Conn.NetConn().Close()
|
||||
}
|
||||
timer := time.AfterFunc(tlsCloseTimeout, func() {
|
||||
c.Conn.NetConn().Close()
|
||||
})
|
||||
|
||||
Reference in New Issue
Block a user