mirror of
https://github.com/XTLS/Xray-core.git
synced 2026-10-01 21:45:44 +00:00
Compare commits
2
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
e0efc13c6e | ||
|
|
518a7efac2 |
@@ -37,7 +37,6 @@ RUN echo '{}' >/tmp/usr/local/etc/xray/08_fakedns.json
|
||||
RUN echo '{}' >/tmp/usr/local/etc/xray/09_metrics.json
|
||||
RUN echo '{}' >/tmp/usr/local/etc/xray/10_observatory.json
|
||||
RUN echo '{}' >/tmp/usr/local/etc/xray/11_geodata.json
|
||||
RUN echo '{}' >/tmp/usr/local/etc/xray/12_env.json
|
||||
RUN echo '{}' >/tmp/usr/local/etc/xray/99_version.json
|
||||
|
||||
# Create log files
|
||||
|
||||
@@ -37,7 +37,6 @@ RUN echo '{}' >/tmp/usr/local/etc/xray/08_fakedns.json
|
||||
RUN echo '{}' >/tmp/usr/local/etc/xray/09_metrics.json
|
||||
RUN echo '{}' >/tmp/usr/local/etc/xray/10_observatory.json
|
||||
RUN echo '{}' >/tmp/usr/local/etc/xray/11_geodata.json
|
||||
RUN echo '{}' >/tmp/usr/local/etc/xray/12_env.json
|
||||
RUN echo '{}' >/tmp/usr/local/etc/xray/99_version.json
|
||||
|
||||
# Create log files
|
||||
|
||||
@@ -64,14 +64,6 @@ jobs:
|
||||
echo "Latest: '$LATEST'."
|
||||
echo "LATEST=$LATEST" >>${GITHUB_ENV}
|
||||
|
||||
NEWEST=false
|
||||
if [[ "${{ github.event_name }}" == "release" ]]; then
|
||||
NEWEST=true
|
||||
fi
|
||||
|
||||
echo "Newest: '$NEWEST'."
|
||||
echo "NEWEST=$NEWEST" >>${GITHUB_ENV}
|
||||
|
||||
- name: Checkout code
|
||||
uses: actions/checkout@v7
|
||||
|
||||
@@ -82,7 +74,7 @@ jobs:
|
||||
uses: docker/setup-buildx-action@v4
|
||||
|
||||
- name: Login to GitHub Container Registry
|
||||
uses: docker/login-action@v4.6.0
|
||||
uses: docker/login-action@v4
|
||||
with:
|
||||
registry: ghcr.io
|
||||
username: ${{ github.repository_owner }}
|
||||
@@ -132,13 +124,6 @@ jobs:
|
||||
${{ env.FULL_IMAGE_NAME }}:${{ env.IMAGE_TAG }}
|
||||
fi
|
||||
|
||||
if [[ "${{ env.NEWEST }}" == "true" ]]; then
|
||||
echo "Adding 'pre-release' tag to manifest: '${{ env.FULL_IMAGE_NAME }}:pre-release'."
|
||||
docker buildx imagetools create \
|
||||
--tag ${{ env.FULL_IMAGE_NAME }}:pre-release \
|
||||
${{ env.FULL_IMAGE_NAME }}:${{ env.IMAGE_TAG }}
|
||||
fi
|
||||
|
||||
- name: Inspect image
|
||||
run: |
|
||||
docker buildx imagetools inspect ${{ env.FULL_IMAGE_NAME }}:${{ env.IMAGE_TAG }}
|
||||
@@ -146,7 +131,3 @@ jobs:
|
||||
if [[ "${{ env.LATEST }}" == "true" ]]; then
|
||||
docker buildx imagetools inspect ${{ env.FULL_IMAGE_NAME }}:latest
|
||||
fi
|
||||
|
||||
if [[ "${{ env.NEWEST }}" == "true" ]]; then
|
||||
docker buildx imagetools inspect ${{ env.FULL_IMAGE_NAME }}:pre-release
|
||||
fi
|
||||
|
||||
@@ -14,13 +14,13 @@ jobs:
|
||||
if: github.event_name != 'pull_request' || github.event.pull_request.head.repo.full_name != github.event.pull_request.base.repo.full_name
|
||||
steps:
|
||||
- name: Restore Geodat Cache
|
||||
uses: actions/cache/restore@v6
|
||||
uses: actions/cache/restore@v5
|
||||
with:
|
||||
path: resources
|
||||
key: xray-geodat-
|
||||
|
||||
- name: Restore Wintun Cache
|
||||
uses: actions/cache/restore@v6
|
||||
uses: actions/cache/restore@v5
|
||||
with:
|
||||
path: resources
|
||||
key: xray-wintun-
|
||||
@@ -92,7 +92,7 @@ jobs:
|
||||
echo "ASSET_NAME=$_NAME" >> $GITHUB_ENV
|
||||
|
||||
- name: Set up Go
|
||||
uses: actions/setup-go@v7
|
||||
uses: actions/setup-go@v6
|
||||
with:
|
||||
go-version-file: go.mod
|
||||
check-latest: true
|
||||
@@ -119,13 +119,13 @@ jobs:
|
||||
# go build -o build_assets/wxray.exe -trimpath -buildvcs=false -gcflags="all=-l=4" -ldflags="-H windowsgui -X github.com/xtls/xray-core/core.build=${COMMID} -s -w -buildid=" -v ./main
|
||||
|
||||
- name: Restore Geodat Cache
|
||||
uses: actions/cache/restore@v6
|
||||
uses: actions/cache/restore@v5
|
||||
with:
|
||||
path: resources
|
||||
key: xray-geodat-
|
||||
|
||||
- name: Restore Wintun Cache
|
||||
uses: actions/cache/restore@v6
|
||||
uses: actions/cache/restore@v5
|
||||
with:
|
||||
path: resources
|
||||
key: xray-wintun-
|
||||
|
||||
@@ -14,13 +14,13 @@ jobs:
|
||||
if: github.event_name != 'pull_request' || github.event.pull_request.head.repo.full_name != github.event.pull_request.base.repo.full_name
|
||||
steps:
|
||||
- name: Restore Geodat Cache
|
||||
uses: actions/cache/restore@v6
|
||||
uses: actions/cache/restore@v5
|
||||
with:
|
||||
path: resources
|
||||
key: xray-geodat-
|
||||
|
||||
- name: Restore Wintun Cache
|
||||
uses: actions/cache/restore@v6
|
||||
uses: actions/cache/restore@v5
|
||||
with:
|
||||
path: resources
|
||||
key: xray-wintun-
|
||||
@@ -193,7 +193,7 @@ jobs:
|
||||
echo "ASSET_NAME=$_NAME" >> $GITHUB_ENV
|
||||
|
||||
- name: Set up Go
|
||||
uses: actions/setup-go@v7
|
||||
uses: actions/setup-go@v6
|
||||
with:
|
||||
go-version-file: go.mod
|
||||
check-latest: true
|
||||
@@ -225,14 +225,14 @@ jobs:
|
||||
fi
|
||||
|
||||
- name: Restore Geodat Cache
|
||||
uses: actions/cache/restore@v6
|
||||
uses: actions/cache/restore@v5
|
||||
with:
|
||||
path: resources
|
||||
key: xray-geodat-
|
||||
|
||||
- name: Restore Wintun Cache
|
||||
if: matrix.goos == 'windows'
|
||||
uses: actions/cache/restore@v6
|
||||
uses: actions/cache/restore@v5
|
||||
with:
|
||||
path: resources
|
||||
key: xray-wintun-
|
||||
|
||||
@@ -26,7 +26,7 @@ jobs:
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- name: Restore Geodat Cache
|
||||
uses: actions/cache/restore@v6
|
||||
uses: actions/cache/restore@v5
|
||||
with:
|
||||
path: resources
|
||||
key: xray-geodat-
|
||||
@@ -59,7 +59,7 @@ jobs:
|
||||
done
|
||||
|
||||
- name: Save Geodat Cache
|
||||
uses: actions/cache/save@v6
|
||||
uses: actions/cache/save@v5
|
||||
if: ${{ steps.update.outputs.unhit }}
|
||||
with:
|
||||
path: resources
|
||||
@@ -73,7 +73,7 @@ jobs:
|
||||
ASSETHASH: 07c256185d6ee3652e09fa55c0b673e2624b565e02c4b9091c79ca7d2f24ef51
|
||||
steps:
|
||||
- name: Restore Wintun Cache
|
||||
uses: actions/cache/restore@v6
|
||||
uses: actions/cache/restore@v5
|
||||
with:
|
||||
path: resources
|
||||
key: xray-wintun-
|
||||
@@ -129,7 +129,7 @@ jobs:
|
||||
fi
|
||||
|
||||
- name: Save Wintun Cache
|
||||
uses: actions/cache/save@v6
|
||||
uses: actions/cache/save@v5
|
||||
if: ${{ steps.update.outputs.unhit }}
|
||||
with:
|
||||
path: resources
|
||||
|
||||
@@ -11,7 +11,7 @@ jobs:
|
||||
if: github.event_name != 'pull_request' || github.event.pull_request.head.repo.full_name != github.event.pull_request.base.repo.full_name
|
||||
steps:
|
||||
- name: Restore Geodat Cache
|
||||
uses: actions/cache/restore@v6
|
||||
uses: actions/cache/restore@v5
|
||||
with:
|
||||
path: resources
|
||||
key: xray-geodat-
|
||||
@@ -61,13 +61,15 @@ jobs:
|
||||
- name: Checkout codebase
|
||||
uses: actions/checkout@v7
|
||||
- name: Set up Go
|
||||
uses: actions/setup-go@v7
|
||||
uses: actions/setup-go@v6
|
||||
with:
|
||||
go-version-file: go.mod
|
||||
check-latest: true
|
||||
cache: false
|
||||
- name: Check Format
|
||||
run: go run ./infra/vformat/main.go -mode check -pwd ./
|
||||
run: |
|
||||
go install -v mvdan.cc/gofumpt@latest
|
||||
go run ./infra/vformat/main.go -mode check -pwd ./
|
||||
|
||||
test:
|
||||
needs: check-assets
|
||||
@@ -83,12 +85,12 @@ jobs:
|
||||
- name: Checkout codebase
|
||||
uses: actions/checkout@v7
|
||||
- name: Set up Go
|
||||
uses: actions/setup-go@v7
|
||||
uses: actions/setup-go@v6
|
||||
with:
|
||||
go-version-file: go.mod
|
||||
check-latest: true
|
||||
- name: Restore Geodat Cache
|
||||
uses: actions/cache/restore@v6
|
||||
uses: actions/cache/restore@v5
|
||||
with:
|
||||
path: resources
|
||||
key: xray-geodat-
|
||||
|
||||
@@ -73,7 +73,6 @@
|
||||
- [Xray_bash_onekey](https://github.com/hello-yunshu/Xray_bash_onekey), [XTool](https://github.com/LordPenguin666/XTool), [VPainLess](https://github.com/vpainless/vpainless)
|
||||
- [v2ray-agent](https://github.com/mack-a/v2ray-agent), [Xray_onekey](https://github.com/wulabing/Xray_onekey), [ProxySU](https://github.com/proxysu/ProxySU)
|
||||
- Magisk
|
||||
- [Magic_V2Ray](https://github.com/vincentng295/Magic_V2Ray)
|
||||
- [Xray_For_Magisk](https://github.com/E7KMbb/Xray_For_Magisk)
|
||||
- Homebrew
|
||||
- `brew install xray`
|
||||
|
||||
@@ -162,7 +162,7 @@ func (d *DefaultDispatcher) getLink(ctx context.Context) (*transport.Link, *tran
|
||||
p := d.policy.ForLevel(user.Level)
|
||||
if p.Stats.UserUplink {
|
||||
name := "user>>>" + user.Email + ">>>traffic>>>uplink"
|
||||
if c, _ := d.stats.GetOrRegisterCounter(name); c != nil {
|
||||
if c, _ := stats.GetOrRegisterCounter(d.stats, name); c != nil {
|
||||
inboundLink.Writer = &SizeStatWriter{
|
||||
Counter: c,
|
||||
Writer: inboundLink.Writer,
|
||||
@@ -171,7 +171,7 @@ func (d *DefaultDispatcher) getLink(ctx context.Context) (*transport.Link, *tran
|
||||
}
|
||||
if p.Stats.UserDownlink {
|
||||
name := "user>>>" + user.Email + ">>>traffic>>>downlink"
|
||||
if c, _ := d.stats.GetOrRegisterCounter(name); c != nil {
|
||||
if c, _ := stats.GetOrRegisterCounter(d.stats, name); c != nil {
|
||||
outboundLink.Writer = &SizeStatWriter{
|
||||
Counter: c,
|
||||
Writer: outboundLink.Writer,
|
||||
@@ -200,13 +200,13 @@ func WrapLink(ctx context.Context, policyManager policy.Manager, statsManager st
|
||||
p := policyManager.ForLevel(user.Level)
|
||||
if p.Stats.UserUplink {
|
||||
name := "user>>>" + user.Email + ">>>traffic>>>uplink"
|
||||
if c, _ := statsManager.GetOrRegisterCounter(name); c != nil {
|
||||
if c, _ := stats.GetOrRegisterCounter(statsManager, name); c != nil {
|
||||
link.Reader.(*buf.TimeoutWrapperReader).Counter = c
|
||||
}
|
||||
}
|
||||
if p.Stats.UserDownlink {
|
||||
name := "user>>>" + user.Email + ">>>traffic>>>downlink"
|
||||
if c, _ := statsManager.GetOrRegisterCounter(name); c != nil {
|
||||
if c, _ := stats.GetOrRegisterCounter(statsManager, name); c != nil {
|
||||
link.Writer = &SizeStatWriter{
|
||||
Counter: c,
|
||||
Writer: link.Writer,
|
||||
@@ -223,7 +223,7 @@ func WrapLink(ctx context.Context, policyManager policy.Manager, statsManager st
|
||||
|
||||
func trackOnlineIP(ctx context.Context, sm stats.Manager, email, ip string) {
|
||||
name := "user>>>" + email + ">>>online"
|
||||
if om, _ := sm.GetOrRegisterOnlineMap(name); om != nil {
|
||||
if om, _ := stats.GetOrRegisterOnlineMap(sm, name); om != nil {
|
||||
om.AddIP(ip)
|
||||
context.AfterFunc(ctx, func() { om.RemoveIP(ip) })
|
||||
}
|
||||
|
||||
@@ -23,7 +23,7 @@ func newFakeDNSSniffer(ctx context.Context) (protocolSnifferWithMetadata, error)
|
||||
}
|
||||
|
||||
if fakeDNSEngine == nil {
|
||||
errNotInit := errors.New("FakeDNSEngine is not initialized, but such a sniffer is used")
|
||||
errNotInit := errors.New("FakeDNSEngine is not initialized, but such a sniffer is used").AtError()
|
||||
return protocolSnifferWithMetadata{}, errNotInit
|
||||
}
|
||||
return protocolSnifferWithMetadata{protocolSniffer: func(ctx context.Context, bytes []byte) (SniffResult, error) {
|
||||
|
||||
+1
-1
@@ -28,7 +28,7 @@ func toNetIP(addrs []net.Address) ([]net.IP, error) {
|
||||
if addr.Family().IsIP() {
|
||||
ips = append(ips, addr.IP())
|
||||
} else {
|
||||
return nil, errors.New("Failed to convert address", addr, "to Net IP.")
|
||||
return nil, errors.New("Failed to convert address", addr, "to Net IP.").AtWarning()
|
||||
}
|
||||
}
|
||||
return ips, nil
|
||||
|
||||
@@ -212,28 +212,6 @@ func (s *DNS) IsOwnLink(ctx context.Context) bool {
|
||||
return false
|
||||
}
|
||||
|
||||
// MayUseSystemResolver reports whether any name server configured here could
|
||||
// still resolve through the system resolver. That is what happens when no name
|
||||
// server is configured at all, and it is also what a name server pointed at
|
||||
// "localhost" does. Callers that are about to redirect the system resolver need
|
||||
// to know, because a resolution path that reaches it would then loop back to
|
||||
// them.
|
||||
//
|
||||
// Any such server is enough: name servers can be selected per domain, so a
|
||||
// single local one makes some query reach the system resolver even when
|
||||
// independent upstreams are configured alongside it.
|
||||
func (s *DNS) MayUseSystemResolver() bool {
|
||||
if len(s.clients) == 0 {
|
||||
return true
|
||||
}
|
||||
for _, client := range s.clients {
|
||||
if _, isLocal := client.server.(*LocalNameServer); isLocal {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// LookupIP implements dns.Client.
|
||||
func (s *DNS) LookupIP(domain string, option dns.IPOption) ([]net.IP, uint32, error) {
|
||||
// Normalize the FQDN form query
|
||||
|
||||
@@ -1,59 +0,0 @@
|
||||
package dns
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
|
||||
"github.com/xtls/xray-core/common/net"
|
||||
feature_dns "github.com/xtls/xray-core/features/dns"
|
||||
)
|
||||
|
||||
// fakeServer stands in for any name server that is not the system resolver.
|
||||
type fakeServer struct{}
|
||||
|
||||
func (fakeServer) Name() string { return "fake" }
|
||||
func (fakeServer) IsDisableCache() bool { return false }
|
||||
func (fakeServer) QueryIP(context.Context, string, feature_dns.IPOption) ([]net.IP, uint32, error) {
|
||||
return nil, 0, nil
|
||||
}
|
||||
|
||||
// Callers that are about to redirect the system resolver rely on this to tell
|
||||
// whether any resolution path could still reach the system resolver, so the
|
||||
// mixed shape has to be reported as reachable: a domain-specific rule can
|
||||
// select the system resolver even when an independent upstream also exists.
|
||||
func TestMayUseSystemResolver(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
clients []*Client
|
||||
want bool
|
||||
}{
|
||||
{
|
||||
name: "no clients at all",
|
||||
want: true,
|
||||
},
|
||||
{
|
||||
name: "only the system resolver",
|
||||
clients: []*Client{{server: NewLocalNameServer()}},
|
||||
want: true,
|
||||
},
|
||||
{
|
||||
name: "the system resolver alongside an independent name server",
|
||||
clients: []*Client{{server: fakeServer{}}, {server: NewLocalNameServer()}},
|
||||
want: true,
|
||||
},
|
||||
{
|
||||
name: "only independent name servers",
|
||||
clients: []*Client{{server: fakeServer{}}},
|
||||
want: false,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
server := &DNS{clients: tt.clients}
|
||||
if got := server.MayUseSystemResolver(); got != tt.want {
|
||||
t.Errorf("MayUseSystemResolver() = %v, want %v", got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -188,10 +188,10 @@ func parseResponse(payload []byte) (*IPRecord, error) {
|
||||
var parser dnsmessage.Parser
|
||||
h, err := parser.Start(payload)
|
||||
if err != nil {
|
||||
return nil, errors.New("failed to parse DNS response").Base(err)
|
||||
return nil, errors.New("failed to parse DNS response").Base(err).AtWarning()
|
||||
}
|
||||
if err := parser.SkipAllQuestions(); err != nil {
|
||||
return nil, errors.New("failed to skip questions in DNS response").Base(err)
|
||||
return nil, errors.New("failed to skip questions in DNS response").Base(err).AtWarning()
|
||||
}
|
||||
|
||||
now := time.Now()
|
||||
|
||||
@@ -58,7 +58,7 @@ func NewFakeDNSHolder() (*Holder, error) {
|
||||
var err error
|
||||
|
||||
if fkdns, err = NewFakeDNSHolderConfigOnly(nil); err != nil {
|
||||
return nil, errors.New("Unable to create Fake Dns Engine").Base(err)
|
||||
return nil, errors.New("Unable to create Fake Dns Engine").Base(err).AtError()
|
||||
}
|
||||
err = fkdns.initialize(dns.FakeIPv4Pool, 65535)
|
||||
if err != nil {
|
||||
@@ -80,13 +80,13 @@ func (fkdns *Holder) initialize(ipPoolCidr string, lruSize int) error {
|
||||
var err error
|
||||
|
||||
if _, ipRange, err = net.ParseCIDR(ipPoolCidr); err != nil {
|
||||
return errors.New("Unable to parse CIDR for Fake DNS IP assignment").Base(err)
|
||||
return errors.New("Unable to parse CIDR for Fake DNS IP assignment").Base(err).AtError()
|
||||
}
|
||||
|
||||
ones, bits := ipRange.Mask.Size()
|
||||
rooms := bits - ones
|
||||
if math.Log2(float64(lruSize)) >= float64(rooms) {
|
||||
return errors.New("LRU size is bigger than subnet size")
|
||||
return errors.New("LRU size is bigger than subnet size").AtError()
|
||||
}
|
||||
fkdns.domainToIP = cache.NewLru(lruSize)
|
||||
fkdns.ipRange = ipRange
|
||||
|
||||
@@ -84,7 +84,7 @@ func NewServer(ctx context.Context, dest net.Destination, dispatcher routing.Dis
|
||||
if dest.Network == net.Network_UDP { // UDP classic DNS mode
|
||||
return NewClassicNameServer(dest, dispatcher, disableCache, serveStale, serveExpiredTTL, clientIP), nil
|
||||
}
|
||||
return nil, errors.New("No available name server could be created from ", dest)
|
||||
return nil, errors.New("No available name server could be created from ", dest).AtWarning()
|
||||
}
|
||||
|
||||
// NewClient creates a DNS client managing a name server with client IP, domain rules and expected IPs.
|
||||
@@ -102,7 +102,7 @@ func NewClient(
|
||||
// Create a new server for each client for now
|
||||
server, err := NewServer(ctx, ns.Address.AsDestination(), dispatcher, disableCache, serveStale, serveExpiredTTL, clientIP)
|
||||
if err != nil {
|
||||
return errors.New("failed to create nameserver").Base(err)
|
||||
return errors.New("failed to create nameserver").Base(err).AtWarning()
|
||||
}
|
||||
|
||||
_, isLocalDNS := server.(*LocalNameServer)
|
||||
@@ -113,7 +113,7 @@ func NewClient(
|
||||
if len(ns.ExpectedIp) > 0 {
|
||||
expectedMatcher, err = geodata.IPReg.BuildIPMatcher(ns.ExpectedIp)
|
||||
if err != nil {
|
||||
return errors.New("failed to create expected ip matcher").Base(err)
|
||||
return errors.New("failed to create expected ip matcher").Base(err).AtWarning()
|
||||
}
|
||||
}
|
||||
|
||||
@@ -122,7 +122,7 @@ func NewClient(
|
||||
if len(ns.UnexpectedIp) > 0 {
|
||||
unexpectedMatcher, err = geodata.IPReg.BuildIPMatcher(ns.UnexpectedIp)
|
||||
if err != nil {
|
||||
return errors.New("failed to create unexpected ip matcher").Base(err)
|
||||
return errors.New("failed to create unexpected ip matcher").Base(err).AtWarning()
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -27,7 +27,7 @@ func (s *FakeDNSServer) IsDisableCache() bool {
|
||||
|
||||
func (f *FakeDNSServer) QueryIP(ctx context.Context, domain string, opt dns.IPOption) ([]net.IP, uint32, error) {
|
||||
if f.fakeDNSEngine == nil {
|
||||
return nil, 0, errors.New("Unable to locate a fake DNS Engine")
|
||||
return nil, 0, errors.New("Unable to locate a fake DNS Engine").AtError()
|
||||
}
|
||||
|
||||
var ips []net.Address
|
||||
@@ -39,7 +39,7 @@ func (f *FakeDNSServer) QueryIP(ctx context.Context, domain string, opt dns.IPOp
|
||||
|
||||
netIP, err := toNetIP(ips)
|
||||
if err != nil {
|
||||
return nil, 0, errors.New("Unable to convert IP to net ip").Base(err)
|
||||
return nil, 0, errors.New("Unable to convert IP to net ip").Base(err).AtError()
|
||||
}
|
||||
|
||||
errors.LogInfo(ctx, f.Name(), " got answer: ", domain, " -> ", ips)
|
||||
|
||||
+2
-6
@@ -89,10 +89,10 @@ func (g *Instance) startInternal() error {
|
||||
g.active = true
|
||||
|
||||
if err := g.initAccessLogger(); err != nil {
|
||||
return errors.New("failed to initialize access logger").Base(err)
|
||||
return errors.New("failed to initialize access logger").Base(err).AtWarning()
|
||||
}
|
||||
if err := g.initErrorLogger(); err != nil {
|
||||
return errors.New("failed to initialize error logger").Base(err)
|
||||
return errors.New("failed to initialize error logger").Base(err).AtWarning()
|
||||
}
|
||||
|
||||
return nil
|
||||
@@ -141,10 +141,6 @@ func (g *Instance) Handle(msg log.Message) {
|
||||
}
|
||||
}
|
||||
|
||||
func (g *Instance) Severity() log.Severity {
|
||||
return g.config.ErrorLogLevel
|
||||
}
|
||||
|
||||
// Close implements common.Closable.Close().
|
||||
func (g *Instance) Close() error {
|
||||
errors.LogDebug(context.Background(), "Logger closing")
|
||||
|
||||
+59
-151
@@ -2,18 +2,15 @@ package metrics
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
stderrors "errors"
|
||||
"expvar"
|
||||
stdnet "net"
|
||||
"net/http"
|
||||
"net/http/pprof"
|
||||
_ "net/http/pprof"
|
||||
"strings"
|
||||
|
||||
"github.com/xtls/xray-core/app/observatory"
|
||||
"github.com/xtls/xray-core/common"
|
||||
"github.com/xtls/xray-core/common/errors"
|
||||
xnet "github.com/xtls/xray-core/common/net"
|
||||
"github.com/xtls/xray-core/common/net"
|
||||
"github.com/xtls/xray-core/common/signal/done"
|
||||
"github.com/xtls/xray-core/core"
|
||||
"github.com/xtls/xray-core/features/extension"
|
||||
@@ -24,17 +21,15 @@ import (
|
||||
type MetricsHandler struct {
|
||||
ohm outbound.Manager
|
||||
statsManager feature_stats.Manager
|
||||
ctx context.Context
|
||||
observatory extension.Observatory
|
||||
tag string
|
||||
listen string
|
||||
tcpListener xnet.Listener
|
||||
listener *OutboundListener
|
||||
tcpListener net.Listener
|
||||
}
|
||||
|
||||
// NewMetricsHandler creates a new MetricsHandler based on the given config.
|
||||
func NewMetricsHandler(ctx context.Context, config *Config) (*MetricsHandler, error) {
|
||||
c := &MetricsHandler{
|
||||
ctx: ctx,
|
||||
tag: config.Tag,
|
||||
listen: config.Listen,
|
||||
}
|
||||
@@ -42,6 +37,46 @@ func NewMetricsHandler(ctx context.Context, config *Config) (*MetricsHandler, er
|
||||
c.statsManager = sm
|
||||
c.ohm = om
|
||||
}))
|
||||
expvar.Publish("stats", expvar.Func(func() interface{} {
|
||||
resp := map[string]map[string]map[string]int64{
|
||||
"inbound": {},
|
||||
"outbound": {},
|
||||
"user": {},
|
||||
}
|
||||
c.statsManager.VisitCounters(func(name string, counter feature_stats.Counter) bool {
|
||||
nameSplit := strings.Split(name, ">>>")
|
||||
typeName, tagOrUser, direction := nameSplit[0], nameSplit[1], nameSplit[3]
|
||||
if item, found := resp[typeName][tagOrUser]; found {
|
||||
item[direction] = counter.Value()
|
||||
} else {
|
||||
resp[typeName][tagOrUser] = map[string]int64{
|
||||
direction: counter.Value(),
|
||||
}
|
||||
}
|
||||
return true
|
||||
})
|
||||
return resp
|
||||
}))
|
||||
expvar.Publish("observatory", expvar.Func(func() interface{} {
|
||||
if c.observatory == nil {
|
||||
common.Must(core.RequireFeatures(ctx, func(observatory extension.Observatory) error {
|
||||
c.observatory = observatory
|
||||
return nil
|
||||
}))
|
||||
if c.observatory == nil {
|
||||
return nil
|
||||
}
|
||||
}
|
||||
resp := map[string]*observatory.OutboundStatus{}
|
||||
if o, err := c.observatory.GetObservation(context.Background()); err != nil {
|
||||
return err
|
||||
} else {
|
||||
for _, x := range o.(*observatory.ObservationResult).GetStatus() {
|
||||
resp[x.OutboundTag] = x
|
||||
}
|
||||
}
|
||||
return resp
|
||||
}))
|
||||
return c, nil
|
||||
}
|
||||
|
||||
@@ -50,172 +85,45 @@ func (p *MetricsHandler) Type() interface{} {
|
||||
}
|
||||
|
||||
func (p *MetricsHandler) Start() error {
|
||||
handler := p.httpHandler()
|
||||
|
||||
// direct listen a port if listen is set
|
||||
if p.listen != "" {
|
||||
TCPlistener, err := xnet.Listen("tcp", p.listen)
|
||||
TCPlistener, err := net.Listen("tcp", p.listen)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
p.tcpListener = TCPlistener
|
||||
errors.LogInfo(context.Background(), "Metrics server listening on ", p.listen)
|
||||
|
||||
go p.serve(TCPlistener, handler)
|
||||
}
|
||||
|
||||
if p.tag == "" {
|
||||
if p.tcpListener == nil {
|
||||
return errors.New("metrics must have a tag or listen address")
|
||||
}
|
||||
return nil
|
||||
go func() {
|
||||
if err := http.Serve(TCPlistener, http.DefaultServeMux); err != nil {
|
||||
errors.LogErrorInner(context.Background(), err, "failed to start metrics server")
|
||||
}
|
||||
}()
|
||||
}
|
||||
|
||||
listener := &OutboundListener{
|
||||
buffer: make(chan xnet.Conn, 4),
|
||||
buffer: make(chan net.Conn, 4),
|
||||
done: done.New(),
|
||||
}
|
||||
p.listener = listener
|
||||
|
||||
go p.serve(listener, handler)
|
||||
go func() {
|
||||
if err := http.Serve(listener, http.DefaultServeMux); err != nil {
|
||||
errors.LogErrorInner(context.Background(), err, "failed to start metrics server")
|
||||
}
|
||||
}()
|
||||
|
||||
if err := p.ohm.RemoveHandler(context.Background(), p.tag); err != nil {
|
||||
errors.LogInfo(context.Background(), "failed to remove existing handler")
|
||||
}
|
||||
|
||||
if err := p.ohm.AddHandler(context.Background(), &Outbound{
|
||||
return p.ohm.AddHandler(context.Background(), &Outbound{
|
||||
tag: p.tag,
|
||||
listener: listener,
|
||||
}); err != nil {
|
||||
if closeErr := p.Close(); closeErr != nil {
|
||||
errors.LogErrorInner(context.Background(), closeErr, "failed to close metrics server after start failure")
|
||||
}
|
||||
return err
|
||||
}
|
||||
|
||||
return nil
|
||||
})
|
||||
}
|
||||
|
||||
func (p *MetricsHandler) Close() error {
|
||||
var errs []error
|
||||
if p.tcpListener != nil {
|
||||
errs = append(errs, p.tcpListener.Close())
|
||||
p.tcpListener = nil
|
||||
}
|
||||
if p.listener != nil {
|
||||
errs = append(errs, p.listener.Close())
|
||||
p.listener = nil
|
||||
}
|
||||
if p.ohm != nil && p.tag != "" {
|
||||
if err := p.ohm.RemoveHandler(context.Background(), p.tag); err != nil {
|
||||
errors.LogInfo(context.Background(), "failed to remove metrics handler")
|
||||
}
|
||||
}
|
||||
return errors.Combine(errs...)
|
||||
}
|
||||
|
||||
func (p *MetricsHandler) serve(listener xnet.Listener, handler http.Handler) {
|
||||
if err := http.Serve(listener, handler); err != nil && !isClosedListenerError(err) {
|
||||
errors.LogErrorInner(context.Background(), err, "failed to start metrics server")
|
||||
}
|
||||
}
|
||||
|
||||
func isClosedListenerError(err error) bool {
|
||||
if err == nil {
|
||||
return true
|
||||
}
|
||||
if stderrors.Is(err, stdnet.ErrClosed) || stderrors.Is(err, http.ErrServerClosed) {
|
||||
return true
|
||||
}
|
||||
errText := err.Error()
|
||||
return strings.Contains(errText, "listen closed") ||
|
||||
strings.Contains(errText, "use of closed network connection")
|
||||
}
|
||||
|
||||
func (p *MetricsHandler) httpHandler() http.Handler {
|
||||
mux := http.NewServeMux()
|
||||
mux.HandleFunc("/debug/vars", p.handleDebugVars)
|
||||
mux.HandleFunc("/debug/pprof/", pprof.Index)
|
||||
mux.HandleFunc("/debug/pprof/cmdline", pprof.Cmdline)
|
||||
mux.HandleFunc("/debug/pprof/profile", pprof.Profile)
|
||||
mux.HandleFunc("/debug/pprof/symbol", pprof.Symbol)
|
||||
mux.HandleFunc("/debug/pprof/trace", pprof.Trace)
|
||||
return mux
|
||||
}
|
||||
|
||||
func (p *MetricsHandler) handleDebugVars(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Type", "application/json; charset=utf-8")
|
||||
vars := map[string]json.RawMessage{}
|
||||
expvar.Do(func(kv expvar.KeyValue) {
|
||||
value := json.RawMessage(kv.Value.String())
|
||||
if !json.Valid(value) {
|
||||
value = json.RawMessage("null")
|
||||
}
|
||||
vars[kv.Key] = value
|
||||
})
|
||||
vars["stats"] = marshalJSON(p.stats())
|
||||
vars["observatory"] = marshalJSON(p.observatoryStatus())
|
||||
|
||||
payload, err := json.Marshal(vars)
|
||||
if err != nil {
|
||||
http.Error(w, err.Error(), http.StatusInternalServerError)
|
||||
return
|
||||
}
|
||||
w.Write(payload)
|
||||
}
|
||||
|
||||
func marshalJSON(value interface{}) json.RawMessage {
|
||||
data, err := json.Marshal(value)
|
||||
if err != nil {
|
||||
return json.RawMessage("null")
|
||||
}
|
||||
return data
|
||||
}
|
||||
|
||||
func (p *MetricsHandler) stats() map[string]map[string]map[string]int64 {
|
||||
resp := map[string]map[string]map[string]int64{
|
||||
"inbound": {},
|
||||
"outbound": {},
|
||||
"user": {},
|
||||
}
|
||||
p.statsManager.VisitCounters(func(name string, counter feature_stats.Counter) bool {
|
||||
nameSplit := strings.Split(name, ">>>")
|
||||
if len(nameSplit) < 4 {
|
||||
return true
|
||||
}
|
||||
typeName, tagOrUser, direction := nameSplit[0], nameSplit[1], nameSplit[3]
|
||||
items, found := resp[typeName]
|
||||
if !found {
|
||||
items = map[string]map[string]int64{}
|
||||
resp[typeName] = items
|
||||
}
|
||||
if item, found := items[tagOrUser]; found {
|
||||
item[direction] = counter.Value()
|
||||
} else {
|
||||
items[tagOrUser] = map[string]int64{
|
||||
direction: counter.Value(),
|
||||
}
|
||||
}
|
||||
return true
|
||||
})
|
||||
return resp
|
||||
}
|
||||
|
||||
func (p *MetricsHandler) observatoryStatus() interface{} {
|
||||
feature := core.MustFromContext(p.ctx).GetFeature(extension.ObservatoryType())
|
||||
if feature == nil {
|
||||
return nil
|
||||
}
|
||||
observatoryFeature := feature.(extension.Observatory)
|
||||
resp := map[string]*observatory.OutboundStatus{}
|
||||
if o, err := observatoryFeature.GetObservation(context.Background()); err != nil {
|
||||
return err
|
||||
} else {
|
||||
for _, x := range o.(*observatory.ObservationResult).GetStatus() {
|
||||
resp[x.OutboundTag] = x
|
||||
}
|
||||
}
|
||||
return resp
|
||||
return nil
|
||||
}
|
||||
|
||||
func init() {
|
||||
|
||||
@@ -1,161 +0,0 @@
|
||||
package metrics
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
stdnet "net"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
|
||||
"github.com/xtls/xray-core/app/dispatcher"
|
||||
"github.com/xtls/xray-core/app/proxyman"
|
||||
_ "github.com/xtls/xray-core/app/proxyman/inbound"
|
||||
_ "github.com/xtls/xray-core/app/proxyman/outbound"
|
||||
appstats "github.com/xtls/xray-core/app/stats"
|
||||
"github.com/xtls/xray-core/common/serial"
|
||||
"github.com/xtls/xray-core/core"
|
||||
feature_outbound "github.com/xtls/xray-core/features/outbound"
|
||||
)
|
||||
|
||||
func TestMetricsCanRestartInSameProcess(t *testing.T) {
|
||||
for i := 0; i < 2; i++ {
|
||||
server := startMetricsTestServer(t)
|
||||
readMetricsVars(t, server)
|
||||
readMetricsPprof(t, server)
|
||||
if err := server.Close(); err != nil {
|
||||
t.Fatalf("failed to close metrics server: %v", err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestMetricsCanRunMultipleInstancesInSameProcess(t *testing.T) {
|
||||
server1 := startMetricsTestServer(t)
|
||||
t.Cleanup(func() {
|
||||
_ = server1.Close()
|
||||
})
|
||||
server2 := startMetricsTestServer(t)
|
||||
t.Cleanup(func() {
|
||||
_ = server2.Close()
|
||||
})
|
||||
|
||||
readMetricsVars(t, server1)
|
||||
readMetricsVars(t, server2)
|
||||
}
|
||||
|
||||
func TestMetricsListenOnlyWithoutTagDoesNotRegisterOutbound(t *testing.T) {
|
||||
listen := pickMetricsListenAddress(t)
|
||||
server := startMetricsTestServerWithMetricsConfig(t, &Config{
|
||||
Listen: listen,
|
||||
})
|
||||
t.Cleanup(func() {
|
||||
_ = server.Close()
|
||||
})
|
||||
|
||||
response, err := http.Get("http://" + listen + "/debug/vars")
|
||||
if err != nil {
|
||||
t.Fatalf("failed to read listen-only metrics: %v", err)
|
||||
}
|
||||
defer response.Body.Close()
|
||||
if response.StatusCode != http.StatusOK {
|
||||
t.Fatalf("unexpected listen-only metrics status: %d", response.StatusCode)
|
||||
}
|
||||
|
||||
outboundManager := server.GetFeature(feature_outbound.ManagerType()).(feature_outbound.Manager)
|
||||
if handlers := outboundManager.ListHandlers(context.Background()); len(handlers) != 0 {
|
||||
t.Fatalf("listen-only metrics registered outbound handlers: got %d, want 0", len(handlers))
|
||||
}
|
||||
}
|
||||
|
||||
func startMetricsTestServer(t *testing.T) *core.Instance {
|
||||
return startMetricsTestServerWithMetricsConfig(t, &Config{
|
||||
Tag: "metrics_out",
|
||||
})
|
||||
}
|
||||
|
||||
func startMetricsTestServerWithMetricsConfig(t *testing.T, metricsConfig *Config) *core.Instance {
|
||||
t.Helper()
|
||||
|
||||
server, err := core.New(metricsTestConfig(metricsConfig))
|
||||
if err != nil {
|
||||
t.Fatalf("failed to create metrics server: %v", err)
|
||||
}
|
||||
if err := server.Start(); err != nil {
|
||||
_ = server.Close()
|
||||
t.Fatalf("failed to start metrics server: %v", err)
|
||||
}
|
||||
return server
|
||||
}
|
||||
|
||||
func metricsTestConfig(metricsConfig *Config) *core.Config {
|
||||
return &core.Config{
|
||||
App: []*serial.TypedMessage{
|
||||
serial.ToTypedMessage(&dispatcher.Config{}),
|
||||
serial.ToTypedMessage(&proxyman.InboundConfig{}),
|
||||
serial.ToTypedMessage(&proxyman.OutboundConfig{}),
|
||||
serial.ToTypedMessage(&appstats.Config{}),
|
||||
serial.ToTypedMessage(metricsConfig),
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func pickMetricsListenAddress(t *testing.T) string {
|
||||
t.Helper()
|
||||
|
||||
listener, err := stdnet.Listen("tcp", "127.0.0.1:0")
|
||||
if err != nil {
|
||||
t.Fatalf("failed to pick metrics listen address: %v", err)
|
||||
}
|
||||
defer listener.Close()
|
||||
return listener.Addr().String()
|
||||
}
|
||||
|
||||
func readMetricsVars(t *testing.T, server *core.Instance) {
|
||||
t.Helper()
|
||||
|
||||
recorder := httptest.NewRecorder()
|
||||
metricsHandler(t, server).httpHandler().ServeHTTP(
|
||||
recorder,
|
||||
httptest.NewRequest(http.MethodGet, "/debug/vars", nil),
|
||||
)
|
||||
|
||||
if recorder.Code != http.StatusOK {
|
||||
t.Fatalf("unexpected metrics vars status: %d", recorder.Code)
|
||||
}
|
||||
|
||||
var payload map[string]interface{}
|
||||
if err := json.NewDecoder(recorder.Body).Decode(&payload); err != nil {
|
||||
t.Fatalf("failed to decode metrics vars: %v", err)
|
||||
}
|
||||
if _, found := payload["stats"]; !found {
|
||||
t.Fatal("metrics vars missing stats")
|
||||
}
|
||||
if _, found := payload["observatory"]; !found {
|
||||
t.Fatal("metrics vars missing observatory")
|
||||
}
|
||||
}
|
||||
|
||||
func readMetricsPprof(t *testing.T, server *core.Instance) {
|
||||
t.Helper()
|
||||
|
||||
recorder := httptest.NewRecorder()
|
||||
metricsHandler(t, server).httpHandler().ServeHTTP(
|
||||
recorder,
|
||||
httptest.NewRequest(http.MethodGet, "/debug/pprof/goroutine?debug=1", nil),
|
||||
)
|
||||
|
||||
if recorder.Code != http.StatusOK {
|
||||
t.Fatalf("unexpected metrics pprof status: %d", recorder.Code)
|
||||
}
|
||||
}
|
||||
|
||||
func metricsHandler(t *testing.T, server *core.Instance) *MetricsHandler {
|
||||
t.Helper()
|
||||
|
||||
feature := server.GetFeature((*MetricsHandler)(nil))
|
||||
handler, ok := feature.(*MetricsHandler)
|
||||
if !ok || handler == nil {
|
||||
t.Fatal("metrics handler not registered")
|
||||
}
|
||||
return handler
|
||||
}
|
||||
@@ -78,12 +78,6 @@ func (o *Observer) background() {
|
||||
sleepTime = time.Duration(o.config.ProbeInterval)
|
||||
}
|
||||
|
||||
if len(outbounds) == 0 {
|
||||
errors.LogWarning(o.ctx, "no outbound matches subjectSelector ", o.config.SubjectSelector)
|
||||
time.Sleep(sleepTime)
|
||||
continue
|
||||
}
|
||||
|
||||
if !o.config.EnableConcurrency {
|
||||
sort.Strings(outbounds)
|
||||
for _, v := range outbounds {
|
||||
|
||||
+22
-11
@@ -330,6 +330,7 @@ type SenderConfig struct {
|
||||
// Send traffic through the given IP. Only IP is allowed.
|
||||
Via *net.IPOrDomain `protobuf:"bytes,1,opt,name=via,proto3" json:"via,omitempty"`
|
||||
StreamSettings *internet.StreamConfig `protobuf:"bytes,2,opt,name=stream_settings,json=streamSettings,proto3" json:"stream_settings,omitempty"`
|
||||
ProxySettings *internet.ProxyConfig `protobuf:"bytes,3,opt,name=proxy_settings,json=proxySettings,proto3" json:"proxy_settings,omitempty"`
|
||||
MultiplexSettings *MultiplexingConfig `protobuf:"bytes,4,opt,name=multiplex_settings,json=multiplexSettings,proto3" json:"multiplex_settings,omitempty"`
|
||||
ViaCidr string `protobuf:"bytes,5,opt,name=via_cidr,json=viaCidr,proto3" json:"via_cidr,omitempty"`
|
||||
TargetStrategy internet.DomainStrategy `protobuf:"varint,6,opt,name=target_strategy,json=targetStrategy,proto3,enum=xray.transport.internet.DomainStrategy" json:"target_strategy,omitempty"`
|
||||
@@ -381,6 +382,13 @@ func (x *SenderConfig) GetStreamSettings() *internet.StreamConfig {
|
||||
return nil
|
||||
}
|
||||
|
||||
func (x *SenderConfig) GetProxySettings() *internet.ProxyConfig {
|
||||
if x != nil {
|
||||
return x.ProxySettings
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (x *SenderConfig) GetMultiplexSettings() *MultiplexingConfig {
|
||||
if x != nil {
|
||||
return x.MultiplexSettings
|
||||
@@ -498,13 +506,14 @@ const file_app_proxyman_config_proto_rawDesc = "" +
|
||||
"\x03tag\x18\x01 \x01(\tR\x03tag\x12M\n" +
|
||||
"\x11receiver_settings\x18\x02 \x01(\v2 .xray.common.serial.TypedMessageR\x10receiverSettings\x12G\n" +
|
||||
"\x0eproxy_settings\x18\x03 \x01(\v2 .xray.common.serial.TypedMessageR\rproxySettings\"\x10\n" +
|
||||
"\x0eOutboundConfig\"\xd6\x02\n" +
|
||||
"\x0eOutboundConfig\"\x9d\x03\n" +
|
||||
"\fSenderConfig\x12-\n" +
|
||||
"\x03via\x18\x01 \x01(\v2\x1b.xray.common.net.IPOrDomainR\x03via\x12N\n" +
|
||||
"\x0fstream_settings\x18\x02 \x01(\v2%.xray.transport.internet.StreamConfigR\x0estreamSettings\x12T\n" +
|
||||
"\x0fstream_settings\x18\x02 \x01(\v2%.xray.transport.internet.StreamConfigR\x0estreamSettings\x12K\n" +
|
||||
"\x0eproxy_settings\x18\x03 \x01(\v2$.xray.transport.internet.ProxyConfigR\rproxySettings\x12T\n" +
|
||||
"\x12multiplex_settings\x18\x04 \x01(\v2%.xray.app.proxyman.MultiplexingConfigR\x11multiplexSettings\x12\x19\n" +
|
||||
"\bvia_cidr\x18\x05 \x01(\tR\aviaCidr\x12P\n" +
|
||||
"\x0ftarget_strategy\x18\x06 \x01(\x0e2'.xray.transport.internet.DomainStrategyR\x0etargetStrategyJ\x04\b\x03\x10\x04\"\xa4\x01\n" +
|
||||
"\x0ftarget_strategy\x18\x06 \x01(\x0e2'.xray.transport.internet.DomainStrategyR\x0etargetStrategy\"\xa4\x01\n" +
|
||||
"\x12MultiplexingConfig\x12\x18\n" +
|
||||
"\aenabled\x18\x01 \x01(\bR\aenabled\x12 \n" +
|
||||
"\vconcurrency\x18\x02 \x01(\x05R\vconcurrency\x12(\n" +
|
||||
@@ -539,7 +548,8 @@ var file_app_proxyman_config_proto_goTypes = []any{
|
||||
(*net.IPOrDomain)(nil), // 10: xray.common.net.IPOrDomain
|
||||
(*internet.StreamConfig)(nil), // 11: xray.transport.internet.StreamConfig
|
||||
(*serial.TypedMessage)(nil), // 12: xray.common.serial.TypedMessage
|
||||
(internet.DomainStrategy)(0), // 13: xray.transport.internet.DomainStrategy
|
||||
(*internet.ProxyConfig)(nil), // 13: xray.transport.internet.ProxyConfig
|
||||
(internet.DomainStrategy)(0), // 14: xray.transport.internet.DomainStrategy
|
||||
}
|
||||
var file_app_proxyman_config_proto_depIdxs = []int32{
|
||||
7, // 0: xray.app.proxyman.SniffingConfig.domains_excluded:type_name -> xray.common.geodata.DomainRule
|
||||
@@ -552,13 +562,14 @@ var file_app_proxyman_config_proto_depIdxs = []int32{
|
||||
12, // 7: xray.app.proxyman.InboundHandlerConfig.proxy_settings:type_name -> xray.common.serial.TypedMessage
|
||||
10, // 8: xray.app.proxyman.SenderConfig.via:type_name -> xray.common.net.IPOrDomain
|
||||
11, // 9: xray.app.proxyman.SenderConfig.stream_settings:type_name -> xray.transport.internet.StreamConfig
|
||||
6, // 10: xray.app.proxyman.SenderConfig.multiplex_settings:type_name -> xray.app.proxyman.MultiplexingConfig
|
||||
13, // 11: xray.app.proxyman.SenderConfig.target_strategy:type_name -> xray.transport.internet.DomainStrategy
|
||||
12, // [12:12] is the sub-list for method output_type
|
||||
12, // [12:12] is the sub-list for method input_type
|
||||
12, // [12:12] is the sub-list for extension type_name
|
||||
12, // [12:12] is the sub-list for extension extendee
|
||||
0, // [0:12] is the sub-list for field type_name
|
||||
13, // 10: xray.app.proxyman.SenderConfig.proxy_settings:type_name -> xray.transport.internet.ProxyConfig
|
||||
6, // 11: xray.app.proxyman.SenderConfig.multiplex_settings:type_name -> xray.app.proxyman.MultiplexingConfig
|
||||
14, // 12: xray.app.proxyman.SenderConfig.target_strategy:type_name -> xray.transport.internet.DomainStrategy
|
||||
13, // [13:13] is the sub-list for method output_type
|
||||
13, // [13:13] is the sub-list for method input_type
|
||||
13, // [13:13] is the sub-list for extension type_name
|
||||
13, // [13:13] is the sub-list for extension extendee
|
||||
0, // [0:13] is the sub-list for field type_name
|
||||
}
|
||||
|
||||
func init() { file_app_proxyman_config_proto_init() }
|
||||
|
||||
@@ -57,7 +57,7 @@ message SenderConfig {
|
||||
// Send traffic through the given IP. Only IP is allowed.
|
||||
xray.common.net.IPOrDomain via = 1;
|
||||
xray.transport.internet.StreamConfig stream_settings = 2;
|
||||
reserved 3;
|
||||
xray.transport.internet.ProxyConfig proxy_settings = 3;
|
||||
MultiplexingConfig multiplex_settings = 4;
|
||||
string via_cidr = 5;
|
||||
xray.transport.internet.DomainStrategy target_strategy = 6;
|
||||
|
||||
@@ -26,7 +26,7 @@ func getStatCounter(v *core.Instance, tag string) (stats.Counter, stats.Counter)
|
||||
if len(tag) > 0 && policy.ForSystem().Stats.InboundUplink {
|
||||
statsManager := v.GetFeature(stats.ManagerType()).(stats.Manager)
|
||||
name := "inbound>>>" + tag + ">>>traffic>>>uplink"
|
||||
c, _ := statsManager.GetOrRegisterCounter(name)
|
||||
c, _ := stats.GetOrRegisterCounter(statsManager, name)
|
||||
if c != nil {
|
||||
uplinkCounter = c
|
||||
}
|
||||
@@ -34,7 +34,7 @@ func getStatCounter(v *core.Instance, tag string) (stats.Counter, stats.Counter)
|
||||
if len(tag) > 0 && policy.ForSystem().Stats.InboundDownlink {
|
||||
statsManager := v.GetFeature(stats.ManagerType()).(stats.Manager)
|
||||
name := "inbound>>>" + tag + ">>>traffic>>>downlink"
|
||||
c, _ := statsManager.GetOrRegisterCounter(name)
|
||||
c, _ := stats.GetOrRegisterCounter(statsManager, name)
|
||||
if c != nil {
|
||||
downlinkCounter = c
|
||||
}
|
||||
@@ -66,7 +66,7 @@ func NewAlwaysOnInboundHandler(ctx context.Context, tag string, receiverConfig *
|
||||
}
|
||||
mss, err := internet.ToMemoryStreamConfig(receiverConfig.StreamSettings)
|
||||
if err != nil {
|
||||
return nil, errors.New("failed to parse stream config").Base(err)
|
||||
return nil, errors.New("failed to parse stream config").Base(err).AtWarning()
|
||||
}
|
||||
|
||||
newCtx := session.ContextWithInbound(ctx, &session.Inbound{Tag: tag, Source: src})
|
||||
|
||||
@@ -165,7 +165,7 @@ func NewHandler(ctx context.Context, config *core.InboundHandlerConfig) (inbound
|
||||
|
||||
receiverSettings, ok := rawReceiverSettings.(*proxyman.ReceiverConfig)
|
||||
if !ok {
|
||||
return nil, errors.New("not a ReceiverConfig")
|
||||
return nil, errors.New("not a ReceiverConfig").AtError()
|
||||
}
|
||||
|
||||
streamSettings := receiverSettings.StreamSettings
|
||||
|
||||
@@ -142,7 +142,7 @@ func (w *tcpWorker) Start() error {
|
||||
go w.callback(conn)
|
||||
})
|
||||
if err != nil {
|
||||
return errors.New("failed to listen TCP on ", w.port).Base(err)
|
||||
return errors.New("failed to listen TCP on ", w.port).AtWarning().Base(err)
|
||||
}
|
||||
w.hub = hub
|
||||
return nil
|
||||
@@ -528,7 +528,7 @@ func (w *dsWorker) Start() error {
|
||||
go w.callback(conn)
|
||||
})
|
||||
if err != nil {
|
||||
return errors.New("failed to listen Unix Domain Socket on ", w.address).Base(err)
|
||||
return errors.New("failed to listen Unix Domain Socket on ", w.address).AtWarning().Base(err)
|
||||
}
|
||||
w.hub = hub
|
||||
return nil
|
||||
|
||||
@@ -15,6 +15,7 @@ import (
|
||||
"github.com/xtls/xray-core/common/errors"
|
||||
"github.com/xtls/xray-core/common/mux"
|
||||
"github.com/xtls/xray-core/common/net"
|
||||
"github.com/xtls/xray-core/common/net/cnc"
|
||||
"github.com/xtls/xray-core/common/serial"
|
||||
"github.com/xtls/xray-core/common/session"
|
||||
"github.com/xtls/xray-core/core"
|
||||
@@ -25,6 +26,8 @@ import (
|
||||
"github.com/xtls/xray-core/transport"
|
||||
"github.com/xtls/xray-core/transport/internet"
|
||||
"github.com/xtls/xray-core/transport/internet/stat"
|
||||
"github.com/xtls/xray-core/transport/internet/tls"
|
||||
"github.com/xtls/xray-core/transport/pipe"
|
||||
"google.golang.org/protobuf/proto"
|
||||
)
|
||||
|
||||
@@ -36,7 +39,7 @@ func getStatCounter(v *core.Instance, tag string) (stats.Counter, stats.Counter)
|
||||
if len(tag) > 0 && policy.ForSystem().Stats.OutboundUplink {
|
||||
statsManager := v.GetFeature(stats.ManagerType()).(stats.Manager)
|
||||
name := "outbound>>>" + tag + ">>>traffic>>>uplink"
|
||||
c, _ := statsManager.GetOrRegisterCounter(name)
|
||||
c, _ := stats.GetOrRegisterCounter(statsManager, name)
|
||||
if c != nil {
|
||||
uplinkCounter = c
|
||||
}
|
||||
@@ -44,7 +47,7 @@ func getStatCounter(v *core.Instance, tag string) (stats.Counter, stats.Counter)
|
||||
if len(tag) > 0 && policy.ForSystem().Stats.OutboundDownlink {
|
||||
statsManager := v.GetFeature(stats.ManagerType()).(stats.Manager)
|
||||
name := "outbound>>>" + tag + ">>>traffic>>>downlink"
|
||||
c, _ := statsManager.GetOrRegisterCounter(name)
|
||||
c, _ := stats.GetOrRegisterCounter(statsManager, name)
|
||||
if c != nil {
|
||||
downlinkCounter = c
|
||||
}
|
||||
@@ -60,6 +63,7 @@ type Handler struct {
|
||||
streamSettings *internet.MemoryStreamConfig
|
||||
proxyConfig proto.Message
|
||||
proxy proxy.Outbound
|
||||
outboundManager outbound.Manager
|
||||
mux *mux.ClientManager
|
||||
xudp *mux.ClientManager
|
||||
udp443 string
|
||||
@@ -73,6 +77,7 @@ func NewHandler(ctx context.Context, config *core.OutboundHandlerConfig) (outbou
|
||||
uplinkCounter, downlinkCounter := getStatCounter(v, config.Tag)
|
||||
h := &Handler{
|
||||
tag: config.Tag,
|
||||
outboundManager: v.GetFeature(outbound.ManagerType()).(outbound.Manager),
|
||||
uplinkCounter: uplinkCounter,
|
||||
downlinkCounter: downlinkCounter,
|
||||
}
|
||||
@@ -87,7 +92,7 @@ func NewHandler(ctx context.Context, config *core.OutboundHandlerConfig) (outbou
|
||||
h.senderSettings = s
|
||||
mss, err := internet.ToMemoryStreamConfig(s.StreamSettings)
|
||||
if err != nil {
|
||||
return nil, errors.New("failed to parse stream settings").Base(err)
|
||||
return nil, errors.New("failed to parse stream settings").Base(err).AtWarning()
|
||||
}
|
||||
h.streamSettings = mss
|
||||
default:
|
||||
@@ -103,11 +108,9 @@ func NewHandler(ctx context.Context, config *core.OutboundHandlerConfig) (outbou
|
||||
|
||||
ctx = session.ContextWithFullHandler(ctx, h)
|
||||
|
||||
if h.streamSettings != nil {
|
||||
ctx = session.ContextWithStreamSettings(ctx, h.streamSettings)
|
||||
}
|
||||
newCtx := session.ContextWithStreamSettings(ctx, h.streamSettings)
|
||||
|
||||
rawProxyHandler, err := common.CreateObject(ctx, proxyConfig)
|
||||
rawProxyHandler, err := common.CreateObject(newCtx, proxyConfig)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
@@ -194,6 +197,7 @@ func (h *Handler) Dispatch(ctx context.Context, link *transport.Link) {
|
||||
common.Interrupt(link.Reader)
|
||||
return
|
||||
}
|
||||
|
||||
} else {
|
||||
unchangedDomain := ob.Target.Address.Domain()
|
||||
ob.Target.Address = net.IPAddress(ips[dice.Roll(len(ips))])
|
||||
@@ -217,7 +221,7 @@ func (h *Handler) Dispatch(ctx context.Context, link *transport.Link) {
|
||||
if ob.Target.Network == net.Network_UDP && ob.Target.Port == 443 {
|
||||
switch h.udp443 {
|
||||
case "reject":
|
||||
test(errors.New("XUDP rejected UDP/443 traffic"))
|
||||
test(errors.New("XUDP rejected UDP/443 traffic").AtInfo())
|
||||
return
|
||||
case "skip":
|
||||
goto out
|
||||
@@ -266,26 +270,66 @@ func (h *Handler) DestIpAddress() net.IP {
|
||||
|
||||
// Dial implements internet.Dialer.
|
||||
func (h *Handler) Dial(ctx context.Context, dest net.Destination) (stat.Connection, error) {
|
||||
if h.senderSettings != nil && h.senderSettings.Via != nil {
|
||||
outbounds := session.OutboundsFromContext(ctx)
|
||||
ob := outbounds[len(outbounds)-1]
|
||||
h.SetOutboundGateway(ctx, ob)
|
||||
if h.senderSettings != nil {
|
||||
|
||||
if h.senderSettings.ProxySettings.HasTag() {
|
||||
|
||||
tag := h.senderSettings.ProxySettings.Tag
|
||||
handler := h.outboundManager.GetHandler(tag)
|
||||
if handler != nil {
|
||||
errors.LogDebug(ctx, "proxying to ", tag, " for dest ", dest)
|
||||
outbounds := session.OutboundsFromContext(ctx)
|
||||
ctx = session.ContextWithOutbounds(ctx, append(outbounds, &session.Outbound{
|
||||
Target: dest,
|
||||
Tag: tag,
|
||||
})) // add another outbound in session ctx
|
||||
opts := pipe.OptionsFromContext(ctx)
|
||||
uplinkReader, uplinkWriter := pipe.New(opts...)
|
||||
downlinkReader, downlinkWriter := pipe.New(opts...)
|
||||
|
||||
go handler.Dispatch(ctx, &transport.Link{Reader: uplinkReader, Writer: downlinkWriter})
|
||||
conn := cnc.NewConnection(cnc.ConnectionInputMulti(uplinkWriter), cnc.ConnectionOutputMulti(downlinkReader))
|
||||
|
||||
if config := tls.ConfigFromStreamSettings(h.streamSettings); config != nil {
|
||||
tlsConfig := config.GetTLSConfig(tls.WithDestination(dest))
|
||||
conn = tls.Client(conn, tlsConfig)
|
||||
}
|
||||
|
||||
return h.getStatCouterConnection(conn), nil
|
||||
}
|
||||
|
||||
errors.LogError(ctx, "failed to get outbound handler with tag: ", tag)
|
||||
return nil, errors.New("failed to get outbound handler with tag: " + tag)
|
||||
}
|
||||
|
||||
if h.senderSettings.Via != nil {
|
||||
outbounds := session.OutboundsFromContext(ctx)
|
||||
ob := outbounds[len(outbounds)-1]
|
||||
h.SetOutboundGateway(ctx, ob)
|
||||
}
|
||||
}
|
||||
|
||||
conn, err := internet.Dial(ctx, dest, h.streamSettings)
|
||||
conn = h.getStatCouterConnection(conn)
|
||||
outbounds := session.OutboundsFromContext(ctx)
|
||||
if outbounds != nil {
|
||||
ob := outbounds[len(outbounds)-1]
|
||||
ob.Conn = conn
|
||||
} else {
|
||||
// for Vision's pre-connect
|
||||
}
|
||||
return conn, err
|
||||
}
|
||||
|
||||
func (h *Handler) SetOutboundGateway(ctx context.Context, ob *session.Outbound) {
|
||||
if ob.Gateway == nil && h.senderSettings != nil && h.senderSettings.Via != nil &&
|
||||
(h.streamSettings.SocketSettings == nil || len(h.streamSettings.SocketSettings.DialerProxy) == 0) {
|
||||
if ob.Gateway == nil && h.senderSettings != nil && h.senderSettings.Via != nil && !h.senderSettings.ProxySettings.HasTag() && (h.streamSettings.SocketSettings == nil || len(h.streamSettings.SocketSettings.DialerProxy) == 0) {
|
||||
var domain string
|
||||
addr := h.senderSettings.Via.AsAddress()
|
||||
domain = h.senderSettings.Via.GetDomain()
|
||||
switch {
|
||||
case h.senderSettings.ViaCidr != "":
|
||||
ob.Gateway = ParseRandomIP(addr, h.senderSettings.ViaCidr)
|
||||
|
||||
case domain == "origin":
|
||||
if inbound := session.InboundFromContext(ctx); inbound != nil {
|
||||
if inbound.Local.IsValid() && inbound.Local.Address.Family().IsIP() {
|
||||
@@ -300,9 +344,12 @@ func (h *Handler) SetOutboundGateway(ctx context.Context, ob *session.Outbound)
|
||||
errors.LogDebug(ctx, "use inbound source ip as sendthrough: ", inbound.Source.Address.String())
|
||||
}
|
||||
}
|
||||
default: // case addr.Family().IsDomain():
|
||||
// case addr.Family().IsDomain():
|
||||
default:
|
||||
ob.Gateway = addr
|
||||
|
||||
}
|
||||
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -68,13 +68,13 @@ func (p *Portal) HandleConnection(ctx context.Context, link *transport.Link) err
|
||||
outbounds := session.OutboundsFromContext(ctx)
|
||||
ob := outbounds[len(outbounds)-1]
|
||||
if ob == nil {
|
||||
return errors.New("outbound metadata not found")
|
||||
return errors.New("outbound metadata not found").AtError()
|
||||
}
|
||||
|
||||
if isDomain(ob.Target, p.domain) {
|
||||
muxClient, err := mux.NewClientWorker(*link, mux.ClientStrategy{})
|
||||
if err != nil {
|
||||
return errors.New("failed to create mux client worker").Base(err)
|
||||
return errors.New("failed to create mux client worker").Base(err).AtWarning()
|
||||
}
|
||||
|
||||
worker, err := NewPortalWorker(muxClient)
|
||||
|
||||
@@ -136,7 +136,7 @@ func (b *Balancer) SelectOutbounds() ([]string, error) {
|
||||
|
||||
// GetPrincipleTarget implements routing.BalancerPrincipleTarget
|
||||
func (r *Router) GetPrincipleTarget(tag string) ([]string, error) {
|
||||
if b, ok := (*r.balancers.Load())[tag]; ok {
|
||||
if b, ok := r.balancers[tag]; ok {
|
||||
if s, ok := b.strategy.(BalancingPrincipleTarget); ok {
|
||||
candidates, err := b.SelectOutbounds()
|
||||
if err != nil {
|
||||
@@ -151,7 +151,7 @@ func (r *Router) GetPrincipleTarget(tag string) ([]string, error) {
|
||||
|
||||
// SetOverrideTarget implements routing.BalancerOverrider
|
||||
func (r *Router) SetOverrideTarget(tag, target string) error {
|
||||
if b, ok := (*r.balancers.Load())[tag]; ok {
|
||||
if b, ok := r.balancers[tag]; ok {
|
||||
b.override.Put(target)
|
||||
return nil
|
||||
}
|
||||
@@ -160,7 +160,7 @@ func (r *Router) SetOverrideTarget(tag, target string) error {
|
||||
|
||||
// GetOverrideTarget implements routing.BalancerOverrider
|
||||
func (r *Router) GetOverrideTarget(tag string) (string, error) {
|
||||
if b, ok := (*r.balancers.Load())[tag]; ok {
|
||||
if b, ok := r.balancers[tag]; ok {
|
||||
return b.override.Get(), nil
|
||||
}
|
||||
return "", errors.New("cannot find tag")
|
||||
|
||||
@@ -2,8 +2,25 @@ package router
|
||||
|
||||
import (
|
||||
sync "sync"
|
||||
|
||||
"github.com/xtls/xray-core/common/errors"
|
||||
)
|
||||
|
||||
func (r *Router) OverrideBalancer(balancer string, target string) error {
|
||||
var b *Balancer
|
||||
for tag, bl := range r.balancers {
|
||||
if tag == balancer {
|
||||
b = bl
|
||||
break
|
||||
}
|
||||
}
|
||||
if b == nil {
|
||||
return errors.New("balancer '", balancer, "' not found")
|
||||
}
|
||||
b.override.Put(target)
|
||||
return nil
|
||||
}
|
||||
|
||||
type overrideSettings struct {
|
||||
target string
|
||||
}
|
||||
|
||||
@@ -5,7 +5,6 @@ import (
|
||||
"os"
|
||||
"path/filepath"
|
||||
"regexp"
|
||||
"runtime"
|
||||
"slices"
|
||||
"strings"
|
||||
|
||||
@@ -394,22 +393,3 @@ func (m *ProcessNameMatcher) Apply(ctx routing.Context) bool {
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// LocalOSMatcher matches the operating system Xray itself is running on. That never
|
||||
// changes while Xray is running, so the result is resolved when the rule is built.
|
||||
type LocalOSMatcher struct {
|
||||
matched bool
|
||||
}
|
||||
|
||||
func NewLocalOSMatcher(names []string) *LocalOSMatcher {
|
||||
return &LocalOSMatcher{
|
||||
matched: slices.ContainsFunc(names, func(name string) bool {
|
||||
return strings.EqualFold(name, runtime.GOOS)
|
||||
}),
|
||||
}
|
||||
}
|
||||
|
||||
// Apply implements Condition.
|
||||
func (m *LocalOSMatcher) Apply(_ routing.Context) bool {
|
||||
return m.matched
|
||||
}
|
||||
|
||||
@@ -2,9 +2,7 @@ package router_test
|
||||
|
||||
import (
|
||||
"path/filepath"
|
||||
"runtime"
|
||||
"strconv"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
. "github.com/xtls/xray-core/app/router"
|
||||
@@ -345,31 +343,6 @@ func TestChinaSites(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestLocalOSRule(t *testing.T) {
|
||||
otherOS := "plan9"
|
||||
if runtime.GOOS == otherOS {
|
||||
otherOS = "linux"
|
||||
}
|
||||
|
||||
cases := []struct {
|
||||
localOS []string
|
||||
output bool
|
||||
}{
|
||||
{localOS: []string{runtime.GOOS}, output: true},
|
||||
{localOS: []string{otherOS}, output: false},
|
||||
{localOS: []string{otherOS, runtime.GOOS}, output: true},
|
||||
{localOS: []string{strings.ToUpper(runtime.GOOS)}, output: true},
|
||||
}
|
||||
|
||||
for _, test := range cases {
|
||||
cond, err := (&RoutingRule{LocalOs: test.localOS}).BuildCondition()
|
||||
common.Must(err)
|
||||
if got := cond.Apply(withBackground()); got != test.output {
|
||||
t.Errorf("for localOS %v on %s: expected %v, got %v", test.localOS, runtime.GOOS, test.output, got)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func BenchmarkMphDomainMatcher(b *testing.B) {
|
||||
b.Setenv("xray.location.asset", filepath.Join("..", "..", "resources"))
|
||||
rules, err := geodata.ParseDomainRules([]string{"geosite:cn"}, geodata.Domain_Substr)
|
||||
|
||||
@@ -33,10 +33,6 @@ func (r *Rule) Apply(ctx routing.Context) bool {
|
||||
func (rr *RoutingRule) BuildCondition() (Condition, error) {
|
||||
conds := NewConditionChan()
|
||||
|
||||
if len(rr.LocalOs) > 0 {
|
||||
conds.Add(NewLocalOSMatcher(rr.LocalOs))
|
||||
}
|
||||
|
||||
if len(rr.InboundTag) > 0 {
|
||||
conds.Add(NewInboundTagMatcher(rr.InboundTag))
|
||||
}
|
||||
@@ -115,7 +111,7 @@ func (rr *RoutingRule) BuildCondition() (Condition, error) {
|
||||
}
|
||||
|
||||
if conds.Len() == 0 {
|
||||
return nil, errors.New("this rule has no effective fields")
|
||||
return nil, errors.New("this rule has no effective fields").AtWarning()
|
||||
}
|
||||
|
||||
return conds, nil
|
||||
@@ -145,7 +141,7 @@ func (br *BalancingRule) Build(ohm outbound.Manager, dispatcher routing.Dispatch
|
||||
}
|
||||
s, ok := i.(*StrategyLeastLoadConfig)
|
||||
if !ok {
|
||||
return nil, errors.New("not a StrategyLeastLoadConfig")
|
||||
return nil, errors.New("not a StrategyLeastLoadConfig").AtError()
|
||||
}
|
||||
leastLoadStrategy := NewLeastLoadStrategy(s)
|
||||
return &Balancer{
|
||||
|
||||
+4
-14
@@ -107,10 +107,8 @@ type RoutingRule struct {
|
||||
VlessRouteList *net.PortList `protobuf:"bytes,20,opt,name=vless_route_list,json=vlessRouteList,proto3" json:"vless_route_list,omitempty"`
|
||||
Process []string `protobuf:"bytes,21,rep,name=process,proto3" json:"process,omitempty"`
|
||||
Webhook *WebhookConfig `protobuf:"bytes,22,opt,name=webhook,proto3" json:"webhook,omitempty"`
|
||||
// List of operating systems for matching the one Xray itself is running on.
|
||||
LocalOs []string `protobuf:"bytes,23,rep,name=local_os,json=localOs,proto3" json:"local_os,omitempty"`
|
||||
unknownFields protoimpl.UnknownFields
|
||||
sizeCache protoimpl.SizeCache
|
||||
unknownFields protoimpl.UnknownFields
|
||||
sizeCache protoimpl.SizeCache
|
||||
}
|
||||
|
||||
func (x *RoutingRule) Reset() {
|
||||
@@ -280,13 +278,6 @@ func (x *RoutingRule) GetWebhook() *WebhookConfig {
|
||||
return nil
|
||||
}
|
||||
|
||||
func (x *RoutingRule) GetLocalOs() []string {
|
||||
if x != nil {
|
||||
return x.LocalOs
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
type isRoutingRule_TargetTag interface {
|
||||
isRoutingRule_TargetTag()
|
||||
}
|
||||
@@ -646,7 +637,7 @@ var File_app_router_config_proto protoreflect.FileDescriptor
|
||||
|
||||
const file_app_router_config_proto_rawDesc = "" +
|
||||
"\n" +
|
||||
"\x17app/router/config.proto\x12\x0fxray.app.router\x1a!common/serial/typed_message.proto\x1a\x15common/net/port.proto\x1a\x18common/net/network.proto\x1a\x1bcommon/geodata/geodat.proto\"\xdc\a\n" +
|
||||
"\x17app/router/config.proto\x12\x0fxray.app.router\x1a!common/serial/typed_message.proto\x1a\x15common/net/port.proto\x1a\x18common/net/network.proto\x1a\x1bcommon/geodata/geodat.proto\"\xc1\a\n" +
|
||||
"\vRoutingRule\x12\x12\n" +
|
||||
"\x03tag\x18\x01 \x01(\tH\x00R\x03tag\x12%\n" +
|
||||
"\rbalancing_tag\x18\f \x01(\tH\x00R\fbalancingTag\x12\x19\n" +
|
||||
@@ -670,8 +661,7 @@ const file_app_router_config_proto_rawDesc = "" +
|
||||
"\x0flocal_port_list\x18\x12 \x01(\v2\x19.xray.common.net.PortListR\rlocalPortList\x12C\n" +
|
||||
"\x10vless_route_list\x18\x14 \x01(\v2\x19.xray.common.net.PortListR\x0evlessRouteList\x12\x18\n" +
|
||||
"\aprocess\x18\x15 \x03(\tR\aprocess\x128\n" +
|
||||
"\awebhook\x18\x16 \x01(\v2\x1e.xray.app.router.WebhookConfigR\awebhook\x12\x19\n" +
|
||||
"\blocal_os\x18\x17 \x03(\tR\alocalOs\x1a=\n" +
|
||||
"\awebhook\x18\x16 \x01(\v2\x1e.xray.app.router.WebhookConfigR\awebhook\x1a=\n" +
|
||||
"\x0fAttributesEntry\x12\x10\n" +
|
||||
"\x03key\x18\x01 \x01(\tR\x03key\x12\x14\n" +
|
||||
"\x05value\x18\x02 \x01(\tR\x05value:\x028\x01B\f\n" +
|
||||
|
||||
@@ -56,9 +56,6 @@ message RoutingRule {
|
||||
|
||||
repeated string process = 21;
|
||||
WebhookConfig webhook = 22;
|
||||
|
||||
// List of operating systems for matching the one Xray itself is running on.
|
||||
repeated string local_os = 23;
|
||||
}
|
||||
|
||||
message WebhookConfig {
|
||||
|
||||
+114
-59
@@ -2,9 +2,7 @@ package router
|
||||
|
||||
import (
|
||||
"context"
|
||||
"maps"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
|
||||
"github.com/xtls/xray-core/common"
|
||||
"github.com/xtls/xray-core/common/errors"
|
||||
@@ -19,8 +17,8 @@ import (
|
||||
// Router is an implementation of routing.Router.
|
||||
type Router struct {
|
||||
domainStrategy Config_DomainStrategy
|
||||
rules atomic.Pointer[[]*Rule]
|
||||
balancers atomic.Pointer[map[string]*Balancer]
|
||||
rules []*Rule
|
||||
balancers map[string]*Balancer
|
||||
dns dns.Client
|
||||
|
||||
ctx context.Context
|
||||
@@ -45,9 +43,52 @@ func (r *Router) Init(ctx context.Context, config *Config, d dns.Client, ohm out
|
||||
r.ohm = ohm
|
||||
r.dispatcher = dispatcher
|
||||
|
||||
r.rules.Store(new([]*Rule))
|
||||
r.balancers.Store(&map[string]*Balancer{})
|
||||
return r.ReloadRules(config, false)
|
||||
r.balancers = make(map[string]*Balancer, len(config.BalancingRule))
|
||||
for _, rule := range config.BalancingRule {
|
||||
balancer, err := rule.Build(ohm, dispatcher)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
balancer.InjectContext(ctx)
|
||||
r.balancers[rule.Tag] = balancer
|
||||
}
|
||||
|
||||
r.rules = make([]*Rule, 0, len(config.Rule))
|
||||
for _, rule := range config.Rule {
|
||||
cond, err := rule.BuildCondition()
|
||||
if err != nil {
|
||||
r.closeWebhooks()
|
||||
return err
|
||||
}
|
||||
rr := &Rule{
|
||||
Condition: cond,
|
||||
Tag: rule.GetTag(),
|
||||
RuleTag: rule.GetRuleTag(),
|
||||
}
|
||||
if wh := rule.GetWebhook(); wh != nil {
|
||||
notifier, err := NewWebhookNotifier(wh)
|
||||
if err != nil {
|
||||
r.closeWebhooks()
|
||||
return err
|
||||
}
|
||||
rr.Webhook = notifier
|
||||
}
|
||||
btag := rule.GetBalancingTag()
|
||||
if len(btag) > 0 {
|
||||
brule, found := r.balancers[btag]
|
||||
if !found {
|
||||
if rr.Webhook != nil {
|
||||
rr.Webhook.Close()
|
||||
}
|
||||
r.closeWebhooks()
|
||||
return errors.New("balancer ", btag, " not found")
|
||||
}
|
||||
rr.Balancer = brule
|
||||
}
|
||||
r.rules = append(r.rules, rr)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// PickRoute implements routing.Router.
|
||||
@@ -83,22 +124,18 @@ func (r *Router) ReloadRules(config *Config, shouldAppend bool) error {
|
||||
r.mu.Lock()
|
||||
defer r.mu.Unlock()
|
||||
|
||||
oldRules := *r.rules.Load()
|
||||
oldBalancers := *r.balancers.Load()
|
||||
|
||||
var newRules []*Rule
|
||||
newBalancers := make(map[string]*Balancer)
|
||||
existTags := make(map[string]bool, len(oldRules)+len(config.Rule))
|
||||
if shouldAppend {
|
||||
newRules = append(newRules, oldRules...)
|
||||
maps.Copy(newBalancers, oldBalancers)
|
||||
for _, rule := range oldRules {
|
||||
existTags[rule.RuleTag] = true
|
||||
if !shouldAppend {
|
||||
for _, rule := range r.rules {
|
||||
if rule.Webhook != nil {
|
||||
rule.Webhook.Close()
|
||||
}
|
||||
}
|
||||
r.balancers = make(map[string]*Balancer, len(config.BalancingRule))
|
||||
r.rules = make([]*Rule, 0, len(config.Rule))
|
||||
}
|
||||
|
||||
for _, rule := range config.BalancingRule {
|
||||
if _, found := newBalancers[rule.Tag]; found {
|
||||
_, found := r.balancers[rule.Tag]
|
||||
if found {
|
||||
return errors.New("duplicate balancer tag")
|
||||
}
|
||||
balancer, err := rule.Build(r.ohm, r.dispatcher)
|
||||
@@ -106,12 +143,27 @@ func (r *Router) ReloadRules(config *Config, shouldAppend bool) error {
|
||||
return err
|
||||
}
|
||||
balancer.InjectContext(r.ctx)
|
||||
newBalancers[rule.Tag] = balancer
|
||||
r.balancers[rule.Tag] = balancer
|
||||
}
|
||||
|
||||
startIdx := len(r.rules)
|
||||
closeNewWebhooks := func() {
|
||||
for i := startIdx; i < len(r.rules); i++ {
|
||||
if r.rules[i].Webhook != nil {
|
||||
r.rules[i].Webhook.Close()
|
||||
}
|
||||
}
|
||||
r.rules = r.rules[:startIdx]
|
||||
}
|
||||
|
||||
for _, rule := range config.Rule {
|
||||
if r.RuleExists(rule.GetRuleTag()) {
|
||||
closeNewWebhooks()
|
||||
return errors.New("duplicate ruleTag ", rule.GetRuleTag())
|
||||
}
|
||||
cond, err := rule.BuildCondition()
|
||||
if err != nil {
|
||||
closeNewWebhooks()
|
||||
return err
|
||||
}
|
||||
rr := &Rule{
|
||||
@@ -119,64 +171,69 @@ func (r *Router) ReloadRules(config *Config, shouldAppend bool) error {
|
||||
Tag: rule.GetTag(),
|
||||
RuleTag: rule.GetRuleTag(),
|
||||
}
|
||||
if rr.RuleTag != "" && existTags[rr.RuleTag] {
|
||||
return errors.New("duplicate ruleTag ", rr.RuleTag)
|
||||
}
|
||||
existTags[rr.RuleTag] = true
|
||||
if wh := rule.GetWebhook(); wh != nil {
|
||||
notifier, err := NewWebhookNotifier(wh)
|
||||
if err != nil {
|
||||
closeNewWebhooks()
|
||||
return err
|
||||
}
|
||||
rr.Webhook = notifier
|
||||
}
|
||||
if btag := rule.GetBalancingTag(); len(btag) > 0 {
|
||||
brule, found := newBalancers[btag]
|
||||
btag := rule.GetBalancingTag()
|
||||
if len(btag) > 0 {
|
||||
brule, found := r.balancers[btag]
|
||||
if !found {
|
||||
if rr.Webhook != nil {
|
||||
rr.Webhook.Close()
|
||||
}
|
||||
closeNewWebhooks()
|
||||
return errors.New("balancer ", btag, " not found")
|
||||
}
|
||||
rr.Balancer = brule
|
||||
}
|
||||
newRules = append(newRules, rr)
|
||||
r.rules = append(r.rules, rr)
|
||||
}
|
||||
|
||||
r.balancers.Store(&newBalancers)
|
||||
r.rules.Store(&newRules)
|
||||
if !shouldAppend {
|
||||
closeWebhooks(oldRules)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (r *Router) RuleExists(tag string) bool {
|
||||
if tag != "" {
|
||||
for _, rule := range r.rules {
|
||||
if rule.RuleTag == tag {
|
||||
return true
|
||||
}
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// RemoveRule implements routing.Router.
|
||||
func (r *Router) RemoveRule(tag string) error {
|
||||
if tag == "" {
|
||||
return errors.New("empty tag name!")
|
||||
}
|
||||
|
||||
r.mu.Lock()
|
||||
defer r.mu.Unlock()
|
||||
|
||||
oldRules := *r.rules.Load()
|
||||
newRules := make([]*Rule, 0, len(oldRules))
|
||||
var removed []*Rule
|
||||
for _, rule := range oldRules {
|
||||
if rule.RuleTag != tag {
|
||||
newRules = append(newRules, rule)
|
||||
} else {
|
||||
removed = append(removed, rule)
|
||||
newRules := []*Rule{}
|
||||
if tag != "" {
|
||||
for _, rule := range r.rules {
|
||||
if rule.RuleTag != tag {
|
||||
newRules = append(newRules, rule)
|
||||
} else if rule.Webhook != nil {
|
||||
rule.Webhook.Close()
|
||||
}
|
||||
}
|
||||
r.rules = newRules
|
||||
return nil
|
||||
}
|
||||
r.rules.Store(&newRules)
|
||||
closeWebhooks(removed)
|
||||
return nil
|
||||
return errors.New("empty tag name!")
|
||||
}
|
||||
|
||||
// ListRule implements routing.Router
|
||||
func (r *Router) ListRule() []routing.Route {
|
||||
rules := *r.rules.Load()
|
||||
ruleList := make([]routing.Route, 0, len(rules))
|
||||
for _, rule := range rules {
|
||||
r.mu.Lock()
|
||||
defer r.mu.Unlock()
|
||||
ruleList := make([]routing.Route, 0)
|
||||
for _, rule := range r.rules {
|
||||
ruleList = append(ruleList, &Route{
|
||||
outboundTag: rule.Tag,
|
||||
ruleTag: rule.RuleTag,
|
||||
@@ -195,9 +252,7 @@ func (r *Router) pickRouteInternal(ctx routing.Context) (*Rule, routing.Context,
|
||||
ctx = routing_dns.ContextWithDNSClient(ctx, r.dns)
|
||||
}
|
||||
|
||||
rules := *r.rules.Load()
|
||||
|
||||
for _, rule := range rules {
|
||||
for _, rule := range r.rules {
|
||||
if rule.Apply(ctx) {
|
||||
return rule, ctx, nil
|
||||
}
|
||||
@@ -210,7 +265,7 @@ func (r *Router) pickRouteInternal(ctx routing.Context) (*Rule, routing.Context,
|
||||
ctx = routing_dns.ContextWithDNSClient(ctx, r.dns)
|
||||
|
||||
// Try applying rules again if we have IPs.
|
||||
for _, rule := range rules {
|
||||
for _, rule := range r.rules {
|
||||
if rule.Apply(ctx) {
|
||||
return rule, ctx, nil
|
||||
}
|
||||
@@ -224,9 +279,9 @@ func (r *Router) Start() error {
|
||||
return nil
|
||||
}
|
||||
|
||||
// closeWebhooks closes all webhook notifiers in the given rule set.
|
||||
func closeWebhooks(rules []*Rule) {
|
||||
for _, rule := range rules {
|
||||
// closeWebhooks closes all webhook notifiers in the current rule set.
|
||||
func (r *Router) closeWebhooks() {
|
||||
for _, rule := range r.rules {
|
||||
if rule.Webhook != nil {
|
||||
rule.Webhook.Close()
|
||||
}
|
||||
@@ -237,7 +292,7 @@ func closeWebhooks(rules []*Rule) {
|
||||
func (r *Router) Close() error {
|
||||
r.mu.Lock()
|
||||
defer r.mu.Unlock()
|
||||
closeWebhooks(*r.rules.Load())
|
||||
r.closeWebhooks()
|
||||
return nil
|
||||
}
|
||||
|
||||
|
||||
+23
-17
@@ -8,7 +8,6 @@ import (
|
||||
"net"
|
||||
"net/http"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"time"
|
||||
|
||||
"github.com/xtls/xray-core/common/errors"
|
||||
@@ -41,7 +40,6 @@ type WebhookNotifier struct {
|
||||
deduplication uint32
|
||||
client *http.Client
|
||||
seen sync.Map
|
||||
lastSweep atomic.Int64
|
||||
done chan struct{}
|
||||
wg sync.WaitGroup
|
||||
closeOnce sync.Once
|
||||
@@ -79,6 +77,11 @@ func NewWebhookNotifier(cfg *WebhookConfig) (*WebhookNotifier, error) {
|
||||
}
|
||||
}
|
||||
|
||||
if h.deduplication > 0 {
|
||||
h.wg.Add(1)
|
||||
go h.cleanupLoop()
|
||||
}
|
||||
|
||||
return h, nil
|
||||
}
|
||||
|
||||
@@ -198,7 +201,6 @@ func (h *WebhookNotifier) isDuplicate(email string) bool {
|
||||
}
|
||||
ttl := time.Duration(h.deduplication) * time.Second
|
||||
now := time.Now()
|
||||
h.maybeSweep(now, ttl)
|
||||
if v, loaded := h.seen.LoadOrStore(email, now); loaded {
|
||||
if now.Sub(v.(time.Time)) < ttl {
|
||||
return true
|
||||
@@ -208,23 +210,27 @@ func (h *WebhookNotifier) isDuplicate(email string) bool {
|
||||
return false
|
||||
}
|
||||
|
||||
func (h *WebhookNotifier) maybeSweep(now time.Time, ttl time.Duration) {
|
||||
last := h.lastSweep.Load()
|
||||
if now.UnixNano()-last < int64(ttl) {
|
||||
return
|
||||
}
|
||||
if !h.lastSweep.CompareAndSwap(last, now.UnixNano()) {
|
||||
return // another goroutine did the sweep
|
||||
}
|
||||
h.seen.Range(func(key, value any) bool {
|
||||
if now.Sub(value.(time.Time)) >= ttl {
|
||||
h.seen.Delete(key)
|
||||
func (h *WebhookNotifier) cleanupLoop() {
|
||||
defer h.wg.Done()
|
||||
ttl := time.Duration(h.deduplication) * time.Second
|
||||
ticker := time.NewTicker(ttl)
|
||||
defer ticker.Stop()
|
||||
for {
|
||||
select {
|
||||
case <-h.done:
|
||||
return
|
||||
case <-ticker.C:
|
||||
now := time.Now()
|
||||
h.seen.Range(func(key, value any) bool {
|
||||
if now.Sub(value.(time.Time)) >= ttl {
|
||||
h.seen.Delete(key)
|
||||
}
|
||||
return true
|
||||
})
|
||||
}
|
||||
return true
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// Only need to call if the Notifier is really used, otherwise GC can clean it
|
||||
func (h *WebhookNotifier) Close() error {
|
||||
h.closeOnce.Do(func() {
|
||||
close(h.done)
|
||||
|
||||
@@ -48,20 +48,6 @@ func (m *Manager) RegisterCounter(name string) (stats.Counter, error) {
|
||||
return c, nil
|
||||
}
|
||||
|
||||
// GetOrRegisterCounter implements stats.Manager.
|
||||
func (m *Manager) GetOrRegisterCounter(name string) (stats.Counter, error) {
|
||||
m.access.Lock()
|
||||
defer m.access.Unlock()
|
||||
|
||||
if c, found := m.counters[name]; found {
|
||||
return c, nil
|
||||
}
|
||||
errors.LogDebug(context.Background(), "create new counter ", name)
|
||||
c := new(Counter)
|
||||
m.counters[name] = c
|
||||
return c, nil
|
||||
}
|
||||
|
||||
// UnregisterCounter implements stats.Manager.
|
||||
func (m *Manager) UnregisterCounter(name string) error {
|
||||
m.access.Lock()
|
||||
@@ -111,20 +97,6 @@ func (m *Manager) RegisterOnlineMap(name string) (stats.OnlineMap, error) {
|
||||
return om, nil
|
||||
}
|
||||
|
||||
// GetOrRegisterOnlineMap implements stats.Manager.
|
||||
func (m *Manager) GetOrRegisterOnlineMap(name string) (stats.OnlineMap, error) {
|
||||
m.access.Lock()
|
||||
defer m.access.Unlock()
|
||||
|
||||
if om, found := m.onlineMaps[name]; found {
|
||||
return om, nil
|
||||
}
|
||||
errors.LogDebug(context.Background(), "create new OnlineMap ", name)
|
||||
om := NewOnlineMap()
|
||||
m.onlineMaps[name] = om
|
||||
return om, nil
|
||||
}
|
||||
|
||||
// UnregisterOnlineMap implements stats.Manager.
|
||||
func (m *Manager) UnregisterOnlineMap(name string) error {
|
||||
m.access.Lock()
|
||||
@@ -177,26 +149,6 @@ func (m *Manager) RegisterChannel(name string) (stats.Channel, error) {
|
||||
return c, nil
|
||||
}
|
||||
|
||||
// GetOrRegisterChannel implements stats.Manager.
|
||||
func (m *Manager) GetOrRegisterChannel(name string) (stats.Channel, error) {
|
||||
m.access.Lock()
|
||||
defer m.access.Unlock()
|
||||
|
||||
if c, found := m.channels[name]; found {
|
||||
return c, nil
|
||||
}
|
||||
errors.LogDebug(context.Background(), "create new channel ", name)
|
||||
c := NewChannel(&ChannelConfig{BufferSize: 64, Blocking: false})
|
||||
if m.running {
|
||||
// Start before publishing so no goroutine can observe an unstarted channel.
|
||||
if err := c.Start(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
m.channels[name] = c
|
||||
return c, nil
|
||||
}
|
||||
|
||||
// UnregisterChannel implements stats.Manager.
|
||||
func (m *Manager) UnregisterChannel(name string) error {
|
||||
m.access.Lock()
|
||||
|
||||
+1
-1
@@ -122,7 +122,7 @@ func NewReader(reader io.Reader) Reader {
|
||||
}
|
||||
|
||||
_, isFile := reader.(*os.File)
|
||||
if !isFile && useReadV() {
|
||||
if !isFile && useReadv {
|
||||
if sc, ok := reader.(syscall.Conn); ok {
|
||||
rawConn, err := sc.SyscallConn()
|
||||
if err != nil {
|
||||
|
||||
@@ -5,7 +5,6 @@ package buf
|
||||
|
||||
import (
|
||||
"io"
|
||||
"sync/atomic"
|
||||
"syscall"
|
||||
|
||||
"github.com/xtls/xray-core/common/platform"
|
||||
@@ -144,24 +143,13 @@ func (r *ReadVReader) ReadMultiBuffer() (MultiBuffer, error) {
|
||||
return mb, nil
|
||||
}
|
||||
|
||||
var useReadv atomic.Bool
|
||||
|
||||
func useReadV() bool {
|
||||
return useReadv.Load()
|
||||
}
|
||||
|
||||
func reloadEnvSettings() error {
|
||||
const defaultFlagValue = "NOT_DEFINED_AT_ALL"
|
||||
value := platform.NewEnvFlag(platform.UseReadV).GetValue(func() string { return defaultFlagValue })
|
||||
enabled := false
|
||||
switch value {
|
||||
case defaultFlagValue, "auto", "enable":
|
||||
enabled = true
|
||||
}
|
||||
useReadv.Store(enabled)
|
||||
return nil
|
||||
}
|
||||
var useReadv bool
|
||||
|
||||
func init() {
|
||||
platform.RegisterEnvReload(reloadEnvSettings)
|
||||
const defaultFlagValue = "NOT_DEFINED_AT_ALL"
|
||||
value := platform.NewEnvFlag(platform.UseReadV).GetValue(func() string { return defaultFlagValue })
|
||||
switch value {
|
||||
case defaultFlagValue, "auto", "enable":
|
||||
useReadv = true
|
||||
}
|
||||
}
|
||||
|
||||
@@ -10,9 +10,7 @@ import (
|
||||
"github.com/xtls/xray-core/features/stats"
|
||||
)
|
||||
|
||||
func useReadV() bool {
|
||||
return false
|
||||
}
|
||||
const useReadv = false
|
||||
|
||||
func NewReadVReader(reader io.Reader, rawConn syscall.RawConn, counter stats.Counter) Reader {
|
||||
panic("not implemented")
|
||||
|
||||
@@ -5,8 +5,7 @@ import (
|
||||
)
|
||||
|
||||
type windowsReader struct {
|
||||
bufs []syscall.WSABuf
|
||||
ready bool
|
||||
bufs []syscall.WSABuf
|
||||
}
|
||||
|
||||
func (r *windowsReader) Init(bs []*Buffer) {
|
||||
@@ -16,7 +15,6 @@ func (r *windowsReader) Init(bs []*Buffer) {
|
||||
for _, b := range bs {
|
||||
r.bufs = append(r.bufs, syscall.WSABuf{Len: uint32(Size), Buf: &b.v[0]})
|
||||
}
|
||||
r.ready = false
|
||||
}
|
||||
|
||||
func (r *windowsReader) Clear() {
|
||||
@@ -27,14 +25,6 @@ func (r *windowsReader) Clear() {
|
||||
}
|
||||
|
||||
func (r *windowsReader) Read(fd uintptr) int32 {
|
||||
// On the first invocation, we return -1 to indicate "not ready"
|
||||
// to make rawConn.Read wait for readability using the runtime's own mechanism
|
||||
// because syscall.WSARecv() is a blocking call when used with nil OVERLAPPED
|
||||
if !r.ready {
|
||||
r.ready = true
|
||||
return -1
|
||||
}
|
||||
|
||||
var nBytes uint32
|
||||
var flags uint32
|
||||
err := syscall.WSARecv(syscall.Handle(fd), &r.bufs[0], uint32(len(r.bufs)), &nBytes, &flags, nil, nil)
|
||||
|
||||
@@ -118,9 +118,7 @@ func (w *BufferedWriter) Write(b []byte) (int, error) {
|
||||
|
||||
nBytes, err := w.buffer.Write(b)
|
||||
totalBytes += nBytes
|
||||
|
||||
// ErrBufferFull means a partial write, so flush below and continue
|
||||
if err != nil && err != ErrBufferFull {
|
||||
if err != nil {
|
||||
return totalBytes, err
|
||||
}
|
||||
if !w.buffered || w.buffer.IsFull() {
|
||||
|
||||
@@ -10,12 +10,12 @@ import (
|
||||
|
||||
// [,)
|
||||
func RandBetween(from int64, to int64) int64 {
|
||||
if from == to {
|
||||
return from
|
||||
}
|
||||
if from > to {
|
||||
from, to = to, from
|
||||
}
|
||||
if d := to - from; d == 0 || d == 1 {
|
||||
return from
|
||||
}
|
||||
bigInt, _ := rand.Int(rand.Reader, big.NewInt(to-from))
|
||||
return from + bigInt.Int64()
|
||||
}
|
||||
|
||||
+65
-13
@@ -18,12 +18,17 @@ type hasInnerError interface {
|
||||
Unwrap() error
|
||||
}
|
||||
|
||||
type hasSeverity interface {
|
||||
Severity() log.Severity
|
||||
}
|
||||
|
||||
// Error is an error object with underlying error.
|
||||
type Error struct {
|
||||
prefix []interface{}
|
||||
message []interface{}
|
||||
caller string
|
||||
inner error
|
||||
prefix []interface{}
|
||||
message []interface{}
|
||||
caller string
|
||||
inner error
|
||||
severity log.Severity
|
||||
}
|
||||
|
||||
// Error implements error.Error().
|
||||
@@ -64,6 +69,46 @@ func (err *Error) Base(e error) *Error {
|
||||
return err
|
||||
}
|
||||
|
||||
func (err *Error) atSeverity(s log.Severity) *Error {
|
||||
err.severity = s
|
||||
return err
|
||||
}
|
||||
|
||||
func (err *Error) Severity() log.Severity {
|
||||
if err.inner == nil {
|
||||
return err.severity
|
||||
}
|
||||
|
||||
if s, ok := err.inner.(hasSeverity); ok {
|
||||
as := s.Severity()
|
||||
if as < err.severity {
|
||||
return as
|
||||
}
|
||||
}
|
||||
|
||||
return err.severity
|
||||
}
|
||||
|
||||
// AtDebug sets the severity to debug.
|
||||
func (err *Error) AtDebug() *Error {
|
||||
return err.atSeverity(log.Severity_Debug)
|
||||
}
|
||||
|
||||
// AtInfo sets the severity to info.
|
||||
func (err *Error) AtInfo() *Error {
|
||||
return err.atSeverity(log.Severity_Info)
|
||||
}
|
||||
|
||||
// AtWarning sets the severity to warning.
|
||||
func (err *Error) AtWarning() *Error {
|
||||
return err.atSeverity(log.Severity_Warning)
|
||||
}
|
||||
|
||||
// AtError sets the severity to error.
|
||||
func (err *Error) AtError() *Error {
|
||||
return err.atSeverity(log.Severity_Error)
|
||||
}
|
||||
|
||||
// String returns the string representation of this error.
|
||||
func (err *Error) String() string {
|
||||
return err.Error()
|
||||
@@ -87,8 +132,9 @@ func New(msg ...interface{}) *Error {
|
||||
details = details[:i]
|
||||
}
|
||||
return &Error{
|
||||
message: msg,
|
||||
caller: details,
|
||||
message: msg,
|
||||
severity: log.Severity_Info,
|
||||
caller: details,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -125,9 +171,6 @@ func LogErrorInner(ctx context.Context, inner error, msg ...interface{}) {
|
||||
}
|
||||
|
||||
func doLog(ctx context.Context, inner error, severity log.Severity, msg ...interface{}) {
|
||||
if log.GetSeverity() < severity {
|
||||
return
|
||||
}
|
||||
pc, _, _, _ := runtime.Caller(2)
|
||||
details := runtime.FuncForPC(pc).Name()
|
||||
if len(details) >= trim {
|
||||
@@ -138,9 +181,10 @@ func doLog(ctx context.Context, inner error, severity log.Severity, msg ...inter
|
||||
details = details[:i]
|
||||
}
|
||||
err := &Error{
|
||||
message: msg,
|
||||
caller: details,
|
||||
inner: inner,
|
||||
message: msg,
|
||||
severity: severity,
|
||||
caller: details,
|
||||
inner: inner,
|
||||
}
|
||||
if ctx != nil && ctx != context.Background() {
|
||||
id := uint32(c.IDFromContext(ctx))
|
||||
@@ -149,7 +193,7 @@ func doLog(ctx context.Context, inner error, severity log.Severity, msg ...inter
|
||||
}
|
||||
}
|
||||
log.Record(&log.GeneralMessage{
|
||||
Severity: severity,
|
||||
Severity: GetSeverity(err),
|
||||
Content: err,
|
||||
})
|
||||
}
|
||||
@@ -173,3 +217,11 @@ L:
|
||||
}
|
||||
return err
|
||||
}
|
||||
|
||||
// GetSeverity returns the actual severity of the error, including inner errors.
|
||||
func GetSeverity(err error) log.Severity {
|
||||
if s, ok := err.(hasSeverity); ok {
|
||||
return s.Severity()
|
||||
}
|
||||
return log.Severity_Info
|
||||
}
|
||||
|
||||
@@ -7,21 +7,30 @@ import (
|
||||
|
||||
"github.com/google/go-cmp/cmp"
|
||||
. "github.com/xtls/xray-core/common/errors"
|
||||
"github.com/xtls/xray-core/common/log"
|
||||
)
|
||||
|
||||
func TestError(t *testing.T) {
|
||||
err := New("TestError")
|
||||
if v := err.Error(); !strings.Contains(v, "TestError") {
|
||||
t.Error("error: ", v)
|
||||
if v := GetSeverity(err); v != log.Severity_Info {
|
||||
t.Error("severity: ", v)
|
||||
}
|
||||
|
||||
err = New("TestError2").Base(io.EOF)
|
||||
if v := err.Error(); !strings.Contains(v, "EOF") {
|
||||
t.Error("error: ", v)
|
||||
if v := GetSeverity(err); v != log.Severity_Info {
|
||||
t.Error("severity: ", v)
|
||||
}
|
||||
|
||||
err = New("TestError3").Base(io.EOF)
|
||||
err = New("TestError4").Base(err)
|
||||
err = New("TestError3").Base(io.EOF).AtWarning()
|
||||
if v := GetSeverity(err); v != log.Severity_Warning {
|
||||
t.Error("severity: ", v)
|
||||
}
|
||||
|
||||
err = New("TestError4").Base(io.EOF).AtWarning()
|
||||
err = New("TestError5").Base(err)
|
||||
if v := GetSeverity(err); v != log.Severity_Warning {
|
||||
t.Error("severity: ", v)
|
||||
}
|
||||
if v := err.Error(); !strings.Contains(v, "EOF") {
|
||||
t.Error("error: ", v)
|
||||
}
|
||||
|
||||
@@ -1,49 +0,0 @@
|
||||
package geodata
|
||||
|
||||
import (
|
||||
"sync"
|
||||
|
||||
"github.com/xtls/xray-core/common"
|
||||
)
|
||||
|
||||
var privateIPMatcher = sync.OnceValue(func() IPMatcher {
|
||||
return common.Must2(IPReg.BuildIPMatcher(common.Must2(ParseIPRules([]string{
|
||||
"0.0.0.0/8",
|
||||
"10.0.0.0/8",
|
||||
"100.64.0.0/10",
|
||||
"127.0.0.0/8",
|
||||
"169.254.0.0/16",
|
||||
"172.16.0.0/12",
|
||||
"192.0.0.0/24",
|
||||
"192.0.2.0/24",
|
||||
"192.88.99.0/24",
|
||||
"192.168.0.0/16",
|
||||
"198.18.0.0/15",
|
||||
"198.51.100.0/24",
|
||||
"203.0.113.0/24",
|
||||
"224.0.0.0/3",
|
||||
"::/127",
|
||||
"fc00::/7",
|
||||
"fe80::/10",
|
||||
"ff00::/8",
|
||||
}))))
|
||||
})
|
||||
|
||||
func GetPrivateIPMatcher() IPMatcher { return privateIPMatcher() }
|
||||
|
||||
var privateDomainMatcher = sync.OnceValue(func() DomainMatcher {
|
||||
return common.Must2(DomainReg.BuildDomainMatcher(common.Must2(ParseDomainRules([]string{
|
||||
"lan",
|
||||
"localdomain",
|
||||
"example",
|
||||
"invalid",
|
||||
"localhost",
|
||||
"test",
|
||||
"local",
|
||||
"home.arpa",
|
||||
"internal",
|
||||
"regexp:^[a-z]([a-z0-9-]{0,61}[a-z0-9])?$", // Dotless domains
|
||||
}, Domain_Domain))))
|
||||
})
|
||||
|
||||
func GetPrivateDomainMatcher() DomainMatcher { return privateDomainMatcher() }
|
||||
@@ -82,10 +82,19 @@ func (f *MphDomainMatcherFactory) BuildMatcher(rules []*DomainRule) (DomainMatch
|
||||
}
|
||||
g.Add(m, uint32(i))
|
||||
case *DomainRule_Geosite:
|
||||
err := loadSiteMatchers(v.Geosite, func(m strmatcher.Matcher) { g.Add(m, uint32(i)) })
|
||||
domains, err := loadSiteWithAttrs(v.Geosite.File, v.Geosite.Code, v.Geosite.Attrs)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
for j, d := range domains {
|
||||
domains[j] = nil // peak mem
|
||||
m, err := parseDomain(d)
|
||||
if err != nil {
|
||||
errors.LogError(context.Background(), "ignore invalid geosite entry in ", v.Geosite.File, ":", v.Geosite.Code, " at index ", j, ", ", err)
|
||||
continue
|
||||
}
|
||||
g.Add(m, uint32(i))
|
||||
}
|
||||
default:
|
||||
panic("unknown domain rule type")
|
||||
}
|
||||
@@ -99,12 +108,12 @@ func (f *MphDomainMatcherFactory) BuildMatcher(rules []*DomainRule) (DomainMatch
|
||||
return g, nil
|
||||
}
|
||||
|
||||
type CompactMphDomainMatcherFactory struct {
|
||||
type CompactDomainMatcherFactory struct {
|
||||
sync.Mutex
|
||||
shared *utils.WeakCacheMap[string, strmatcher.MphValueMatcher]
|
||||
shared *utils.WeakCacheMap[string, strmatcher.LinearAnyMatcher]
|
||||
}
|
||||
|
||||
func (f *CompactMphDomainMatcherFactory) getOrCreateFrom(rule *GeoSiteRule) (*strmatcher.MphValueMatcher, error) {
|
||||
func (f *CompactDomainMatcherFactory) getOrCreateFrom(rule *GeoSiteRule) (strmatcher.MatcherSet, error) {
|
||||
key := rule.File + ":" + rule.Code + "@" + rule.Attrs
|
||||
|
||||
f.Lock()
|
||||
@@ -116,23 +125,33 @@ func (f *CompactMphDomainMatcherFactory) getOrCreateFrom(rule *GeoSiteRule) (*st
|
||||
}
|
||||
errors.LogDebug(context.Background(), "geodata geosite matcher cache MISS ", key)
|
||||
|
||||
s := strmatcher.NewMphValueMatcher()
|
||||
if err := loadSiteMatchers(rule, func(m strmatcher.Matcher) { s.Add(m, 0) }); err != nil {
|
||||
s := strmatcher.NewLinearAnyMatcher()
|
||||
domains, err := loadSiteWithAttrs(rule.File, rule.Code, rule.Attrs)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := s.Build(); err != nil {
|
||||
return nil, err
|
||||
for i, d := range domains {
|
||||
domains[i] = nil // peak mem
|
||||
m, err := parseDomain(d)
|
||||
if err != nil {
|
||||
errors.LogError(context.Background(), "ignore invalid geosite entry in ", rule.File, ":", rule.Code, " at index ", i, ", ", err)
|
||||
continue
|
||||
}
|
||||
s.Add(m)
|
||||
}
|
||||
f.shared.Store(key, s)
|
||||
return s, nil
|
||||
return s, err
|
||||
}
|
||||
|
||||
// BuildMatcher implements DomainMatcherFactory.
|
||||
func (f *CompactMphDomainMatcherFactory) BuildMatcher(rules []*DomainRule) (DomainMatcher, error) {
|
||||
func (f *CompactDomainMatcherFactory) BuildMatcher(rules []*DomainRule) (DomainMatcher, error) {
|
||||
if len(rules) == 0 {
|
||||
return nil, errors.New("empty domain rule list")
|
||||
}
|
||||
compact := new(CompactMphDomainMatcher)
|
||||
compact := &CompactDomainMatcher{
|
||||
matchers: make([]strmatcher.MatcherSet, 0, len(rules)),
|
||||
values: make([]uint32, 0, len(rules)),
|
||||
}
|
||||
for i, r := range rules {
|
||||
switch v := r.Value.(type) {
|
||||
case *DomainRule_Custom:
|
||||
@@ -149,7 +168,8 @@ func (f *CompactMphDomainMatcherFactory) BuildMatcher(rules []*DomainRule) (Doma
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
compact.combiner.Add(m, uint32(i))
|
||||
compact.matchers = append(compact.matchers, m)
|
||||
compact.values = append(compact.values, uint32(i))
|
||||
default:
|
||||
panic("unknown domain rule type")
|
||||
}
|
||||
@@ -157,40 +177,37 @@ func (f *CompactMphDomainMatcherFactory) BuildMatcher(rules []*DomainRule) (Doma
|
||||
return compact, nil
|
||||
}
|
||||
|
||||
type CompactMphDomainMatcher struct {
|
||||
type CompactDomainMatcher struct {
|
||||
custom strmatcher.ValueMatcher
|
||||
combiner strmatcher.MphValueMatcherCombiner
|
||||
matchers []strmatcher.MatcherSet
|
||||
values []uint32
|
||||
}
|
||||
|
||||
// Match implements DomainMatcher.
|
||||
func (c *CompactMphDomainMatcher) Match(input string) []uint32 {
|
||||
result := c.combiner.Match(input)
|
||||
func (c *CompactDomainMatcher) Match(input string) []uint32 {
|
||||
var result []uint32
|
||||
if c.custom != nil {
|
||||
result = append(c.custom.Match(input), result...)
|
||||
result = append(result, c.custom.Match(input)...)
|
||||
}
|
||||
for i, m := range c.matchers {
|
||||
if m.MatchAny(input) {
|
||||
result = append(result, c.values[i])
|
||||
}
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
// MatchAny implements DomainMatcher.
|
||||
func (c *CompactMphDomainMatcher) MatchAny(input string) bool {
|
||||
func (c *CompactDomainMatcher) MatchAny(input string) bool {
|
||||
if c.custom != nil && c.custom.MatchAny(input) {
|
||||
return true
|
||||
}
|
||||
return c.combiner.MatchAny(input)
|
||||
}
|
||||
|
||||
// loadSiteMatchers calls add with a matcher for every domain of the geosite rule and logs the invalid ones.
|
||||
func loadSiteMatchers(rule *GeoSiteRule, add func(strmatcher.Matcher)) error {
|
||||
i := 0
|
||||
return loadSite(rule.File, rule.Code, rule.Attrs, func(t Domain_Type, value []byte) {
|
||||
m, err := parseDomain(&Domain{Type: t, Value: string(value)})
|
||||
if err != nil {
|
||||
errors.LogError(context.Background(), "ignore invalid geosite entry in ", rule.File, ":", rule.Code, " at index ", i, ", ", err)
|
||||
} else {
|
||||
add(m)
|
||||
for _, m := range c.matchers {
|
||||
if m.MatchAny(input) {
|
||||
return true
|
||||
}
|
||||
i++
|
||||
})
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func parseDomain(d *Domain) (strmatcher.Matcher, error) {
|
||||
@@ -214,7 +231,7 @@ func parseDomain(d *Domain) (strmatcher.Matcher, error) {
|
||||
func newDomainMatcherFactory() DomainMatcherFactory {
|
||||
switch runtime.GOOS {
|
||||
case "ios", "android":
|
||||
return &CompactMphDomainMatcherFactory{shared: utils.NewWeakCacheMap[string, strmatcher.MphValueMatcher]()}
|
||||
return &CompactDomainMatcherFactory{shared: utils.NewWeakCacheMap[string, strmatcher.LinearAnyMatcher]()}
|
||||
default:
|
||||
return &MphDomainMatcherFactory{shared: utils.NewWeakCacheMap[string, strmatcher.MphValueMatcher]()}
|
||||
}
|
||||
|
||||
@@ -4,7 +4,6 @@ import (
|
||||
"path/filepath"
|
||||
"reflect"
|
||||
"slices"
|
||||
"sync"
|
||||
"testing"
|
||||
|
||||
"github.com/xtls/xray-core/common/geodata/strmatcher"
|
||||
@@ -12,7 +11,7 @@ import (
|
||||
)
|
||||
|
||||
func TestCompactDomainMatcher_PreservesCustomRuleIndices(t *testing.T) {
|
||||
factory := &CompactMphDomainMatcherFactory{shared: utils.NewWeakCacheMap[string, strmatcher.MphValueMatcher]()}
|
||||
factory := &CompactDomainMatcherFactory{shared: utils.NewWeakCacheMap[string, strmatcher.LinearAnyMatcher]()}
|
||||
matcher, err := factory.BuildMatcher([]*DomainRule{
|
||||
{Value: &DomainRule_Custom{Custom: &Domain{Type: Domain_Full, Value: "example.com"}}},
|
||||
{Value: &DomainRule_Custom{Custom: &Domain{Type: Domain_Domain, Value: "example.com"}}},
|
||||
@@ -33,7 +32,7 @@ func TestCompactDomainMatcher_PreservesCustomRuleIndices(t *testing.T) {
|
||||
func TestCompactDomainMatcher_PreservesMixedRuleIndices(t *testing.T) {
|
||||
t.Setenv("xray.location.asset", filepath.Join("..", "..", "resources"))
|
||||
|
||||
factory := &CompactMphDomainMatcherFactory{shared: utils.NewWeakCacheMap[string, strmatcher.MphValueMatcher]()}
|
||||
factory := &CompactDomainMatcherFactory{shared: utils.NewWeakCacheMap[string, strmatcher.LinearAnyMatcher]()}
|
||||
matcher, err := factory.BuildMatcher([]*DomainRule{
|
||||
{Value: &DomainRule_Geosite{Geosite: &GeoSiteRule{File: DefaultGeoSiteDat, Code: "CN"}}},
|
||||
{Value: &DomainRule_Custom{Custom: &Domain{Type: Domain_Full, Value: "163.com"}}},
|
||||
@@ -73,76 +72,3 @@ func TestMphDomainMatcher_MatchReturnsDetachedSlice(t *testing.T) {
|
||||
t.Fatalf("Match() after caller mutation = %v, want %v", gotAgain, []uint32{0, 1})
|
||||
}
|
||||
}
|
||||
|
||||
// DNS sorts every Match result in place, so a matcher must never hand out a
|
||||
// slice it keeps, also when only its keyword or regex part matches.
|
||||
func TestDomainMatcher_MatchResultsCanBeSortedConcurrently(t *testing.T) {
|
||||
t.Setenv("xray.location.asset", filepath.Join("..", "..", "resources"))
|
||||
|
||||
rules := []*DomainRule{
|
||||
{Value: &DomainRule_Custom{Custom: &Domain{Type: Domain_Full, Value: "example.com"}}},
|
||||
{Value: &DomainRule_Custom{Custom: &Domain{Type: Domain_Domain, Value: "example.com"}}},
|
||||
{Value: &DomainRule_Custom{Custom: &Domain{Type: Domain_Substr, Value: "exam"}}},
|
||||
{Value: &DomainRule_Custom{Custom: &Domain{Type: Domain_Regex, Value: `^ex.*\.org$`}}},
|
||||
{Value: &DomainRule_Custom{Custom: &Domain{Type: Domain_Substr, Value: "exam"}}},
|
||||
{Value: &DomainRule_Geosite{Geosite: &GeoSiteRule{File: DefaultGeoSiteDat, Code: "CN"}}},
|
||||
{Value: &DomainRule_Custom{Custom: &Domain{Type: Domain_Full, Value: "only.full.test"}}},
|
||||
}
|
||||
cases := []struct {
|
||||
input string
|
||||
want []uint32
|
||||
}{
|
||||
{"example.com", []uint32{0, 1, 2, 4}},
|
||||
{"www.example.com", []uint32{1, 2, 4}},
|
||||
{"exam.net", []uint32{2, 4}}, // keyword part only
|
||||
{"example.org", []uint32{2, 3, 4}},
|
||||
{"163.com", []uint32{5}},
|
||||
{"www.163.com", []uint32{5}},
|
||||
{"only.full.test", []uint32{6}}, // full part only
|
||||
{"nomatch.test", nil},
|
||||
}
|
||||
factories := map[string]DomainMatcherFactory{
|
||||
"mph": &MphDomainMatcherFactory{shared: utils.NewWeakCacheMap[string, strmatcher.MphValueMatcher]()},
|
||||
"compact": &CompactMphDomainMatcherFactory{shared: utils.NewWeakCacheMap[string, strmatcher.MphValueMatcher]()},
|
||||
}
|
||||
for name, factory := range factories {
|
||||
t.Run(name, func(t *testing.T) {
|
||||
matcher, err := factory.BuildMatcher(rules)
|
||||
if err != nil {
|
||||
t.Fatalf("BuildMatcher() failed: %v", err)
|
||||
}
|
||||
for _, c := range cases {
|
||||
got := matcher.Match(c.input)
|
||||
if sorted := slices.Sorted(slices.Values(got)); !slices.Equal(sorted, c.want) {
|
||||
t.Fatalf("Match(%q) = %v, want %v", c.input, sorted, c.want)
|
||||
}
|
||||
got = got[:cap(got)]
|
||||
for j := range got {
|
||||
got[j] = ^uint32(0)
|
||||
}
|
||||
if again := slices.Sorted(slices.Values(matcher.Match(c.input))); !slices.Equal(again, c.want) {
|
||||
t.Fatalf("Match(%q) after caller mutation = %v, want %v", c.input, again, c.want)
|
||||
}
|
||||
}
|
||||
|
||||
var wg sync.WaitGroup
|
||||
for range 8 {
|
||||
wg.Add(1)
|
||||
go func() {
|
||||
defer wg.Done()
|
||||
for range 500 {
|
||||
for _, c := range cases {
|
||||
got := matcher.Match(c.input)
|
||||
slices.Sort(got)
|
||||
if !slices.Equal(got, c.want) {
|
||||
t.Errorf("Match(%q) = %v, want %v", c.input, got, c.want)
|
||||
return
|
||||
}
|
||||
}
|
||||
}
|
||||
}()
|
||||
}
|
||||
wg.Wait()
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
+62
-213
@@ -5,14 +5,11 @@ import (
|
||||
"bytes"
|
||||
"io"
|
||||
"runtime"
|
||||
"slices"
|
||||
"strings"
|
||||
"unicode/utf8"
|
||||
|
||||
"github.com/xtls/xray-core/common/errors"
|
||||
"github.com/xtls/xray-core/common/platform/filesystem"
|
||||
|
||||
"google.golang.org/protobuf/encoding/protowire"
|
||||
"google.golang.org/protobuf/proto"
|
||||
)
|
||||
|
||||
@@ -55,56 +52,17 @@ func loadIP(file, code string) ([]*CIDR, error) {
|
||||
return geoip.Cidr, nil
|
||||
}
|
||||
|
||||
// loadSite calls fn, in file order, with the type and value of every domain of the geosite code
|
||||
// that has all the "@"-separated attrs. It decodes the entry while reading the file instead of
|
||||
// unmarshalling it into a []*Domain, so value is only valid during fn.
|
||||
func loadSite(file, code, attrs string, fn func(Domain_Type, []byte)) error {
|
||||
runtime.GC() // peak mem
|
||||
r, err := filesystem.OpenAsset(file)
|
||||
func loadSite(file, code string) ([]*Domain, error) {
|
||||
bs, err := loadFile(file, code)
|
||||
if err != nil {
|
||||
return errors.New("failed to open ", file).Base(err)
|
||||
return nil, err
|
||||
}
|
||||
defer r.Close()
|
||||
br := bufio.NewReaderSize(r, 64*1024)
|
||||
n, err := seek(br, []byte(code))
|
||||
if err != nil {
|
||||
return errors.New("failed to load code ", code, " from ", file).Base(err)
|
||||
defer runtime.GC() // peak mem
|
||||
var geosite GeoSite
|
||||
if err := proto.Unmarshal(bs, &geosite); err != nil {
|
||||
return nil, errors.New("error unmarshal Site in ", file, ":", code).Base(err)
|
||||
}
|
||||
loadErr := func(err error) error {
|
||||
if err == io.EOF {
|
||||
err = io.ErrUnexpectedEOF
|
||||
}
|
||||
return errors.New("failed to load code ", code, " from ", file).Base(err)
|
||||
}
|
||||
unmarshalErr := func(err error) error {
|
||||
return errors.New("error unmarshal Site in ", file, ":", code).Base(err)
|
||||
}
|
||||
d := newSiteDecoder(attrs, fn)
|
||||
for n > 0 {
|
||||
w, err := br.Peek(min(n, br.Size()))
|
||||
if err != nil {
|
||||
return loadErr(err)
|
||||
}
|
||||
used, err := d.decode(w, len(w) < n)
|
||||
if err != nil {
|
||||
return unmarshalErr(err)
|
||||
}
|
||||
if used == 0 {
|
||||
break // a field longer than the buffer
|
||||
}
|
||||
br.Discard(used)
|
||||
n -= used
|
||||
}
|
||||
if n > 0 {
|
||||
w := make([]byte, n)
|
||||
if _, err := io.ReadFull(br, w); err != nil {
|
||||
return loadErr(err)
|
||||
}
|
||||
if _, err := d.decode(w, false); err != nil {
|
||||
return unmarshalErr(err)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
return geosite.Domain, nil
|
||||
}
|
||||
|
||||
func decodeVarint(br *bufio.Reader) (uint64, error) {
|
||||
@@ -124,63 +82,68 @@ func decodeVarint(br *bufio.Reader) (uint64, error) {
|
||||
}
|
||||
|
||||
func find(r io.Reader, code []byte, readBody bool) ([]byte, error) {
|
||||
br := bufio.NewReaderSize(r, 64*1024)
|
||||
bodyL, err := seek(br, code)
|
||||
if err != nil || !readBody {
|
||||
return nil, err
|
||||
}
|
||||
out := make([]byte, bodyL)
|
||||
if _, err := io.ReadFull(br, out); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
|
||||
// seek advances br to the body of the entry for code and returns the body length.
|
||||
func seek(br *bufio.Reader, code []byte) (int, error) {
|
||||
codeL := len(code)
|
||||
if codeL == 0 {
|
||||
return 0, errors.New("empty code")
|
||||
return nil, errors.New("empty code")
|
||||
}
|
||||
|
||||
br := bufio.NewReaderSize(r, 64*1024)
|
||||
need := 2 + codeL // TODO: if code too long
|
||||
prefixBuf := make([]byte, need)
|
||||
|
||||
for {
|
||||
if _, err := br.ReadByte(); err != nil {
|
||||
return 0, err
|
||||
return nil, err
|
||||
}
|
||||
|
||||
x, err := decodeVarint(br)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
return nil, err
|
||||
}
|
||||
bodyL := int(x)
|
||||
if bodyL <= 0 {
|
||||
return 0, errors.New("invalid body length: ", bodyL)
|
||||
return nil, errors.New("invalid body length: ", bodyL)
|
||||
}
|
||||
|
||||
// Peek no more than the buffer holds: a code longer than the buffer cannot match a single
|
||||
// length byte anyway, so a short peek only skips it, as base find (io.ReadFull) does.
|
||||
prefix, err := br.Peek(min(bodyL, need, br.Size()))
|
||||
if err != nil {
|
||||
if err == io.EOF && len(prefix) > 0 {
|
||||
err = io.ErrUnexpectedEOF // as io.ReadFull
|
||||
prefixL := bodyL
|
||||
if prefixL > need {
|
||||
prefixL = need
|
||||
}
|
||||
prefix := prefixBuf[:prefixL]
|
||||
if _, err := io.ReadFull(br, prefix); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
match := false
|
||||
if bodyL >= need {
|
||||
if int(prefix[1]) == codeL && bytes.Equal(prefix[2:need], code) {
|
||||
if !readBody {
|
||||
return nil, nil
|
||||
}
|
||||
match = true
|
||||
}
|
||||
return 0, err
|
||||
}
|
||||
if bodyL >= need && len(prefix) >= need && int(prefix[1]) == codeL && bytes.Equal(prefix[2:], code) {
|
||||
return bodyL, nil
|
||||
|
||||
remain := bodyL - prefixL
|
||||
if match {
|
||||
out := make([]byte, bodyL)
|
||||
copy(out, prefix)
|
||||
if remain > 0 {
|
||||
if _, err := io.ReadFull(br, out[prefixL:]); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
if _, err := br.Discard(bodyL); err != nil {
|
||||
return 0, err
|
||||
|
||||
if remain > 0 {
|
||||
if _, err := br.Discard(remain); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// AttributeMatcher, HasAttrMatcher, AllAttrsMatcher and NewAllAttrsMatcher are the exported
|
||||
// attribute helpers that have been part of this package's API since #5814. The streaming loader
|
||||
// above filters attributes itself without building a *Domain, so it does not use them, but they
|
||||
// are kept for external callers. Their behaviour is unchanged.
|
||||
|
||||
type AttributeMatcher interface {
|
||||
Match(*Domain) bool
|
||||
}
|
||||
@@ -222,137 +185,23 @@ func NewAllAttrsMatcher(attrs string) AttributeMatcher {
|
||||
return m
|
||||
}
|
||||
|
||||
var errInvalidUTF8 = errors.New("string field contains invalid UTF-8")
|
||||
func loadSiteWithAttrs(file, code, attrs string) ([]*Domain, error) {
|
||||
domains, err := loadSite(file, code)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
type siteDecoder struct {
|
||||
want []string
|
||||
has []bool
|
||||
fn func(Domain_Type, []byte)
|
||||
}
|
||||
matcher := NewAllAttrsMatcher(attrs)
|
||||
if matcher == nil {
|
||||
return domains, nil
|
||||
}
|
||||
|
||||
func newSiteDecoder(attrs string, fn func(Domain_Type, []byte)) *siteDecoder {
|
||||
d := &siteDecoder{fn: fn}
|
||||
if attrs != "" {
|
||||
d.want = strings.Split(attrs, "@")
|
||||
d.has = make([]bool, len(d.want))
|
||||
filtered := make([]*Domain, 0, len(domains))
|
||||
for _, d := range domains {
|
||||
if matcher.Match(d) {
|
||||
filtered = append(filtered, d)
|
||||
}
|
||||
}
|
||||
return d
|
||||
}
|
||||
|
||||
// decode walks the whole fields at the start of b, a part of an encoded GeoSite (see geodat.proto),
|
||||
// calls fn for every domain that has all attrs and returns how many bytes it used. A field cut off
|
||||
// by the end of b is an error unless more is set. It accepts and rejects what proto.Unmarshal does.
|
||||
func (d *siteDecoder) decode(b []byte, more bool) (int, error) {
|
||||
used := 0
|
||||
for used < len(b) {
|
||||
f, n, err := consumeField(b[used:])
|
||||
if err == io.ErrUnexpectedEOF && more {
|
||||
break
|
||||
}
|
||||
if err != nil {
|
||||
return used, err
|
||||
}
|
||||
used += n
|
||||
if f.typ != protowire.BytesType {
|
||||
continue
|
||||
}
|
||||
switch f.num {
|
||||
case 1: // code
|
||||
if !utf8.Valid(f.v) {
|
||||
return used, errInvalidUTF8
|
||||
}
|
||||
case 2: // domain
|
||||
t, value, err := decodeDomain(f.v, d.want, d.has)
|
||||
if err != nil {
|
||||
return used, err
|
||||
}
|
||||
if !slices.Contains(d.has, false) {
|
||||
d.fn(t, value)
|
||||
}
|
||||
}
|
||||
}
|
||||
return used, nil
|
||||
}
|
||||
|
||||
// decodeDomain decodes an encoded Domain and sets has[i] if one of its attributes has the key want[i].
|
||||
func decodeDomain(b []byte, want []string, has []bool) (t Domain_Type, value []byte, err error) {
|
||||
clear(has)
|
||||
for len(b) > 0 {
|
||||
f, n, err := consumeField(b)
|
||||
if err != nil {
|
||||
return 0, nil, err
|
||||
}
|
||||
b = b[n:]
|
||||
switch {
|
||||
case f.num == 1 && f.typ == protowire.VarintType: // type
|
||||
t = Domain_Type(f.x)
|
||||
case f.num == 2 && f.typ == protowire.BytesType: // value
|
||||
if !utf8.Valid(f.v) {
|
||||
return 0, nil, errInvalidUTF8
|
||||
}
|
||||
value = f.v
|
||||
case f.num == 3 && f.typ == protowire.BytesType: // attribute
|
||||
key, err := decodeAttributeKey(f.v)
|
||||
if err != nil {
|
||||
return 0, nil, err
|
||||
}
|
||||
for i, w := range want {
|
||||
if string(key) == w {
|
||||
has[i] = true
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
return t, value, nil
|
||||
}
|
||||
|
||||
// decodeAttributeKey returns the key of an encoded Domain.Attribute.
|
||||
func decodeAttributeKey(b []byte) ([]byte, error) {
|
||||
var key []byte
|
||||
for len(b) > 0 {
|
||||
f, n, err := consumeField(b)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
b = b[n:]
|
||||
if f.num == 1 && f.typ == protowire.BytesType {
|
||||
if !utf8.Valid(f.v) {
|
||||
return nil, errInvalidUTF8
|
||||
}
|
||||
key = f.v
|
||||
}
|
||||
}
|
||||
return key, nil
|
||||
}
|
||||
|
||||
type protoField struct {
|
||||
num protowire.Number
|
||||
typ protowire.Type
|
||||
v []byte // payload of a length-delimited field
|
||||
x uint64 // value of a varint field
|
||||
}
|
||||
|
||||
// consumeField parses the first field of an encoded message and returns it with its length.
|
||||
func consumeField(b []byte) (protoField, int, error) {
|
||||
num, typ, n := protowire.ConsumeTag(b)
|
||||
if n < 0 {
|
||||
return protoField{}, 0, protowire.ParseError(n)
|
||||
}
|
||||
if num > protowire.MaxValidNumber {
|
||||
return protoField{}, 0, errors.New("invalid field number ", num)
|
||||
}
|
||||
f := protoField{num: num, typ: typ}
|
||||
var m int
|
||||
switch typ {
|
||||
case protowire.BytesType:
|
||||
f.v, m = protowire.ConsumeBytes(b[n:])
|
||||
case protowire.VarintType:
|
||||
f.x, m = protowire.ConsumeVarint(b[n:])
|
||||
default:
|
||||
m = protowire.ConsumeFieldValue(num, typ, b[n:])
|
||||
}
|
||||
if m < 0 {
|
||||
return protoField{}, 0, protowire.ParseError(m)
|
||||
}
|
||||
return f, n + m, nil
|
||||
return filtered, nil
|
||||
}
|
||||
|
||||
@@ -1,283 +0,0 @@
|
||||
package geodata
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"slices"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"google.golang.org/protobuf/encoding/protowire"
|
||||
"google.golang.org/protobuf/proto"
|
||||
)
|
||||
|
||||
type siteEntry struct {
|
||||
Type Domain_Type
|
||||
Value string
|
||||
}
|
||||
|
||||
// unmarshalSite is what loadSite used to do: proto.Unmarshal, then keep the domains that have all attrs.
|
||||
func unmarshalSite(b []byte, attrs string) ([]siteEntry, error) {
|
||||
var site GeoSite
|
||||
if err := proto.Unmarshal(b, &site); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
var entries []siteEntry
|
||||
for _, d := range site.Domain {
|
||||
ok := true
|
||||
for _, key := range strings.Split(attrs, "@") {
|
||||
ok = ok && (attrs == "" || slices.ContainsFunc(d.Attribute, func(a *Domain_Attribute) bool { return a.Key == key }))
|
||||
}
|
||||
if ok {
|
||||
entries = append(entries, siteEntry{d.Type, d.Value})
|
||||
}
|
||||
}
|
||||
return entries, nil
|
||||
}
|
||||
|
||||
func checkDecodeSite(t *testing.T, name string, b []byte, attrs string) {
|
||||
t.Helper()
|
||||
want, wantErr := unmarshalSite(b, attrs)
|
||||
var got []siteEntry
|
||||
_, err := newSiteDecoder(attrs, func(typ Domain_Type, value []byte) {
|
||||
got = append(got, siteEntry{typ, string(value)})
|
||||
}).decode(b, false)
|
||||
if (err == nil) != (wantErr == nil) {
|
||||
t.Fatalf("%s@%s: error %v, proto.Unmarshal: %v", name, attrs, err, wantErr)
|
||||
}
|
||||
if err == nil && !slices.Equal(got, want) {
|
||||
t.Fatalf("%s@%s: got %v, want %v", name, attrs, got, want)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDecodeSiteMatchesUnmarshal(t *testing.T) {
|
||||
bs, err := os.ReadFile(filepath.Join("..", "..", "resources", DefaultGeoSiteDat))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
for len(bs) > 0 {
|
||||
num, typ, n := protowire.ConsumeTag(bs)
|
||||
if n < 0 || num != 1 || typ != protowire.BytesType {
|
||||
t.Fatal("unexpected GeoSiteList field")
|
||||
}
|
||||
entry, m := protowire.ConsumeBytes(bs[n:])
|
||||
if m < 0 {
|
||||
t.Fatal(protowire.ParseError(m))
|
||||
}
|
||||
bs = bs[n+m:]
|
||||
|
||||
var site GeoSite
|
||||
if err := proto.Unmarshal(entry, &site); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
queries := []string{"", "none"}
|
||||
for _, d := range site.Domain {
|
||||
for _, a := range d.Attribute {
|
||||
if !slices.Contains(queries, a.Key) {
|
||||
queries = append(queries, a.Key, a.Key+"@none")
|
||||
}
|
||||
}
|
||||
}
|
||||
for _, attrs := range queries {
|
||||
checkDecodeSite(t, site.Code, entry, attrs)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestDecodeSiteUnusualEncodings(t *testing.T) {
|
||||
field := func(num protowire.Number, v []byte) []byte {
|
||||
return protowire.AppendBytes(protowire.AppendTag(nil, num, protowire.BytesType), v)
|
||||
}
|
||||
typ := func(v Domain_Type) []byte {
|
||||
return protowire.AppendVarint(protowire.AppendTag(nil, 1, protowire.VarintType), uint64(v))
|
||||
}
|
||||
value := func(s string) []byte { return field(2, []byte(s)) }
|
||||
attr := func(keys ...string) []byte {
|
||||
var b []byte
|
||||
for _, k := range keys {
|
||||
b = append(b, field(1, []byte(k))...)
|
||||
}
|
||||
return field(3, b)
|
||||
}
|
||||
domain := func(fields ...[]byte) []byte { return field(2, slices.Concat(fields...)) }
|
||||
unknown := protowire.AppendFixed32(protowire.AppendTag(nil, 9, protowire.Fixed32Type), 1)
|
||||
|
||||
for name, b := range map[string][]byte{
|
||||
"unknown field": domain(typ(Domain_Full), unknown, value("example.com")),
|
||||
"repeated value": domain(value("a.com"), typ(Domain_Full), value("b.com")),
|
||||
"repeated type": domain(typ(Domain_Full), value("a.com"), typ(Domain_Regex)),
|
||||
"repeated key": domain(value("a.com"), attr("cn", "ads")),
|
||||
"type as bytes": domain(field(1, []byte("x")), value("a.com")),
|
||||
"no value": domain(typ(Domain_Domain), attr("cn")),
|
||||
"truncated": domain(typ(Domain_Full), value("example.com"))[:10],
|
||||
"invalid utf8": domain(value("example.\xff")),
|
||||
"invalid key": domain(value("a.com"), attr("\xff")),
|
||||
"bad field": protowire.AppendVarint(protowire.AppendTag(nil, protowire.MaxValidNumber+1, protowire.VarintType), 1),
|
||||
"stray end group": protowire.AppendTag(nil, 5, protowire.EndGroupType),
|
||||
} {
|
||||
for _, attrs := range []string{"", "cn", "ads", "cn@ads"} {
|
||||
checkDecodeSite(t, name, b, attrs)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestLoadSiteReadsInPieces covers what real lists never do: an entry far longer than the read
|
||||
// buffer, with a field longer than the buffer in the middle, and a file cut short.
|
||||
func TestLoadSiteReadsInPieces(t *testing.T) {
|
||||
site := &GeoSite{Code: "BIG"}
|
||||
for i := range 5000 {
|
||||
d := &Domain{Type: Domain_Domain, Value: strings.Repeat("x", i%40) + ".example.com"}
|
||||
if i%3 == 0 {
|
||||
d.Attribute = []*Domain_Attribute{{Key: "cn"}}
|
||||
}
|
||||
if i == 2500 {
|
||||
d = &Domain{Type: Domain_Regex, Value: strings.Repeat("a", 100_000)}
|
||||
}
|
||||
site.Domain = append(site.Domain, d)
|
||||
}
|
||||
list := &GeoSiteList{Entry: []*GeoSite{{Code: "SMALL", Domain: []*Domain{{Type: Domain_Full, Value: "a.com"}}}, site}}
|
||||
bs, err := proto.Marshal(list)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
entry, err := proto.Marshal(site)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
dir := t.TempDir()
|
||||
t.Setenv("xray.location.asset", dir)
|
||||
write := func(b []byte) {
|
||||
if err := os.WriteFile(filepath.Join(dir, "big.dat"), b, 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
for _, attrs := range []string{"", "cn"} {
|
||||
want, _ := unmarshalSite(entry, attrs)
|
||||
var got []siteEntry
|
||||
write(bs)
|
||||
err := loadSite("big.dat", "BIG", attrs, func(typ Domain_Type, value []byte) {
|
||||
got = append(got, siteEntry{typ, string(value)})
|
||||
})
|
||||
if err != nil || !slices.Equal(got, want) {
|
||||
t.Fatalf("attrs %q: %d entries, want %d, error %v", attrs, len(got), len(want), err)
|
||||
}
|
||||
for _, cut := range []int{30_000, len(bs) - 150_000, len(bs) - 1} {
|
||||
write(bs[:cut])
|
||||
if err := loadSite("big.dat", "BIG", attrs, func(Domain_Type, []byte) {}); err == nil {
|
||||
t.Fatalf("file cut at %d of %d: no error", cut, len(bs))
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// oneEntryGeoSiteFile wraps an encoded GeoSite as a one-entry GeoSiteList, the file loadSite reads.
|
||||
func oneEntryGeoSiteFile(entry []byte) []byte {
|
||||
return protowire.AppendBytes(protowire.AppendTag(nil, 1, protowire.BytesType), entry)
|
||||
}
|
||||
|
||||
// TestLoadSiteWindowedMatchesSingleShot checks that the windowed reader in loadSite (its Peek/Discard
|
||||
// loop, the more-break when a field is cut by a window edge, the used==0 fallback for a field longer
|
||||
// than the buffer, and the tail path) reaches exactly the same result as decoding the whole entry at
|
||||
// once, for a category several 64 KiB windows long, valid and then mutated near a window edge and
|
||||
// early in the file: same error-or-not, and the same emitted (type, value) sequence when both accept.
|
||||
func TestLoadSiteWindowedMatchesSingleShot(t *testing.T) {
|
||||
const window = 64 * 1024
|
||||
site := &GeoSite{Code: "BIG"}
|
||||
for i := range 12000 { // ~250 KiB, four windows
|
||||
d := &Domain{Type: Domain_Domain, Value: fmt.Sprintf("host%d.%s.example.com", i, strings.Repeat("y", i%30))}
|
||||
if i%3 == 0 {
|
||||
d.Attribute = []*Domain_Attribute{{Key: "cn"}}
|
||||
}
|
||||
site.Domain = append(site.Domain, d)
|
||||
}
|
||||
// a field longer than the buffer, straddling the third window, to force the used==0 fallback
|
||||
site.Domain = slices.Insert(site.Domain, 8000, &Domain{Type: Domain_Regex, Value: strings.Repeat("a", 90_000)})
|
||||
entry, err := proto.Marshal(site)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
dir := t.TempDir()
|
||||
t.Setenv("xray.location.asset", dir)
|
||||
|
||||
// mutations of the encoded entry: unchanged, a byte flipped at several offsets (early windows and
|
||||
// either side of a window edge), and truncations at the same places.
|
||||
type mut struct {
|
||||
name string
|
||||
make func([]byte) []byte
|
||||
}
|
||||
muts := []mut{{"valid", func(b []byte) []byte { return b }}}
|
||||
for _, off := range []int{3, 40, 4000, window - 1, window, window + 1, 2*window - 2, 2 * window} {
|
||||
if off < len(entry) {
|
||||
off := off
|
||||
muts = append(muts, mut{fmt.Sprintf("flip@%d", off), func(b []byte) []byte {
|
||||
c := slices.Clone(b)
|
||||
c[off] ^= 0xff
|
||||
return c
|
||||
}})
|
||||
muts = append(muts, mut{fmt.Sprintf("cut@%d", off), func(b []byte) []byte { return slices.Clone(b[:off]) }})
|
||||
}
|
||||
}
|
||||
|
||||
for _, attrs := range []string{"", "cn"} {
|
||||
for _, m := range muts {
|
||||
e := m.make(entry)
|
||||
// single-shot reference: decode the whole entry in one call
|
||||
var want []siteEntry
|
||||
_, wantErr := newSiteDecoder(attrs, func(typ Domain_Type, value []byte) {
|
||||
want = append(want, siteEntry{typ, string(value)})
|
||||
}).decode(e, false)
|
||||
// windowed: loadSite reads the file 64 KiB at a time
|
||||
if err := os.WriteFile(filepath.Join(dir, "w.dat"), oneEntryGeoSiteFile(e), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
var got []siteEntry
|
||||
gotErr := loadSite("w.dat", "BIG", attrs, func(typ Domain_Type, value []byte) {
|
||||
got = append(got, siteEntry{typ, string(value)})
|
||||
})
|
||||
if (gotErr == nil) != (wantErr == nil) {
|
||||
t.Fatalf("%s attrs=%q: windowed err %v, single-shot err %v", m.name, attrs, gotErr, wantErr)
|
||||
}
|
||||
if gotErr == nil && !slices.Equal(got, want) {
|
||||
t.Fatalf("%s attrs=%q: windowed got %d entries, single-shot %d", m.name, attrs, len(got), len(want))
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestLoadSiteLongCode covers a geosite entry whose code is longer than the 64 KiB read buffer. seek
|
||||
// must skip it (find compares a single length byte, so it never matches such a code) and still find a
|
||||
// later entry, and looking the long code up must fail cleanly, like a missing code, not panic.
|
||||
func TestLoadSiteLongCode(t *testing.T) {
|
||||
longCode := strings.Repeat("Z", 70000)
|
||||
list := &GeoSiteList{Entry: []*GeoSite{
|
||||
{Code: "FIRST", Domain: []*Domain{{Type: Domain_Full, Value: "first.com"}}},
|
||||
{Code: longCode, Domain: []*Domain{{Type: Domain_Full, Value: "huge.com"}}},
|
||||
{Code: "AFTER", Domain: []*Domain{{Type: Domain_Domain, Value: "after.com"}}},
|
||||
}}
|
||||
bs, err := proto.Marshal(list)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
dir := t.TempDir()
|
||||
t.Setenv("xray.location.asset", dir)
|
||||
if err := os.WriteFile(filepath.Join(dir, "lc.dat"), bs, 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
collect := func(code string) ([]siteEntry, error) {
|
||||
var got []siteEntry
|
||||
err := loadSite("lc.dat", code, "", func(typ Domain_Type, value []byte) {
|
||||
got = append(got, siteEntry{typ, string(value)})
|
||||
})
|
||||
return got, err
|
||||
}
|
||||
if got, err := collect("FIRST"); err != nil || !slices.Equal(got, []siteEntry{{Domain_Full, "first.com"}}) {
|
||||
t.Fatalf("FIRST: %v %v", got, err)
|
||||
}
|
||||
if got, err := collect("AFTER"); err != nil || !slices.Equal(got, []siteEntry{{Domain_Domain, "after.com"}}) {
|
||||
t.Fatalf("AFTER (past the oversized entry): %v %v", got, err)
|
||||
}
|
||||
if _, err := collect(longCode); err == nil {
|
||||
t.Fatal("oversized code: expected a not-found error, got nil")
|
||||
}
|
||||
}
|
||||
@@ -1,7 +1,6 @@
|
||||
package strmatcher_test
|
||||
|
||||
import (
|
||||
"regexp"
|
||||
"strconv"
|
||||
"testing"
|
||||
|
||||
@@ -73,64 +72,6 @@ func BenchmarkSubstrMatcher(b *testing.B) {
|
||||
})
|
||||
}
|
||||
|
||||
func BenchmarkRegexMatcher(b *testing.B) {
|
||||
patterns := []string{ // taken from geosite
|
||||
`(^|\.)91porn\.(best|com|cool|fun|group|party|plus|site|tw|work)$`,
|
||||
`(^|\.)91porn[0-9]{3}\.me$`,
|
||||
`(^|\.)apiproxy-device-prod-nlb-.+\.amazonaws\.com$`,
|
||||
`(^|\.)dualstack\.apiproxy-.+\.amazonaws\.com$`,
|
||||
`(^|\.)aqdk[0-9]{3}\.com$`,
|
||||
`(^|\.)bilibili3(0[1-9]|1[0-2])\.xyz$`,
|
||||
`(^|\.)byyum([3589]|2[235689]|3[34]|4[1-9]|5[1-79]|6[0134679])?\.com$`,
|
||||
`(^|\.)fiftymvapi\..+$`,
|
||||
`(^|\.)gossipfuli[0-9]{3,4}\.xyz$`,
|
||||
`(^|\.)kpkuang\.(bond|fun|info|one|us)$`,
|
||||
`(^|\.)rule34\.(asia|us|world|xxx|xyz)$`,
|
||||
`(^|\.)[a-z][1-9][0-9][a-z]\.com$`,
|
||||
`.+\.awsdns-[0-9][0-9]\.(co\.uk|com|net|org)$`,
|
||||
`.+\.dkr\.ecr\.[^\.]+\.amazonaws\.com$`,
|
||||
`^(.+\.)*zh\.okaapps\.com$`,
|
||||
`^cdn\d-epicgames-\d+\.file\.myqcloud\.com$`,
|
||||
`^chatgpt-async-webps-prod-\S+-\d+\.webpubsub\.azure\.com$`,
|
||||
`^r+[0-9]+(---|\.)sn-(2x3|ni5|j5o)\w{5}\.googlevideo\.com$`,
|
||||
`^speed\.(coe|open)\.ad\.[a-z]{2,6}\.prod\.hosts\.ooklaserver\.net$`,
|
||||
`javdb\d+\.com$`,
|
||||
}
|
||||
domains := []string{
|
||||
"www.google.com", "rr3---sn-4g5edndy.googlevideo.com", "r1---sn-2x3abcde.googlevideo.com", "i.ytimg.com",
|
||||
"graph.facebook.com", "api.twitter.com", "www.baidu.com", "github.com", "objects.githubusercontent.com",
|
||||
"login.microsoftonline.com", "e1234.dscb.akamaiedge.net", "d1a2b3c4d5e6f7.cloudfront.net",
|
||||
"s3.us-east-1.amazonaws.com", "123456789012.dkr.ecr.us-east-1.amazonaws.com", "www.wikipedia.org",
|
||||
"discord.com", "telegram.org", "store.steampowered.com", "www.91porn.com", "ns-1234.awsdns-12.org",
|
||||
}
|
||||
bench := func(b *testing.B, ctor func(pattern string) func(string) bool) {
|
||||
var matchers []func(string) bool
|
||||
for _, p := range patterns {
|
||||
matchers = append(matchers, ctor(p))
|
||||
}
|
||||
b.ResetTimer()
|
||||
for i := 0; i < b.N; i++ {
|
||||
for _, d := range domains {
|
||||
for _, match := range matchers {
|
||||
_ = match(d)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
b.Run("regexp", func(b *testing.B) {
|
||||
bench(b, func(pattern string) func(string) bool {
|
||||
return regexp.MustCompile(pattern).MatchString
|
||||
})
|
||||
})
|
||||
b.Run("prefilter", func(b *testing.B) {
|
||||
bench(b, func(pattern string) func(string) bool {
|
||||
m, err := Regex.New(pattern)
|
||||
common.Must(err)
|
||||
return m.Match
|
||||
})
|
||||
})
|
||||
}
|
||||
|
||||
// Utility functions for benchmark
|
||||
|
||||
func benchmarkMatcherType(b *testing.B, t Type, ctor func() MatcherGroup) {
|
||||
|
||||
@@ -52,9 +52,7 @@ func (g *MphIndexMatcher) Add(matcher Matcher) uint32 {
|
||||
func (g *MphIndexMatcher) Build() error {
|
||||
if g.mph != nil {
|
||||
runtime.GC() // peak mem
|
||||
if err := g.mph.Build(); err != nil {
|
||||
return err
|
||||
}
|
||||
g.mph.Build()
|
||||
}
|
||||
runtime.GC() // peak mem
|
||||
if g.ac != nil {
|
||||
@@ -66,17 +64,23 @@ func (g *MphIndexMatcher) Build() error {
|
||||
|
||||
// Match implements IndexMatcher.Match.
|
||||
func (g *MphIndexMatcher) Match(input string) []uint32 {
|
||||
var result []uint32
|
||||
result := make([][]uint32, 0, 5)
|
||||
if g.mph != nil {
|
||||
result = g.mph.Match(input) // a new slice, returned without another copy
|
||||
if matches := g.mph.Match(input); len(matches) > 0 {
|
||||
result = append(result, matches)
|
||||
}
|
||||
}
|
||||
if g.ac != nil {
|
||||
result = append(result, g.ac.Match(input)...)
|
||||
if matches := g.ac.Match(input); len(matches) > 0 {
|
||||
result = append(result, matches)
|
||||
}
|
||||
}
|
||||
if g.regex != nil {
|
||||
result = append(result, g.regex.Match(input)...)
|
||||
if matches := g.regex.Match(input); len(matches) > 0 {
|
||||
result = append(result, matches)
|
||||
}
|
||||
}
|
||||
return result
|
||||
return CompositeMatches(result)
|
||||
}
|
||||
|
||||
// MatchAny implements IndexMatcher.MatchAny.
|
||||
|
||||
@@ -78,10 +78,6 @@ func TestMphIndexMatcher(t *testing.T) {
|
||||
Input: "example.com",
|
||||
Output: []uint32{10, 4},
|
||||
},
|
||||
{
|
||||
Input: "apis.org",
|
||||
Output: []uint32{2, 6},
|
||||
},
|
||||
}
|
||||
matcherGroup := NewMphIndexMatcher()
|
||||
for _, rule := range rules {
|
||||
@@ -91,13 +87,8 @@ func TestMphIndexMatcher(t *testing.T) {
|
||||
}
|
||||
matcherGroup.Build()
|
||||
for _, test := range cases {
|
||||
m := matcherGroup.Match(test.Input)
|
||||
if !reflect.DeepEqual(m, test.Output) {
|
||||
t.Error("unexpected output: ", m, " for test case ", test)
|
||||
}
|
||||
clear(m) // the caller owns the result, so this must not change the next one
|
||||
if m := matcherGroup.Match(test.Input); !reflect.DeepEqual(m, test.Output) {
|
||||
t.Error("unexpected output after clearing the previous one: ", m, " for test case ", test)
|
||||
t.Error("unexpected output: ", m, " for test case ", test)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,440 +1,198 @@
|
||||
package strmatcher
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"cmp"
|
||||
"encoding/binary"
|
||||
"errors"
|
||||
"math"
|
||||
"slices"
|
||||
"math/bits"
|
||||
"runtime"
|
||||
"sort"
|
||||
"strings"
|
||||
"unsafe"
|
||||
)
|
||||
|
||||
// Flags of a level1 slot, stored above the record offset.
|
||||
const (
|
||||
mphDomain = 1 << 31 // matches the pattern and its subdomains
|
||||
mphFull = 1 << 30 // matches the pattern only
|
||||
mphParent = 1 << 29 // matches subdomains only, from a pattern with a leading dot
|
||||
mphOffMask = mphParent - 1
|
||||
)
|
||||
// PrimeRK is the prime base used in Rabin-Karp algorithm.
|
||||
const PrimeRK = 16777619
|
||||
|
||||
// Kinds of an added pattern, indexes of mphKinds.
|
||||
const (
|
||||
mphKindFull = iota
|
||||
mphKindParent
|
||||
mphKindDomain
|
||||
)
|
||||
|
||||
// mphKinds are the slot flags in the order Match reports their values.
|
||||
var mphKinds = [...]uint32{mphFull, mphParent, mphDomain}
|
||||
|
||||
// mphMultipliers are odd multipliers for the suffix hash. Build moves to the next one if two patterns collide.
|
||||
var mphMultipliers = [...]uint64{0x9e3779b97f4a7c15, 0xc2b2ae3d27d4eb4f, 0x165667b19e3779f9, 0x27d4eb2f165667c5}
|
||||
|
||||
var (
|
||||
errMphCollision = errors.New("strmatcher: suffix hash collision in MphMatcherGroup")
|
||||
errMphBuilt = errors.New("strmatcher: MphMatcherGroup is already built")
|
||||
)
|
||||
|
||||
type mphEntry struct {
|
||||
off uint32 // pattern start in buf
|
||||
value uint32
|
||||
n uint32 // pattern length
|
||||
kind uint8
|
||||
// RollingHash calculates the rolling murmurHash of given string based on a provided suffix hash.
|
||||
func RollingHash(hash uint32, input string) uint32 {
|
||||
for i := len(input) - 1; i >= 0; i-- {
|
||||
hash = hash*PrimeRK + uint32(input[i])
|
||||
}
|
||||
return hash
|
||||
}
|
||||
|
||||
// MphMatcherGroup is an implementation of MatcherGroup for Full and Domain matchers.
|
||||
// Each distinct pattern is stored once as a record in arena: its length (255 means a uvarint length follows),
|
||||
// its bytes and, if the group holds more than one distinct value, its values. A minimal perfect hash table
|
||||
// built with hash, displace and compress (http://cmph.sourceforge.net/papers/esa09.pdf) maps a pattern to its
|
||||
// record. Patterns are hashed from the right, so one pass over the input hashes all its parent domains.
|
||||
type MphMatcherGroup struct {
|
||||
arena string
|
||||
level0 []uint16 // bucket -> seed
|
||||
level1 []uint32 // slot -> flags | record offset
|
||||
fp []uint8 // slot -> low byte of its pattern's hash, rejects most misses without reading arena
|
||||
n0, n1 uint32
|
||||
mul uint64 // multiplier of the suffix hash
|
||||
single uint32 // the only value if !multi
|
||||
multi bool
|
||||
// MemHash is the hash function used by go map, it utilizes available hardware instructions(behaves
|
||||
// as aeshash if aes instruction is available).
|
||||
// With different seed, each MemHash<seed> performs as distinct hash functions.
|
||||
func MemHash(seed uint32, input string) uint32 {
|
||||
return uint32(strhash(unsafe.Pointer(&input), uintptr(seed))) // nosemgrep
|
||||
}
|
||||
|
||||
buf []byte // build only, patterns in Add order
|
||||
entries []mphEntry
|
||||
const (
|
||||
mphMatchTypeCount = 2 // Full and Domain
|
||||
)
|
||||
|
||||
type mphRuleInfo struct {
|
||||
rollingHash uint32
|
||||
matchers [mphMatchTypeCount][]uint32
|
||||
}
|
||||
|
||||
// MphMatcherGroup is an implementation of MatcherGroup.
|
||||
// It implements Rabin-Karp algorithm and minimal perfect hash table for Full and Domain matcher.
|
||||
type MphMatcherGroup struct {
|
||||
rules []string // RuleIdx -> pattern string, index 0 reserved for failed lookup
|
||||
values [][]uint32 // RuleIdx -> registered matcher values for the pattern (Full Matcher takes precedence)
|
||||
level0 []uint32 // RollingHash & Mask -> seed for Memhash
|
||||
level0Mask uint32 // Mask restricting RollingHash to 0 ~ len(level0)
|
||||
level1 []uint32 // Memhash<seed> & Mask -> stored index for rules
|
||||
level1Mask uint32 // Mask for restricting Memhash<seed> to 0 ~ len(level1)
|
||||
ruleInfos *map[string]mphRuleInfo
|
||||
}
|
||||
|
||||
func NewMphMatcherGroup() *MphMatcherGroup {
|
||||
return new(MphMatcherGroup)
|
||||
return &MphMatcherGroup{
|
||||
rules: []string{""},
|
||||
values: [][]uint32{nil},
|
||||
level0: nil,
|
||||
level0Mask: 0,
|
||||
level1: nil,
|
||||
level1Mask: 0,
|
||||
ruleInfos: &map[string]mphRuleInfo{}, // Only used for building, destroyed after build complete
|
||||
}
|
||||
}
|
||||
|
||||
// AddFullMatcher implements MatcherGroupForFull.
|
||||
func (g *MphMatcherGroup) AddFullMatcher(matcher FullMatcher, value uint32) {
|
||||
g.add(matcher.Pattern(), mphKindFull, value)
|
||||
pattern := strings.ToLower(matcher.Pattern())
|
||||
g.addPattern(0, "", pattern, matcher.Type(), value)
|
||||
}
|
||||
|
||||
// AddDomainMatcher implements MatcherGroupForDomain.
|
||||
func (g *MphMatcherGroup) AddDomainMatcher(matcher DomainMatcher, value uint32) {
|
||||
g.add(matcher.Pattern(), mphKindDomain, value)
|
||||
pattern := strings.ToLower(matcher.Pattern())
|
||||
hash := g.addPattern(0, "", pattern, matcher.Type(), value) // For full domain match
|
||||
g.addPattern(hash, pattern, ".", matcher.Type(), value) // For partial domain match
|
||||
}
|
||||
|
||||
func (g *MphMatcherGroup) add(pattern string, kind uint8, value uint32) {
|
||||
if g.arena != "" {
|
||||
panic(errMphBuilt)
|
||||
}
|
||||
pattern = strings.ToLower(pattern)
|
||||
off := uint32(len(g.buf))
|
||||
g.buf = append(g.buf, pattern...)
|
||||
g.entries = append(g.entries, mphEntry{off: off, value: value, n: uint32(len(pattern)), kind: kind})
|
||||
if len(pattern) > 0 && pattern[0] == '.' {
|
||||
// ".x" has always matched "*.x" as well, so it also gets a parent-only record for "x"
|
||||
g.entries = append(g.entries, mphEntry{off: off + 1, value: value, n: uint32(len(pattern) - 1), kind: mphKindParent})
|
||||
func (g *MphMatcherGroup) addPattern(suffixHash uint32, suffixPattern string, pattern string, matcherType Type, value uint32) uint32 {
|
||||
fullPattern := pattern + suffixPattern
|
||||
info, found := (*g.ruleInfos)[fullPattern]
|
||||
if !found {
|
||||
info = mphRuleInfo{rollingHash: RollingHash(suffixHash, pattern)}
|
||||
g.rules = append(g.rules, fullPattern)
|
||||
g.values = append(g.values, nil)
|
||||
}
|
||||
info.matchers[matcherType] = append(info.matchers[matcherType], value)
|
||||
(*g.ruleInfos)[fullPattern] = info
|
||||
return info.rollingHash
|
||||
}
|
||||
|
||||
func (g *MphMatcherGroup) key(i uint32) []byte {
|
||||
e := &g.entries[i]
|
||||
return g.buf[e.off : e.off+e.n]
|
||||
}
|
||||
|
||||
// Build builds the hash table. It must be called once, after the last Add.
|
||||
// Build builds a minimal perfect hash table for insert rules.
|
||||
// Algorithm used: Hash, displace, and compress. See http://cmph.sourceforge.net/papers/esa09.pdf
|
||||
func (g *MphMatcherGroup) Build() error {
|
||||
if g.arena != "" {
|
||||
return errMphBuilt
|
||||
}
|
||||
if uint64(len(g.buf)) > math.MaxUint32 {
|
||||
return errors.New("too many rules for MphMatcherGroup")
|
||||
}
|
||||
recs := g.writeRecords()
|
||||
if len(g.arena) > mphOffMask {
|
||||
return errors.New("too many rules for MphMatcherGroup")
|
||||
}
|
||||
hashes := make([]uint64, len(recs))
|
||||
for _, mul := range mphMultipliers {
|
||||
for i, rec := range recs {
|
||||
hashes[i] = mphMix(mphHash(mul, g.recKey(rec)))
|
||||
}
|
||||
g.mul = mul
|
||||
if err := g.place(recs, hashes); err != errMphCollision {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return errMphCollision
|
||||
}
|
||||
ruleCount := len(*g.ruleInfos)
|
||||
g.level0 = make([]uint32, nextPow2(ruleCount/4))
|
||||
g.level0Mask = uint32(len(g.level0) - 1)
|
||||
g.level1 = make([]uint32, nextPow2(ruleCount))
|
||||
g.level1Mask = uint32(len(g.level1) - 1)
|
||||
|
||||
// writeRecords writes one record per distinct pattern to arena and returns flags | offset of each.
|
||||
func (g *MphMatcherGroup) writeRecords() []uint32 {
|
||||
g.multi = false
|
||||
if len(g.entries) > 0 {
|
||||
g.single = g.entries[0].value
|
||||
for _, e := range g.entries {
|
||||
if e.value != g.single {
|
||||
g.multi = true
|
||||
break
|
||||
}
|
||||
}
|
||||
// Create buckets based on all rule's rolling hash
|
||||
buckets := make([][]uint32, len(g.level0))
|
||||
for ruleIdx := 1; ruleIdx < len(g.rules); ruleIdx++ { // Traverse rules starting from index 1 (0 reserved for failed lookup)
|
||||
ruleInfo := (*g.ruleInfos)[g.rules[ruleIdx]]
|
||||
bucketIdx := ruleInfo.rollingHash & g.level0Mask
|
||||
buckets[bucketIdx] = append(buckets[bucketIdx], uint32(ruleIdx))
|
||||
g.values[ruleIdx] = append(ruleInfo.matchers[Full], ruleInfo.matchers[Domain]...) // nolint:gocritic
|
||||
}
|
||||
// Equal patterns become neighbours in Add order, so their values keep their priority
|
||||
order := make([]uint32, len(g.entries))
|
||||
for i := range order {
|
||||
order[i] = uint32(i)
|
||||
}
|
||||
slices.SortFunc(order, func(a, b uint32) int {
|
||||
return cmp.Or(bytes.Compare(g.key(a), g.key(b)), cmp.Compare(a, b))
|
||||
})
|
||||
g.ruleInfos = nil // Set ruleInfos nil to release memory
|
||||
runtime.GC() // peak mem
|
||||
|
||||
size := len(g.buf) + len(g.entries) + 2
|
||||
if g.multi {
|
||||
size += 3 * len(g.entries)
|
||||
// Sort buckets in descending order with respect to each bucket's size
|
||||
bucketIdxs := make([]int, len(buckets))
|
||||
for bucketIdx := range buckets {
|
||||
bucketIdxs[bucketIdx] = bucketIdx
|
||||
}
|
||||
arena := make([]byte, 0, size)
|
||||
recs := make([]uint32, 0, len(order))
|
||||
var vals [len(mphKinds)][]uint32
|
||||
for i := 0; i < len(order); {
|
||||
k := g.key(order[i])
|
||||
for t := range vals {
|
||||
vals[t] = vals[t][:0]
|
||||
}
|
||||
for ; i < len(order) && bytes.Equal(g.key(order[i]), k); i++ {
|
||||
e := &g.entries[order[i]]
|
||||
if !slices.Contains(vals[e.kind], e.value) {
|
||||
vals[e.kind] = append(vals[e.kind], e.value)
|
||||
}
|
||||
}
|
||||
rec := uint32(len(arena))
|
||||
if len(k) < 255 {
|
||||
arena = append(arena, byte(len(k)))
|
||||
} else {
|
||||
arena = binary.AppendUvarint(append(arena, 255), uint64(len(k)))
|
||||
}
|
||||
arena = append(arena, k...)
|
||||
for t, v := range vals {
|
||||
if len(v) == 0 {
|
||||
continue
|
||||
}
|
||||
rec |= mphKinds[t]
|
||||
if g.multi {
|
||||
arena = binary.AppendUvarint(arena, uint64(len(v)))
|
||||
for _, x := range v {
|
||||
arena = binary.AppendUvarint(arena, uint64(x))
|
||||
sort.Slice(bucketIdxs, func(i, j int) bool { return len(buckets[bucketIdxs[i]]) > len(buckets[bucketIdxs[j]]) })
|
||||
|
||||
// Exercise Hash, Displace, and Compress algorithm to construct minimal perfect hash table
|
||||
occupied := make([]bool, len(g.level1)) // Whether a second-level hash has been already used
|
||||
hashedBucket := make([]uint32, 0, 4) // Second-level hashes for each rule in a specific bucket
|
||||
for _, bucketIdx := range bucketIdxs {
|
||||
bucket := buckets[bucketIdx]
|
||||
hashedBucket = hashedBucket[:0]
|
||||
seed := uint32(0)
|
||||
for len(hashedBucket) != len(bucket) {
|
||||
for _, ruleIdx := range bucket {
|
||||
memHash := MemHash(seed, g.rules[ruleIdx]) & g.level1Mask
|
||||
if occupied[memHash] { // Collision occurred with this seed
|
||||
for _, hash := range hashedBucket { // Revert all values in this hashed bucket
|
||||
occupied[hash] = false
|
||||
g.level1[hash] = 0
|
||||
}
|
||||
hashedBucket = hashedBucket[:0]
|
||||
seed++ // Try next seed
|
||||
break
|
||||
}
|
||||
occupied[memHash] = true
|
||||
g.level1[memHash] = ruleIdx // The final value in the hash table
|
||||
hashedBucket = append(hashedBucket, memHash)
|
||||
}
|
||||
}
|
||||
recs = append(recs, rec)
|
||||
}
|
||||
// Lookups may point one byte past a pattern, and an empty group needs a record at offset 0 for empty slots
|
||||
arena = append(arena, 0)
|
||||
if len(recs) == 0 {
|
||||
arena = append(arena, 0)
|
||||
}
|
||||
g.buf, g.entries = nil, nil
|
||||
if cap(arena)-len(arena) > len(arena)/32 {
|
||||
arena = slices.Clone(arena)
|
||||
}
|
||||
g.arena = unsafe.String(unsafe.SliceData(arena), len(arena)) // arena is not written after this
|
||||
return recs
|
||||
}
|
||||
|
||||
// place fills level0, level1 and fp: records are bucketed by hash, and each bucket, largest first, gets
|
||||
// the first seed that puts all its records in free slots.
|
||||
func (g *MphMatcherGroup) place(recs []uint32, hashes []uint64) error {
|
||||
r := len(recs)
|
||||
n0, n1 := max(1, r/3), max(1, r+r/99)
|
||||
g.n0, g.n1 = uint32(n0), uint32(n1)
|
||||
g.level0 = make([]uint16, n0)
|
||||
g.level1 = make([]uint32, n1)
|
||||
g.fp = make([]uint8, n1)
|
||||
|
||||
start := make([]uint32, n0+1)
|
||||
for _, h := range hashes {
|
||||
start[g.bucket(h)+1]++
|
||||
}
|
||||
for b := range n0 {
|
||||
start[b+1] += start[b]
|
||||
}
|
||||
members := make([]uint32, r)
|
||||
fill := slices.Clone(start[:n0])
|
||||
for i, h := range hashes {
|
||||
b := g.bucket(h)
|
||||
members[fill[b]] = uint32(i)
|
||||
fill[b]++
|
||||
}
|
||||
fill = nil
|
||||
buckets := make([]uint32, n0)
|
||||
for b := range buckets {
|
||||
buckets[b] = uint32(b)
|
||||
}
|
||||
slices.SortStableFunc(buckets, func(a, b uint32) int {
|
||||
return cmp.Compare(start[b+1]-start[b], start[a+1]-start[a])
|
||||
})
|
||||
|
||||
occupied := make([]uint64, (n1+63)/64)
|
||||
var slots []uint32
|
||||
next:
|
||||
for _, b := range buckets {
|
||||
m := members[start[b]:start[b+1]]
|
||||
if len(m) == 0 {
|
||||
break
|
||||
}
|
||||
for i := range m {
|
||||
for j := range i {
|
||||
if hashes[m[i]] == hashes[m[j]] {
|
||||
return errMphCollision // no seed can separate them
|
||||
}
|
||||
}
|
||||
}
|
||||
search:
|
||||
for seed := range math.MaxUint16 + 1 {
|
||||
slots = slots[:0]
|
||||
for _, ri := range m {
|
||||
s := g.slot(hashes[ri], uint16(seed))
|
||||
if occupied[s/64]&(1<<(s%64)) != 0 || slices.Contains(slots, s) {
|
||||
continue search
|
||||
}
|
||||
slots = append(slots, s)
|
||||
}
|
||||
for k, ri := range m {
|
||||
s := slots[k]
|
||||
occupied[s/64] |= 1 << (s % 64)
|
||||
g.level1[s] = recs[ri]
|
||||
g.fp[s] = uint8(hashes[ri])
|
||||
}
|
||||
g.level0[b] = uint16(seed)
|
||||
continue next
|
||||
}
|
||||
return errors.New("strmatcher: no seed found for a bucket in MphMatcherGroup")
|
||||
g.level0[bucketIdx] = seed // Displacement value for this bucket
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// mphHash is the suffix hash of s, taken from the right: the hash of s[i:] is the state after reading s[i].
|
||||
func mphHash(mul uint64, s string) uint64 {
|
||||
h := uint64(0)
|
||||
for i := len(s) - 1; i >= 0; i-- {
|
||||
h = h*mul + uint64(s[i])
|
||||
}
|
||||
return h
|
||||
}
|
||||
|
||||
// mphMix spreads the weak low bits of a suffix hash.
|
||||
func mphMix(h uint64) uint64 {
|
||||
h ^= h >> 32
|
||||
h *= 0xd6e8feb86659fd93
|
||||
return h ^ h>>32
|
||||
}
|
||||
|
||||
func (g *MphMatcherGroup) bucket(f uint64) uint32 {
|
||||
return uint32(((f >> 32) * uint64(g.n0)) >> 32)
|
||||
}
|
||||
|
||||
func (g *MphMatcherGroup) slot(f uint64, seed uint16) uint32 {
|
||||
x := ((f ^ uint64(seed)*0x9e3779b97f4a7c15) * 0xc4ceb9fe1a85ec53) >> 32
|
||||
return uint32((x * uint64(g.n1)) >> 32)
|
||||
}
|
||||
|
||||
func (g *MphMatcherGroup) uvarint(p uint32) (x, next uint32) {
|
||||
for shift := 0; ; shift += 7 {
|
||||
c := g.arena[p]
|
||||
p++
|
||||
x |= uint32(c&0x7f) << shift
|
||||
if c < 0x80 {
|
||||
return x, p
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// recSpan returns where the pattern of the record at off starts and how long it is.
|
||||
func (g *MphMatcherGroup) recSpan(off uint32) (p, n uint32) {
|
||||
n, p = uint32(g.arena[off]), off+1
|
||||
if n == 255 {
|
||||
n, p = g.uvarint(p)
|
||||
}
|
||||
return p, n
|
||||
}
|
||||
|
||||
func (g *MphMatcherGroup) recKey(rec uint32) string {
|
||||
p, n := g.recSpan(rec & mphOffMask)
|
||||
return g.arena[p : p+n]
|
||||
}
|
||||
|
||||
// lookup returns the level1 entry of s, or 0 if s is not a pattern. h is the suffix hash of s.
|
||||
func (g *MphMatcherGroup) lookup(h uint64, s string) uint32 {
|
||||
f := mphMix(h)
|
||||
// bucket < n0 == len(level0) and slot < n1 == len(level1) == len(fp), skip the bounds checks
|
||||
seed := *(*uint16)(unsafe.Add(unsafe.Pointer(unsafe.SliceData(g.level0)), uintptr(g.bucket(f))*2))
|
||||
slot := uintptr(g.slot(f, seed))
|
||||
if *(*uint8)(unsafe.Add(unsafe.Pointer(unsafe.SliceData(g.fp)), slot)) != uint8(f) {
|
||||
return 0
|
||||
}
|
||||
e := *(*uint32)(unsafe.Add(unsafe.Pointer(unsafe.SliceData(g.level1)), slot*4))
|
||||
if len(s) < 255 {
|
||||
// A record whose length byte is len(s) has len(s) pattern bytes after it
|
||||
p := unsafe.Add(unsafe.Pointer(unsafe.StringData(g.arena)), e&mphOffMask)
|
||||
if int(*(*byte)(p)) == len(s) && unsafe.String((*byte)(unsafe.Add(p, 1)), len(s)) == s {
|
||||
return e
|
||||
}
|
||||
return 0
|
||||
}
|
||||
if g.recKey(e) == s {
|
||||
return e
|
||||
// Lookup searches for input in minimal perfect hash table and returns its index. 0 indicates not found.
|
||||
func (g *MphMatcherGroup) Lookup(rollingHash uint32, input string) uint32 {
|
||||
i0 := rollingHash & g.level0Mask
|
||||
seed := g.level0[i0]
|
||||
i1 := MemHash(seed, input) & g.level1Mask
|
||||
if n := g.level1[i1]; g.rules[n] == input {
|
||||
return n
|
||||
}
|
||||
return 0
|
||||
}
|
||||
|
||||
// appendValues appends the values of record e for the flags in want, in mphKinds order.
|
||||
func (g *MphMatcherGroup) appendValues(dst []uint32, e, want uint32) []uint32 {
|
||||
if !g.multi {
|
||||
for _, flag := range mphKinds {
|
||||
if e&want&flag != 0 {
|
||||
dst = append(dst, g.single)
|
||||
}
|
||||
}
|
||||
return dst
|
||||
}
|
||||
if e&want == 0 {
|
||||
return dst
|
||||
}
|
||||
p, n := g.recSpan(e & mphOffMask)
|
||||
p += n
|
||||
for _, flag := range mphKinds {
|
||||
if e&flag == 0 {
|
||||
continue
|
||||
}
|
||||
var count, v uint32
|
||||
for count, p = g.uvarint(p); count > 0; count-- {
|
||||
v, p = g.uvarint(p)
|
||||
if want&flag != 0 {
|
||||
dst = append(dst, v)
|
||||
}
|
||||
}
|
||||
}
|
||||
return dst
|
||||
}
|
||||
|
||||
// Match implements MatcherGroup.Match. Values of an exact match come first (Full, then Domain), then those of
|
||||
// the parent domains, nearest first.
|
||||
// Match implements MatcherGroup.Match.
|
||||
func (g *MphMatcherGroup) Match(input string) []uint32 {
|
||||
var stack [8]uint32
|
||||
parents := stack[:0] // TLD side first
|
||||
h, mul := uint64(0), g.mul
|
||||
matches := make([][]uint32, 0, 5)
|
||||
hash := uint32(0)
|
||||
for i := len(input) - 1; i >= 0; i-- {
|
||||
hash = hash*PrimeRK + uint32(input[i])
|
||||
if input[i] == '.' {
|
||||
if e := g.lookup(h, input[i+1:]); e&(mphDomain|mphParent) != 0 {
|
||||
parents = append(parents, e)
|
||||
if mphIdx := g.Lookup(hash, input[i:]); mphIdx != 0 {
|
||||
matches = append(matches, g.values[mphIdx])
|
||||
}
|
||||
}
|
||||
h = h*mul + uint64(input[i])
|
||||
}
|
||||
exact := g.lookup(h, input)
|
||||
if exact&(mphFull|mphDomain) == 0 && len(parents) == 0 {
|
||||
return nil
|
||||
if mphIdx := g.Lookup(hash, input); mphIdx != 0 {
|
||||
matches = append(matches, g.values[mphIdx])
|
||||
}
|
||||
result := g.appendValues(make([]uint32, 0, len(parents)+1), exact, mphFull|mphDomain)
|
||||
for k := len(parents) - 1; k >= 0; k-- {
|
||||
result = g.appendValues(result, parents[k], mphParent|mphDomain)
|
||||
}
|
||||
return result
|
||||
return CompositeMatchesReverse(matches)
|
||||
}
|
||||
|
||||
// MatchAny implements MatcherGroup.MatchAny.
|
||||
func (g *MphMatcherGroup) MatchAny(input string) bool {
|
||||
h, mul := uint64(0), g.mul
|
||||
for i := len(input) - 1; i >= 0; i-- {
|
||||
if input[i] == '.' && g.lookup(h, input[i+1:])&(mphDomain|mphParent) != 0 {
|
||||
return true
|
||||
}
|
||||
h = h*mul + uint64(input[i])
|
||||
}
|
||||
return g.lookup(h, input)&(mphFull|mphDomain) != 0
|
||||
}
|
||||
|
||||
// mphSuffix is the suffix hash of input[off:], a parent domain of the input.
|
||||
type mphSuffix struct {
|
||||
h uint64
|
||||
off int
|
||||
}
|
||||
|
||||
// mphSuffixes appends the suffix hashes of the parent domains of input to dst, TLD side first, and returns them
|
||||
// with the hash of input itself: what MatchAny computes, computed once for several groups.
|
||||
func mphSuffixes(dst []mphSuffix, mul uint64, input string) ([]mphSuffix, uint64) {
|
||||
h := uint64(0)
|
||||
hash := uint32(0)
|
||||
for i := len(input) - 1; i >= 0; i-- {
|
||||
hash = hash*PrimeRK + uint32(input[i])
|
||||
if input[i] == '.' {
|
||||
dst = append(dst, mphSuffix{h, i + 1})
|
||||
if g.Lookup(hash, input[i:]) != 0 {
|
||||
return true
|
||||
}
|
||||
}
|
||||
h = h*mul + uint64(input[i])
|
||||
}
|
||||
return dst, h
|
||||
return g.Lookup(hash, input) != 0
|
||||
}
|
||||
|
||||
// matchAnyHashed is MatchAny with parents and h from mphSuffixes(_, mul, input).
|
||||
func (g *MphMatcherGroup) matchAnyHashed(input string, parents []mphSuffix, h, mul uint64) bool {
|
||||
if g.mul != mul {
|
||||
return g.MatchAny(input) // built with a later multiplier after a collision
|
||||
func nextPow2(v int) int {
|
||||
if v <= 1 {
|
||||
return 1
|
||||
}
|
||||
for _, p := range parents {
|
||||
if g.lookup(p.h, input[p.off:])&(mphDomain|mphParent) != 0 {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return g.lookup(h, input)&(mphFull|mphDomain) != 0
|
||||
const MaxUInt = ^uint(0)
|
||||
n := (MaxUInt >> bits.LeadingZeros(uint(v))) + 1
|
||||
return int(n)
|
||||
}
|
||||
|
||||
//go:noescape
|
||||
//go:linkname strhash runtime.strhash
|
||||
func strhash(p unsafe.Pointer, h uintptr) uintptr
|
||||
|
||||
@@ -1,108 +0,0 @@
|
||||
package strmatcher
|
||||
|
||||
import (
|
||||
"slices"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestMphMatcherGroupHashCollision(t *testing.T) {
|
||||
saved := mphMultipliers
|
||||
defer func() { mphMultipliers = saved }()
|
||||
|
||||
mphMultipliers[0] = 1 // anagrams collide
|
||||
g := NewMphMatcherGroup()
|
||||
g.AddFullMatcher(FullMatcher("ab.com"), 1)
|
||||
g.AddDomainMatcher(DomainMatcher("ba.com"), 2)
|
||||
g.AddDomainMatcher(DomainMatcher("com"), 3)
|
||||
if err := g.Build(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if g.mul != saved[1] {
|
||||
t.Errorf("multiplier %#x, want the second one %#x", g.mul, saved[1])
|
||||
}
|
||||
for input, want := range map[string][]uint32{"ab.com": {1, 3}, "x.ba.com": {2, 3}, "x.ab.com": {3}, "ba.com": {2, 3}} {
|
||||
if m := g.Match(input); !slices.Equal(m, want) {
|
||||
t.Errorf("Match(%q) = %v, want %v", input, m, want)
|
||||
}
|
||||
}
|
||||
|
||||
// Thue-Morse strings of 2048 bytes and their complements collide for every odd multiplier
|
||||
mphMultipliers = saved
|
||||
a, b := make([]byte, 2048), make([]byte, 2048)
|
||||
for i := range a {
|
||||
a[i], b[i] = "ab"[bitsOnes(i)%2], "ba"[bitsOnes(i)%2]
|
||||
}
|
||||
g = NewMphMatcherGroup()
|
||||
g.AddFullMatcher(FullMatcher(a), 1)
|
||||
g.AddFullMatcher(FullMatcher(b), 1)
|
||||
if err := g.Build(); err != errMphCollision {
|
||||
t.Errorf("Build() = %v, want %v", err, errMphCollision)
|
||||
}
|
||||
}
|
||||
|
||||
func bitsOnes(i int) int {
|
||||
n := 0
|
||||
for ; i > 0; i &= i - 1 {
|
||||
n++
|
||||
}
|
||||
return n
|
||||
}
|
||||
|
||||
func TestMphValueMatcherCombiner(t *testing.T) {
|
||||
build := func(matchers ...Matcher) *MphValueMatcher {
|
||||
m := NewMphValueMatcher()
|
||||
for _, x := range matchers {
|
||||
m.Add(x, 0)
|
||||
}
|
||||
if err := m.Build(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
return m
|
||||
}
|
||||
regex, err := Regex.New(`^a\d+\.net$`)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
saved := mphMultipliers
|
||||
t.Cleanup(func() { mphMultipliers = saved })
|
||||
mphMultipliers[0] = 1 // anagrams collide, so this one falls back to its own hash pass
|
||||
collided := build(FullMatcher("ab.com"), DomainMatcher("ba.com"))
|
||||
mphMultipliers = saved
|
||||
if collided.mph.mul == mphMultipliers[0] {
|
||||
t.Fatal("collided matcher uses the first multiplier")
|
||||
}
|
||||
matchers := []*MphValueMatcher{
|
||||
build(DomainMatcher("example.com"), FullMatcher("full.org"), DomainMatcher(".dot.io")),
|
||||
collided,
|
||||
build(regex, SubstrMatcher("keyword")),
|
||||
build(),
|
||||
build(DomainMatcher("com"), DomainMatcher("a.b.c.d.e.f.g.h.i.j.k.l.m.n.o.p.q.r.s")),
|
||||
}
|
||||
var s MphValueMatcherCombiner
|
||||
for i, m := range matchers {
|
||||
s.Add(m, uint32(10+i))
|
||||
}
|
||||
inputs := []string{
|
||||
"", ".", "..", "com", "example.com", "www.example.com", "xexample.com", "example.com.", "full.org", "x.full.org",
|
||||
"dot.io", "x.dot.io", ".dot.io", "ab.com", "x.ab.com", "ba.com", "x.ba.com", "a12.net", "a12.net.x", "my-keyword.org",
|
||||
"a.b.c.d.e.f.g.h.i.j.k.l.m.n.o.p.q.r.s", "0.a.b.c.d.e.f.g.h.i.j.k.l.m.n.o.p.q.r.s", "b.c.d.e.f.g.h.i.j.k.l.m.n.o.p.q.r.s",
|
||||
"x.y.z.1.2.3.4.5.6.7.8.9.10.11.12.13.14.15.16.17.ab.com", "x.y.z.1.2.3.4.5.6.7.8.9.10.11.12.13.14.15.16.17.org",
|
||||
}
|
||||
for _, input := range inputs {
|
||||
var want []uint32
|
||||
for i, m := range matchers {
|
||||
if m.MatchAny(input) {
|
||||
want = append(want, uint32(10+i))
|
||||
}
|
||||
}
|
||||
if got := s.Match(input); !slices.Equal(got, want) {
|
||||
t.Errorf("Match(%q) = %v, want %v", input, got, want)
|
||||
}
|
||||
if got := s.MatchAny(input); got != (len(want) > 0) {
|
||||
t.Errorf("MatchAny(%q) = %v", input, got)
|
||||
}
|
||||
}
|
||||
if n := testing.AllocsPerRun(100, func() { s.MatchAny("www.a.b.c.example.org") }); n != 0 {
|
||||
t.Errorf("MatchAny allocates %v times", n)
|
||||
}
|
||||
}
|
||||
@@ -1,10 +1,7 @@
|
||||
package strmatcher_test
|
||||
|
||||
import (
|
||||
"math/rand"
|
||||
"reflect"
|
||||
"slices"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/xtls/xray-core/common"
|
||||
@@ -279,142 +276,3 @@ func TestEmptyMphMatcherGroup(t *testing.T) {
|
||||
t.Error("Expect [], but ", r)
|
||||
}
|
||||
}
|
||||
|
||||
func TestMphMatcherGroupRandom(t *testing.T) {
|
||||
inputs := []string{""} // All strings over "ab." up to 7 bytes
|
||||
for i := 0; len(inputs[i]) < 7; i++ {
|
||||
for _, c := range []string{"a", "b", "."} {
|
||||
inputs = append(inputs, inputs[i]+c)
|
||||
}
|
||||
}
|
||||
for seed := int64(0); seed < 300; seed++ {
|
||||
r := rand.New(rand.NewSource(seed))
|
||||
g := NewMphMatcherGroup()
|
||||
full, domain := map[string][]uint32{}, map[string][]uint32{} // Stored pattern -> values
|
||||
for value := uint32(r.Intn(200)); value > 0; value-- {
|
||||
pattern := make([]byte, r.Intn(8))
|
||||
for i := range pattern {
|
||||
pattern[i] = "ab."[r.Intn(3)]
|
||||
}
|
||||
if p := string(pattern); r.Intn(2) == 0 {
|
||||
g.AddFullMatcher(FullMatcher(p), value)
|
||||
full[p] = append(full[p], value)
|
||||
} else {
|
||||
g.AddDomainMatcher(DomainMatcher(p), value)
|
||||
domain[p] = append(domain[p], value)
|
||||
domain["."+p] = append(domain["."+p], value)
|
||||
}
|
||||
}
|
||||
common.Must(g.Build())
|
||||
for _, input := range inputs {
|
||||
keys := []string{input} // Whole input first, then "." suffixes from longest to shortest
|
||||
for i := range len(input) {
|
||||
if input[i] == '.' {
|
||||
keys = append(keys, input[i:])
|
||||
}
|
||||
}
|
||||
var want []uint32
|
||||
for _, k := range keys {
|
||||
want = append(append(want, full[k]...), domain[k]...)
|
||||
}
|
||||
// Compared as sets: Match reports a value once per matching pattern, and orders them differently
|
||||
// from want for patterns and inputs with a leading dot
|
||||
m := g.Match(input)
|
||||
if !slices.Equal(sortedSet(m), sortedSet(want)) {
|
||||
t.Fatalf("seed %d: Match(%q) = %v, want %v", seed, input, m, want)
|
||||
}
|
||||
if m := g.MatchAny(input); m != (len(want) > 0) {
|
||||
t.Fatalf("seed %d: MatchAny(%q) = %v", seed, input, m)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestMphMatcherGroupAppend(t *testing.T) {
|
||||
g := NewMphMatcherGroup()
|
||||
g.AddFullMatcher(FullMatcher("a.com"), 1)
|
||||
g.AddFullMatcher(FullMatcher("b.com"), 2)
|
||||
g.Build()
|
||||
if m := append(g.Match("a.com"), 3); !slices.Equal(m, []uint32{1, 3}) {
|
||||
t.Error("expect [1 3], but ", m)
|
||||
}
|
||||
if m := g.Match("b.com"); !slices.Equal(m, []uint32{2}) {
|
||||
t.Error("expect [2], but ", m)
|
||||
}
|
||||
}
|
||||
|
||||
func sortedSet(v []uint32) []uint32 {
|
||||
v = slices.Clone(v)
|
||||
slices.Sort(v)
|
||||
return slices.Compact(v)
|
||||
}
|
||||
|
||||
func TestMphMatcherGroupLongPattern(t *testing.T) {
|
||||
long := strings.Repeat("a", 300) + ".com"
|
||||
for _, values := range [][4]uint32{{1, 2, 3, 4}, {7, 7, 7, 7}} {
|
||||
g := NewMphMatcherGroup()
|
||||
g.AddDomainMatcher(DomainMatcher(long), values[0])
|
||||
g.AddFullMatcher(FullMatcher("x."+long), values[1])
|
||||
g.AddFullMatcher(FullMatcher(long[:255]), values[2]) // the shortest pattern stored with a long length
|
||||
g.AddFullMatcher(FullMatcher(long[:254]), values[3])
|
||||
common.Must(g.Build())
|
||||
cases := []struct {
|
||||
input string
|
||||
want []uint32
|
||||
}{
|
||||
{long, []uint32{values[0]}},
|
||||
{"www." + long, []uint32{values[0]}},
|
||||
{"x." + long, []uint32{values[1], values[0]}},
|
||||
{long[1:], nil},
|
||||
{"a" + long, nil},
|
||||
{long[:255], []uint32{values[2]}},
|
||||
{long[:254], []uint32{values[3]}},
|
||||
{long[:256], nil},
|
||||
{long[:253], nil},
|
||||
}
|
||||
for _, c := range cases {
|
||||
if m := g.Match(c.input); !slices.Equal(m, c.want) {
|
||||
t.Errorf("Match(%d bytes) = %v, want %v", len(c.input), m, c.want)
|
||||
}
|
||||
if m := g.MatchAny(c.input); m != (c.want != nil) {
|
||||
t.Errorf("MatchAny(%d bytes) = %v", len(c.input), m)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// A pattern longer than 65535 bytes builds and matches: a record's length is a uvarint,
|
||||
// so the only cap was the build-time length field, now widened to uint32.
|
||||
huge := strings.Repeat("a", 70000)
|
||||
g := NewMphMatcherGroup()
|
||||
g.AddFullMatcher(FullMatcher(strings.Repeat("a", 65535)), 1)
|
||||
g.AddDomainMatcher(DomainMatcher(huge+".com"), 2)
|
||||
g.AddFullMatcher(FullMatcher("a.com"), 3)
|
||||
common.Must(g.Build())
|
||||
if !g.MatchAny(strings.Repeat("a", 65535)) || g.MatchAny(strings.Repeat("a", 65534)) {
|
||||
t.Error("wrong answer for a 65535-byte pattern")
|
||||
}
|
||||
if m := g.Match(huge + ".com"); !slices.Equal(m, []uint32{2}) {
|
||||
t.Errorf("Match(%d-byte input) = %v, want [2]", len(huge)+4, m)
|
||||
}
|
||||
if m := g.Match("x." + huge + ".com"); !slices.Equal(m, []uint32{2}) {
|
||||
t.Errorf("Match(subdomain of a %d-byte pattern) = %v, want [2]", len(huge)+4, m)
|
||||
}
|
||||
if g.MatchAny(huge) { // the 70000-byte label on its own is not a rule
|
||||
t.Error("unexpected match for the bare 70000-byte label")
|
||||
}
|
||||
}
|
||||
|
||||
func TestMphMatcherGroupBuildOnce(t *testing.T) {
|
||||
g := NewMphMatcherGroup()
|
||||
g.AddFullMatcher(FullMatcher("a.com"), 1)
|
||||
common.Must(g.Build())
|
||||
if err := g.Build(); err == nil || !g.MatchAny("a.com") {
|
||||
t.Errorf("second Build() = %v, MatchAny(a.com) = %v", err, g.MatchAny("a.com"))
|
||||
}
|
||||
defer func() {
|
||||
if recover() == nil {
|
||||
t.Error("Add after Build did not panic")
|
||||
}
|
||||
}()
|
||||
g.AddDomainMatcher(DomainMatcher("b.com"), 2)
|
||||
}
|
||||
|
||||
@@ -2,12 +2,9 @@ package strmatcher
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"math/bits"
|
||||
"regexp"
|
||||
"regexp/syntax"
|
||||
"slices"
|
||||
"strings"
|
||||
"unicode"
|
||||
"unicode/utf8"
|
||||
|
||||
"golang.org/x/net/idna"
|
||||
@@ -76,274 +73,7 @@ func (m SubstrMatcher) Match(s string) bool {
|
||||
|
||||
// RegexMatcher is an implementation of Matcher.
|
||||
type RegexMatcher struct {
|
||||
pattern *regexp.Regexp
|
||||
literals []string // every match contains all of them, longest first
|
||||
tail []byteSet // tail[i] holds the bytes a matching input can have i bytes before its end
|
||||
rest *byteSet // the bytes it can have further before, nil if any
|
||||
}
|
||||
|
||||
func newRegexMatcher(pattern string) (Matcher, error) {
|
||||
regex, err := regexp.Compile(pattern)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
m := &RegexMatcher{pattern: regex}
|
||||
if re, err := syntax.Parse(pattern, syntax.Perl); err == nil { // same flags as regexp.Compile
|
||||
m.literals = requiredLiterals(re, nil)
|
||||
slices.SortStableFunc(m.literals, func(a, b string) int { return len(b) - len(a) })
|
||||
m.tail, m.rest = tailGuard(re)
|
||||
}
|
||||
return m, nil
|
||||
}
|
||||
|
||||
// byteSet is a set of bytes. The bytes >= 0x80 share one bit with 0x7f.
|
||||
type byteSet [4]uint32
|
||||
|
||||
func (s *byteSet) add(c byte) { c = min(c, 0x7f); s[c>>5] |= 1 << (c & 31) }
|
||||
func (s *byteSet) has(c byte) bool { c = min(c, 0x7f); return s[c>>5]&(1<<(c&31)) != 0 }
|
||||
func (s *byteSet) or(t *byteSet) {
|
||||
for i := range s {
|
||||
s[i] |= t[i]
|
||||
}
|
||||
}
|
||||
|
||||
var allBytes = byteSet{^uint32(0), ^uint32(0), ^uint32(0), ^uint32(0)}
|
||||
|
||||
// tailLen is how many positions before the end of the input tailGuard tells apart.
|
||||
const tailLen = 8
|
||||
|
||||
// tailBudget caps how many repetition steps tailGuard walks. Only nested repeats can make the
|
||||
// walk explode, so only they are charged: a flat pattern, however long, is walked once and keeps
|
||||
// its guard.
|
||||
const tailBudget = 100000
|
||||
|
||||
// tailWalk is a set of positions in the input, counted in bytes before its end.
|
||||
type tailWalk struct {
|
||||
at uint32 // bit i: exactly i bytes before the end, for i < tailLen
|
||||
far bool // tailLen or more bytes before the end
|
||||
free bool // not tied to the end of the input yet
|
||||
}
|
||||
|
||||
func (w tailWalk) union(v tailWalk) tailWalk {
|
||||
return tailWalk{w.at | v.at, w.far || v.far, w.free || v.free}
|
||||
}
|
||||
|
||||
type tailBuilder struct {
|
||||
tail [tailLen]byteSet
|
||||
rest byteSet
|
||||
void bool
|
||||
work int
|
||||
}
|
||||
|
||||
// tailGuard walks re backwards from the end of the input and collects the bytes an input
|
||||
// matching re can have at each position before its end. It returns nil, nil when a branch
|
||||
// of re does not end with $ or when nested repeats push the walk past tailBudget.
|
||||
func tailGuard(re *syntax.Regexp) ([]byteSet, *byteSet) {
|
||||
var b tailBuilder
|
||||
w := b.walk(re, tailWalk{free: true})
|
||||
b.stop(w)
|
||||
if b.void {
|
||||
return nil, nil
|
||||
}
|
||||
if w.at != 0 { // a match can start here, so any bytes can come before
|
||||
for i := bits.TrailingZeros32(w.at); i < tailLen; i++ {
|
||||
b.tail[i] = allBytes
|
||||
}
|
||||
}
|
||||
if w.at != 0 || w.far {
|
||||
b.rest = allBytes
|
||||
}
|
||||
n := tailLen
|
||||
for n > 0 && b.tail[n-1] == b.rest {
|
||||
n--
|
||||
}
|
||||
var tail []byteSet
|
||||
if n > 0 {
|
||||
tail = slices.Clone(b.tail[:n])
|
||||
}
|
||||
if b.rest != allBytes {
|
||||
rest := b.rest
|
||||
return tail, &rest
|
||||
}
|
||||
return tail, nil
|
||||
}
|
||||
|
||||
// stop ends the paths of w. One that never met $ lets its match be followed by anything.
|
||||
func (b *tailBuilder) stop(w tailWalk) {
|
||||
if w.free {
|
||||
b.void = true
|
||||
}
|
||||
}
|
||||
|
||||
func (b *tailBuilder) walk(re *syntax.Regexp, w tailWalk) tailWalk {
|
||||
if w == (tailWalk{}) || b.void {
|
||||
return w
|
||||
}
|
||||
switch re.Op {
|
||||
case syntax.OpNoMatch:
|
||||
return tailWalk{}
|
||||
case syntax.OpLiteral:
|
||||
for i := len(re.Rune) - 1; i >= 0; i-- {
|
||||
var set byteSet
|
||||
set.add(byte(min(re.Rune[i], utf8.RuneSelf)))
|
||||
if re.Flags&syntax.FoldCase != 0 {
|
||||
for f := unicode.SimpleFold(re.Rune[i]); f != re.Rune[i]; f = unicode.SimpleFold(f) {
|
||||
set.add(byte(min(f, utf8.RuneSelf)))
|
||||
}
|
||||
}
|
||||
w = b.step(w, &set)
|
||||
}
|
||||
return w
|
||||
case syntax.OpCharClass:
|
||||
var set byteSet
|
||||
for i := 0; i+1 < len(re.Rune); i += 2 {
|
||||
for r := min(re.Rune[i], utf8.RuneSelf); r <= min(re.Rune[i+1], utf8.RuneSelf); r++ {
|
||||
set.add(byte(r))
|
||||
}
|
||||
}
|
||||
return b.step(w, &set)
|
||||
case syntax.OpAnyChar, syntax.OpAnyCharNotNL: // a domain has no \n to reject
|
||||
return b.step(w, &allBytes)
|
||||
case syntax.OpBeginText: // nothing comes before
|
||||
b.stop(w)
|
||||
return tailWalk{}
|
||||
case syntax.OpEndText:
|
||||
out := tailWalk{at: w.at & 1}
|
||||
if w.free {
|
||||
out.at = 1
|
||||
}
|
||||
return out
|
||||
case syntax.OpCapture:
|
||||
return b.walk(re.Sub[0], w)
|
||||
case syntax.OpConcat:
|
||||
for i := len(re.Sub) - 1; i >= 0; i-- {
|
||||
w = b.walk(re.Sub[i], w)
|
||||
}
|
||||
return w
|
||||
case syntax.OpAlternate:
|
||||
var out tailWalk
|
||||
for _, sub := range re.Sub {
|
||||
out = out.union(b.walk(sub, w))
|
||||
}
|
||||
return out
|
||||
case syntax.OpQuest:
|
||||
return b.repeat(re.Sub[0], w, 1)
|
||||
case syntax.OpStar:
|
||||
return b.repeat(re.Sub[0], w, -1)
|
||||
case syntax.OpPlus:
|
||||
return b.repeat(re.Sub[0], b.walk(re.Sub[0], w), -1)
|
||||
case syntax.OpRepeat:
|
||||
for i := 0; i < re.Min; i++ {
|
||||
if b.charge() {
|
||||
return w
|
||||
}
|
||||
w = b.walk(re.Sub[0], w)
|
||||
}
|
||||
if re.Max < 0 {
|
||||
return b.repeat(re.Sub[0], w, -1)
|
||||
}
|
||||
return b.repeat(re.Sub[0], w, re.Max-re.Min)
|
||||
}
|
||||
return w // empty match, line and word boundaries: no constraint
|
||||
}
|
||||
|
||||
// charge counts one repetition step and reports whether the walk has run out of budget. Only
|
||||
// repeats re-walk their body, so charging them alone bounds the blow-up of nested repeats while
|
||||
// leaving a single linear pass, of any length, free.
|
||||
func (b *tailBuilder) charge() bool {
|
||||
b.work++
|
||||
if b.work > tailBudget {
|
||||
b.void = true
|
||||
}
|
||||
return b.void
|
||||
}
|
||||
|
||||
// repeat walks back over up to n more repetitions of re, any number if n < 0.
|
||||
func (b *tailBuilder) repeat(re *syntax.Regexp, w tailWalk, n int) tailWalk {
|
||||
for ; n != 0; n-- {
|
||||
if b.charge() {
|
||||
return w
|
||||
}
|
||||
next := w.union(b.walk(re, w))
|
||||
if next == w {
|
||||
break
|
||||
}
|
||||
w = next
|
||||
}
|
||||
return w
|
||||
}
|
||||
|
||||
// step walks back over one character whose last byte is in set. A character that can be
|
||||
// non-ASCII can take up to 4 bytes, all >= 0x80; regexp matches an invalid byte as U+FFFD.
|
||||
func (b *tailBuilder) step(w tailWalk, set *byteSet) tailWalk {
|
||||
out := tailWalk{far: w.far, free: w.free}
|
||||
if w.far {
|
||||
b.rest.or(set)
|
||||
}
|
||||
width := 1
|
||||
if set.has(0x80) {
|
||||
width = utf8.UTFMax
|
||||
}
|
||||
for i := 0; i < tailLen; i++ {
|
||||
if w.at&(1<<i) == 0 {
|
||||
continue
|
||||
}
|
||||
b.tail[i].or(set)
|
||||
for n := 1; n <= width; n++ {
|
||||
if j := i + n; j < tailLen {
|
||||
out.at |= 1 << j
|
||||
if n < width {
|
||||
b.tail[j].add(0x80)
|
||||
}
|
||||
} else {
|
||||
out.far = true
|
||||
if n < width {
|
||||
b.rest.add(0x80)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
// mayMatch reports whether s passes the tail guard.
|
||||
func (m *RegexMatcher) mayMatch(s string) bool {
|
||||
n := len(s)
|
||||
if m.rest == nil {
|
||||
n = min(n, len(m.tail))
|
||||
}
|
||||
for i := 0; i < n; i++ {
|
||||
set := m.rest
|
||||
if i < len(m.tail) {
|
||||
set = &m.tail[i]
|
||||
}
|
||||
if !set.has(s[len(s)-1-i]) {
|
||||
return false
|
||||
}
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
// requiredLiterals appends to dst the case-sensitive strings that every match of re contains.
|
||||
func requiredLiterals(re *syntax.Regexp, dst []string) []string {
|
||||
switch re.Op {
|
||||
case syntax.OpLiteral:
|
||||
// regexp matches U+FFFD against invalid UTF-8 bytes, strings.Contains does not
|
||||
if re.Flags&syntax.FoldCase == 0 && !slices.Contains(re.Rune, utf8.RuneError) {
|
||||
dst = append(dst, string(re.Rune))
|
||||
}
|
||||
case syntax.OpCapture, syntax.OpPlus:
|
||||
dst = requiredLiterals(re.Sub[0], dst)
|
||||
case syntax.OpRepeat:
|
||||
if re.Min > 0 {
|
||||
dst = requiredLiterals(re.Sub[0], dst)
|
||||
}
|
||||
case syntax.OpConcat:
|
||||
for _, sub := range re.Sub {
|
||||
dst = requiredLiterals(sub, dst)
|
||||
}
|
||||
}
|
||||
return dst
|
||||
pattern *regexp.Regexp
|
||||
}
|
||||
|
||||
func (*RegexMatcher) Type() Type {
|
||||
@@ -359,14 +89,6 @@ func (m *RegexMatcher) String() string {
|
||||
}
|
||||
|
||||
func (m *RegexMatcher) Match(s string) bool {
|
||||
if !m.mayMatch(s) {
|
||||
return false
|
||||
}
|
||||
for _, l := range m.literals {
|
||||
if !strings.Contains(s, l) {
|
||||
return false
|
||||
}
|
||||
}
|
||||
return m.pattern.MatchString(s)
|
||||
}
|
||||
|
||||
@@ -380,7 +102,11 @@ func (t Type) New(pattern string) (Matcher, error) {
|
||||
case Domain:
|
||||
return DomainMatcher(pattern), nil
|
||||
case Regex: // 1. regex matching is case-sensitive
|
||||
return newRegexMatcher(pattern)
|
||||
regex, err := regexp.Compile(pattern)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &RegexMatcher{pattern: regex}, nil
|
||||
default:
|
||||
return nil, errors.New("unknown matcher type")
|
||||
}
|
||||
@@ -409,7 +135,11 @@ func (t Type) NewDomainPattern(pattern string) (Matcher, error) {
|
||||
}
|
||||
return DomainMatcher(pattern), nil
|
||||
case Regex: // Regex's charset not in LDH subset
|
||||
return newRegexMatcher(pattern)
|
||||
regex, err := regexp.Compile(pattern)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &RegexMatcher{pattern: regex}, nil
|
||||
default:
|
||||
return nil, errors.New("unknown matcher type")
|
||||
}
|
||||
|
||||
@@ -1,233 +0,0 @@
|
||||
package strmatcher
|
||||
|
||||
import (
|
||||
"hash/fnv"
|
||||
"math/rand/v2"
|
||||
"regexp"
|
||||
"regexp/syntax"
|
||||
"slices"
|
||||
"strconv"
|
||||
"strings"
|
||||
"testing"
|
||||
"unicode"
|
||||
"unicode/utf8"
|
||||
)
|
||||
|
||||
var regexLiteralCases = []struct {
|
||||
pattern string
|
||||
literals []string
|
||||
}{
|
||||
{`(^|\.)91porn\.(best|com)$`, []string{"91porn."}},
|
||||
{`.+\.awsdns-cn-[0-9][0-9]\.(biz|com|net|top)$`, []string{".awsdns-cn-", "."}},
|
||||
{`^r+[0-9]+(---|\.)sn-(2x3|ni5|j5o)\w{5}\.googlevideo\.com$`, []string{".googlevideo.com", "sn-", "r"}},
|
||||
{`(?i)abc`, nil},
|
||||
{`ab(?i:CD)ef`, []string{"ab", "ef"}},
|
||||
{`(abc)?x`, []string{"x"}},
|
||||
{`(abc)*x`, []string{"x"}},
|
||||
{`x{0,3}yy`, []string{"yy"}},
|
||||
{`(ab)+c{2}`, []string{"ab", "c"}},
|
||||
{`abc|abd`, []string{"ab"}},
|
||||
{`\Qa.b\E`, []string{"a.b"}},
|
||||
{`a\x{FFFD}b`, nil},
|
||||
{`^[^.]+$`, nil},
|
||||
}
|
||||
|
||||
func TestRegexRequiredLiterals(t *testing.T) {
|
||||
for _, test := range regexLiteralCases {
|
||||
m, err := newRegexMatcher(test.pattern)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if got := m.(*RegexMatcher).literals; !slices.Equal(got, test.literals) {
|
||||
t.Errorf("%s: got %q, want %q", test.pattern, got, test.literals)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
var regexTailCases = []struct {
|
||||
pattern string
|
||||
guard bool
|
||||
match []string // inputs the pattern matches
|
||||
reject []string // inputs the tail guard alone rejects
|
||||
}{
|
||||
{`^[a-z]([a-z0-9-]{0,61}[a-z0-9])?$`, true, []string{"a", "localhost", "x-1"}, []string{"www.example.com", "localhost.", "LOCALHOST", "a b"}},
|
||||
{`(^|\.)[a-z][1-9][0-9][a-z]\.com$`, true, []string{"a12b.com", "x.q10z.com"}, []string{"google.com", "a12b.co", "a12b.com.", "ab12.com"}},
|
||||
{`^hses[1-7]?\.akamaized\.net$`, true, []string{"hses.akamaized.net", "hses3.akamaized.net"}, []string{"xhses.akamaized.net", "www.hses.akamaized.net"}},
|
||||
{`(?i)k\.net$`, true, []string{"k.net", "K.NET", "\u212a.net"}, []string{"x.net", "k.nex"}},
|
||||
{`[^.]+\.cn$`, true, []string{"a.cn", "\xff.cn", "\u4e2d.cn"}, []string{"a.cnn", "a.c"}},
|
||||
{`\x{FFFD}$`, true, []string{"\xff", "a\xc3", "\uFFFD"}, []string{"a", "\xff."}},
|
||||
{`^.\.cn$`, true, []string{"a.cn", "\u4E2D.cn", "\xff.cn"}, []string{"ab.cn"}},
|
||||
{`^$`, true, []string{""}, []string{"a"}},
|
||||
{`(^|\.)youyuapi\..+$`, false, []string{"youyuapi.com"}, nil},
|
||||
{`abc`, false, []string{"abc", "xabcx"}, nil},
|
||||
{`^ab`, false, []string{"ab", "abc"}, nil},
|
||||
{`a$|b`, false, []string{"a", "bx"}, nil},
|
||||
{`(?m)a$`, false, []string{"a", "a\nb"}, nil},
|
||||
{strings.Repeat(`(?:abcdefgh(?:a`, 20) + strings.Repeat(`)*)*`, 20) + `\.com$`, false, []string{".com", "abcdefgha.com"}, nil}, // over tailBudget
|
||||
}
|
||||
|
||||
func TestRegexTailGuard(t *testing.T) {
|
||||
for _, test := range regexTailCases {
|
||||
m, err := newRegexMatcher(test.pattern)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
rm := m.(*RegexMatcher)
|
||||
if guard := rm.tail != nil || rm.rest != nil; guard != test.guard {
|
||||
t.Errorf("%s: guard %v, want %v", test.pattern, guard, test.guard)
|
||||
}
|
||||
for _, s := range test.match {
|
||||
if !rm.pattern.MatchString(s) || !rm.Match(s) {
|
||||
t.Errorf("%s: %q does not match", test.pattern, s)
|
||||
}
|
||||
}
|
||||
for _, s := range test.reject {
|
||||
if rm.pattern.MatchString(s) || rm.mayMatch(s) {
|
||||
t.Errorf("%s: %q passes the guard", test.pattern, s)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestRegexTailGuardFlatAlternation checks that a long but non-recursive pattern keeps its
|
||||
// guard. Only nested repeats are charged against tailBudget, so a flat alternation of many
|
||||
// names, however large, is walked once and guarded; its guard is checked against regexp.
|
||||
func TestRegexTailGuardFlatAlternation(t *testing.T) {
|
||||
var sb strings.Builder
|
||||
sb.WriteString("(?:")
|
||||
for i := 0; i < 20000; i++ {
|
||||
if i > 0 {
|
||||
sb.WriteByte('|')
|
||||
}
|
||||
sb.WriteString("name")
|
||||
sb.WriteString(strconv.Itoa(i))
|
||||
}
|
||||
sb.WriteString(`)\.example\.com$`)
|
||||
m, err := newRegexMatcher(sb.String())
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
rm := m.(*RegexMatcher)
|
||||
if rm.tail == nil && rm.rest == nil {
|
||||
t.Fatal("flat alternation of 20000 names lost its guard")
|
||||
}
|
||||
for _, s := range []string{"name0.example.com", "name19999.example.com", "x.name12345.example.com"} {
|
||||
if !rm.pattern.MatchString(s) || !rm.Match(s) {
|
||||
t.Errorf("%q should match", s)
|
||||
}
|
||||
}
|
||||
for _, s := range []string{"name0.example.org", "name0.example.com.", "name0.example.con", "google.com"} {
|
||||
if rm.pattern.MatchString(s) {
|
||||
t.Fatalf("test bug: %q matches the pattern", s)
|
||||
}
|
||||
if rm.mayMatch(s) {
|
||||
t.Errorf("%q should be rejected by the guard", s)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// sampleMatch appends a string that re matches, assertions aside, unless it runs out of
|
||||
// budget, which it spends one per call so that nested repeats stay cheap.
|
||||
func sampleMatch(sb *strings.Builder, re *syntax.Regexp, rnd *rand.Rand, budget *int) {
|
||||
if *budget <= 0 {
|
||||
return
|
||||
}
|
||||
*budget--
|
||||
switch re.Op {
|
||||
case syntax.OpLiteral:
|
||||
for _, r := range re.Rune {
|
||||
if re.Flags&syntax.FoldCase != 0 {
|
||||
for n := rnd.IntN(4); n > 0; n-- {
|
||||
r = unicode.SimpleFold(r)
|
||||
}
|
||||
}
|
||||
sampleRune(sb, r, rnd)
|
||||
}
|
||||
case syntax.OpCharClass:
|
||||
if len(re.Rune) > 0 {
|
||||
i := rnd.IntN(len(re.Rune)/2) * 2
|
||||
sampleRune(sb, re.Rune[i]+rnd.Int32N(min(re.Rune[i+1]-re.Rune[i]+1, 300)), rnd)
|
||||
}
|
||||
case syntax.OpAnyChar, syntax.OpAnyCharNotNL:
|
||||
sampleRune(sb, []rune{'a', '.', '\n', 0xe9, 0x212a, utf8.RuneError}[rnd.IntN(6)], rnd)
|
||||
case syntax.OpCapture:
|
||||
sampleMatch(sb, re.Sub[0], rnd, budget)
|
||||
case syntax.OpConcat:
|
||||
for _, sub := range re.Sub {
|
||||
sampleMatch(sb, sub, rnd, budget)
|
||||
}
|
||||
case syntax.OpAlternate:
|
||||
sampleMatch(sb, re.Sub[rnd.IntN(len(re.Sub))], rnd, budget)
|
||||
case syntax.OpQuest, syntax.OpStar, syntax.OpPlus, syntax.OpRepeat:
|
||||
lo, hi := 0, 3
|
||||
switch re.Op {
|
||||
case syntax.OpQuest:
|
||||
hi = 1
|
||||
case syntax.OpPlus:
|
||||
lo = 1
|
||||
case syntax.OpRepeat:
|
||||
lo, hi = re.Min, re.Min+3
|
||||
if re.Max >= 0 {
|
||||
hi = min(hi, re.Max)
|
||||
}
|
||||
}
|
||||
for n := lo + rnd.IntN(hi-lo+1); n > 0; n-- {
|
||||
sampleMatch(sb, re.Sub[0], rnd, budget)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func sampleRune(sb *strings.Builder, r rune, rnd *rand.Rand) {
|
||||
if r == utf8.RuneError && rnd.IntN(2) == 0 {
|
||||
sb.WriteByte(0x80 | byte(rnd.IntN(0x80))) // regexp matches an invalid byte as U+FFFD
|
||||
return
|
||||
}
|
||||
sb.WriteRune(r)
|
||||
}
|
||||
|
||||
func FuzzRegexMatcher(f *testing.F) {
|
||||
inputs := []string{
|
||||
"", "x", "yy", "abd", "ccc", "ABC", "abCDef", "abcdef", "abababcc", "a.b", "a\xffb", "a\uFFFDb",
|
||||
"www.91porn.com", "ns1.awsdns-cn-01.top", "r1---sn-2x3abcde.googlevideo.com",
|
||||
}
|
||||
for _, test := range regexLiteralCases {
|
||||
for _, s := range inputs {
|
||||
f.Add(test.pattern, s)
|
||||
}
|
||||
}
|
||||
for _, test := range regexTailCases {
|
||||
for _, s := range append(test.match, test.reject...) {
|
||||
f.Add(test.pattern, s)
|
||||
}
|
||||
}
|
||||
f.Fuzz(func(t *testing.T, pattern, s string) {
|
||||
re, err := regexp.Compile(pattern)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
m, _ := newRegexMatcher(pattern)
|
||||
check := func(s string) {
|
||||
if got, want := m.Match(s), re.MatchString(s); got != want {
|
||||
t.Errorf("pattern %q, input %q: got %v, want %v", pattern, s, got, want)
|
||||
}
|
||||
}
|
||||
check(s)
|
||||
// random inputs seldom match, so also try strings built from the pattern
|
||||
parsed, _ := syntax.Parse(pattern, syntax.Perl)
|
||||
h := fnv.New64a()
|
||||
h.Write([]byte(s))
|
||||
rnd := rand.New(rand.NewPCG(h.Sum64(), 1))
|
||||
for range 8 {
|
||||
var sb strings.Builder
|
||||
budget := 256
|
||||
sampleMatch(&sb, parsed, rnd, &budget)
|
||||
sample := sb.String()
|
||||
check(sample)
|
||||
check(s + sample)
|
||||
if len(sample) > 0 && len(s) > 0 {
|
||||
i := rnd.IntN(len(sample))
|
||||
check(sample[:i] + s[:1] + sample[i+1:])
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
@@ -46,9 +46,7 @@ func (g *MphValueMatcher) Add(matcher Matcher, value uint32) {
|
||||
func (g *MphValueMatcher) Build() error {
|
||||
if g.mph != nil {
|
||||
runtime.GC() // peak mem
|
||||
if err := g.mph.Build(); err != nil {
|
||||
return err
|
||||
}
|
||||
g.mph.Build()
|
||||
}
|
||||
runtime.GC() // peak mem
|
||||
if g.ac != nil {
|
||||
@@ -60,17 +58,23 @@ func (g *MphValueMatcher) Build() error {
|
||||
|
||||
// Match implements ValueMatcher.Match.
|
||||
func (g *MphValueMatcher) Match(input string) []uint32 {
|
||||
var result []uint32
|
||||
result := make([][]uint32, 0, 5)
|
||||
if g.mph != nil {
|
||||
result = g.mph.Match(input) // a new slice, returned without another copy
|
||||
if matches := g.mph.Match(input); len(matches) > 0 {
|
||||
result = append(result, matches)
|
||||
}
|
||||
}
|
||||
if g.ac != nil {
|
||||
result = append(result, g.ac.Match(input)...)
|
||||
if matches := g.ac.Match(input); len(matches) > 0 {
|
||||
result = append(result, matches)
|
||||
}
|
||||
}
|
||||
if g.regex != nil {
|
||||
result = append(result, g.regex.Match(input)...)
|
||||
if matches := g.regex.Match(input); len(matches) > 0 {
|
||||
result = append(result, matches)
|
||||
}
|
||||
}
|
||||
return result
|
||||
return CompositeMatches(result)
|
||||
}
|
||||
|
||||
// MatchAny implements ValueMatcher.MatchAny.
|
||||
@@ -83,62 +87,3 @@ func (g *MphValueMatcher) MatchAny(input string) bool {
|
||||
}
|
||||
return g.regex != nil && g.regex.MatchAny(input)
|
||||
}
|
||||
|
||||
func (g *MphValueMatcher) matchAnyHashed(input string, parents []mphSuffix, h, mul uint64) bool {
|
||||
if g.mph != nil && g.mph.matchAnyHashed(input, parents, h, mul) {
|
||||
return true
|
||||
}
|
||||
if g.ac != nil && g.ac.MatchAny(input) {
|
||||
return true
|
||||
}
|
||||
return g.regex != nil && g.regex.MatchAny(input)
|
||||
}
|
||||
|
||||
// MphValueMatcherCombiner combines several built MphValueMatchers, each bound to one value, and matches an input
|
||||
// against them as their MatchAny would, hashing the input once for all of them.
|
||||
type MphValueMatcherCombiner struct {
|
||||
matchers []*MphValueMatcher
|
||||
values []uint32
|
||||
}
|
||||
|
||||
// Add adds a built matcher that stands for value.
|
||||
func (s *MphValueMatcherCombiner) Add(m *MphValueMatcher, value uint32) {
|
||||
s.matchers = append(s.matchers, m)
|
||||
s.values = append(s.values, value)
|
||||
}
|
||||
|
||||
// Match returns the values of the matchers that match input, in Add order.
|
||||
func (s *MphValueMatcherCombiner) Match(input string) []uint32 {
|
||||
if len(s.matchers) == 0 {
|
||||
return nil
|
||||
}
|
||||
var stack [16]mphSuffix
|
||||
mul := mphMultipliers[0]
|
||||
parents, h := mphSuffixes(stack[:0], mul, input)
|
||||
var result []uint32
|
||||
for i, m := range s.matchers {
|
||||
if m.matchAnyHashed(input, parents, h, mul) {
|
||||
result = append(result, s.values[i])
|
||||
}
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
// MatchAny returns true as soon as one matcher matches input.
|
||||
func (s *MphValueMatcherCombiner) MatchAny(input string) bool {
|
||||
switch len(s.matchers) {
|
||||
case 0:
|
||||
return false
|
||||
case 1:
|
||||
return s.matchers[0].MatchAny(input) // nothing to share, and it stops at the first matching suffix
|
||||
}
|
||||
var stack [16]mphSuffix
|
||||
mul := mphMultipliers[0]
|
||||
parents, h := mphSuffixes(stack[:0], mul, input)
|
||||
for _, m := range s.matchers {
|
||||
if m.matchAnyHashed(input, parents, h, mul) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
+25
-21
@@ -1,7 +1,7 @@
|
||||
package log // import "github.com/xtls/xray-core/common/log"
|
||||
|
||||
import (
|
||||
"sync/atomic"
|
||||
"sync"
|
||||
|
||||
"github.com/xtls/xray-core/common/serial"
|
||||
)
|
||||
@@ -29,32 +29,36 @@ func (m *GeneralMessage) String() string {
|
||||
|
||||
// Record writes a message into log stream.
|
||||
func Record(msg Message) {
|
||||
if h := logHandler.Load(); h != nil {
|
||||
(*h).Handle(msg)
|
||||
}
|
||||
logHandler.Handle(msg)
|
||||
}
|
||||
|
||||
type SeverityLogger interface {
|
||||
Handler
|
||||
Severity() Severity
|
||||
}
|
||||
|
||||
func GetSeverity() Severity {
|
||||
if h := logHandler.Load(); h != nil {
|
||||
if sh, ok := (*h).(SeverityLogger); ok {
|
||||
return sh.Severity()
|
||||
}
|
||||
}
|
||||
// log everything by default
|
||||
return Severity_Debug
|
||||
}
|
||||
|
||||
var logHandler atomic.Pointer[Handler]
|
||||
var logHandler syncHandler
|
||||
|
||||
// RegisterHandler registers a new handler as current log handler. Previous registered handler will be discarded.
|
||||
func RegisterHandler(handler Handler) {
|
||||
if handler == nil {
|
||||
panic("Log handler is nil")
|
||||
}
|
||||
logHandler.Store(&handler)
|
||||
logHandler.Set(handler)
|
||||
}
|
||||
|
||||
type syncHandler struct {
|
||||
sync.RWMutex
|
||||
Handler
|
||||
}
|
||||
|
||||
func (h *syncHandler) Handle(msg Message) {
|
||||
h.RLock()
|
||||
defer h.RUnlock()
|
||||
|
||||
if h.Handler != nil {
|
||||
h.Handler.Handle(msg)
|
||||
}
|
||||
}
|
||||
|
||||
func (h *syncHandler) Set(handler Handler) {
|
||||
h.Lock()
|
||||
defer h.Unlock()
|
||||
|
||||
h.Handler = handler
|
||||
}
|
||||
|
||||
@@ -68,10 +68,6 @@ func (l *serverityLogger) Handle(msg Message) {
|
||||
}
|
||||
}
|
||||
|
||||
func (l *serverityLogger) Severity() Severity {
|
||||
return l.logLevel
|
||||
}
|
||||
|
||||
func (l *generalLogger) run() {
|
||||
defer l.access.Signal()
|
||||
|
||||
|
||||
@@ -38,7 +38,7 @@ func (m *ClientManager) Dispatch(ctx context.Context, link *transport.Link) erro
|
||||
}
|
||||
}
|
||||
|
||||
return errors.New("unable to find an available mux client")
|
||||
return errors.New("unable to find an available mux client").AtWarning()
|
||||
}
|
||||
|
||||
type WorkerPicker interface {
|
||||
|
||||
+1
-1
@@ -117,7 +117,7 @@ func (f *FrameMetadata) Unmarshal(reader io.Reader, readSourceAndLocal bool) err
|
||||
return err
|
||||
}
|
||||
if metaLen > 512 {
|
||||
return errors.New("invalid metalen ", metaLen)
|
||||
return errors.New("invalid metalen ", metaLen).AtError()
|
||||
}
|
||||
|
||||
b := buf.New()
|
||||
|
||||
@@ -351,7 +351,7 @@ func (w *ServerWorker) handleFrame(ctx context.Context, reader *buf.BufferedRead
|
||||
err = w.handleStatusKeep(&meta, reader)
|
||||
default:
|
||||
status := meta.SessionStatus
|
||||
return errors.New("unknown status: ", status)
|
||||
return errors.New("unknown status: ", status).AtError()
|
||||
}
|
||||
|
||||
if err != nil {
|
||||
|
||||
@@ -1,350 +0,0 @@
|
||||
//go:build darwin && !ios
|
||||
|
||||
package net
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"net"
|
||||
"net/netip"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"syscall"
|
||||
"unsafe"
|
||||
|
||||
"golang.org/x/sys/unix"
|
||||
|
||||
"github.com/xtls/xray-core/common/errors"
|
||||
)
|
||||
|
||||
const (
|
||||
darwinProcPIDListFDs = 1
|
||||
darwinProcPIDFDSocketInfo = 3
|
||||
darwinProcFDTypeSocket = 2
|
||||
darwinProcFDInfoSize = 8
|
||||
darwinSocketFDInfoSize = 792
|
||||
darwinSocketFDInfoPSIOff = 24
|
||||
darwinSocketInfoProtoOff = darwinSocketFDInfoPSIOff + 156
|
||||
darwinSocketInfoFamilyOff = darwinSocketFDInfoPSIOff + 160
|
||||
darwinSocketInfoKindOff = darwinSocketFDInfoPSIOff + 232
|
||||
darwinSocketInfoInSockOff = darwinSocketFDInfoPSIOff + 240
|
||||
darwinInSockInfoFPortOff = darwinSocketInfoInSockOff
|
||||
darwinInSockInfoLPortOff = darwinSocketInfoInSockOff + 4
|
||||
darwinInSockInfoVFlagOff = darwinSocketInfoInSockOff + 24
|
||||
darwinInSockInfoFAddrOff = darwinSocketInfoInSockOff + 32
|
||||
darwinInSockInfoLAddrOff = darwinSocketInfoInSockOff + 48
|
||||
darwinInSockInfoSize = 80
|
||||
darwinInSockInfoIPv4 = 0x1
|
||||
darwinInSockInfoIPv6 = 0x2
|
||||
darwinSockInfoIN = 1
|
||||
darwinSockInfoTCP = 2
|
||||
)
|
||||
|
||||
type darwinSocketMatchLevel int
|
||||
|
||||
const (
|
||||
darwinSocketNoMatch darwinSocketMatchLevel = iota
|
||||
darwinSocketPortMatch
|
||||
darwinSocketRemoteMatch
|
||||
darwinSocketLocalMatch
|
||||
darwinSocketExactMatch
|
||||
)
|
||||
|
||||
func FindProcess(network, srcIP string, srcPort uint16, destIP string, destPort uint16) (PID int, Name string, AbsolutePath string, err error) {
|
||||
isLocal, err := IsLocal(net.ParseIP(srcIP))
|
||||
if err != nil {
|
||||
return 0, "", "", errors.New("failed to determine if address is local: ", err)
|
||||
}
|
||||
if !isLocal {
|
||||
return 0, "", "", ErrNotLocal
|
||||
}
|
||||
if network != "tcp" && network != "udp" {
|
||||
panic("Unsupported network type for process lookup.")
|
||||
}
|
||||
|
||||
srcAddr, err := netip.ParseAddr(srcIP)
|
||||
if err != nil {
|
||||
return 0, "", "", errors.New("invalid source IP address: ", srcIP)
|
||||
}
|
||||
srcAddr = srcAddr.Unmap()
|
||||
|
||||
var dstAddr netip.Addr
|
||||
hasDstAddr := false
|
||||
if destIP != "" && destPort != 0 {
|
||||
dstAddr, err = netip.ParseAddr(destIP)
|
||||
if err != nil {
|
||||
return 0, "", "", errors.New("invalid destination IP address: ", destIP)
|
||||
}
|
||||
dstAddr = dstAddr.Unmap()
|
||||
hasDstAddr = true
|
||||
}
|
||||
|
||||
processes, err := unix.SysctlKinfoProcSlice("kern.proc.all")
|
||||
if err != nil {
|
||||
return 0, "", "", errors.New("failed to list processes").Base(err)
|
||||
}
|
||||
|
||||
var bestPID int32
|
||||
bestLevel := darwinSocketNoMatch
|
||||
ambiguousBest := false
|
||||
|
||||
for _, process := range processes {
|
||||
pid := process.Proc.P_pid
|
||||
if pid <= 0 {
|
||||
continue
|
||||
}
|
||||
|
||||
matchLevel, err := darwinProcessSocketMatchLevel(pid, network, srcAddr, srcPort, dstAddr, destPort, hasDstAddr)
|
||||
if err != nil || matchLevel == darwinSocketNoMatch {
|
||||
continue
|
||||
}
|
||||
if matchLevel == darwinSocketExactMatch {
|
||||
bestPID = pid
|
||||
bestLevel = matchLevel
|
||||
ambiguousBest = false
|
||||
break
|
||||
}
|
||||
if matchLevel > bestLevel {
|
||||
bestPID = pid
|
||||
bestLevel = matchLevel
|
||||
ambiguousBest = false
|
||||
continue
|
||||
}
|
||||
if matchLevel == bestLevel {
|
||||
ambiguousBest = true
|
||||
}
|
||||
}
|
||||
|
||||
if bestLevel == darwinSocketNoMatch {
|
||||
return 0, "", "", errors.New("process not found for ", network, " connection from ", srcIP, ":", srcPort, " to ", destIP, ":", destPort)
|
||||
}
|
||||
if ambiguousBest {
|
||||
return 0, "", "", errors.New("ambiguous process match for ", network, " connection from ", srcIP, ":", srcPort, " to ", destIP, ":", destPort)
|
||||
}
|
||||
|
||||
absPath, err := darwinProcessPath(bestPID)
|
||||
if err != nil {
|
||||
return 0, "", "", errors.New("could not get process path for PID ", bestPID, ": ", err)
|
||||
}
|
||||
|
||||
absPath = filepath.ToSlash(absPath)
|
||||
return int(bestPID), filepath.Base(absPath), absPath, nil
|
||||
}
|
||||
|
||||
func darwinProcessSocketMatchLevel(pid int32, network string, srcAddr netip.Addr, srcPort uint16, dstAddr netip.Addr, dstPort uint16, hasDstAddr bool) (darwinSocketMatchLevel, error) {
|
||||
fds, err := darwinProcessFDs(pid)
|
||||
if err != nil {
|
||||
return darwinSocketNoMatch, err
|
||||
}
|
||||
|
||||
bestLevel := darwinSocketNoMatch
|
||||
info := make([]byte, darwinSocketFDInfoSize)
|
||||
for fd := 0; fd+darwinProcFDInfoSize <= len(fds); fd += darwinProcFDInfoSize {
|
||||
fdNumber := int32(darwinReadNativeUint32(fds[fd : fd+4]))
|
||||
fdType := darwinReadNativeUint32(fds[fd+4 : fd+8])
|
||||
if fdType != darwinProcFDTypeSocket {
|
||||
continue
|
||||
}
|
||||
|
||||
n, err := darwinProcPIDFDInfo(pid, fdNumber, darwinProcPIDFDSocketInfo, info)
|
||||
if err != nil || n < darwinSocketInfoInSockOff+darwinInSockInfoSize {
|
||||
continue
|
||||
}
|
||||
level := darwinSocketInfoMatchLevel(info[:n], network, srcAddr, srcPort, dstAddr, dstPort, hasDstAddr)
|
||||
if level == darwinSocketExactMatch {
|
||||
return level, nil
|
||||
}
|
||||
if level > bestLevel {
|
||||
bestLevel = level
|
||||
}
|
||||
}
|
||||
|
||||
return bestLevel, nil
|
||||
}
|
||||
|
||||
func darwinProcessFDs(pid int32) ([]byte, error) {
|
||||
n, err := darwinProcPIDInfo(pid, darwinProcPIDListFDs, 0, nil)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if n <= 0 {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
buf := make([]byte, n)
|
||||
n, err = darwinProcPIDInfo(pid, darwinProcPIDListFDs, 0, buf)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return buf[:n], nil
|
||||
}
|
||||
|
||||
func darwinSocketInfoMatchLevel(info []byte, network string, srcAddr netip.Addr, srcPort uint16, dstAddr netip.Addr, dstPort uint16, hasDstAddr bool) darwinSocketMatchLevel {
|
||||
protocol := int(darwinReadNativeUint32(info[darwinSocketInfoProtoOff : darwinSocketInfoProtoOff+4]))
|
||||
family := int(darwinReadNativeUint32(info[darwinSocketInfoFamilyOff : darwinSocketInfoFamilyOff+4]))
|
||||
kind := int(darwinReadNativeUint32(info[darwinSocketInfoKindOff : darwinSocketInfoKindOff+4]))
|
||||
|
||||
switch network {
|
||||
case "tcp":
|
||||
if protocol != unix.IPPROTO_TCP || kind != darwinSockInfoTCP {
|
||||
return darwinSocketNoMatch
|
||||
}
|
||||
case "udp":
|
||||
if protocol != unix.IPPROTO_UDP || kind != darwinSockInfoIN {
|
||||
return darwinSocketNoMatch
|
||||
}
|
||||
default:
|
||||
return darwinSocketNoMatch
|
||||
}
|
||||
|
||||
vflag := info[darwinInSockInfoVFlagOff]
|
||||
if srcAddr.Is4() {
|
||||
// Dual-stack sockets expose IPv4-mapped connections as AF_INET6
|
||||
// while marking the endpoint as IPv4 in ini_vflag.
|
||||
if (family != unix.AF_INET && family != unix.AF_INET6) || vflag&darwinInSockInfoIPv4 == 0 {
|
||||
return darwinSocketNoMatch
|
||||
}
|
||||
} else {
|
||||
if family != unix.AF_INET6 || vflag&darwinInSockInfoIPv6 == 0 {
|
||||
return darwinSocketNoMatch
|
||||
}
|
||||
}
|
||||
|
||||
localPort := int32(darwinReadNativeUint32(info[darwinInSockInfoLPortOff : darwinInSockInfoLPortOff+4]))
|
||||
if !darwinPortMatches(localPort, srcPort) {
|
||||
return darwinSocketNoMatch
|
||||
}
|
||||
localAddrMatches := darwinAddrMatchesOrUnspecified(info[darwinInSockInfoLAddrOff:darwinInSockInfoLAddrOff+16], srcAddr)
|
||||
|
||||
foreignAddrRaw := info[darwinInSockInfoFAddrOff : darwinInSockInfoFAddrOff+16]
|
||||
foreignPort := int32(darwinReadNativeUint32(info[darwinInSockInfoFPortOff : darwinInSockInfoFPortOff+4]))
|
||||
|
||||
if !hasDstAddr {
|
||||
if localAddrMatches {
|
||||
return darwinSocketExactMatch
|
||||
}
|
||||
return darwinSocketNoMatch
|
||||
}
|
||||
|
||||
remoteMatches := darwinPortMatches(foreignPort, dstPort) && darwinAddrMatches(foreignAddrRaw, dstAddr)
|
||||
if network == "udp" && darwinEndpointIsZero(foreignAddrRaw, foreignPort) && localAddrMatches {
|
||||
return darwinSocketExactMatch
|
||||
}
|
||||
switch {
|
||||
case localAddrMatches && remoteMatches:
|
||||
return darwinSocketExactMatch
|
||||
case localAddrMatches:
|
||||
return darwinSocketLocalMatch
|
||||
case remoteMatches:
|
||||
return darwinSocketRemoteMatch
|
||||
default:
|
||||
return darwinSocketPortMatch
|
||||
}
|
||||
}
|
||||
|
||||
func darwinPortMatches(value int32, port uint16) bool {
|
||||
raw := uint16(value)
|
||||
return raw == port || darwinNtohs(raw) == port
|
||||
}
|
||||
|
||||
func darwinNtohs(value uint16) uint16 {
|
||||
return value<<8 | value>>8
|
||||
}
|
||||
|
||||
func darwinAddrMatches(raw []byte, addr netip.Addr) bool {
|
||||
if addr.Is4() {
|
||||
ip := addr.As4()
|
||||
return bytes.Equal(raw[12:16], ip[:])
|
||||
}
|
||||
ip := addr.As16()
|
||||
return bytes.Equal(raw, ip[:])
|
||||
}
|
||||
|
||||
func darwinAddrMatchesOrUnspecified(raw []byte, addr netip.Addr) bool {
|
||||
if darwinAddrMatches(raw, addr) {
|
||||
return true
|
||||
}
|
||||
if addr.Is4() {
|
||||
return darwinBytesAreZero(raw[12:16])
|
||||
}
|
||||
return darwinBytesAreZero(raw)
|
||||
}
|
||||
|
||||
func darwinEndpointIsZero(rawAddr []byte, port int32) bool {
|
||||
return uint32(port) == 0 && darwinBytesAreZero(rawAddr)
|
||||
}
|
||||
|
||||
func darwinBytesAreZero(raw []byte) bool {
|
||||
for _, value := range raw {
|
||||
if value != 0 {
|
||||
return false
|
||||
}
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
func darwinReadNativeUint32(b []byte) uint32 {
|
||||
return *(*uint32)(unsafe.Pointer(&b[0]))
|
||||
}
|
||||
|
||||
func darwinProcessPath(pid int32) (string, error) {
|
||||
buf := make([]byte, unix.PathMax)
|
||||
n, err := darwinProcPIDPath(pid, buf)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
if n <= 0 {
|
||||
return "", errors.New("empty process path")
|
||||
}
|
||||
return strings.TrimRight(string(buf[:n]), "\x00"), nil
|
||||
}
|
||||
|
||||
func darwinProcPIDInfo(pid int32, flavor int, arg uint64, buf []byte) (int, error) {
|
||||
var ptr unsafe.Pointer
|
||||
if len(buf) > 0 {
|
||||
ptr = unsafe.Pointer(&buf[0])
|
||||
}
|
||||
|
||||
r0, _, errno := syscall_syscall6(libc_proc_pidinfo_trampoline_addr, uintptr(pid), uintptr(flavor), uintptr(arg), uintptr(ptr), uintptr(len(buf)), 0)
|
||||
if errno != 0 {
|
||||
return 0, errno
|
||||
}
|
||||
return int(r0), nil
|
||||
}
|
||||
|
||||
func darwinProcPIDFDInfo(pid int32, fd int32, flavor int, buf []byte) (int, error) {
|
||||
var ptr unsafe.Pointer
|
||||
if len(buf) > 0 {
|
||||
ptr = unsafe.Pointer(&buf[0])
|
||||
}
|
||||
|
||||
r0, _, errno := syscall_syscall6(libc_proc_pidfdinfo_trampoline_addr, uintptr(pid), uintptr(fd), uintptr(flavor), uintptr(ptr), uintptr(len(buf)), 0)
|
||||
if errno != 0 {
|
||||
return 0, errno
|
||||
}
|
||||
return int(r0), nil
|
||||
}
|
||||
|
||||
func darwinProcPIDPath(pid int32, buf []byte) (int, error) {
|
||||
r0, _, errno := syscall_syscall6(libc_proc_pidpath_trampoline_addr, uintptr(pid), uintptr(unsafe.Pointer(&buf[0])), uintptr(len(buf)), 0, 0, 0)
|
||||
if errno != 0 {
|
||||
return 0, errno
|
||||
}
|
||||
return int(r0), nil
|
||||
}
|
||||
|
||||
var libc_proc_pidinfo_trampoline_addr uintptr
|
||||
|
||||
//go:cgo_import_dynamic libc_proc_pidinfo proc_pidinfo "/usr/lib/libproc.dylib"
|
||||
|
||||
var libc_proc_pidfdinfo_trampoline_addr uintptr
|
||||
|
||||
//go:cgo_import_dynamic libc_proc_pidfdinfo proc_pidfdinfo "/usr/lib/libproc.dylib"
|
||||
|
||||
var libc_proc_pidpath_trampoline_addr uintptr
|
||||
|
||||
//go:cgo_import_dynamic libc_proc_pidpath proc_pidpath "/usr/lib/libproc.dylib"
|
||||
|
||||
// Implemented in the runtime package (runtime/sys_darwin.go).
|
||||
func syscall_syscall6(fn, a1, a2, a3, a4, a5, a6 uintptr) (r1, r2 uintptr, err syscall.Errno)
|
||||
|
||||
//go:linkname syscall_syscall6 syscall.syscall6
|
||||
@@ -1,18 +0,0 @@
|
||||
//go:build darwin && !ios
|
||||
|
||||
#include "textflag.h"
|
||||
|
||||
TEXT libc_proc_pidinfo_trampoline<>(SB),NOSPLIT,$0-0
|
||||
JMP libc_proc_pidinfo(SB)
|
||||
GLOBL ·libc_proc_pidinfo_trampoline_addr(SB), RODATA, $8
|
||||
DATA ·libc_proc_pidinfo_trampoline_addr(SB)/8, $libc_proc_pidinfo_trampoline<>(SB)
|
||||
|
||||
TEXT libc_proc_pidfdinfo_trampoline<>(SB),NOSPLIT,$0-0
|
||||
JMP libc_proc_pidfdinfo(SB)
|
||||
GLOBL ·libc_proc_pidfdinfo_trampoline_addr(SB), RODATA, $8
|
||||
DATA ·libc_proc_pidfdinfo_trampoline_addr(SB)/8, $libc_proc_pidfdinfo_trampoline<>(SB)
|
||||
|
||||
TEXT libc_proc_pidpath_trampoline<>(SB),NOSPLIT,$0-0
|
||||
JMP libc_proc_pidpath(SB)
|
||||
GLOBL ·libc_proc_pidpath_trampoline_addr(SB), RODATA, $8
|
||||
DATA ·libc_proc_pidpath_trampoline_addr(SB)/8, $libc_proc_pidpath_trampoline<>(SB)
|
||||
@@ -1,356 +0,0 @@
|
||||
//go:build darwin && !ios
|
||||
|
||||
package net
|
||||
|
||||
import (
|
||||
stdnet "net"
|
||||
"net/netip"
|
||||
"os"
|
||||
"testing"
|
||||
"time"
|
||||
"unsafe"
|
||||
|
||||
"golang.org/x/sys/unix"
|
||||
)
|
||||
|
||||
func TestFindProcessDarwinTCP(t *testing.T) {
|
||||
listener, err := stdnet.Listen("tcp", "127.0.0.1:0")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer listener.Close()
|
||||
|
||||
accepted := make(chan stdnet.Conn, 1)
|
||||
go func() {
|
||||
conn, err := listener.Accept()
|
||||
if err == nil {
|
||||
accepted <- conn
|
||||
return
|
||||
}
|
||||
close(accepted)
|
||||
}()
|
||||
|
||||
conn, err := stdnet.Dial("tcp", listener.Addr().String())
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer conn.Close()
|
||||
|
||||
serverConn := <-accepted
|
||||
if serverConn == nil {
|
||||
t.Fatal("server did not accept tcp connection")
|
||||
}
|
||||
defer serverConn.Close()
|
||||
|
||||
local := conn.LocalAddr().(*stdnet.TCPAddr)
|
||||
remote := conn.RemoteAddr().(*stdnet.TCPAddr)
|
||||
|
||||
pid, name, path, err := FindProcess("tcp", local.IP.String(), uint16(local.Port), remote.IP.String(), uint16(remote.Port))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
assertCurrentProcess(t, pid, name, path)
|
||||
}
|
||||
|
||||
func TestFindProcessDarwinTCPIPv4Mapped(t *testing.T) {
|
||||
listener, err := stdnet.Listen("tcp4", "127.0.0.1:0")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer listener.Close()
|
||||
|
||||
accepted := make(chan stdnet.Conn, 1)
|
||||
go func() {
|
||||
conn, err := listener.Accept()
|
||||
if err == nil {
|
||||
accepted <- conn
|
||||
return
|
||||
}
|
||||
close(accepted)
|
||||
}()
|
||||
|
||||
listenerAddr := listener.Addr().(*stdnet.TCPAddr)
|
||||
fd, err := unix.Socket(unix.AF_INET6, unix.SOCK_STREAM, unix.IPPROTO_TCP)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer unix.Close(fd)
|
||||
|
||||
mappedAddr := [16]byte{10: 0xff, 11: 0xff, 12: 127, 15: 1}
|
||||
if err := unix.Connect(fd, &unix.SockaddrInet6{
|
||||
Port: listenerAddr.Port,
|
||||
Addr: mappedAddr,
|
||||
}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
serverConn := <-accepted
|
||||
if serverConn == nil {
|
||||
t.Fatal("server did not accept tcp connection")
|
||||
}
|
||||
defer serverConn.Close()
|
||||
|
||||
local, err := unix.Getsockname(fd)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
localPort := local.(*unix.SockaddrInet6).Port
|
||||
|
||||
pid, name, path, err := FindProcess("tcp", "127.0.0.1", uint16(localPort), "127.0.0.1", uint16(listenerAddr.Port))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
assertCurrentProcess(t, pid, name, path)
|
||||
}
|
||||
|
||||
func TestFindProcessDarwinUDP(t *testing.T) {
|
||||
conn, err := stdnet.ListenUDP("udp", &stdnet.UDPAddr{IP: stdnet.ParseIP("127.0.0.1")})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer conn.Close()
|
||||
|
||||
local := conn.LocalAddr().(*stdnet.UDPAddr)
|
||||
|
||||
pid, name, path, err := FindProcess("udp", local.IP.String(), uint16(local.Port), "", 0)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
assertCurrentProcess(t, pid, name, path)
|
||||
}
|
||||
|
||||
func TestFindProcessDarwinNonLocal(t *testing.T) {
|
||||
_, _, _, err := FindProcess("tcp", "203.0.113.1", 80, "", 0)
|
||||
if err != ErrNotLocal {
|
||||
t.Fatalf("expected ErrNotLocal, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestFindProcessDarwinUnsupportedNetwork(t *testing.T) {
|
||||
defer func() {
|
||||
if recover() == nil {
|
||||
t.Fatal("expected panic")
|
||||
}
|
||||
}()
|
||||
|
||||
_, _, _, _ = FindProcess("icmp", "127.0.0.1", 0, "", 0)
|
||||
}
|
||||
|
||||
func assertCurrentProcess(t *testing.T, pid int, name string, path string) {
|
||||
t.Helper()
|
||||
|
||||
if pid != os.Getpid() {
|
||||
t.Fatalf("expected pid %d, got %d (%s, %s)", os.Getpid(), pid, name, path)
|
||||
}
|
||||
|
||||
executable, err := os.Executable()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if path == "" || name == "" {
|
||||
t.Fatalf("expected process path and name, got name=%q path=%q", name, path)
|
||||
}
|
||||
if sameFile(executable, path) {
|
||||
return
|
||||
}
|
||||
t.Fatalf("expected executable %q, got %q", executable, path)
|
||||
}
|
||||
|
||||
func sameFile(left string, right string) bool {
|
||||
leftInfo, leftErr := os.Stat(left)
|
||||
rightInfo, rightErr := os.Stat(right)
|
||||
if leftErr != nil || rightErr != nil {
|
||||
return false
|
||||
}
|
||||
return os.SameFile(leftInfo, rightInfo)
|
||||
}
|
||||
|
||||
func TestFindProcessDarwinTCPWithoutDestination(t *testing.T) {
|
||||
listener, err := stdnet.Listen("tcp", "127.0.0.1:0")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer listener.Close()
|
||||
|
||||
accepted := make(chan stdnet.Conn, 1)
|
||||
go func() {
|
||||
conn, err := listener.Accept()
|
||||
if err == nil {
|
||||
accepted <- conn
|
||||
return
|
||||
}
|
||||
close(accepted)
|
||||
}()
|
||||
|
||||
conn, err := stdnet.DialTimeout("tcp", listener.Addr().String(), time.Second)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer conn.Close()
|
||||
|
||||
serverConn := <-accepted
|
||||
if serverConn == nil {
|
||||
t.Fatal("server did not accept tcp connection")
|
||||
}
|
||||
defer serverConn.Close()
|
||||
|
||||
local := conn.LocalAddr().(*stdnet.TCPAddr)
|
||||
pid, name, path, err := FindProcess("tcp", local.IP.String(), uint16(local.Port), "", 0)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
assertCurrentProcess(t, pid, name, path)
|
||||
}
|
||||
|
||||
func TestFindProcessDarwinTCPWithDifferentDestination(t *testing.T) {
|
||||
listener, err := stdnet.Listen("tcp", "127.0.0.1:0")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer listener.Close()
|
||||
|
||||
accepted := make(chan stdnet.Conn, 1)
|
||||
go func() {
|
||||
conn, err := listener.Accept()
|
||||
if err == nil {
|
||||
accepted <- conn
|
||||
return
|
||||
}
|
||||
close(accepted)
|
||||
}()
|
||||
|
||||
conn, err := stdnet.DialTimeout("tcp", listener.Addr().String(), time.Second)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer conn.Close()
|
||||
|
||||
serverConn := <-accepted
|
||||
if serverConn == nil {
|
||||
t.Fatal("server did not accept tcp connection")
|
||||
}
|
||||
defer serverConn.Close()
|
||||
|
||||
local := conn.LocalAddr().(*stdnet.TCPAddr)
|
||||
pid, name, path, err := FindProcess("tcp", local.IP.String(), uint16(local.Port), "203.0.113.10", 443)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
assertCurrentProcess(t, pid, name, path)
|
||||
}
|
||||
|
||||
func TestDarwinSocketInfoMatchLevelFallbacks(t *testing.T) {
|
||||
src := netip.MustParseAddr("198.18.0.2")
|
||||
dst := netip.MustParseAddr("203.0.113.10")
|
||||
otherLocal := netip.MustParseAddr("192.168.1.10")
|
||||
otherRemote := netip.MustParseAddr("198.51.100.10")
|
||||
unspecifiedLocal := netip.MustParseAddr("0.0.0.0")
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
local netip.Addr
|
||||
remote netip.Addr
|
||||
hasDst bool
|
||||
wantLevel darwinSocketMatchLevel
|
||||
}{
|
||||
{
|
||||
name: "exact",
|
||||
local: src,
|
||||
remote: dst,
|
||||
hasDst: true,
|
||||
wantLevel: darwinSocketExactMatch,
|
||||
},
|
||||
{
|
||||
name: "unspecified local with matching remote",
|
||||
local: unspecifiedLocal,
|
||||
remote: dst,
|
||||
hasDst: true,
|
||||
wantLevel: darwinSocketExactMatch,
|
||||
},
|
||||
{
|
||||
name: "unspecified local without destination",
|
||||
local: unspecifiedLocal,
|
||||
remote: otherRemote,
|
||||
hasDst: false,
|
||||
wantLevel: darwinSocketExactMatch,
|
||||
},
|
||||
{
|
||||
name: "local match with different remote",
|
||||
local: src,
|
||||
remote: otherRemote,
|
||||
hasDst: true,
|
||||
wantLevel: darwinSocketLocalMatch,
|
||||
},
|
||||
{
|
||||
name: "remote match with different local",
|
||||
local: otherLocal,
|
||||
remote: dst,
|
||||
hasDst: true,
|
||||
wantLevel: darwinSocketRemoteMatch,
|
||||
},
|
||||
{
|
||||
name: "port only with destination",
|
||||
local: otherLocal,
|
||||
remote: otherRemote,
|
||||
hasDst: true,
|
||||
wantLevel: darwinSocketPortMatch,
|
||||
},
|
||||
{
|
||||
name: "different local without destination",
|
||||
local: otherLocal,
|
||||
remote: otherRemote,
|
||||
hasDst: false,
|
||||
wantLevel: darwinSocketNoMatch,
|
||||
},
|
||||
}
|
||||
|
||||
for _, test := range tests {
|
||||
t.Run(test.name, func(t *testing.T) {
|
||||
info := newDarwinSocketInfo("tcp", test.local, 12345, test.remote, 443)
|
||||
level := darwinSocketInfoMatchLevel(info, "tcp", src, 12345, dst, 443, test.hasDst)
|
||||
if level != test.wantLevel {
|
||||
t.Fatalf("unexpected match level: got %d, want %d", level, test.wantLevel)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestDarwinSocketInfoMatchLevelIPv4Mapped(t *testing.T) {
|
||||
src := netip.MustParseAddr("127.0.0.1")
|
||||
dst := netip.MustParseAddr("203.0.113.10")
|
||||
info := newDarwinSocketInfo("tcp", src, 12345, dst, 443)
|
||||
writeDarwinNativeUint32(info, darwinSocketInfoFamilyOff, uint32(unix.AF_INET6))
|
||||
|
||||
level := darwinSocketInfoMatchLevel(info, "tcp", src, 12345, dst, 443, true)
|
||||
if level != darwinSocketExactMatch {
|
||||
t.Fatalf("unexpected match level: got %d, want %d", level, darwinSocketExactMatch)
|
||||
}
|
||||
}
|
||||
|
||||
func newDarwinSocketInfo(network string, local netip.Addr, localPort uint16, remote netip.Addr, remotePort uint16) []byte {
|
||||
info := make([]byte, darwinSocketFDInfoSize)
|
||||
switch network {
|
||||
case "tcp":
|
||||
writeDarwinNativeUint32(info, darwinSocketInfoProtoOff, uint32(unix.IPPROTO_TCP))
|
||||
writeDarwinNativeUint32(info, darwinSocketInfoKindOff, uint32(darwinSockInfoTCP))
|
||||
case "udp":
|
||||
writeDarwinNativeUint32(info, darwinSocketInfoProtoOff, uint32(unix.IPPROTO_UDP))
|
||||
writeDarwinNativeUint32(info, darwinSocketInfoKindOff, uint32(darwinSockInfoIN))
|
||||
}
|
||||
writeDarwinNativeUint32(info, darwinSocketInfoFamilyOff, uint32(unix.AF_INET))
|
||||
info[darwinInSockInfoVFlagOff] = darwinInSockInfoIPv4
|
||||
writeDarwinNativeUint32(info, darwinInSockInfoLPortOff, uint32(localPort))
|
||||
writeDarwinNativeUint32(info, darwinInSockInfoFPortOff, uint32(remotePort))
|
||||
copyDarwinIPv4(info[darwinInSockInfoLAddrOff:darwinInSockInfoLAddrOff+16], local)
|
||||
copyDarwinIPv4(info[darwinInSockInfoFAddrOff:darwinInSockInfoFAddrOff+16], remote)
|
||||
return info
|
||||
}
|
||||
|
||||
func writeDarwinNativeUint32(b []byte, offset int, value uint32) {
|
||||
*(*uint32)(unsafe.Pointer(&b[offset])) = value
|
||||
}
|
||||
|
||||
func copyDarwinIPv4(dst []byte, addr netip.Addr) {
|
||||
ip := addr.As4()
|
||||
copy(dst[12:16], ip[:])
|
||||
}
|
||||
@@ -1,11 +0,0 @@
|
||||
//go:build ios
|
||||
|
||||
package net
|
||||
|
||||
import (
|
||||
"github.com/xtls/xray-core/common/errors"
|
||||
)
|
||||
|
||||
func FindProcess(network, srcIP string, srcPort uint16, destIP string, destPort uint16) (int, string, string, error) {
|
||||
return 0, "", "", errors.New("process lookup is not supported on this platform")
|
||||
}
|
||||
@@ -1,4 +1,4 @@
|
||||
//go:build !windows && !linux && !android && !darwin
|
||||
//go:build !windows && !linux && !android
|
||||
|
||||
package net
|
||||
|
||||
|
||||
@@ -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,40 +0,0 @@
|
||||
package platform
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"sync"
|
||||
)
|
||||
|
||||
var envReloadRegistry = struct {
|
||||
sync.RWMutex
|
||||
handlers []func() error
|
||||
}{}
|
||||
|
||||
// RegisterEnvReload registers an environment reload handler and runs it once
|
||||
// immediately so package defaults keep the same behavior as init-time reads.
|
||||
func RegisterEnvReload(handler func() error) {
|
||||
if handler == nil {
|
||||
return
|
||||
}
|
||||
envReloadRegistry.Lock()
|
||||
envReloadRegistry.handlers = append(envReloadRegistry.handlers, handler)
|
||||
envReloadRegistry.Unlock()
|
||||
if err := handler(); err != nil {
|
||||
panic(err)
|
||||
}
|
||||
}
|
||||
|
||||
// ReloadEnvSettings refreshes all registered environment-backed package state.
|
||||
func ReloadEnvSettings() error {
|
||||
envReloadRegistry.RLock()
|
||||
handlers := append([]func() error{}, envReloadRegistry.handlers...)
|
||||
envReloadRegistry.RUnlock()
|
||||
|
||||
var errs []error
|
||||
for _, handler := range handlers {
|
||||
if err := handler(); err != nil {
|
||||
errs = append(errs, err)
|
||||
}
|
||||
}
|
||||
return errors.Join(errs...)
|
||||
}
|
||||
@@ -3,8 +3,11 @@ package bittorrent
|
||||
import (
|
||||
"encoding/binary"
|
||||
"errors"
|
||||
"math"
|
||||
"time"
|
||||
|
||||
"github.com/xtls/xray-core/common"
|
||||
"github.com/xtls/xray-core/common/buf"
|
||||
)
|
||||
|
||||
type SniffHeader struct{}
|
||||
@@ -36,44 +39,50 @@ func SniffUTP(b []byte) (*SniffHeader, error) {
|
||||
return nil, common.ErrNoClue
|
||||
}
|
||||
|
||||
// type 4 (ST_SYN), version 1
|
||||
if b[0] != 0x41 {
|
||||
buffer := buf.FromBytes(b)
|
||||
|
||||
var typeAndVersion uint8
|
||||
|
||||
if binary.Read(buffer, binary.BigEndian, &typeAndVersion) != nil {
|
||||
return nil, common.ErrNoClue
|
||||
} else if b[0]>>4&0xF > 4 || b[0]&0xF != 1 {
|
||||
return nil, errNotBittorrent
|
||||
}
|
||||
|
||||
// timestamp_difference is always 0 in new connections
|
||||
if binary.BigEndian.Uint32(b[8:12]) != 0 {
|
||||
var extension uint8
|
||||
|
||||
if binary.Read(buffer, binary.BigEndian, &extension) != nil {
|
||||
return nil, common.ErrNoClue
|
||||
} else if extension != 0 && extension != 1 {
|
||||
return nil, errNotBittorrent
|
||||
}
|
||||
|
||||
// Walk the extension chain. Selective ack (1) and extension bits (2)
|
||||
extension, offset := b[1], 20
|
||||
for extension != 0 {
|
||||
if len(b) < offset+2 {
|
||||
if extension != 1 {
|
||||
return nil, errNotBittorrent
|
||||
}
|
||||
length := int(b[offset+1])
|
||||
switch extension {
|
||||
case 1: // selective ack
|
||||
if length < 4 || length%4 != 0 {
|
||||
return nil, errNotBittorrent
|
||||
}
|
||||
case 2: // extension bits: fixed 8 bytes, sent in ST_SYN by µTorrent
|
||||
if length != 8 {
|
||||
return nil, errNotBittorrent
|
||||
}
|
||||
default:
|
||||
return nil, errNotBittorrent
|
||||
if binary.Read(buffer, binary.BigEndian, &extension) != nil {
|
||||
return nil, common.ErrNoClue
|
||||
}
|
||||
if len(b) < offset+2+length {
|
||||
return nil, errNotBittorrent
|
||||
|
||||
var length uint8
|
||||
if err := binary.Read(buffer, binary.BigEndian, &length); err != nil {
|
||||
return nil, common.ErrNoClue
|
||||
}
|
||||
if common.Error2(buffer.ReadBytes(int32(length))) != nil {
|
||||
return nil, common.ErrNoClue
|
||||
}
|
||||
extension = b[offset]
|
||||
offset += 2 + length
|
||||
}
|
||||
|
||||
// extensions should consume all ST_SYN payload
|
||||
if len(b) != offset {
|
||||
if common.Error2(buffer.ReadBytes(2)) != nil {
|
||||
return nil, common.ErrNoClue
|
||||
}
|
||||
|
||||
var timestamp uint32
|
||||
if err := binary.Read(buffer, binary.BigEndian, ×tamp); err != nil {
|
||||
return nil, common.ErrNoClue
|
||||
}
|
||||
if math.Abs(float64(time.Now().UnixMicro()-int64(timestamp))) > float64(24*time.Hour) {
|
||||
return nil, errNotBittorrent
|
||||
}
|
||||
|
||||
|
||||
@@ -1,67 +0,0 @@
|
||||
package bittorrent
|
||||
|
||||
import (
|
||||
"encoding/binary"
|
||||
"testing"
|
||||
|
||||
"github.com/xtls/xray-core/common"
|
||||
)
|
||||
|
||||
// utpPacket builds the fixed 20-byte header defined by BEP 29.
|
||||
func utpPacket(packetType, extension byte, tsDiff uint32, payload ...byte) []byte {
|
||||
b := make([]byte, 20)
|
||||
b[0] = packetType<<4 | 1
|
||||
b[1] = extension
|
||||
binary.BigEndian.PutUint16(b[2:4], 0x4a3f) // connection_id, random
|
||||
binary.BigEndian.PutUint32(b[4:8], 0x8c3a91d2) // timestamp_microseconds, sender's clock
|
||||
binary.BigEndian.PutUint32(b[8:12], tsDiff)
|
||||
binary.BigEndian.PutUint32(b[12:16], 0x00100000) // wnd_size
|
||||
binary.BigEndian.PutUint16(b[16:18], 0x71ee) // seq_nr, random in libutp/libtorrent
|
||||
binary.BigEndian.PutUint16(b[18:20], 0x0000) // ack_nr
|
||||
return append(b, payload...)
|
||||
}
|
||||
|
||||
func TestSniffUTP(t *testing.T) {
|
||||
selectiveAck := []byte{0, 4, 0xff, 0x00, 0xff, 0x00}
|
||||
wrongVersion := utpPacket(4, 0, 0)
|
||||
wrongVersion[0] = 4<<4 | 2
|
||||
|
||||
cases := []struct {
|
||||
name string
|
||||
payload []byte
|
||||
err error
|
||||
}{
|
||||
{"syn", utpPacket(4, 0, 0), nil},
|
||||
{"syn with selective ack", append(utpPacket(4, 1, 0), selectiveAck...), nil},
|
||||
{"syn with extension bits", append(utpPacket(4, 2, 0), 0, 8, 1, 2, 3, 4, 5, 6, 7, 8), nil},
|
||||
{"extension bits with wrong length", append(utpPacket(4, 2, 0), 0, 4, 1, 2, 3, 4), errNotBittorrent},
|
||||
{"syn with nonzero timestamp_difference", utpPacket(4, 0, 0x1234), errNotBittorrent},
|
||||
{"syn with trailing payload", utpPacket(4, 0, 0, 'x'), errNotBittorrent},
|
||||
// txid 0x4100, no EDNS0: the worst case colliding with the uTP header
|
||||
{"dns query", []byte{
|
||||
0x41, 0x00, 0x01, 0x00, 0x00, 0x01, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00,
|
||||
0x01, 'a', 0x02, 'c', 'o', 0x00, 0x00, 0x01, 0x00, 0x01,
|
||||
}, errNotBittorrent},
|
||||
{"established connection packets", utpPacket(0, 0, 0x5678, 'x', 'y', 'z'), errNotBittorrent},
|
||||
{"state", utpPacket(2, 0, 0x5678), errNotBittorrent},
|
||||
{"fin", utpPacket(1, 0, 0x5678), errNotBittorrent},
|
||||
{"wrong version", wrongVersion, errNotBittorrent},
|
||||
{"unknown packet type", utpPacket(5, 0, 0), errNotBittorrent},
|
||||
{"unknown extension", utpPacket(4, 2, 0), errNotBittorrent},
|
||||
{"extension chain past the datagram", utpPacket(4, 1, 0, 0, 8, 0xff), errNotBittorrent},
|
||||
{"selective ack not in multiples of 4", append(utpPacket(4, 1, 0), 0, 3, 0xff, 0x00, 0xff), errNotBittorrent},
|
||||
{"shorter than the header", utpPacket(4, 0, 0)[:19], common.ErrNoClue},
|
||||
}
|
||||
|
||||
for _, c := range cases {
|
||||
t.Run(c.name, func(t *testing.T) {
|
||||
h, err := SniffUTP(c.payload)
|
||||
if err != c.err {
|
||||
t.Fatalf("expected error %v, got %v", c.err, err)
|
||||
}
|
||||
if err == nil && h == nil {
|
||||
t.Fatal("expected a sniff header, got nil")
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -28,6 +28,8 @@ const (
|
||||
SecurityType_AUTO SecurityType = 2
|
||||
SecurityType_AES128_GCM SecurityType = 3
|
||||
SecurityType_CHACHA20_POLY1305 SecurityType = 4
|
||||
SecurityType_NONE SecurityType = 5 // [DEPRECATED 2023-06]
|
||||
SecurityType_ZERO SecurityType = 6
|
||||
)
|
||||
|
||||
// Enum value maps for SecurityType.
|
||||
@@ -37,12 +39,16 @@ var (
|
||||
2: "AUTO",
|
||||
3: "AES128_GCM",
|
||||
4: "CHACHA20_POLY1305",
|
||||
5: "NONE",
|
||||
6: "ZERO",
|
||||
}
|
||||
SecurityType_value = map[string]int32{
|
||||
"UNKNOWN": 0,
|
||||
"AUTO": 2,
|
||||
"AES128_GCM": 3,
|
||||
"CHACHA20_POLY1305": 4,
|
||||
"NONE": 5,
|
||||
"ZERO": 6,
|
||||
}
|
||||
)
|
||||
|
||||
@@ -123,13 +129,15 @@ const file_common_protocol_headers_proto_rawDesc = "" +
|
||||
"\n" +
|
||||
"\x1dcommon/protocol/headers.proto\x12\x14xray.common.protocol\"H\n" +
|
||||
"\x0eSecurityConfig\x126\n" +
|
||||
"\x04type\x18\x01 \x01(\x0e2\".xray.common.protocol.SecurityTypeR\x04type*L\n" +
|
||||
"\x04type\x18\x01 \x01(\x0e2\".xray.common.protocol.SecurityTypeR\x04type*`\n" +
|
||||
"\fSecurityType\x12\v\n" +
|
||||
"\aUNKNOWN\x10\x00\x12\b\n" +
|
||||
"\x04AUTO\x10\x02\x12\x0e\n" +
|
||||
"\n" +
|
||||
"AES128_GCM\x10\x03\x12\x15\n" +
|
||||
"\x11CHACHA20_POLY1305\x10\x04B^\n" +
|
||||
"\x11CHACHA20_POLY1305\x10\x04\x12\b\n" +
|
||||
"\x04NONE\x10\x05\x12\b\n" +
|
||||
"\x04ZERO\x10\x06B^\n" +
|
||||
"\x18com.xray.common.protocolP\x01Z)github.com/xtls/xray-core/common/protocol\xaa\x02\x14Xray.Common.Protocolb\x06proto3"
|
||||
|
||||
var (
|
||||
|
||||
@@ -11,6 +11,8 @@ enum SecurityType {
|
||||
AUTO = 2;
|
||||
AES128_GCM = 3;
|
||||
CHACHA20_POLY1305 = 4;
|
||||
NONE = 5; // [DEPRECATED 2023-06]
|
||||
ZERO = 6;
|
||||
}
|
||||
|
||||
message SecurityConfig {
|
||||
|
||||
@@ -1,10 +1,18 @@
|
||||
package quic
|
||||
|
||||
import (
|
||||
"crypto"
|
||||
"crypto/cipher"
|
||||
_ "crypto/tls"
|
||||
_ "unsafe"
|
||||
)
|
||||
|
||||
type CipherSuiteTLS13 struct {
|
||||
ID uint16
|
||||
KeyLen int
|
||||
AEAD func(key, fixedNonce []byte) cipher.AEAD
|
||||
Hash crypto.Hash
|
||||
}
|
||||
|
||||
//go:linkname AEADAESGCMTLS13 crypto/tls.aeadAESGCMTLS13
|
||||
func AEADAESGCMTLS13(key, nonceMask []byte) cipher.AEAD
|
||||
|
||||
@@ -3,6 +3,7 @@ package quic
|
||||
import (
|
||||
"crypto"
|
||||
"crypto/aes"
|
||||
"crypto/tls"
|
||||
"encoding/binary"
|
||||
"io"
|
||||
|
||||
@@ -27,43 +28,22 @@ func (s SniffHeader) Domain() string {
|
||||
return s.domain
|
||||
}
|
||||
|
||||
var (
|
||||
errNotQUIC = errors.New("not quic")
|
||||
errNotQUICInitial = errors.New("not initial packet")
|
||||
const (
|
||||
versionDraft29 uint32 = 0xff00001d
|
||||
version1 uint32 = 0x1
|
||||
)
|
||||
|
||||
type quicVersionSpec struct {
|
||||
ver uint32
|
||||
typeInitial byte
|
||||
initialSalt []byte
|
||||
labelPrefix string
|
||||
}
|
||||
|
||||
var (
|
||||
quicDraft29 = quicVersionSpec{
|
||||
ver: 0xff00001d,
|
||||
typeInitial: 0b00,
|
||||
initialSalt: []byte{0xaf, 0xbf, 0xec, 0x28, 0x99, 0x93, 0xd2, 0x4c, 0x9e, 0x97, 0x86, 0xf1, 0x9c, 0x61, 0x11, 0xe0, 0x43, 0x90, 0xa8, 0x99},
|
||||
labelPrefix: "quic",
|
||||
}
|
||||
quicV1 = quicVersionSpec{
|
||||
ver: 0x1,
|
||||
typeInitial: 0b00,
|
||||
initialSalt: []byte{0x38, 0x76, 0x2c, 0xf7, 0xf5, 0x59, 0x34, 0xb3, 0x4d, 0x17, 0x9a, 0xe6, 0xa4, 0xc8, 0x0c, 0xad, 0xcc, 0xbb, 0x7f, 0x0a},
|
||||
labelPrefix: "quic",
|
||||
}
|
||||
quicV2 = quicVersionSpec{
|
||||
ver: 0x6b3343cf,
|
||||
typeInitial: 0b01,
|
||||
initialSalt: []byte{0x0d, 0xed, 0xe3, 0xde, 0xf7, 0x00, 0xa6, 0xdb, 0x81, 0x93, 0x81, 0xbe, 0x6e, 0x26, 0x9d, 0xcb, 0xf9, 0xbd, 0x2e, 0xd9},
|
||||
labelPrefix: "quicv2",
|
||||
}
|
||||
|
||||
quicVersionSpecMap = map[uint32]*quicVersionSpec{
|
||||
quicDraft29.ver: &quicDraft29,
|
||||
quicV1.ver: &quicV1,
|
||||
quicV2.ver: &quicV2,
|
||||
quicSaltOld = []byte{0xaf, 0xbf, 0xec, 0x28, 0x99, 0x93, 0xd2, 0x4c, 0x9e, 0x97, 0x86, 0xf1, 0x9c, 0x61, 0x11, 0xe0, 0x43, 0x90, 0xa8, 0x99}
|
||||
quicSalt = []byte{0x38, 0x76, 0x2c, 0xf7, 0xf5, 0x59, 0x34, 0xb3, 0x4d, 0x17, 0x9a, 0xe6, 0xa4, 0xc8, 0x0c, 0xad, 0xcc, 0xbb, 0x7f, 0x0a}
|
||||
initialSuite = &CipherSuiteTLS13{
|
||||
ID: tls.TLS_AES_128_GCM_SHA256,
|
||||
KeyLen: 16,
|
||||
AEAD: AEADAESGCMTLS13,
|
||||
Hash: crypto.SHA256,
|
||||
}
|
||||
errNotQuic = errors.New("not quic")
|
||||
errNotQuicInitial = errors.New("not initial packet")
|
||||
)
|
||||
|
||||
func SniffQUIC(b []byte) (*SniffHeader, error) {
|
||||
@@ -83,61 +63,60 @@ func SniffQUIC(b []byte) (*SniffHeader, error) {
|
||||
buffer := buf.FromBytes(b)
|
||||
typeByte, err := buffer.ReadByte()
|
||||
if err != nil {
|
||||
return nil, errNotQUIC
|
||||
return nil, errNotQuic
|
||||
}
|
||||
|
||||
isLongHeader := typeByte&0x80 > 0
|
||||
if !isLongHeader || typeByte&0x40 == 0 {
|
||||
return nil, errNotQUICInitial
|
||||
return nil, errNotQuicInitial
|
||||
}
|
||||
|
||||
vb, err := buffer.ReadBytes(4)
|
||||
if err != nil {
|
||||
return nil, errNotQUIC
|
||||
return nil, errNotQuic
|
||||
}
|
||||
|
||||
versionNumber := binary.BigEndian.Uint32(vb)
|
||||
var s *quicVersionSpec
|
||||
if v, ok := quicVersionSpecMap[versionNumber]; ok {
|
||||
s = v
|
||||
} else {
|
||||
return nil, errNotQUIC
|
||||
}
|
||||
|
||||
var destConnID []byte
|
||||
if l, err := buffer.ReadByte(); err != nil {
|
||||
return nil, errNotQUIC
|
||||
} else if destConnID, err = buffer.ReadBytes(int32(l)); err != nil {
|
||||
return nil, errNotQUIC
|
||||
}
|
||||
|
||||
if l, err := buffer.ReadByte(); err != nil {
|
||||
return nil, errNotQUIC
|
||||
} else if common.Error2(buffer.ReadBytes(int32(l))) != nil {
|
||||
return nil, errNotQUIC
|
||||
if versionNumber != 0 && typeByte&0x40 == 0 {
|
||||
return nil, errNotQuic
|
||||
} else if versionNumber != versionDraft29 && versionNumber != version1 {
|
||||
return nil, errNotQuic
|
||||
}
|
||||
|
||||
packetType := (typeByte & 0x30) >> 4
|
||||
isQUICInitial := packetType == s.typeInitial
|
||||
isQuicInitial := packetType == 0x0
|
||||
|
||||
if isQUICInitial { // Only initial packets have token, see https://datatracker.ietf.org/doc/html/rfc9000#section-17.2.2
|
||||
tokenLen, err := readShortQUICVarint(buffer)
|
||||
var destConnID []byte
|
||||
if l, err := buffer.ReadByte(); err != nil {
|
||||
return nil, errNotQuic
|
||||
} else if destConnID, err = buffer.ReadBytes(int32(l)); err != nil {
|
||||
return nil, errNotQuic
|
||||
}
|
||||
|
||||
if l, err := buffer.ReadByte(); err != nil {
|
||||
return nil, errNotQuic
|
||||
} else if common.Error2(buffer.ReadBytes(int32(l))) != nil {
|
||||
return nil, errNotQuic
|
||||
}
|
||||
|
||||
if isQuicInitial { // Only initial packets have token, see https://datatracker.ietf.org/doc/html/rfc9000#section-17.2.2
|
||||
tokenLen, err := readShortQuicVarint(buffer)
|
||||
if err != nil || tokenLen > int32(len(b)) {
|
||||
return nil, errNotQUIC
|
||||
return nil, errNotQuic
|
||||
}
|
||||
|
||||
if _, err = buffer.ReadBytes(tokenLen); err != nil {
|
||||
return nil, errNotQUIC
|
||||
return nil, errNotQuic
|
||||
}
|
||||
}
|
||||
|
||||
packetLen, err := readShortQUICVarint(buffer)
|
||||
packetLen, err := readShortQuicVarint(buffer)
|
||||
if err != nil {
|
||||
return nil, errNotQUIC
|
||||
return nil, errNotQuic
|
||||
}
|
||||
// packetLen is impossible to be shorter than this
|
||||
if packetLen < 4 {
|
||||
return nil, errNotQUIC
|
||||
return nil, errNotQuic
|
||||
}
|
||||
|
||||
hdrLen := len(b) - int(buffer.Len())
|
||||
@@ -146,23 +125,25 @@ func SniffQUIC(b []byte) (*SniffHeader, error) {
|
||||
}
|
||||
|
||||
restPayload := b[hdrLen+int(packetLen):]
|
||||
if !isQUICInitial { // Skip this packet if it's not initial packet
|
||||
if !isQuicInitial { // Skip this packet if it's not initial packet
|
||||
b = restPayload
|
||||
continue
|
||||
}
|
||||
|
||||
salt := s.initialSalt
|
||||
label := s.labelPrefix
|
||||
var salt []byte
|
||||
if versionNumber == version1 {
|
||||
salt = quicSalt
|
||||
} else {
|
||||
salt = quicSaltOld
|
||||
}
|
||||
initialSecret := hkdf.Extract(crypto.SHA256.New, destConnID, salt)
|
||||
secret := hkdfExpandLabel(initialSecret, "client in", crypto.SHA256.Size())
|
||||
hpKey := hkdfExpandLabel(secret, label+" hp", 16)
|
||||
secret := hkdfExpandLabel(crypto.SHA256, initialSecret, []byte{}, "client in", crypto.SHA256.Size())
|
||||
hpKey := hkdfExpandLabel(initialSuite.Hash, secret, []byte{}, "quic hp", initialSuite.KeyLen)
|
||||
block, err := aes.NewCipher(hpKey)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if len(b) < hdrLen+4+block.BlockSize() {
|
||||
return nil, errNotQUIC
|
||||
}
|
||||
|
||||
cache.Clear()
|
||||
mask := cache.Extend(int32(block.BlockSize()))
|
||||
block.Encrypt(mask, b[hdrLen+4:hdrLen+4+len(mask)])
|
||||
@@ -172,8 +153,8 @@ func SniffQUIC(b []byte) (*SniffHeader, error) {
|
||||
b[hdrLen+i] ^= mask[i+1]
|
||||
}
|
||||
|
||||
key := hkdfExpandLabel(secret, label+" key", 16)
|
||||
iv := hkdfExpandLabel(secret, label+" iv", 12)
|
||||
key := hkdfExpandLabel(crypto.SHA256, secret, []byte{}, "quic key", 16)
|
||||
iv := hkdfExpandLabel(crypto.SHA256, secret, []byte{}, "quic iv", 12)
|
||||
cipher := AEADAESGCMTLS13(key, iv)
|
||||
|
||||
nonce := cache.Extend(int32(cipher.NonceSize()))
|
||||
@@ -198,44 +179,44 @@ func SniffQUIC(b []byte) (*SniffHeader, error) {
|
||||
case 0x00: // PADDING frame
|
||||
case 0x01: // PING frame
|
||||
case 0x02, 0x03: // ACK frame
|
||||
if _, err = readShortQUICVarint(buffer); err != nil { // Field: Largest Acknowledged
|
||||
if _, err = readShortQuicVarint(buffer); err != nil { // Field: Largest Acknowledged
|
||||
return nil, io.ErrUnexpectedEOF
|
||||
}
|
||||
if _, err = readShortQUICVarint(buffer); err != nil { // Field: ACK Delay
|
||||
if _, err = readShortQuicVarint(buffer); err != nil { // Field: ACK Delay
|
||||
return nil, io.ErrUnexpectedEOF
|
||||
}
|
||||
ackRangeCount, err := readShortQUICVarint(buffer) // Field: ACK Range Count
|
||||
ackRangeCount, err := readShortQuicVarint(buffer) // Field: ACK Range Count
|
||||
if err != nil {
|
||||
return nil, io.ErrUnexpectedEOF
|
||||
}
|
||||
if _, err = readShortQUICVarint(buffer); err != nil { // Field: First ACK Range
|
||||
if _, err = readShortQuicVarint(buffer); err != nil { // Field: First ACK Range
|
||||
return nil, io.ErrUnexpectedEOF
|
||||
}
|
||||
for i := 0; i < int(ackRangeCount); i++ { // Field: ACK Range
|
||||
if _, err = readShortQUICVarint(buffer); err != nil { // Field: ACK Range -> Gap
|
||||
if _, err = readShortQuicVarint(buffer); err != nil { // Field: ACK Range -> Gap
|
||||
return nil, io.ErrUnexpectedEOF
|
||||
}
|
||||
if _, err = readShortQUICVarint(buffer); err != nil { // Field: ACK Range -> ACK Range Length
|
||||
if _, err = readShortQuicVarint(buffer); err != nil { // Field: ACK Range -> ACK Range Length
|
||||
return nil, io.ErrUnexpectedEOF
|
||||
}
|
||||
}
|
||||
if frameType == 0x03 {
|
||||
if _, err = readShortQUICVarint(buffer); err != nil { // Field: ECN Counts -> ECT0 Count
|
||||
if _, err = readShortQuicVarint(buffer); err != nil { // Field: ECN Counts -> ECT0 Count
|
||||
return nil, io.ErrUnexpectedEOF
|
||||
}
|
||||
if _, err = readShortQUICVarint(buffer); err != nil { // Field: ECN Counts -> ECT1 Count
|
||||
if _, err = readShortQuicVarint(buffer); err != nil { // Field: ECN Counts -> ECT1 Count
|
||||
return nil, io.ErrUnexpectedEOF
|
||||
}
|
||||
if _, err = readShortQUICVarint(buffer); err != nil { //nolint:misspell // Field: ECN Counts -> ECT-CE Count
|
||||
if _, err = readShortQuicVarint(buffer); err != nil { //nolint:misspell // Field: ECN Counts -> ECT-CE Count
|
||||
return nil, io.ErrUnexpectedEOF
|
||||
}
|
||||
}
|
||||
case 0x06: // CRYPTO frame, we will use this frame
|
||||
offset, err := readShortQUICVarint(buffer) // Field: Offset
|
||||
offset, err := readShortQuicVarint(buffer) // Field: Offset
|
||||
if err != nil {
|
||||
return nil, io.ErrUnexpectedEOF
|
||||
}
|
||||
length, err := readShortQUICVarint(buffer) // Field: Length
|
||||
length, err := readShortQuicVarint(buffer) // Field: Length
|
||||
if err != nil || length > buffer.Len() {
|
||||
return nil, io.ErrUnexpectedEOF
|
||||
}
|
||||
@@ -251,13 +232,13 @@ func SniffQUIC(b []byte) (*SniffHeader, error) {
|
||||
return nil, io.ErrUnexpectedEOF
|
||||
}
|
||||
case 0x1c: // CONNECTION_CLOSE frame, only 0x1c is permitted in initial packet
|
||||
if _, err = readShortQUICVarint(buffer); err != nil { // Field: Error Code
|
||||
if _, err = readShortQuicVarint(buffer); err != nil { // Field: Error Code
|
||||
return nil, io.ErrUnexpectedEOF
|
||||
}
|
||||
if _, err = readShortQUICVarint(buffer); err != nil { // Field: Frame Type
|
||||
if _, err = readShortQuicVarint(buffer); err != nil { // Field: Frame Type
|
||||
return nil, io.ErrUnexpectedEOF
|
||||
}
|
||||
length, err := readShortQUICVarint(buffer) // Field: Reason Phrase Length
|
||||
length, err := readShortQuicVarint(buffer) // Field: Reason Phrase Length
|
||||
if err != nil {
|
||||
return nil, io.ErrUnexpectedEOF
|
||||
}
|
||||
@@ -267,7 +248,7 @@ func SniffQUIC(b []byte) (*SniffHeader, error) {
|
||||
default:
|
||||
// Only above frame types are permitted in initial packet.
|
||||
// See https://www.rfc-editor.org/rfc/rfc9000.html#section-17.2.2-8
|
||||
return nil, errNotQUICInitial
|
||||
return nil, errNotQuicInitial
|
||||
}
|
||||
}
|
||||
|
||||
@@ -285,33 +266,35 @@ func SniffQUIC(b []byte) (*SniffHeader, error) {
|
||||
return nil, protocol.ErrProtoNeedMoreData
|
||||
}
|
||||
|
||||
func hkdfExpandLabel(secret []byte, label string, length int) []byte {
|
||||
b := make([]byte, 0, 2+1+6+len(label)+1)
|
||||
b = binary.BigEndian.AppendUint16(b, uint16(length))
|
||||
b = append(b, byte(6+len(label)))
|
||||
b = append(b, "tls13 "...)
|
||||
b = append(b, label...)
|
||||
b = append(b, 0) // context
|
||||
func hkdfExpandLabel(hash crypto.Hash, secret, context []byte, label string, length int) []byte {
|
||||
b := make([]byte, 3, 3+6+len(label)+1+len(context))
|
||||
binary.BigEndian.PutUint16(b, uint16(length))
|
||||
b[2] = uint8(6 + len(label))
|
||||
b = append(b, []byte("tls13 ")...)
|
||||
b = append(b, []byte(label)...)
|
||||
b = b[:3+6+len(label)+1]
|
||||
b[3+6+len(label)] = uint8(len(context))
|
||||
b = append(b, context...)
|
||||
|
||||
out := make([]byte, length)
|
||||
n, err := hkdf.Expand(crypto.SHA256.New, secret, b).Read(out)
|
||||
n, err := hkdf.Expand(hash.New, secret, b).Read(out)
|
||||
if err != nil || n != length {
|
||||
panic("quic: HKDF-Expand-Label invocation failed unexpectedly")
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
// readShortQUICVarint wraps quicvarint.Read with a max limit for length related fields.
|
||||
// readShortQuicVarint wraps quicvarint.Read with a max limit for length related fields.
|
||||
// we only handle QUIC Initial so these numbers should not exceed 65535
|
||||
// returns int32 to reduce type conversion
|
||||
func readShortQUICVarint(reader io.ByteReader) (int32, error) {
|
||||
func readShortQuicVarint(reader io.ByteReader) (int32, error) {
|
||||
v, err := quicvarint.Read(reader)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
if v > 65535 {
|
||||
// not used(
|
||||
return 0, errNotQUICInitial
|
||||
return 0, errNotQuicInitial
|
||||
}
|
||||
return int32(v), nil
|
||||
}
|
||||
|
||||
File diff suppressed because one or more lines are too long
@@ -7,7 +7,7 @@ import (
|
||||
|
||||
func (u *User) GetTypedAccount() (Account, error) {
|
||||
if u.GetAccount() == nil {
|
||||
return nil, errors.New("Account is missing")
|
||||
return nil, errors.New("Account is missing").AtWarning()
|
||||
}
|
||||
|
||||
rawAccount, err := u.Account.GetInstance()
|
||||
|
||||
@@ -207,7 +207,6 @@ func getConfig() string {
|
||||
"tag": "XHTTP_IN",
|
||||
"streamSettings": {
|
||||
"network": "xhttp",
|
||||
"security": "tls",
|
||||
"xhttpSettings": {
|
||||
"host": "bing.com",
|
||||
"path": "/xhttp_client_upload",
|
||||
|
||||
@@ -70,6 +70,8 @@ type Outbound struct {
|
||||
Tag string
|
||||
// Name of the outbound proxy that handles the connection.
|
||||
Name string
|
||||
// Unused. Conn is actually internet.Connection. May be nil. It is currently nil for outbound with proxySettings
|
||||
Conn net.Conn
|
||||
// CanSpliceCopy is a property for this connection
|
||||
// 1 = can, 2 = after processing protocol info should be able to, 3 = cannot
|
||||
CanSpliceCopy int
|
||||
|
||||
@@ -0,0 +1,53 @@
|
||||
package singbridge
|
||||
|
||||
import (
|
||||
M "github.com/sagernet/sing/common/metadata"
|
||||
N "github.com/sagernet/sing/common/network"
|
||||
"github.com/xtls/xray-core/common/errors"
|
||||
"github.com/xtls/xray-core/common/net"
|
||||
)
|
||||
|
||||
func ToNetwork(network string) net.Network {
|
||||
switch N.NetworkName(network) {
|
||||
case N.NetworkTCP:
|
||||
return net.Network_TCP
|
||||
case N.NetworkUDP:
|
||||
return net.Network_UDP
|
||||
default:
|
||||
return net.Network_Unknown
|
||||
}
|
||||
}
|
||||
|
||||
func ToDestination(socksaddr M.Socksaddr, network net.Network) (net.Destination, error) {
|
||||
// IsFqdn() implicitly checks if the domain name is valid
|
||||
if socksaddr.IsFqdn() {
|
||||
return net.Destination{
|
||||
Network: network,
|
||||
Address: net.DomainAddress(socksaddr.Fqdn),
|
||||
Port: net.Port(socksaddr.Port),
|
||||
}, nil
|
||||
}
|
||||
|
||||
// IsIP() implicitly checks if the IP address is valid
|
||||
if socksaddr.IsIP() {
|
||||
return net.Destination{
|
||||
Network: network,
|
||||
Address: net.IPAddress(socksaddr.Addr.AsSlice()),
|
||||
Port: net.Port(socksaddr.Port),
|
||||
}, nil
|
||||
}
|
||||
|
||||
return net.Destination{}, errors.New("invalid socks address: ", socksaddr)
|
||||
}
|
||||
|
||||
func ToSocksaddr(destination net.Destination) M.Socksaddr {
|
||||
var addr M.Socksaddr
|
||||
switch destination.Address.Family() {
|
||||
case net.AddressFamilyDomain:
|
||||
addr.Fqdn = destination.Address.Domain()
|
||||
default:
|
||||
addr.Addr = M.AddrFromIP(destination.Address.IP())
|
||||
}
|
||||
addr.Port = uint16(destination.Port)
|
||||
return addr
|
||||
}
|
||||
@@ -0,0 +1,72 @@
|
||||
package singbridge
|
||||
|
||||
import (
|
||||
"context"
|
||||
"os"
|
||||
|
||||
M "github.com/sagernet/sing/common/metadata"
|
||||
N "github.com/sagernet/sing/common/network"
|
||||
"github.com/xtls/xray-core/common/net"
|
||||
"github.com/xtls/xray-core/common/net/cnc"
|
||||
"github.com/xtls/xray-core/common/session"
|
||||
"github.com/xtls/xray-core/proxy"
|
||||
"github.com/xtls/xray-core/transport"
|
||||
"github.com/xtls/xray-core/transport/internet"
|
||||
"github.com/xtls/xray-core/transport/pipe"
|
||||
)
|
||||
|
||||
var _ N.Dialer = (*XrayDialer)(nil)
|
||||
|
||||
type XrayDialer struct {
|
||||
internet.Dialer
|
||||
}
|
||||
|
||||
func NewDialer(dialer internet.Dialer) *XrayDialer {
|
||||
return &XrayDialer{dialer}
|
||||
}
|
||||
|
||||
func (d *XrayDialer) DialContext(ctx context.Context, network string, destination M.Socksaddr) (net.Conn, error) {
|
||||
dest, err := ToDestination(destination, ToNetwork(network))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return d.Dialer.Dial(ctx, dest)
|
||||
}
|
||||
|
||||
func (d *XrayDialer) ListenPacket(ctx context.Context, destination M.Socksaddr) (net.PacketConn, error) {
|
||||
return nil, os.ErrInvalid
|
||||
}
|
||||
|
||||
type XrayOutboundDialer struct {
|
||||
outbound proxy.Outbound
|
||||
dialer internet.Dialer
|
||||
}
|
||||
|
||||
func NewOutboundDialer(outbound proxy.Outbound, dialer internet.Dialer) *XrayOutboundDialer {
|
||||
return &XrayOutboundDialer{outbound, dialer}
|
||||
}
|
||||
|
||||
func (d *XrayOutboundDialer) DialContext(ctx context.Context, network string, destination M.Socksaddr) (net.Conn, error) {
|
||||
dest, err := ToDestination(destination, ToNetwork(network))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
outbounds := session.OutboundsFromContext(ctx)
|
||||
if len(outbounds) == 0 {
|
||||
outbounds = []*session.Outbound{{}}
|
||||
ctx = session.ContextWithOutbounds(ctx, outbounds)
|
||||
}
|
||||
ob := outbounds[len(outbounds)-1]
|
||||
ob.Target = dest
|
||||
|
||||
opts := []pipe.Option{pipe.WithSizeLimit(64 * 1024)}
|
||||
uplinkReader, uplinkWriter := pipe.New(opts...)
|
||||
downlinkReader, downlinkWriter := pipe.New(opts...)
|
||||
conn := cnc.NewConnection(cnc.ConnectionInputMulti(downlinkWriter), cnc.ConnectionOutputMulti(uplinkReader))
|
||||
go d.outbound.Process(ctx, &transport.Link{Reader: downlinkReader, Writer: uplinkWriter}, d.dialer)
|
||||
return conn, nil
|
||||
}
|
||||
|
||||
func (d *XrayOutboundDialer) ListenPacket(ctx context.Context, destination M.Socksaddr) (net.PacketConn, error) {
|
||||
return nil, os.ErrInvalid
|
||||
}
|
||||
@@ -0,0 +1,10 @@
|
||||
package singbridge
|
||||
|
||||
import E "github.com/sagernet/sing/common/exceptions"
|
||||
|
||||
func ReturnError(err error) error {
|
||||
if E.IsClosedOrCanceled(err) {
|
||||
return nil
|
||||
}
|
||||
return err
|
||||
}
|
||||
@@ -0,0 +1,58 @@
|
||||
package singbridge
|
||||
|
||||
import (
|
||||
"context"
|
||||
"io"
|
||||
|
||||
M "github.com/sagernet/sing/common/metadata"
|
||||
N "github.com/sagernet/sing/common/network"
|
||||
"github.com/xtls/xray-core/common/buf"
|
||||
"github.com/xtls/xray-core/common/errors"
|
||||
"github.com/xtls/xray-core/common/net"
|
||||
"github.com/xtls/xray-core/features/routing"
|
||||
"github.com/xtls/xray-core/transport"
|
||||
)
|
||||
|
||||
var (
|
||||
_ N.TCPConnectionHandler = (*Dispatcher)(nil)
|
||||
_ N.UDPConnectionHandler = (*Dispatcher)(nil)
|
||||
)
|
||||
|
||||
type Dispatcher struct {
|
||||
upstream routing.Dispatcher
|
||||
newErrorFunc func(values ...any) *errors.Error
|
||||
}
|
||||
|
||||
func NewDispatcher(dispatcher routing.Dispatcher, newErrorFunc func(values ...any) *errors.Error) *Dispatcher {
|
||||
return &Dispatcher{
|
||||
upstream: dispatcher,
|
||||
newErrorFunc: newErrorFunc,
|
||||
}
|
||||
}
|
||||
|
||||
func (d *Dispatcher) NewConnection(ctx context.Context, conn net.Conn, metadata M.Metadata) error {
|
||||
dest, err := ToDestination(metadata.Destination, net.Network_TCP)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
xConn := NewConn(conn)
|
||||
return d.upstream.DispatchLink(ctx, dest, &transport.Link{
|
||||
Reader: xConn,
|
||||
Writer: xConn,
|
||||
})
|
||||
}
|
||||
|
||||
func (d *Dispatcher) NewPacketConnection(ctx context.Context, conn N.PacketConn, metadata M.Metadata) error {
|
||||
dest, err := ToDestination(metadata.Destination, net.Network_UDP)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return d.upstream.DispatchLink(ctx, dest, &transport.Link{
|
||||
Reader: buf.NewPacketReader(conn.(io.Reader)),
|
||||
Writer: buf.NewWriter(conn.(io.Writer)),
|
||||
})
|
||||
}
|
||||
|
||||
func (d *Dispatcher) NewError(ctx context.Context, err error) {
|
||||
errors.LogInfo(ctx, err.Error())
|
||||
}
|
||||
@@ -0,0 +1,70 @@
|
||||
package singbridge
|
||||
|
||||
import (
|
||||
"context"
|
||||
|
||||
"github.com/sagernet/sing/common/logger"
|
||||
"github.com/xtls/xray-core/common/errors"
|
||||
)
|
||||
|
||||
var _ logger.ContextLogger = (*XrayLogger)(nil)
|
||||
|
||||
type XrayLogger struct {
|
||||
newError func(values ...any) *errors.Error
|
||||
}
|
||||
|
||||
func NewLogger(newErrorFunc func(values ...any) *errors.Error) *XrayLogger {
|
||||
return &XrayLogger{
|
||||
newErrorFunc,
|
||||
}
|
||||
}
|
||||
|
||||
func (l *XrayLogger) Trace(args ...any) {
|
||||
}
|
||||
|
||||
func (l *XrayLogger) Debug(args ...any) {
|
||||
errors.LogDebug(context.Background(), args...)
|
||||
}
|
||||
|
||||
func (l *XrayLogger) Info(args ...any) {
|
||||
errors.LogInfo(context.Background(), args...)
|
||||
}
|
||||
|
||||
func (l *XrayLogger) Warn(args ...any) {
|
||||
errors.LogWarning(context.Background(), args...)
|
||||
}
|
||||
|
||||
func (l *XrayLogger) Error(args ...any) {
|
||||
errors.LogError(context.Background(), args...)
|
||||
}
|
||||
|
||||
func (l *XrayLogger) Fatal(args ...any) {
|
||||
}
|
||||
|
||||
func (l *XrayLogger) Panic(args ...any) {
|
||||
}
|
||||
|
||||
func (l *XrayLogger) TraceContext(ctx context.Context, args ...any) {
|
||||
}
|
||||
|
||||
func (l *XrayLogger) DebugContext(ctx context.Context, args ...any) {
|
||||
errors.LogDebug(ctx, args...)
|
||||
}
|
||||
|
||||
func (l *XrayLogger) InfoContext(ctx context.Context, args ...any) {
|
||||
errors.LogInfo(ctx, args...)
|
||||
}
|
||||
|
||||
func (l *XrayLogger) WarnContext(ctx context.Context, args ...any) {
|
||||
errors.LogWarning(ctx, args...)
|
||||
}
|
||||
|
||||
func (l *XrayLogger) ErrorContext(ctx context.Context, args ...any) {
|
||||
errors.LogError(ctx, args...)
|
||||
}
|
||||
|
||||
func (l *XrayLogger) FatalContext(ctx context.Context, args ...any) {
|
||||
}
|
||||
|
||||
func (l *XrayLogger) PanicContext(ctx context.Context, args ...any) {
|
||||
}
|
||||
@@ -0,0 +1,107 @@
|
||||
package singbridge
|
||||
|
||||
import (
|
||||
"context"
|
||||
"time"
|
||||
|
||||
B "github.com/sagernet/sing/common/buf"
|
||||
"github.com/sagernet/sing/common/bufio"
|
||||
M "github.com/sagernet/sing/common/metadata"
|
||||
"github.com/xtls/xray-core/common"
|
||||
"github.com/xtls/xray-core/common/buf"
|
||||
"github.com/xtls/xray-core/common/net"
|
||||
"github.com/xtls/xray-core/common/signal"
|
||||
"github.com/xtls/xray-core/transport"
|
||||
)
|
||||
|
||||
func CopyPacketConn(ctx context.Context, inboundConn net.Conn, link *transport.Link, destination net.Destination, serverConn net.PacketConn) error {
|
||||
cancel := func() {
|
||||
common.Interrupt(link.Reader)
|
||||
common.Interrupt(serverConn)
|
||||
}
|
||||
conn := &PacketConnWrapper{
|
||||
Reader: link.Reader,
|
||||
Writer: link.Writer,
|
||||
Dest: destination,
|
||||
Conn: inboundConn,
|
||||
T: signal.CancelAfterInactivity(ctx, cancel, 300*time.Second),
|
||||
}
|
||||
return ReturnError(bufio.CopyPacketConn(ctx, conn, bufio.NewPacketConn(serverConn)))
|
||||
}
|
||||
|
||||
type PacketConnWrapper struct {
|
||||
buf.Reader
|
||||
buf.Writer
|
||||
net.Conn
|
||||
Dest net.Destination
|
||||
cached buf.MultiBuffer
|
||||
|
||||
// A simple patch to avoid goroutine leak since sing infra cannot awake read block by write err
|
||||
T *signal.ActivityTimer
|
||||
}
|
||||
|
||||
func (w *PacketConnWrapper) ReadPacket(buffer *B.Buffer) (addr M.Socksaddr, err error) {
|
||||
w.T.Update()
|
||||
defer func() {
|
||||
if err != nil {
|
||||
// uplinkonly
|
||||
w.T.SetTimeout(2 * time.Second)
|
||||
}
|
||||
}()
|
||||
if w.cached != nil {
|
||||
mb, bb := buf.SplitFirst(w.cached)
|
||||
if bb == nil {
|
||||
w.cached = nil
|
||||
} else {
|
||||
buffer.Write(bb.Bytes())
|
||||
w.cached = mb
|
||||
var destination net.Destination
|
||||
if bb.UDP != nil {
|
||||
destination = *bb.UDP
|
||||
} else {
|
||||
destination = w.Dest
|
||||
}
|
||||
bb.Release()
|
||||
return ToSocksaddr(destination), nil
|
||||
}
|
||||
}
|
||||
mb, err := w.ReadMultiBuffer()
|
||||
nb, bb := buf.SplitFirst(mb)
|
||||
if bb == nil {
|
||||
return M.Socksaddr{}, nil
|
||||
} else {
|
||||
buffer.Write(bb.Bytes())
|
||||
w.cached = nb
|
||||
var destination net.Destination
|
||||
if bb.UDP != nil {
|
||||
destination = *bb.UDP
|
||||
} else {
|
||||
destination = w.Dest
|
||||
}
|
||||
bb.Release()
|
||||
return ToSocksaddr(destination), nil
|
||||
}
|
||||
}
|
||||
|
||||
func (w *PacketConnWrapper) WritePacket(buffer *B.Buffer, destination M.Socksaddr) (err error) {
|
||||
w.T.Update()
|
||||
defer func() {
|
||||
if err != nil {
|
||||
// downlinkonly
|
||||
w.T.SetTimeout(5 * time.Second)
|
||||
}
|
||||
}()
|
||||
endpoint, err := ToDestination(destination, net.Network_UDP)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
vBuf := buf.New()
|
||||
vBuf.Write(buffer.Bytes())
|
||||
vBuf.UDP = &endpoint
|
||||
return w.WriteMultiBuffer(buf.MultiBuffer{vBuf})
|
||||
}
|
||||
|
||||
func (w *PacketConnWrapper) Close() error {
|
||||
buf.ReleaseMulti(w.cached)
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,81 @@
|
||||
package singbridge
|
||||
|
||||
import (
|
||||
"context"
|
||||
"io"
|
||||
"net"
|
||||
"time"
|
||||
|
||||
"github.com/sagernet/sing/common/bufio"
|
||||
"github.com/xtls/xray-core/common"
|
||||
"github.com/xtls/xray-core/common/buf"
|
||||
"github.com/xtls/xray-core/common/signal"
|
||||
"github.com/xtls/xray-core/transport"
|
||||
)
|
||||
|
||||
func CopyConn(ctx context.Context, inboundConn net.Conn, link *transport.Link, serverConn net.Conn) error {
|
||||
conn := &PipeConnWrapper{
|
||||
W: link.Writer,
|
||||
Conn: inboundConn,
|
||||
}
|
||||
if ir, ok := link.Reader.(io.Reader); ok {
|
||||
conn.R = ir
|
||||
} else {
|
||||
conn.R = &buf.BufferedReader{Reader: link.Reader}
|
||||
}
|
||||
cancel := func() {
|
||||
common.Interrupt(link.Reader)
|
||||
common.Interrupt(serverConn)
|
||||
}
|
||||
conn.T = signal.CancelAfterInactivity(ctx, cancel, 300*time.Second)
|
||||
return ReturnError(bufio.CopyConn(ctx, conn, serverConn))
|
||||
}
|
||||
|
||||
type PipeConnWrapper struct {
|
||||
R io.Reader
|
||||
W buf.Writer
|
||||
net.Conn
|
||||
|
||||
// A simple patch to avoid goroutine leak since sing infra cannot awake read block by write err
|
||||
T *signal.ActivityTimer
|
||||
}
|
||||
|
||||
func (w *PipeConnWrapper) Close() error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func (w *PipeConnWrapper) Read(b []byte) (n int, err error) {
|
||||
w.T.Update()
|
||||
n, err = w.R.Read(b)
|
||||
if err != nil {
|
||||
// uplinkonly
|
||||
w.T.SetTimeout(2 * time.Second)
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
func (w *PipeConnWrapper) Write(p []byte) (n int, err error) {
|
||||
w.T.Update()
|
||||
n = len(p)
|
||||
var mb buf.MultiBuffer
|
||||
pLen := len(p)
|
||||
for pLen > 0 {
|
||||
buffer := buf.New()
|
||||
if pLen > buf.Size {
|
||||
_, err = buffer.Write(p[:buf.Size])
|
||||
p = p[buf.Size:]
|
||||
} else {
|
||||
buffer.Write(p)
|
||||
}
|
||||
pLen -= int(buffer.Len())
|
||||
mb = append(mb, buffer)
|
||||
}
|
||||
err = w.W.WriteMultiBuffer(mb)
|
||||
if err != nil {
|
||||
n = 0
|
||||
buf.ReleaseMulti(mb)
|
||||
// downlinkonly
|
||||
w.T.SetTimeout(5 * time.Second)
|
||||
}
|
||||
return
|
||||
}
|
||||
@@ -0,0 +1,66 @@
|
||||
package singbridge
|
||||
|
||||
import (
|
||||
"time"
|
||||
|
||||
"github.com/sagernet/sing/common"
|
||||
"github.com/sagernet/sing/common/bufio"
|
||||
N "github.com/sagernet/sing/common/network"
|
||||
"github.com/xtls/xray-core/common/buf"
|
||||
"github.com/xtls/xray-core/common/net"
|
||||
)
|
||||
|
||||
var (
|
||||
_ buf.Reader = (*Conn)(nil)
|
||||
_ buf.TimeoutReader = (*Conn)(nil)
|
||||
_ buf.Writer = (*Conn)(nil)
|
||||
)
|
||||
|
||||
type Conn struct {
|
||||
net.Conn
|
||||
writer N.VectorisedWriter
|
||||
}
|
||||
|
||||
func NewConn(conn net.Conn) *Conn {
|
||||
writer, _ := bufio.CreateVectorisedWriter(conn)
|
||||
return &Conn{
|
||||
Conn: conn,
|
||||
writer: writer,
|
||||
}
|
||||
}
|
||||
|
||||
func (c *Conn) ReadMultiBuffer() (buf.MultiBuffer, error) {
|
||||
buffer, err := buf.ReadBuffer(c.Conn)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return buf.MultiBuffer{buffer}, nil
|
||||
}
|
||||
|
||||
func (c *Conn) ReadMultiBufferTimeout(duration time.Duration) (buf.MultiBuffer, error) {
|
||||
err := c.SetReadDeadline(time.Now().Add(duration))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer c.SetReadDeadline(time.Time{})
|
||||
return c.ReadMultiBuffer()
|
||||
}
|
||||
|
||||
func (c *Conn) WriteMultiBuffer(bufferList buf.MultiBuffer) error {
|
||||
defer buf.ReleaseMulti(bufferList)
|
||||
if c.writer != nil {
|
||||
bytesList := make([][]byte, len(bufferList))
|
||||
for i, buffer := range bufferList {
|
||||
bytesList[i] = buffer.Bytes()
|
||||
}
|
||||
return common.Error(bufio.WriteVectorised(c.writer, bytesList))
|
||||
}
|
||||
// Since this conn is only used by tun, we don't force buffer writes to merge.
|
||||
for _, buffer := range bufferList {
|
||||
_, err := c.Conn.Write(buffer.Bytes())
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
+2
-2
@@ -16,7 +16,7 @@ var typeCreatorRegistry = make(map[reflect.Type]ConfigCreator)
|
||||
func RegisterConfig(config interface{}, configCreator ConfigCreator) error {
|
||||
configType := reflect.TypeOf(config)
|
||||
if _, found := typeCreatorRegistry[configType]; found {
|
||||
return errors.New(configType.Name() + " is already registered")
|
||||
return errors.New(configType.Name() + " is already registered").AtError()
|
||||
}
|
||||
typeCreatorRegistry[configType] = configCreator
|
||||
return nil
|
||||
@@ -27,7 +27,7 @@ func CreateObject(ctx context.Context, config interface{}) (interface{}, error)
|
||||
configType := reflect.TypeOf(config)
|
||||
creator, found := typeCreatorRegistry[configType]
|
||||
if !found {
|
||||
return nil, errors.New(configType.String() + " is not registered")
|
||||
return nil, errors.New(configType.String() + " is not registered").AtError()
|
||||
}
|
||||
return creator(ctx, config)
|
||||
}
|
||||
|
||||
+19
-33
@@ -8,7 +8,7 @@ import (
|
||||
"io"
|
||||
"strconv"
|
||||
"strings"
|
||||
"sync/atomic"
|
||||
"time"
|
||||
|
||||
"github.com/xtls/xray-core/common/buf"
|
||||
"github.com/xtls/xray-core/common/errors"
|
||||
@@ -27,39 +27,25 @@ var AddrParser = protocol.NewAddressParser(
|
||||
)
|
||||
|
||||
var (
|
||||
Show atomic.Bool
|
||||
baseKey atomic.Value
|
||||
Show bool
|
||||
BaseKey []byte
|
||||
)
|
||||
|
||||
func reloadEnvSettings() error {
|
||||
Show.Store(strings.ToLower(platform.NewEnvFlag(platform.XUDPLog).GetValue(func() string { return "" })) == "true")
|
||||
raw := platform.NewEnvFlag(platform.XUDPBaseKey).GetValue(func() string { return "" })
|
||||
if raw == "" {
|
||||
ensureBaseKey()
|
||||
return nil
|
||||
}
|
||||
key, _ := base64.RawURLEncoding.DecodeString(raw)
|
||||
if len(key) != 32 {
|
||||
return errors.New(platform.XUDPBaseKey + ": invalid value (BaseKey must be 32 bytes): " + raw + " len " + strconv.Itoa(len(key)))
|
||||
}
|
||||
baseKey.Store(append([]byte(nil), key...))
|
||||
return nil
|
||||
}
|
||||
|
||||
func ensureBaseKey() []byte {
|
||||
if key := baseKey.Load(); key != nil {
|
||||
return key.([]byte)
|
||||
}
|
||||
key := make([]byte, 32)
|
||||
if _, err := rand.Read(key); err != nil {
|
||||
panic(err)
|
||||
}
|
||||
baseKey.Store(key)
|
||||
return key
|
||||
}
|
||||
|
||||
func init() {
|
||||
platform.RegisterEnvReload(reloadEnvSettings)
|
||||
if strings.ToLower(platform.NewEnvFlag(platform.XUDPLog).GetValue(func() string { return "" })) == "true" {
|
||||
Show = true
|
||||
}
|
||||
BaseKey = make([]byte, 32)
|
||||
rand.Read(BaseKey)
|
||||
go func() {
|
||||
time.Sleep(100 * time.Millisecond) // this is not nice, but need to give some time for Android to setup ENV
|
||||
if raw := platform.NewEnvFlag(platform.XUDPBaseKey).GetValue(func() string { return "" }); raw != "" {
|
||||
if BaseKey, _ = base64.RawURLEncoding.DecodeString(raw); len(BaseKey) == 32 {
|
||||
return
|
||||
}
|
||||
panic(platform.XUDPBaseKey + ": invalid value (BaseKey must be 32 bytes): " + raw + " len " + strconv.Itoa(len(BaseKey)))
|
||||
}
|
||||
}()
|
||||
}
|
||||
|
||||
func GetGlobalID(ctx context.Context) (globalID [8]byte) {
|
||||
@@ -68,10 +54,10 @@ func GetGlobalID(ctx context.Context) (globalID [8]byte) {
|
||||
}
|
||||
if inbound := session.InboundFromContext(ctx); inbound != nil && inbound.Source.Network == net.Network_UDP &&
|
||||
(inbound.Name == "dokodemo-door" || inbound.Name == "socks" || inbound.Name == "shadowsocks" || inbound.Name == "tun") {
|
||||
h := blake3.New(8, ensureBaseKey())
|
||||
h := blake3.New(8, BaseKey)
|
||||
h.Write([]byte(inbound.Source.String()))
|
||||
copy(globalID[:], h.Sum(nil))
|
||||
if Show.Load() {
|
||||
if Show {
|
||||
errors.LogInfo(ctx, fmt.Sprintf("XUDP inbound.Source.String(): %v\tglobalID: %v\n", inbound.Source.String(), globalID))
|
||||
}
|
||||
}
|
||||
|
||||
+4
-4
@@ -125,7 +125,7 @@ func LoadConfig(formatName string, input interface{}) (*Config, error) {
|
||||
}
|
||||
|
||||
if f == "" {
|
||||
return nil, errors.New("Failed to get format of ", file)
|
||||
return nil, errors.New("Failed to get format of ", file).AtWarning()
|
||||
}
|
||||
|
||||
if f == "protobuf" {
|
||||
@@ -142,7 +142,7 @@ func LoadConfig(formatName string, input interface{}) (*Config, error) {
|
||||
if len(v) == 1 {
|
||||
return configLoaderByName["protobuf"].Loader(v)
|
||||
} else {
|
||||
return nil, errors.New("Only one protobuf config file is allowed")
|
||||
return nil, errors.New("Only one protobuf config file is allowed").AtWarning()
|
||||
}
|
||||
}
|
||||
|
||||
@@ -152,11 +152,11 @@ func LoadConfig(formatName string, input interface{}) (*Config, error) {
|
||||
if f, found := configLoaderByName[formatName]; found {
|
||||
return f.Loader(v)
|
||||
} else {
|
||||
return nil, errors.New("Unable to load config in", formatName)
|
||||
return nil, errors.New("Unable to load config in", formatName).AtWarning()
|
||||
}
|
||||
}
|
||||
|
||||
return nil, errors.New("Unable to load config")
|
||||
return nil, errors.New("Unable to load config").AtWarning()
|
||||
}
|
||||
|
||||
func loadProtobufConfig(data []byte) (*Config, error) {
|
||||
|
||||
+2
-2
@@ -19,8 +19,8 @@ import (
|
||||
|
||||
var (
|
||||
Version_x byte = 26
|
||||
Version_y byte = 9
|
||||
Version_z byte = 30
|
||||
Version_y byte = 6
|
||||
Version_z byte = 22
|
||||
)
|
||||
|
||||
var (
|
||||
|
||||
@@ -187,9 +187,6 @@ func NewWithContext(ctx context.Context, config *Config) (*Instance, error) {
|
||||
}
|
||||
|
||||
func initInstanceWithConfig(config *Config, server *Instance) (bool, error) {
|
||||
if err := platform.ReloadEnvSettings(); err != nil {
|
||||
return true, errors.New("failed to reload environment settings").Base(err)
|
||||
}
|
||||
server.ctx = context.WithValue(server.ctx, "cone",
|
||||
platform.NewEnvFlag(platform.UseCone).GetValue(func() string { return "" }) != "true")
|
||||
|
||||
|
||||
@@ -13,7 +13,7 @@ type FakeDNSEngine interface {
|
||||
|
||||
var (
|
||||
FakeIPv4Pool = "198.18.0.0/15"
|
||||
FakeIPv6Pool = "2001:2::/48"
|
||||
FakeIPv6Pool = "fc00::/18"
|
||||
)
|
||||
|
||||
type FakeDNSEngineRev0 interface {
|
||||
|
||||
@@ -97,9 +97,6 @@ func New() *Client {
|
||||
r := &net.Resolver{
|
||||
PreferGo: true,
|
||||
Dial: func(ctx context.Context, network, address string) (net.Conn, error) {
|
||||
if internet.IsSkippedDNSServer(address) {
|
||||
return nil, errors.New("skipped DNS server ", address)
|
||||
}
|
||||
return d.DialContext(ctx, network, address)
|
||||
},
|
||||
}
|
||||
|
||||
@@ -1,23 +0,0 @@
|
||||
package localdns
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net/netip"
|
||||
"testing"
|
||||
|
||||
"github.com/xtls/xray-core/transport/internet"
|
||||
)
|
||||
|
||||
func TestSkippedDNSServers(t *testing.T) {
|
||||
internet.SkipDNSServers([]netip.Addr{netip.MustParseAddr("203.0.113.53")})
|
||||
t.Cleanup(func() { internet.SkipDNSServers(nil) })
|
||||
c := New()
|
||||
if _, err := c.r.Dial(context.Background(), "udp", "203.0.113.53:53"); err == nil {
|
||||
t.Error("a skipped DNS server was dialed")
|
||||
}
|
||||
conn, err := c.r.Dial(context.Background(), "udp", "127.0.0.1:53")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
conn.Close()
|
||||
}
|
||||
@@ -3,7 +3,6 @@ package policy
|
||||
import (
|
||||
"context"
|
||||
"runtime"
|
||||
"sync/atomic"
|
||||
"time"
|
||||
|
||||
"github.com/xtls/xray-core/common/platform"
|
||||
@@ -83,41 +82,32 @@ func ManagerType() interface{} {
|
||||
return (*Manager)(nil)
|
||||
}
|
||||
|
||||
var defaultBufferSize atomic.Int32
|
||||
var defaultBufferSize int32
|
||||
|
||||
func reloadEnvSettings() error {
|
||||
defaultBufferSize.Store(readDefaultBufferSize())
|
||||
return nil
|
||||
}
|
||||
|
||||
func readDefaultBufferSize() int32 {
|
||||
func init() {
|
||||
const defaultValue = -17
|
||||
size := platform.NewEnvFlag(platform.BufferSize).GetValueAsInt(defaultValue)
|
||||
|
||||
switch size {
|
||||
case 0:
|
||||
return -1 // For pipe to use unlimited size
|
||||
defaultBufferSize = -1 // For pipe to use unlimited size
|
||||
case defaultValue: // Env flag not defined. Use default values per CPU-arch.
|
||||
switch runtime.GOARCH {
|
||||
case "arm", "mips", "mipsle":
|
||||
return 0
|
||||
defaultBufferSize = 0
|
||||
case "arm64", "mips64", "mips64le":
|
||||
return 4 * 1024 // 4k cache for low-end devices
|
||||
defaultBufferSize = 4 * 1024 // 4k cache for low-end devices
|
||||
default:
|
||||
return 512 * 1024
|
||||
defaultBufferSize = 512 * 1024
|
||||
}
|
||||
default:
|
||||
return int32(size) * 1024 * 1024
|
||||
defaultBufferSize = int32(size) * 1024 * 1024
|
||||
}
|
||||
}
|
||||
|
||||
func init() {
|
||||
platform.RegisterEnvReload(reloadEnvSettings)
|
||||
}
|
||||
|
||||
func defaultBufferPolicy() Buffer {
|
||||
return Buffer{
|
||||
PerConnection: defaultBufferSize.Load(),
|
||||
PerConnection: defaultBufferSize,
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
+30
-21
@@ -81,8 +81,6 @@ type Manager interface {
|
||||
|
||||
// RegisterCounter registers a new counter to the manager. The identifier string must not be empty, and unique among other counters.
|
||||
RegisterCounter(string) (Counter, error)
|
||||
// GetOrRegisterCounter returns the counter by its identifier, atomically creating and registering it if absent.
|
||||
GetOrRegisterCounter(string) (Counter, error)
|
||||
// UnregisterCounter unregisters a counter from the manager by its identifier.
|
||||
UnregisterCounter(string) error
|
||||
// GetCounter returns a counter by its identifier.
|
||||
@@ -93,8 +91,6 @@ type Manager interface {
|
||||
|
||||
// RegisterOnlineMap registers a new OnlineMap to the manager. The identifier string must not be empty, and unique among other OnlineMaps.
|
||||
RegisterOnlineMap(string) (OnlineMap, error)
|
||||
// GetOrRegisterOnlineMap returns the OnlineMap by its identifier, atomically creating and registering it if absent.
|
||||
GetOrRegisterOnlineMap(string) (OnlineMap, error)
|
||||
// UnregisterOnlineMap unregisters an OnlineMap from the manager by its identifier.
|
||||
UnregisterOnlineMap(string) error
|
||||
// GetOnlineMap returns an OnlineMap by its identifier.
|
||||
@@ -105,8 +101,6 @@ type Manager interface {
|
||||
|
||||
// RegisterChannel registers a new channel to the manager. The identifier string must not be empty, and unique among other channels.
|
||||
RegisterChannel(string) (Channel, error)
|
||||
// GetOrRegisterChannel returns the channel by its identifier, atomically creating and registering it if absent.
|
||||
GetOrRegisterChannel(string) (Channel, error)
|
||||
// UnregisterChannel unregisters a channel from the manager by its identifier.
|
||||
UnregisterChannel(string) error
|
||||
// GetChannel returns a channel by its identifier.
|
||||
@@ -116,6 +110,36 @@ type Manager interface {
|
||||
GetAllOnlineUsers() []string
|
||||
}
|
||||
|
||||
// GetOrRegisterCounter tries to get the StatCounter first. If not exist, it then tries to create a new counter.
|
||||
func GetOrRegisterCounter(m Manager, name string) (Counter, error) {
|
||||
counter := m.GetCounter(name)
|
||||
if counter != nil {
|
||||
return counter, nil
|
||||
}
|
||||
|
||||
return m.RegisterCounter(name)
|
||||
}
|
||||
|
||||
// GetOrRegisterOnlineMap tries to get the OnlineMap first. If not exist, it then tries to create a new OnlineMap.
|
||||
func GetOrRegisterOnlineMap(m Manager, name string) (OnlineMap, error) {
|
||||
onlineMap := m.GetOnlineMap(name)
|
||||
if onlineMap != nil {
|
||||
return onlineMap, nil
|
||||
}
|
||||
|
||||
return m.RegisterOnlineMap(name)
|
||||
}
|
||||
|
||||
// GetOrRegisterChannel tries to get the StatChannel first. If not exist, it then tries to create a new channel.
|
||||
func GetOrRegisterChannel(m Manager, name string) (Channel, error) {
|
||||
channel := m.GetChannel(name)
|
||||
if channel != nil {
|
||||
return channel, nil
|
||||
}
|
||||
|
||||
return m.RegisterChannel(name)
|
||||
}
|
||||
|
||||
// ManagerType returns the type of Manager interface. Can be used to implement common.HasType.
|
||||
//
|
||||
// xray:api:stable
|
||||
@@ -136,11 +160,6 @@ func (NoopManager) RegisterCounter(string) (Counter, error) {
|
||||
return nil, errors.New("not implemented")
|
||||
}
|
||||
|
||||
// GetOrRegisterCounter implements Manager.
|
||||
func (NoopManager) GetOrRegisterCounter(string) (Counter, error) {
|
||||
return nil, errors.New("not implemented")
|
||||
}
|
||||
|
||||
// UnregisterCounter implements Manager.
|
||||
func (NoopManager) UnregisterCounter(string) error {
|
||||
return nil
|
||||
@@ -159,11 +178,6 @@ func (NoopManager) RegisterOnlineMap(string) (OnlineMap, error) {
|
||||
return nil, errors.New("not implemented")
|
||||
}
|
||||
|
||||
// GetOrRegisterOnlineMap implements Manager.
|
||||
func (NoopManager) GetOrRegisterOnlineMap(string) (OnlineMap, error) {
|
||||
return nil, errors.New("not implemented")
|
||||
}
|
||||
|
||||
// UnregisterOnlineMap implements Manager.
|
||||
func (NoopManager) UnregisterOnlineMap(string) error {
|
||||
return nil
|
||||
@@ -182,11 +196,6 @@ func (NoopManager) RegisterChannel(string) (Channel, error) {
|
||||
return nil, errors.New("not implemented")
|
||||
}
|
||||
|
||||
// GetOrRegisterChannel implements Manager.
|
||||
func (NoopManager) GetOrRegisterChannel(string) (Channel, error) {
|
||||
return nil, errors.New("not implemented")
|
||||
}
|
||||
|
||||
// UnregisterChannel implements Manager.
|
||||
func (NoopManager) UnregisterChannel(string) error {
|
||||
return nil
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user