Compare commits

..
10 Commits
Author SHA1 Message Date
patternihaandClaude Opus 5.5 747b153333 TUN inbound: Error out when autoSystemWfpBlockLeak or autoSystemDnsToGateway cannot apply
As asked in review, rather than run without them:

- The config is rejected, also by xray -test, for autoSystemWfpBlockLeak
  without autoSystemRoutingTable, or with "dns" but without dns, on
  Windows, and for autoSystemDnsToGateway without gateway on Linux.
- Xray does not start when the filters cannot be added, now on every
  Windows version, or when the system DNS cannot be set on Linux,
  instead of logging it and running without them.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
2026-09-30 08:46:45 +03:30
patternihaandClaude Opus 5.5 94cd83ccb4 TUN inbound: Rename "misconfig" to "misconfigtun"
The autoSystemWfpBlockLeak value that blocks an IP version not routed to
the TUN, as asked in review; "misconfig" was never released.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
2026-09-30 07:59:39 +03:30
patternihaandClaude Opus 5.5 1c225a8041 TUN inbound: Make autoSystemWfpBlockLeak a list of the leaks to block
autoSystemWfpBlockLeak now takes ["dns", "misconfig"] instead of true:
"dns" keeps DNS inside the TUN, and "misconfig" blocks an IP version
that no route leads to the TUN, the leak of a configuration that routes
only one of them. Either can be used alone, e.g. ["dns"] to block DNS
leaks while an IP version stays out of the TUN on purpose. Unknown
values are rejected. The config field becomes a repeated string with
the same number; it was never released.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
2026-09-30 07:46:51 +03:30
patternihaandClaude Opus 5.5 3a6bdb19ba TUN inbound: autoSystemDnsToGateway falls back to an IPv6 gateway on Linux
Without an IPv4 address in gateway, the system DNS now points at the
first IPv6 gateway plus one (e.g. fc00::1/64 -> fc00::2) instead of
nothing, and the routing check before the takeover accepts IPv6
addresses for it.

The README also says what each system does without gateway: Xray
assigns no address on Linux, Windows gives the TUN link-local ones
itself, and macOS and FreeBSD use 169.254.10.1/30.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
2026-09-30 07:18:55 +03:30
patternihaandClaude Opus 5.5 de02da553a TUN inbound: Block IPv4 too when it is not routed to the TUN on Windows
autoSystemWfpBlockLeak blocked IPv6 when the TUN had no IPv6 address or
no IPv6 route, and IPv4 never: with only IPv6 routed to the TUN, IPv4
went around it. Now each IP version is blocked when no route of it leads
to the TUN, except for loopback, DHCP, IPv6 neighbor and multicast
listener discovery, and Xray itself.

Addresses no longer count: without one of a version in gateway, Windows
gives the TUN a link-local one itself (fe80:: at once, 169.254.x.x after
some seconds), and what is routed to the TUN goes through it with that,
so a TUN with IPv6 routes but no IPv6 address had its IPv6 blocked for
nothing.

Tested on Windows 11, elevated: with only IPv6 routed, other programs'
IPv4 is denied, while loopback, a DHCP renew of Wi-Fi and Xray still
work; without gateway, IPv4 and IPv6 routed to the TUN enter it from
169.254.x.x and fe80::.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
2026-09-30 03:07:14 +03:30
patternihaandClaude Opus 5.5 4ec4fb8aab TUN inbound: Rename to autoSystemWfpBlockLeak and autoSystemDnsToGateway
autoSystemWFP becomes autoSystemWfpBlockLeak, saying that the WFP filters
block leaks, and autoSystemDNS becomes autoSystemDnsToGateway, saying
where it points the system DNS, so that pointing the system DNS at the
gateway on Windows later would fit the same name. Their config fields
keep their numbers; neither was released.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
2026-09-30 02:31:00 +03:30
patternihaandClaude Opus 5.5 63de6135cb TUN inbound: Rename strictRoute to autoSystemWFP
It only turns on the Windows Filtering Platform filters, along with Xray
resolving its own lookups while they restrict DNS, so it is named after
what it sets up in the system, like autoSystemRoutingTable and
autoSystemDNS. The config field becomes auto_system_wfp, with the same
number; strictRoute was never released.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
2026-09-30 01:22:22 +03:30
patternihaandClaude Opus 5.5 edd916b08e TUN inbound: Keep Windows' DNS Client from sending DoH/DoT outside the TUN
With strictRoute, only port 53 was kept inside the TUN. But Windows' DNS
Client service sends the queries for an interface's DNS servers out
through that interface, whatever the routes say, and since Windows 11
and Server 2022 it may send them over HTTPS or TLS, when that is set up
for the interface (as Windows Settings does) or for the server. Those
left through the physical link.

On those versions, the DNS Client service may now only connect through
the TUN, except for its mDNS and LLMNR. The filters recognize the
service by its SID in the token of its process, as Windows Firewall's
own rules for it do. Earlier versions only query port 53, and may run
the service in one process with others, so they get no such filters.
The port 53 rule stays, for the programs that query a resolver on the
local network themselves, and for those earlier versions.

Tested on Windows 11, elevated: with DoH set on Wi-Fi per adapter, per
network profile or by global auto-upgrade, none of the DNS Client's
connections left through Wi-Fi (WFP logged the drops by the new filter),
names still resolved through the TUN, mDNS and LLMNR still went out, and
other programs were unaffected, also in a real Xray run.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
2026-09-30 01:09:52 +03:30
patternihaandClaude Opus 5.5 eb29a4e3de TUN inbound: Make strictRoute false by default
Like sing-box's strict_route, strictRoute is now false by default, so the
Windows Filtering Platform filters are only added when it is set to true
(together with autoSystemRoutingTable). With unset meaning false,
strict_route becomes a plain bool field.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
2026-09-29 14:51:17 +03:30
patternihaandClaude Opus 5.5 2db099b34b TUN inbound: Block DNS and IPv6 leaks outside the TUN on Windows; Add strictRoute
Windows sends name queries to the DNS servers of all interfaces, and a
resolver on the local network (e.g. 192.168.1.1 from DHCP) is reached
through its more specific LAN route instead of the TUN, so DNS leaks
past it. IPv6 bypasses a TUN that cannot carry it.

With autoSystemRoutingTable set, the Windows TUN now adds Windows
Filtering Platform filters, all in one transaction and in a dynamic
session, so that they are removed when Xray exits, even if it crashes:
- DNS (port 53) only goes through the TUN, in both directions: its local
  address, and the interface it leaves or arrives by, must be the TUN's.
- IPv6 is blocked in both directions when the TUN has no IPv6 address or
  no IPv6 route, except loopback, neighbor and multicast listener
  discovery, and DHCPv6.
- Xray's own traffic is exempt: its connections out with a hard permit,
  which Windows Firewall rules do not override (like sing-box's
  strict_route), connections to its inbounds with an ordinary one.
If the filters cannot be added, the TUN does not start on Windows 10 and
later (only a warning on 7/8). The new `strictRoute` option (true by
default) turns them off.

Also on Windows:
- A warning for `dns` servers outside gateway and autoSystemRoutingTable,
  as queries to them cannot go through the TUN and are blocked.
- While DNS is restricted and autoOutboundsInterface is in use, Xray
  resolves the names it would ask Windows for itself (Go's resolver on
  its own sockets). Those lookups and the `localhost` DNS server skip the
  TUN's DNS servers, unless another interface uses them too, instead of
  looping back into the TUN.
- The DNS cache is flushed when the TUN starts and stops, and DNS
  registration is turned off on the TUN (through netsh before Windows 10
  1809).
- Close no longer panics when registering the route or interface change
  callbacks failed.
The README's Windows section describes all of it.

Tested on Windows 11, elevated, amd64 and 386: the filters, DNS arriving
through a real Wintun adapter and blocked outside it, the IPv6 block,
Windows Firewall rules, and a real Xray run. Windows 7/8 and Windows 10
before 1809 are untested.

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