Compare commits

..
Author SHA1 Message Date
Fangliding e0efc13c6e Fix h1 and h2 2026-06-24 16:16:52 +08:00
Fangliding 518a7efac2 Apply utls in geodata 2026-06-24 14:08:56 +08:00
443 changed files with 10493 additions and 44796 deletions
-1
View File
@@ -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
-1
View File
@@ -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
+1 -20
View File
@@ -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
+5 -5
View File
@@ -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-
+5 -5
View File
@@ -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
+7 -5
View File
@@ -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-
-1
View File
@@ -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`
+5 -8
View File
@@ -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) })
}
@@ -470,9 +470,6 @@ func (d *DefaultDispatcher) routedDispatch(ctx context.Context, link *transport.
return // DO NOT CHANGE: the traffic shouldn't be processed by default outbound if the specified outbound tag doesn't exist (yet), e.g., VLESS Reverse Proxy
}
} else {
if err != common.ErrNoClue {
errors.LogErrorInner(ctx, err, "failed to pick route for ", destination)
}
errors.LogInfo(ctx, "default route for ", destination)
}
}
+1 -1
View File
@@ -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
View File
@@ -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
+6 -25
View File
@@ -93,7 +93,6 @@ type NameServer struct {
UnexpectedIp []*geodata.IPRule `protobuf:"bytes,13,rep,name=unexpected_ip,json=unexpectedIp,proto3" json:"unexpected_ip,omitempty"`
ActUnprior bool `protobuf:"varint,14,opt,name=actUnprior,proto3" json:"actUnprior,omitempty"`
PolicyID uint32 `protobuf:"varint,17,opt,name=policyID,proto3" json:"policyID,omitempty"`
Id string `protobuf:"bytes,18,opt,name=id,proto3" json:"id,omitempty"`
unknownFields protoimpl.UnknownFields
sizeCache protoimpl.SizeCache
}
@@ -240,13 +239,6 @@ func (x *NameServer) GetPolicyID() uint32 {
return 0
}
func (x *NameServer) GetId() string {
if x != nil {
return x.Id
}
return ""
}
type Config struct {
state protoimpl.MessageState `protogen:"open.v1"`
// NameServer list used by this DNS client.
@@ -266,10 +258,8 @@ type Config struct {
DisableFallback bool `protobuf:"varint,10,opt,name=disableFallback,proto3" json:"disableFallback,omitempty"`
DisableFallbackIfMatch bool `protobuf:"varint,11,opt,name=disableFallbackIfMatch,proto3" json:"disableFallbackIfMatch,omitempty"`
EnableParallelQuery bool `protobuf:"varint,14,opt,name=enableParallelQuery,proto3" json:"enableParallelQuery,omitempty"`
// Absolute path to the Lua DNS query script.
Script string `protobuf:"bytes,15,opt,name=script,proto3" json:"script,omitempty"`
unknownFields protoimpl.UnknownFields
sizeCache protoimpl.SizeCache
unknownFields protoimpl.UnknownFields
sizeCache protoimpl.SizeCache
}
func (x *Config) Reset() {
@@ -379,13 +369,6 @@ func (x *Config) GetEnableParallelQuery() bool {
return false
}
func (x *Config) GetScript() string {
if x != nil {
return x.Script
}
return ""
}
type Config_HostMapping struct {
state protoimpl.MessageState `protogen:"open.v1"`
Domain *geodata.DomainRule `protobuf:"bytes,2,opt,name=domain,proto3" json:"domain,omitempty"`
@@ -452,7 +435,7 @@ var File_app_dns_config_proto protoreflect.FileDescriptor
const file_app_dns_config_proto_rawDesc = "" +
"\n" +
"\x14app/dns/config.proto\x12\fxray.app.dns\x1a\x1ccommon/net/destination.proto\x1a\x1bcommon/geodata/geodat.proto\"\xee\x05\n" +
"\x14app/dns/config.proto\x12\fxray.app.dns\x1a\x1ccommon/net/destination.proto\x1a\x1bcommon/geodata/geodat.proto\"\xde\x05\n" +
"\n" +
"NameServer\x123\n" +
"\aaddress\x18\x01 \x01(\v2\x19.xray.common.net.EndpointR\aaddress\x12\x1b\n" +
@@ -478,11 +461,10 @@ const file_app_dns_config_proto_rawDesc = "" +
"\n" +
"actUnprior\x18\x0e \x01(\bR\n" +
"actUnprior\x12\x1a\n" +
"\bpolicyID\x18\x11 \x01(\rR\bpolicyID\x12\x0e\n" +
"\x02id\x18\x12 \x01(\tR\x02idB\x0f\n" +
"\bpolicyID\x18\x11 \x01(\rR\bpolicyIDB\x0f\n" +
"\r_disableCacheB\r\n" +
"\v_serveStaleB\x12\n" +
"\x10_serveExpiredTTLJ\x04\b\x04\x10\x05\"\x9a\x05\n" +
"\x10_serveExpiredTTLJ\x04\b\x04\x10\x05\"\x82\x05\n" +
"\x06Config\x129\n" +
"\vname_server\x18\x05 \x03(\v2\x18.xray.app.dns.NameServerR\n" +
"nameServer\x12\x1b\n" +
@@ -498,8 +480,7 @@ const file_app_dns_config_proto_rawDesc = "" +
"\x0fdisableFallback\x18\n" +
" \x01(\bR\x0fdisableFallback\x126\n" +
"\x16disableFallbackIfMatch\x18\v \x01(\bR\x16disableFallbackIfMatch\x120\n" +
"\x13enableParallelQuery\x18\x0e \x01(\bR\x13enableParallelQuery\x12\x16\n" +
"\x06script\x18\x0f \x01(\tR\x06script\x1a}\n" +
"\x13enableParallelQuery\x18\x0e \x01(\bR\x13enableParallelQuery\x1a}\n" +
"\vHostMapping\x127\n" +
"\x06domain\x18\x02 \x01(\v2\x1f.xray.common.geodata.DomainRuleR\x06domain\x12\x0e\n" +
"\x02ip\x18\x03 \x03(\fR\x02ip\x12%\n" +
-4
View File
@@ -27,7 +27,6 @@ message NameServer {
repeated xray.common.geodata.IPRule unexpected_ip = 13;
bool actUnprior = 14;
uint32 policyID = 17;
string id = 18;
}
enum QueryStrategy {
@@ -74,7 +73,4 @@ message Config {
bool disableFallbackIfMatch = 11;
bool enableParallelQuery = 14;
// Absolute path to the Lua DNS query script.
string script = 15;
}
-38
View File
@@ -31,8 +31,6 @@ type DNS struct {
domainMatcher geodata.DomainMatcher
matcherInfos []*DomainMatcherInfo
checkSystem bool
script *scriptEngine
scriptPath string
}
// DomainMatcherInfo contains information attached to index returned by Server.domainMatcher.
@@ -182,7 +180,6 @@ func New(ctx context.Context, config *Config) (*DNS, error) {
disableFallbackIfMatch: config.DisableFallbackIfMatch,
enableParallelQuery: config.EnableParallelQuery,
checkSystem: checkSystem,
scriptPath: config.Script,
}, nil
}
@@ -193,21 +190,11 @@ func (*DNS) Type() interface{} {
// Start implements common.Runnable.
func (s *DNS) Start() error {
if s.scriptPath != "" {
engine, err := newScriptEngine(s.scriptPath, s)
if err != nil {
return errors.New("failed to initialize DNS script").Base(err)
}
s.script = engine
}
return nil
}
// Close implements common.Closable.
func (s *DNS) Close() error {
if s.script != nil {
s.script.close()
}
return nil
}
@@ -225,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
@@ -292,9 +257,6 @@ func (s *DNS) LookupIP(domain string, option dns.IPOption) ([]net.IP, uint32, er
}
// Name servers lookup
if s.script != nil {
return s.script.query(domain, option)
}
if s.enableParallelQuery {
return s.parallelQuery(domain, option)
} else {
-59
View File
@@ -1,59 +0,0 @@
package dns
import (
"context"
"testing"
"github.com/xtls/xray-core/common/net"
feature_dns "github.com/xtls/xray-core/features/dns"
)
// fakeServer stands in for any name server that is not the system resolver.
type fakeServer struct{}
func (fakeServer) Name() string { return "fake" }
func (fakeServer) IsDisableCache() bool { return false }
func (fakeServer) QueryIP(context.Context, string, feature_dns.IPOption) ([]net.IP, uint32, error) {
return nil, 0, nil
}
// Callers that are about to redirect the system resolver rely on this to tell
// whether any resolution path could still reach the system resolver, so the
// mixed shape has to be reported as reachable: a domain-specific rule can
// select the system resolver even when an independent upstream also exists.
func TestMayUseSystemResolver(t *testing.T) {
tests := []struct {
name string
clients []*Client
want bool
}{
{
name: "no clients at all",
want: true,
},
{
name: "only the system resolver",
clients: []*Client{{server: NewLocalNameServer()}},
want: true,
},
{
name: "the system resolver alongside an independent name server",
clients: []*Client{{server: fakeServer{}}, {server: NewLocalNameServer()}},
want: true,
},
{
name: "only independent name servers",
clients: []*Client{{server: fakeServer{}}},
want: false,
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
server := &DNS{clients: tt.clients}
if got := server.MayUseSystemResolver(); got != tt.want {
t.Errorf("MayUseSystemResolver() = %v, want %v", got, tt.want)
}
})
}
}
+2 -2
View File
@@ -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()
+3 -3
View File
@@ -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
-165
View File
@@ -1,165 +0,0 @@
package dns
import (
"context"
"strings"
"github.com/xtls/xray-core/common/errors"
xlua "github.com/xtls/xray-core/common/lua"
"github.com/xtls/xray-core/common/net"
featureDNS "github.com/xtls/xray-core/features/dns"
"github.com/xtls/xray-core/features/dns/localdns"
lua "github.com/yuin/gopher-lua"
)
// luaDNSServer adapts configured and local DNS to the same Lua API.
type luaDNSServer struct {
id string
name string
query func(context.Context, string, featureDNS.IPOption) ([]net.IP, uint32, error)
}
// RegisterLua makes xray.dns available to scripts backed by client.
func RegisterLua(L *lua.LState, client featureDNS.Client) {
var servers []luaDNSServer
switch client := client.(type) {
case *DNS:
servers = luaServers(client)
case *localdns.Client:
servers = []luaDNSServer{{
id: "localhost",
name: "localhost",
query: func(_ context.Context, domain string, option featureDNS.IPOption) ([]net.IP, uint32, error) {
return client.LookupIP(domain, option)
},
}}
}
registerLua(L, servers, client)
}
// registerLua makes xray.dns available to DNS scripts.
func (s *DNS) registerLua(L *lua.LState) {
registerLua(L, luaServers(s), nil)
}
func luaServers(s *DNS) []luaDNSServer {
servers := make([]luaDNSServer, len(s.clients))
for i, client := range s.clients {
servers[i] = luaDNSServer{id: client.id, name: client.Name(), query: client.QueryIP}
}
return servers
}
func registerLua(L *lua.LState, servers []luaDNSServer, client featureDNS.Client) {
L.PreloadModule("xray.dns", func(L *lua.LState) int {
serverList := L.NewTable()
for i, client := range servers {
server := L.NewTable()
server.RawSetString("ID", lua.LString(client.id))
server.RawSetString("Query", L.NewFunction(func(L *lua.LState) int {
domain, ok := L.Get(2).(lua.LString)
if !ok {
L.RaiseError("server:Query requires a domain")
return 0
}
option := featureDNS.IPOption{
IPv4Enable: L.CheckBool(3),
IPv6Enable: L.CheckBool(4),
FakeEnable: L.CheckBool(5),
}
ctx := L.Context()
if ctx == nil {
L.RaiseError("server:Query requires an active DNS query")
return 0
}
var ips []net.IP
var ttl uint32
var err error
if !option.FakeEnable && strings.EqualFold(client.name, "FakeDNS") {
err = featureDNS.ErrEmptyResponse
} else {
ips, ttl, err = client.query(ctx, string(domain), option)
}
xlua.PushUserData(L, ips)
xlua.PushNumber(L, ttl)
xlua.PushError(L, err)
return 3
}))
serverList.RawSetInt(i+1, server)
}
module := L.NewTable()
if servers != nil {
module.RawSetString("Servers", serverList)
}
if client != nil {
module.RawSetString("Query", newLuaClientQuery(L, client))
}
L.Push(module)
return 1
})
}
func newLuaClientQuery(L *lua.LState, client featureDNS.Client) *lua.LFunction {
return L.NewFunction(func(L *lua.LState) int {
domain, ok := L.Get(1).(lua.LString)
if !ok {
L.RaiseError("dns.Query requires a domain")
return 0
}
option := featureDNS.IPOption{
IPv4Enable: L.CheckBool(2),
IPv6Enable: L.CheckBool(3),
FakeEnable: L.CheckBool(4),
}
if L.Context() == nil {
L.RaiseError("dns.Query requires an active DNS query")
return 0
}
ips, ttl, err := client.LookupIP(string(domain), option)
xlua.PushUserData(L, ips)
xlua.PushNumber(L, ttl)
xlua.PushError(L, err)
return 3
})
}
// callLuaHook invokes HandleDNSQuery in the supplied state.
// Returned slices and IP bytes may share storage with DNS caches or matcher inputs.
func (s *DNS) callLuaHook(L *lua.LState, domain string, option featureDNS.IPOption) ([]net.IP, uint32, error) {
top := L.GetTop()
defer L.SetTop(top)
fn := L.GetGlobal("HandleDNSQuery")
if fn.Type() != lua.LTFunction {
return nil, 0, errors.New("DNS script must define HandleDNSQuery(...)")
}
if err := L.CallByParam(lua.P{Fn: fn, NRet: 3, Protect: true},
lua.LString(strings.ToLower(domain)), lua.LBool(option.IPv4Enable),
lua.LBool(option.IPv6Enable), lua.LBool(option.FakeEnable)); err != nil {
return nil, 0, err
}
return readLuaDNSResult(L.Get(-3), L.Get(-2), L.Get(-1))
}
func readLuaDNSResult(addresses, ttlValue, errorValue lua.LValue) ([]net.IP, uint32, error) {
if err := xlua.ReadError(errorValue, "DNS script error must be an error or string"); err != nil {
return nil, 0, err
}
ttl, err := xlua.ReadUint32(ttlValue, "DNS script returned invalid TTL")
if err != nil {
return nil, 0, err
}
if addresses == lua.LNil {
return nil, 0, featureDNS.ErrEmptyResponse
}
ips, err := xlua.ReadUserData[[]net.IP](addresses, "DNS script IPs must be native IP slice userdata")
if err != nil {
return nil, 0, err
}
if len(ips) == 0 {
return nil, 0, featureDNS.ErrEmptyResponse
}
return ips, ttl, nil
}
-294
View File
@@ -1,294 +0,0 @@
package dns
import (
"context"
go_errors "errors"
"math"
"strings"
"testing"
"time"
"github.com/xtls/xray-core/common/geodata"
"github.com/xtls/xray-core/common/net"
featureDNS "github.com/xtls/xray-core/features/dns"
"github.com/xtls/xray-core/features/dns/localdns"
lua "github.com/yuin/gopher-lua"
)
func TestReadLuaDNSResult(t *testing.T) {
L := lua.NewState()
defer L.Close()
want := []net.IP{net.ParseIP("8.8.8.8"), {127, 0, 0, 1}, net.ParseIP("::1")}
addresses := L.NewUserData()
addresses.Value = want
ips, ttl, err := readLuaDNSResult(addresses, lua.LNumber(45), lua.LNil)
if err != nil || ttl != 45 || len(ips) != len(want) {
t.Fatalf("readLuaDNSResult() = %v, TTL %d, %v", ips, ttl, err)
}
for i := range want {
if !ips[i].Equal(want[i]) {
t.Fatalf("IP %d = %v, want %v", i, ips[i], want[i])
}
}
}
func TestReadLuaDNSResultValidation(t *testing.T) {
L := lua.NewState()
defer L.Close()
for _, tc := range []struct {
name string
change func(*[3]lua.LValue)
want string
}{
{"fractional TTL", func(v *[3]lua.LValue) { v[1] = lua.LNumber(1.5) }, "invalid TTL"},
{"oversized TTL", func(v *[3]lua.LValue) { v[1] = lua.LNumber(4294967296) }, "invalid TTL"},
{"negative TTL", func(v *[3]lua.LValue) { v[1] = lua.LNumber(-1) }, "invalid TTL"},
{"NaN TTL", func(v *[3]lua.LValue) { v[1] = lua.LNumber(math.NaN()) }, "invalid TTL"},
{"missing TTL", func(v *[3]lua.LValue) { v[1] = lua.LNil }, "invalid TTL"},
{"string IPs", func(v *[3]lua.LValue) { v[0] = lua.LString("127.0.0.1") }, "native IP slice"},
{"wrong userdata", func(v *[3]lua.LValue) { v[0].(*lua.LUserData).Value = net.ParseIP("127.0.0.1") }, "native IP slice"},
{"script error", func(v *[3]lua.LValue) { v[2] = lua.LString("blocked by script") }, "blocked by script"},
{"invalid error", func(v *[3]lua.LValue) { v[2] = lua.LTrue }, "error or string"},
} {
t.Run(tc.name, func(t *testing.T) {
addresses := L.NewUserData()
addresses.Value = []net.IP{net.ParseIP("127.0.0.1")}
values := [3]lua.LValue{addresses, lua.LNumber(60), lua.LNil}
tc.change(&values)
_, _, err := readLuaDNSResult(values[0], values[1], values[2])
if err == nil || !strings.Contains(err.Error(), tc.want) {
t.Fatalf("readLuaDNSResult error = %v, want %q", err, tc.want)
}
})
}
addresses := L.NewUserData()
addresses.Value = []net.IP(nil)
for _, empty := range []lua.LValue{addresses, lua.LNil} {
if _, _, err := readLuaDNSResult(empty, lua.LNumber(0), lua.LNil); !go_errors.Is(err, featureDNS.ErrEmptyResponse) {
t.Fatalf("empty result error = %v, want ErrEmptyResponse", err)
}
}
wantErr := go_errors.New("upstream failed")
errorValue := L.NewUserData()
errorValue.Value = wantErr
if _, _, err := readLuaDNSResult(lua.LNil, lua.LNil, errorValue); err != wantErr {
t.Fatalf("upstream error = %v, want original error %v", err, wantErr)
}
}
func TestCallLuaHookCancellation(t *testing.T) {
L := lua.NewState()
defer L.Close()
if err := L.DoString(`function HandleDNSQuery(domain, ipv4, ipv6, fake) while true do end end`); err != nil {
t.Fatal(err)
}
ctx, cancel := context.WithCancel(context.Background())
cancel()
L.SetContext(ctx)
_, _, err := (&DNS{}).callLuaHook(L, "example.com", featureDNS.IPOption{IPv4Enable: true})
if err == nil {
t.Fatal("CallLuaHook did not stop after context cancellation")
}
if L.Context() != ctx {
t.Fatal("CallLuaHook changed the Lua state's context")
}
}
func TestCallLuaHookNormalizesDomain(t *testing.T) {
L := lua.NewState()
defer L.Close()
addresses := L.NewUserData()
addresses.Value = []net.IP{net.ParseIP("127.0.0.1")}
L.SetGlobal("ips", addresses)
if err := L.DoString(`
function HandleDNSQuery(domain, ipv4, ipv6, fake)
assert(domain == "example.com")
assert(ipv4 and not ipv6 and not fake)
return ips, 60, nil
end
`); err != nil {
t.Fatal(err)
}
s := &DNS{}
if _, _, err := s.callLuaHook(L, "ExAmPlE.CoM", featureDNS.IPOption{IPv4Enable: true}); err != nil {
t.Fatal(err)
}
}
func TestCallLuaHookRestoresStack(t *testing.T) {
for _, tc := range []struct {
name string
body string
wantErr bool
}{
{"success", `return ips, 60`, false},
{"error", `error("failed")`, true},
} {
t.Run(tc.name, func(t *testing.T) {
L := lua.NewState()
defer L.Close()
addresses := L.NewUserData()
addresses.Value = []net.IP{net.ParseIP("127.0.0.1")}
L.SetGlobal("ips", addresses)
if err := L.DoString("function HandleDNSQuery() " + tc.body + " end"); err != nil {
t.Fatal(err)
}
L.Push(lua.LTrue)
_, _, err := (&DNS{}).callLuaHook(L, "example.com", featureDNS.IPOption{IPv4Enable: true})
if (err != nil) != tc.wantErr {
t.Fatalf("hook error = %v, want error %t", err, tc.wantErr)
}
if L.GetTop() != 1 || L.Get(1) != lua.LTrue {
t.Fatal("hook did not restore the stack")
}
})
}
}
func TestLuaDNSServerQuery(t *testing.T) {
L := lua.NewState()
defer L.Close()
geodata.RegisterLua(L)
option := featureDNS.IPOption{IPv4Enable: true}
ips := []net.IP{net.ParseIP("127.0.0.1"), net.ParseIP("8.8.8.8")}
server := &DNS{clients: []*Client{{server: &benchmarkLuaNameServer{ips: ips}, ipOption: &option, timeoutMs: time.Second}}}
server.registerLua(L)
if err := L.DoString(`
local server = require("xray.dns").Servers[1]
local matcher = require("xray.geodata").BuildIPMatcher("127.0.0.0/8")
function HandleDNSQuery(domain, ipv4, ipv6, fake)
local ips, ttl, err = server:Query(domain, ipv4, ipv6, fake)
assert(type(ips) == "userdata" and not err)
assert(matcher:AnyMatch(ips))
local matched = matcher:FilterIPs(ips)
return matched, ttl, err
end
`); err != nil {
t.Fatal(err)
}
L.SetContext(context.Background())
got, ttl, err := server.callLuaHook(L, "example.com", option)
if err != nil || ttl != 60 || len(got) != 1 || !got[0].Equal(ips[0]) {
t.Fatalf("server query = %v, TTL %d, %v", got, ttl, err)
}
}
type luaDNSClient struct {
featureDNS.Client
lookup func(string, featureDNS.IPOption) ([]net.IP, uint32, error)
}
func (c *luaDNSClient) LookupIP(domain string, option featureDNS.IPOption) ([]net.IP, uint32, error) {
return c.lookup(domain, option)
}
func TestLuaDNSClientQuery(t *testing.T) {
L := lua.NewState()
defer L.Close()
L.SetContext(context.Background())
geodata.RegisterLua(L)
want := []net.IP{{127, 0, 0, 1}}
client := &luaDNSClient{lookup: func(domain string, option featureDNS.IPOption) ([]net.IP, uint32, error) {
if domain != "MiXeD.Example." || !option.IPv4Enable || option.IPv6Enable || !option.FakeEnable {
t.Fatalf("dns.Query arguments = %q, %+v", domain, option)
}
return want, 42, nil
}}
RegisterLua(L, client)
if err := L.DoString(`
local dns = require("xray.dns")
local matcher = require("xray.geodata").BuildIPMatcher("127.0.0.1")
assert(dns.Servers == nil)
ips, ttl, err = dns.Query("MiXeD.Example.", true, false, true)
assert(not err and ttl == 42 and matcher:AnyMatch(ips))
`); err != nil {
t.Fatal(err)
}
got := L.GetGlobal("ips").(*lua.LUserData).Value.([]net.IP)
if &got[0] != &want[0] {
t.Fatal("dns.Query copied the IP slice")
}
}
func TestLuaDNSLocalClient(t *testing.T) {
L := lua.NewState()
defer L.Close()
L.SetContext(context.Background())
RegisterLua(L, localdns.New())
if err := L.DoString(`
local dns = require("xray.dns")
assert(dns.Servers[1].ID == "localhost")
serverIPs, _, serverErr = dns.Servers[1]:Query("127.0.0.1", true, false, false)
clientIPs, _, clientErr = dns.Query("127.0.0.1", true, false, false)
assert(not serverErr and not clientErr)
`); err != nil {
t.Fatal(err)
}
for _, name := range []string{"serverIPs", "clientIPs"} {
ips := L.GetGlobal(name).(*lua.LUserData).Value.([]net.IP)
if len(ips) != 1 || !ips[0].Equal(net.ParseIP("127.0.0.1")) {
t.Fatalf("%s = %v", name, ips)
}
}
}
type benchmarkLuaNameServer struct {
ips []net.IP
}
func (*benchmarkLuaNameServer) Name() string { return "benchmark" }
func (*benchmarkLuaNameServer) IsDisableCache() bool { return true }
func (s *benchmarkLuaNameServer) QueryIP(context.Context, string, featureDNS.IPOption) ([]net.IP, uint32, error) {
return s.ips, 60, nil
}
// BenchmarkLuaDNSHookCall isolates a preloaded Lua hook and its server:Query bridge.
// The direct case measures the same DNS client without Lua.
func BenchmarkLuaDNSHookCall(b *testing.B) {
option := featureDNS.IPOption{IPv4Enable: true}
ip := net.ParseIP("127.0.0.1")
upstream := &benchmarkLuaNameServer{ips: []net.IP{ip}}
client := &Client{server: upstream, ipOption: &option, timeoutMs: time.Second}
server := &DNS{clients: []*Client{client}}
L := lua.NewState()
defer L.Close()
server.registerLua(L)
if err := L.DoString(`
local server = require("xray.dns").Servers[1]
function HandleDNSQuery(domain, ipv4, ipv6, fake)
return server:Query(domain, ipv4, ipv6, fake)
end
`); err != nil {
b.Fatal(err)
}
ctx := context.Background()
L.SetContext(ctx)
for _, bench := range []struct {
name string
query func() ([]net.IP, uint32, error)
}{
{"direct", func() ([]net.IP, uint32, error) { return client.QueryIP(ctx, "example.com", option) }},
{"lua_hook", func() ([]net.IP, uint32, error) { return server.callLuaHook(L, "example.com", option) }},
} {
b.Run(bench.name, func(b *testing.B) {
b.ReportAllocs()
b.ResetTimer()
var ips []net.IP
var ttl uint32
var err error
for i := 0; i < b.N; i++ {
ips, ttl, err = bench.query()
if err != nil {
b.Fatal(err)
}
}
b.StopTimer()
if ttl != 60 || len(ips) != 1 || !ips[0].Equal(ip) {
b.Fatalf("query() = %v, TTL %d; want %v, TTL 60", ips, ttl, ip)
}
})
}
}
+5 -6
View File
@@ -29,7 +29,6 @@ type Server interface {
// Client is the interface for DNS client.
type Client struct {
id string
server Server
skipFallback bool
expectedIPs geodata.IPMatcher
@@ -85,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.
@@ -98,12 +97,12 @@ func NewClient(
ipOption dns.IPOption,
updateRules func(bool),
) (*Client, error) {
client := &Client{id: ns.Id}
client := &Client{}
err := core.RequireFeatures(ctx, func(dispatcher routing.Dispatcher) error {
// 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)
@@ -114,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()
}
}
@@ -123,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()
}
}
+2 -2
View File
@@ -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)
+1 -1
View File
@@ -49,5 +49,5 @@ func NewLocalNameServer() *LocalNameServer {
// NewLocalDNSClient creates localdns client object for directly lookup in system DNS.
func NewLocalDNSClient(ipOption dns.IPOption) *Client {
return &Client{id: "localhost", server: NewLocalNameServer(), ipOption: &ipOption}
return &Client{server: NewLocalNameServer(), ipOption: &ipOption}
}
-59
View File
@@ -1,59 +0,0 @@
package dns
import (
"time"
"github.com/xtls/xray-core/common/errors"
"github.com/xtls/xray-core/common/geodata"
"github.com/xtls/xray-core/common/log"
xlua "github.com/xtls/xray-core/common/lua"
"github.com/xtls/xray-core/common/net"
"github.com/xtls/xray-core/features/dns"
lua "github.com/yuin/gopher-lua"
)
const scriptExecutionTimeout = 6 * time.Second
type scriptEngine struct {
dns *DNS
pool *xlua.Pool
}
func newScriptEngine(path string, server *DNS) (*scriptEngine, error) {
program, err := xlua.CompileFile(path)
if err != nil {
return nil, err
}
e := &scriptEngine{dns: server}
e.pool, err = xlua.NewPool(server.ctx, scriptExecutionTimeout, program.NewStateFactory(
scriptExecutionTimeout*20,
func(L *lua.LState) {
geodata.RegisterLua(L)
log.RegisterLua(L)
server.registerLua(L)
},
func(L *lua.LState) error {
if L.GetGlobal("HandleDNSQuery").Type() != lua.LTFunction {
return errors.New("DNS script must define HandleDNSQuery(...)")
}
return nil
}))
if err != nil {
return nil, err
}
errors.LogInfo(server.ctx, "DNS script initialized from ", path)
return e, nil
}
func (e *scriptEngine) close() {
e.pool.Close()
}
func (e *scriptEngine) query(domain string, option dns.IPOption) (ips []net.IP, ttl uint32, err error) {
err = e.pool.WithState(nil, 0, func(L *lua.LState) error {
var hookErr error
ips, ttl, hookErr = e.dns.callLuaHook(L, domain, option)
return hookErr
})
return
}
-197
View File
@@ -1,197 +0,0 @@
package dns
import (
"context"
"os"
"path/filepath"
"strings"
"testing"
"time"
"github.com/xtls/xray-core/common/net"
featureDNS "github.com/xtls/xray-core/features/dns"
)
type geoIPScriptNameServer struct {
name string
answers map[string]net.IP
ttl uint32
calls int
}
func (s *geoIPScriptNameServer) Name() string { return s.name }
func (s *geoIPScriptNameServer) IsDisableCache() bool { return true }
func (s *geoIPScriptNameServer) QueryIP(ctx context.Context, domain string, _ featureDNS.IPOption) ([]net.IP, uint32, error) {
if err := ctx.Err(); err != nil {
return nil, 0, err
}
s.calls++
ip, ok := s.answers[domain]
if !ok {
return nil, 0, featureDNS.ErrEmptyResponse
}
return []net.IP{ip}, s.ttl, nil
}
func TestDNSScriptGeoIPFallback(t *testing.T) {
t.Setenv("xray.location.asset", filepath.Join("..", "..", "resources"))
script := `
local servers = require("xray.dns").Servers
local us_ips = require("xray.geodata").BuildIPMatcher("geoip:us")
local by_id = {}
for _, server in ipairs(servers) do
by_id[server.ID] = server
end
assert(by_id.primary and by_id.fallback, "primary and fallback DNS servers are required")
function HandleDNSQuery(domain, ipv4, ipv6, fake)
local ips, ttl, err = by_id.primary:Query(domain, ipv4, ipv6, fake)
if not err and us_ips:AnyMatch(ips) then
return ips, ttl, nil
end
return by_id.fallback:Query(domain, ipv4, ipv6, fake)
end
`
scriptPath := filepath.Join(t.TempDir(), "geoip_fallback.lua")
if err := os.WriteFile(scriptPath, []byte(script), 0o600); err != nil {
t.Fatal(err)
}
primary := &geoIPScriptNameServer{
name: "primary",
answers: map[string]net.IP{
"us.example": net.ParseIP("2001:4860:4860::8888"),
"other.example": net.ParseIP("127.0.0.1"),
},
ttl: 30,
}
fallback := &geoIPScriptNameServer{
name: "fallback",
answers: map[string]net.IP{"other.example": net.ParseIP("9.9.9.9")},
ttl: 60,
}
option := featureDNS.IPOption{IPv4Enable: true, IPv6Enable: true}
hosts, err := NewStaticHosts(nil)
if err != nil {
t.Fatal(err)
}
server := &DNS{
ctx: context.Background(),
hosts: hosts,
ipOption: &option,
scriptPath: scriptPath,
clients: []*Client{
{id: "primary", server: primary, ipOption: &option, timeoutMs: 2 * time.Second},
{id: "fallback", server: fallback, ipOption: &option, timeoutMs: 2 * time.Second},
},
}
if err := server.Start(); err != nil {
t.Fatal(err)
}
defer server.Close()
for _, tc := range []struct {
domain string
ip net.IP
ttl uint32
}{
{"Us.Example.", net.ParseIP("2001:4860:4860::8888"), 30},
{"other.example", net.ParseIP("9.9.9.9"), 60},
} {
ips, ttl, err := server.LookupIP(tc.domain, option)
if err != nil {
t.Fatalf("LookupIP(%q): %v", tc.domain, err)
}
if ttl != tc.ttl || len(ips) != 1 || !ips[0].Equal(tc.ip) {
t.Fatalf("LookupIP(%q) = %v, TTL %d; want %v, TTL %d", tc.domain, ips, ttl, tc.ip, tc.ttl)
}
}
if primary.calls != 2 || fallback.calls != 1 {
t.Fatalf("upstream calls: primary %d, fallback %d; want 2 and 1", primary.calls, fallback.calls)
}
}
func TestDNSScriptRejectsInvalidStartup(t *testing.T) {
for _, tc := range []struct {
name string
script string
}{
{"syntax", "function HandleDNSQuery("},
{"missing hook", "value = 1"},
{"top-level error", `error("setup failed")`},
} {
t.Run(tc.name, func(t *testing.T) {
path := filepath.Join(t.TempDir(), "script.lua")
if err := os.WriteFile(path, []byte(tc.script), 0o600); err != nil {
t.Fatal(err)
}
server := &DNS{ctx: context.Background(), scriptPath: path}
if err := server.Start(); err == nil {
t.Fatal("Start accepted an invalid DNS script")
}
if server.script != nil {
t.Fatal("Start retained a script engine after failure")
}
})
}
}
func TestDNSScriptHookErrorAndFakeDNSOption(t *testing.T) {
path := filepath.Join(t.TempDir(), "script.lua")
script := `
local server = require("xray.dns").Servers[1]
local log = require("xray.log")
log.Info("DNS script loaded")
function HandleDNSQuery(domain, ipv4, ipv6, fake)
log.Debug("DNS query: ", domain)
if domain == "bad.example" then error("script failure") end
local ips, ttl, err = server:Query(domain, ipv4, ipv6, fake)
if err then log.Error("DNS failed: ", err) end
return ips, ttl, err
end
`
if err := os.WriteFile(path, []byte(script), 0o600); err != nil {
t.Fatal(err)
}
option := featureDNS.IPOption{IPv4Enable: true}
hosts, err := NewStaticHosts(nil)
if err != nil {
t.Fatal(err)
}
upstream := &geoIPScriptNameServer{
name: "FakeDNS",
answers: map[string]net.IP{"good.example": net.ParseIP("198.18.0.1")},
ttl: 30,
}
server := &DNS{
ctx: context.Background(),
hosts: hosts,
ipOption: &option,
scriptPath: path,
clients: []*Client{{id: "fake", server: upstream, ipOption: &option, timeoutMs: time.Second}},
}
if err := server.Start(); err != nil {
t.Fatal(err)
}
defer server.Close()
if _, _, err := server.LookupIP("bad.example", option); err == nil || !strings.Contains(err.Error(), "script failure") {
t.Fatalf("hook failure = %v, want script failure", err)
}
if _, _, err := server.LookupIP("good.example", option); err != featureDNS.ErrEmptyResponse {
t.Fatalf("FakeDNS without FakeEnable = %v, want ErrEmptyResponse", err)
}
if upstream.calls != 0 {
t.Fatalf("FakeDNS was queried without FakeEnable: %d calls", upstream.calls)
}
withFake := featureDNS.IPOption{IPv4Enable: true, FakeEnable: true}
ips, ttl, err := server.LookupIP("good.example", withFake)
if err != nil || ttl != 30 || len(ips) != 1 || !ips[0].Equal(net.ParseIP("198.18.0.1")) {
t.Fatalf("FakeDNS with FakeEnable = %v, TTL %d, %v", ips, ttl, err)
}
if upstream.calls != 1 {
t.Fatalf("FakeDNS query count = %d, want 1", upstream.calls)
}
}
+2 -6
View File
@@ -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
View File
@@ -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() {
-161
View File
@@ -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
}
-6
View File
@@ -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
View File
@@ -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() }
+1 -1
View File
@@ -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;
+3 -3
View File
@@ -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})
+1 -1
View File
@@ -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
+2 -2
View File
@@ -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
+62 -15
View File
@@ -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
}
}
}
+2 -2
View File
@@ -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)
+3 -3
View File
@@ -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")
+17
View File
@@ -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
}
-20
View File
@@ -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
}
-27
View File
@@ -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)
+2 -6
View File
@@ -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{
+8 -28
View File
@@ -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()
}
@@ -587,10 +578,8 @@ type Config struct {
DomainStrategy Config_DomainStrategy `protobuf:"varint,1,opt,name=domain_strategy,json=domainStrategy,proto3,enum=xray.app.router.Config_DomainStrategy" json:"domain_strategy,omitempty"`
Rule []*RoutingRule `protobuf:"bytes,2,rep,name=rule,proto3" json:"rule,omitempty"`
BalancingRule []*BalancingRule `protobuf:"bytes,3,rep,name=balancing_rule,json=balancingRule,proto3" json:"balancing_rule,omitempty"`
// Absolute path to the Lua routing script.
Script string `protobuf:"bytes,4,opt,name=script,proto3" json:"script,omitempty"`
unknownFields protoimpl.UnknownFields
sizeCache protoimpl.SizeCache
unknownFields protoimpl.UnknownFields
sizeCache protoimpl.SizeCache
}
func (x *Config) Reset() {
@@ -644,18 +633,11 @@ func (x *Config) GetBalancingRule() []*BalancingRule {
return nil
}
func (x *Config) GetScript() string {
if x != nil {
return x.Script
}
return ""
}
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" +
@@ -679,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" +
@@ -708,12 +689,11 @@ const file_app_router_config_proto_rawDesc = "" +
"\tbaselines\x18\x03 \x03(\x03R\tbaselines\x12\x1a\n" +
"\bexpected\x18\x04 \x01(\x05R\bexpected\x12\x16\n" +
"\x06maxRTT\x18\x05 \x01(\x03R\x06maxRTT\x12\x1c\n" +
"\ttolerance\x18\x06 \x01(\x02R\ttolerance\"\xae\x02\n" +
"\ttolerance\x18\x06 \x01(\x02R\ttolerance\"\x96\x02\n" +
"\x06Config\x12O\n" +
"\x0fdomain_strategy\x18\x01 \x01(\x0e2&.xray.app.router.Config.DomainStrategyR\x0edomainStrategy\x120\n" +
"\x04rule\x18\x02 \x03(\v2\x1c.xray.app.router.RoutingRuleR\x04rule\x12E\n" +
"\x0ebalancing_rule\x18\x03 \x03(\v2\x1e.xray.app.router.BalancingRuleR\rbalancingRule\x12\x16\n" +
"\x06script\x18\x04 \x01(\tR\x06script\"B\n" +
"\x0ebalancing_rule\x18\x03 \x03(\v2\x1e.xray.app.router.BalancingRuleR\rbalancingRule\"B\n" +
"\x0eDomainStrategy\x12\b\n" +
"\x04AsIs\x10\x00\x12\x10\n" +
"\fIpIfNonMatch\x10\x02\x12\x0e\n" +
-5
View File
@@ -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 {
@@ -110,6 +107,4 @@ message Config {
DomainStrategy domain_strategy = 1;
repeated RoutingRule rule = 2;
repeated BalancingRule balancing_rule = 3;
// Absolute path to the Lua routing script.
string script = 4;
}
-167
View File
@@ -1,167 +0,0 @@
package router
import (
"runtime"
"strings"
"github.com/xtls/xray-core/common/errors"
xlua "github.com/xtls/xray-core/common/lua"
"github.com/xtls/xray-core/common/net"
"github.com/xtls/xray-core/features/routing"
lua "github.com/yuin/gopher-lua"
)
const (
luaContextType = "xray.router.Context"
luaAttributesType = "xray.router.Attributes"
)
// RegisterLua makes xray.router available to routing scripts.
func (r *Router) RegisterLua(L *lua.LState) {
registerLuaContext(L)
L.PreloadModule("xray.router", func(L *lua.LState) int {
module := L.NewTable()
module.RawSetString("NetworkUnknown", lua.LNumber(net.Network_Unknown))
module.RawSetString("NetworkTCP", lua.LNumber(net.Network_TCP))
module.RawSetString("NetworkUDP", lua.LNumber(net.Network_UDP))
module.RawSetString("NetworkUNIX", lua.LNumber(net.Network_UNIX))
module.RawSetString("LocalOS", lua.LString(runtime.GOOS))
module.RawSetString("PickOutbound", L.NewFunction(func(L *lua.LState) int {
tag, ok := L.Get(2).(lua.LString)
if !ok {
L.ArgError(2, "balancer tag must be a string")
return 0
}
balancer, found := (*r.balancers.Load())[string(tag)]
if !found {
xlua.PushNil(L)
xlua.PushError(L, errors.New("balancer ", tag, " not found"))
return 2
}
outboundTag, err := balancer.PickOutbound()
xlua.PushString(L, outboundTag)
xlua.PushError(L, err)
return 2
}))
module.RawSetString("FindProcess", L.NewFunction(func(L *lua.LState) int {
pid, name, path, err := findProcess(checkLuaContext(L), net.FindProcess)
xlua.PushNumber(L, pid)
xlua.PushString(L, name)
xlua.PushString(L, path)
xlua.PushError(L, err)
return 4
}))
L.Push(module)
return 1
})
}
func registerLuaContext(L *lua.LState) {
attributes := L.NewTypeMetatable(luaAttributesType)
L.SetField(attributes, "__index", L.NewFunction(func(L *lua.LState) int {
values := L.CheckUserData(1).Value.(map[string]string)
key := L.CheckString(2)
if value, found := values[key]; found {
xlua.PushString(L, value)
} else {
xlua.PushNil(L)
}
return 1
}))
methods := L.NewTable()
L.SetFuncs(methods, map[string]lua.LGFunction{
"GetSourceIPs": func(L *lua.LState) int {
xlua.PushUserData(L, checkLuaContext(L).GetSourceIPs())
return 1
},
"GetTargetIPs": func(L *lua.LState) int {
xlua.PushUserData(L, checkLuaContext(L).GetTargetIPs())
return 1
},
"GetLocalIPs": func(L *lua.LState) int {
xlua.PushUserData(L, checkLuaContext(L).GetLocalIPs())
return 1
},
"GetAttributes": func(L *lua.LState) int {
values := L.NewUserData()
values.Value = checkLuaContext(L).GetAttributes()
L.SetMetatable(values, attributes)
L.Push(values)
return 1
},
})
L.SetField(L.NewTypeMetatable(luaContextType), "__index", methods)
}
func checkLuaContext(L *lua.LState) routing.Context {
ctx, ok := L.CheckUserData(1).Value.(routing.Context)
if !ok {
L.ArgError(1, "routing context expected")
}
return ctx
}
// callLuaHook invokes HandleRoute in the supplied state.
func (r *Router) callLuaHook(L *lua.LState, routeCtx routing.Context) (string, string, error) {
top := L.GetTop()
defer L.SetTop(top)
fn := L.GetGlobal("HandleRoute")
if fn.Type() != lua.LTFunction {
return "", "", errors.New("routing script must define HandleRoute(...)")
}
value := L.NewUserData()
value.Value = routeCtx
L.SetMetatable(value, L.GetTypeMetatable(luaContextType))
if err := L.CallByParam(lua.P{Fn: fn, NRet: 3, Protect: true},
value, lua.LString(routeCtx.GetInboundTag()), lua.LNumber(routeCtx.GetSourcePort()),
lua.LNumber(routeCtx.GetTargetPort()), lua.LNumber(routeCtx.GetLocalPort()),
lua.LString(strings.ToLower(routeCtx.GetTargetDomain())), lua.LNumber(routeCtx.GetNetwork()),
lua.LString(routeCtx.GetProtocol()), lua.LString(routeCtx.GetUser()),
lua.LNumber(routeCtx.GetVlessRoute()), lua.LBool(routeCtx.GetSkipDNSResolve())); err != nil {
return "", "", err
}
return readLuaRouteResult(L.Get(-3), L.Get(-2), L.Get(-1))
}
func readLuaRouteResult(tagValue, ruleValue, errorValue lua.LValue) (string, string, error) {
if err := xlua.ReadError(errorValue, "routing script error must be an error or string"); err != nil {
return "", "", err
}
tag, err := xlua.ReadOptionalString(tagValue, "routing script outboundTag must be a string or nil")
if err != nil || tag == "" {
return "", "", err
}
ruleTag, err := xlua.ReadOptionalString(ruleValue, "routing script ruleTag must be a string")
if err != nil {
return "", "", err
}
return tag, ruleTag, nil
}
type processFinder func(string, string, uint16, string, uint16) (int, string, string, error)
func findProcess(ctx routing.Context, finder processFinder) (int, string, string, error) {
sources := ctx.GetSourceIPs()
if len(sources) == 0 {
return 0, "", "", errors.New("process lookup requires a source IP")
}
var network string
switch ctx.GetNetwork() {
case net.Network_TCP:
network = "tcp"
case net.Network_UDP:
network = "udp"
default:
return 0, "", "", errors.New("process lookup requires TCP or UDP")
}
targetIP, targetPort := "", uint16(0)
if targets := ctx.GetTargetIPs(); len(targets) > 0 {
targetIP, targetPort = targets[0].String(), uint16(ctx.GetTargetPort())
}
return finder(network, sources[0].String(), uint16(ctx.GetSourcePort()), targetIP, targetPort)
}
-300
View File
@@ -1,300 +0,0 @@
package router
import (
"context"
go_errors "errors"
"runtime"
"strings"
"testing"
"github.com/xtls/xray-core/common/geodata"
"github.com/xtls/xray-core/common/net"
"github.com/xtls/xray-core/common/protocol"
"github.com/xtls/xray-core/common/session"
"github.com/xtls/xray-core/features/routing"
routing_session "github.com/xtls/xray-core/features/routing/session"
lua "github.com/yuin/gopher-lua"
)
type luaRouteTestContext struct {
*routing_session.Context
sourceIPs, targetIPs, localIPs []net.IP
}
func (c *luaRouteTestContext) GetSourceIPs() []net.IP { return c.sourceIPs }
func (c *luaRouteTestContext) GetTargetIPs() []net.IP { return c.targetIPs }
func (c *luaRouteTestContext) GetLocalIPs() []net.IP { return c.localIPs }
func newLuaRouteTestContext() *luaRouteTestContext {
return &luaRouteTestContext{
Context: &routing_session.Context{
Inbound: &session.Inbound{
Tag: "in", VlessRoute: 4321,
Source: net.TCPDestination(net.LocalHostIP, 1234),
Local: net.TCPDestination(net.LocalHostIP, 5678),
User: &protocol.MemoryUser{Email: "user@example.com"},
},
Outbound: &session.Outbound{
Target: net.TCPDestination(net.LocalHostIP, 443),
RouteTarget: net.TCPDestination(net.DomainAddress("MiXeD.Example."), 443),
},
Content: &session.Content{
Protocol: "tls", Attributes: map[string]string{"key": "value"}, SkipDNSResolve: true,
},
},
sourceIPs: []net.IP{{127, 0, 0, 2}},
targetIPs: []net.IP{{127, 0, 0, 3}},
localIPs: []net.IP{{127, 0, 0, 1}},
}
}
func newLuaRouterState(t *testing.T, script string) (*Router, *lua.LState) {
t.Helper()
r := new(Router)
if err := r.Init(context.Background(), &Config{}, nil, nil, nil); err != nil {
t.Fatal(err)
}
L := lua.NewState()
t.Cleanup(L.Close)
r.RegisterLua(L)
geodata.RegisterLua(L)
if err := L.DoString(script); err != nil {
t.Fatal(err)
}
return r, L
}
func TestLuaRouteBinding(t *testing.T) {
r, L := newLuaRouterState(t, `
local router = require("xray.router")
local matcher = require("xray.geodata").BuildIPMatcher("127.0.0.0/8")
assert(router.NetworkUnknown == 0 and router.NetworkTCP == 2)
assert(router.NetworkUDP == 3 and router.NetworkUNIX == 4)
assert(router.BuildIPMatcher == nil and router.BuildDomainMatcher == nil)
function HandleRoute(ctx, inboundTag, sourcePort, targetPort, localPort,
targetDomain, network, protocol, user, vlessRoute, skipDNSResolve, ...)
assert(select("#", ...) == 0)
assert(inboundTag == "in" and sourcePort == 1234 and targetPort == 443 and localPort == 5678)
assert(targetDomain == "mixed.example." and network == router.NetworkTCP)
assert(protocol == "tls" and user == "user@example.com" and vlessRoute == 4321 and skipDNSResolve)
assert(ctx.GetNetwork == nil and ctx.Context == nil)
savedContext = ctx
sourceIPs, targetIPs, localIPs = ctx:GetSourceIPs(), ctx:GetTargetIPs(), ctx:GetLocalIPs()
attributes = ctx:GetAttributes()
assert(matcher:AnyMatch(sourceIPs) and matcher:AnyMatch(targetIPs) and matcher:AnyMatch(localIPs))
assert(attributes.key == "value" and attributes.missing == nil)
assert(not pcall(function() attributes.key = "changed" end))
return "out", "rule"
end`)
ctx := newLuaRouteTestContext()
tag, rule, err := r.callLuaHook(L, ctx)
if err != nil || tag != "out" || rule != "rule" {
t.Fatalf("hook = %q, %q, %v", tag, rule, err)
}
if L.GetGlobal("savedContext").(*lua.LUserData).Value != ctx {
t.Fatal("routing context was copied")
}
for _, tc := range []struct {
name string
want []net.IP
}{
{"sourceIPs", ctx.sourceIPs},
{"targetIPs", ctx.targetIPs},
{"localIPs", ctx.localIPs},
} {
got := L.GetGlobal(tc.name).(*lua.LUserData).Value.([]net.IP)
if &got[0] != &tc.want[0] {
t.Fatalf("%s storage was copied", tc.name)
}
}
ctx.Content.Attributes["key"] = "updated"
L.SetGlobal("expectedOS", lua.LString(runtime.GOOS))
if err := L.DoString(`
assert(attributes.key == "updated")
assert(require("xray.router").LocalOS == expectedOS)`); err != nil {
t.Fatal(err)
}
}
func TestLuaRouteResult(t *testing.T) {
nativeErr := go_errors.New("native failure")
for _, tc := range []struct {
name, body, tag, rule, wantErr string
native bool
}{
{name: "route", body: `return "out", "rule"`, tag: "out", rule: "rule"},
{name: "no match", body: `return nil`},
{name: "empty tag", body: `return ""`},
{name: "no match ignores rule", body: `return nil, false`},
{name: "empty tag ignores rule", body: `return "", false`},
{name: "missing rule", body: `return "out"`, tag: "out"},
{name: "invalid tag", body: `return 1`, wantErr: "outboundTag"},
{name: "invalid rule", body: `return "out", false`, wantErr: "ruleTag"},
{name: "string error", body: `return nil, nil, "script failure"`, wantErr: "script failure"},
{name: "native error", body: `return nil, nil, nativeError`, native: true},
{name: "error overrides invalid tags", body: `return false, false, nativeError`, native: true},
{name: "invalid error", body: `return "out", "rule", false`, wantErr: "error or string"},
{name: "wrong error userdata", body: `return "out", "rule", wrongError`, wantErr: "error or string"},
{name: "runtime error", body: `error("runtime failure")`, wantErr: "runtime failure"},
} {
t.Run(tc.name, func(t *testing.T) {
r, L := newLuaRouterState(t, "function HandleRoute() "+tc.body+" end")
value := L.NewUserData()
value.Value = nativeErr
L.SetGlobal("nativeError", value)
wrong := L.NewUserData()
wrong.Value = "not a native error"
L.SetGlobal("wrongError", wrong)
L.Push(lua.LTrue)
tag, rule, err := r.callLuaHook(L, &routing_session.Context{})
if tag != tc.tag || rule != tc.rule {
t.Fatalf("result = %q, %q, %v", tag, rule, err)
}
switch {
case tc.native:
if err != nativeErr {
t.Fatalf("error = %v, want original error", err)
}
case tc.wantErr != "":
if err == nil || !strings.Contains(err.Error(), tc.wantErr) {
t.Fatalf("error = %v, want %q", err, tc.wantErr)
}
case err != nil:
t.Fatal(err)
}
if L.GetTop() != 1 || L.Get(1) != lua.LTrue {
t.Fatal("hook did not restore the stack")
}
})
}
}
func TestLuaRouteCancellation(t *testing.T) {
r, L := newLuaRouterState(t, `function HandleRoute() while true do end end`)
ctx, cancel := context.WithCancel(context.Background())
cancel()
L.SetContext(ctx)
if _, _, err := r.callLuaHook(L, &routing_session.Context{}); err == nil {
t.Fatal("CallLuaHook did not stop after context cancellation")
}
if L.Context() != ctx || L.GetTop() != 0 {
t.Fatal("CallLuaHook did not restore the Lua state")
}
}
func TestFindProcess(t *testing.T) {
for _, tc := range []struct {
name, network, target string
targetPort uint16
modify func(*luaRouteTestContext)
wantErr bool
}{
{name: "TCP", network: "tcp", target: "127.0.0.3", targetPort: 443},
{name: "UDP", network: "udp", target: "127.0.0.3", targetPort: 443, modify: func(c *luaRouteTestContext) {
c.Outbound.Target.Network = net.Network_UDP
}},
{name: "domain target", network: "tcp", modify: func(c *luaRouteTestContext) { c.targetIPs = nil }},
{name: "missing source", modify: func(c *luaRouteTestContext) { c.sourceIPs = nil }, wantErr: true},
{name: "unsupported network", modify: func(c *luaRouteTestContext) {
c.Outbound.Target.Network = net.Network_UNIX
}, wantErr: true},
} {
t.Run(tc.name, func(t *testing.T) {
ctx := newLuaRouteTestContext()
if tc.modify != nil {
tc.modify(ctx)
}
called := false
pid, name, path, err := findProcess(ctx, func(network, source string, sourcePort uint16, target string, targetPort uint16) (int, string, string, error) {
called = true
if network != tc.network || source != "127.0.0.2" || sourcePort != 1234 || target != tc.target || targetPort != tc.targetPort {
t.Fatalf("endpoints = %s %s:%d -> %s:%d", network, source, sourcePort, target, targetPort)
}
return 42, "process", "/path/process", nil
})
if tc.wantErr {
if err == nil || called {
t.Fatalf("findProcess = %d, %q, %q, %v", pid, name, path, err)
}
return
}
if err != nil || !called || pid != 42 || name != "process" || path != "/path/process" {
t.Fatalf("findProcess = %d, %q, %q, %v", pid, name, path, err)
}
})
}
}
// BenchmarkLuaRouteHookCall isolates a preloaded Lua hook and its routing context bridge.
// The direct case runs an equivalent native routing rule.
func BenchmarkLuaRouteHookCall(b *testing.B) {
r := new(Router)
if err := r.Init(context.Background(), &Config{Rule: []*RoutingRule{{
TargetTag: &RoutingRule_Tag{Tag: "out"},
RuleTag: "rule",
InboundTag: []string{"in"},
Networks: []net.Network{net.Network_TCP},
Ip: []*geodata.IPRule{{
Value: &geodata.IPRule_Custom{Custom: &geodata.CIDRRule{
Cidr: &geodata.CIDR{Ip: []byte{127, 0, 0, 0}, Prefix: 8},
}},
}},
}}}, nil, nil, nil); err != nil {
b.Fatal(err)
}
L := lua.NewState()
defer L.Close()
r.RegisterLua(L)
geodata.RegisterLua(L)
if err := L.DoString(`
local router = require("xray.router")
local matcher = require("xray.geodata").BuildIPMatcher("127.0.0.0/8")
function HandleRoute(ctx, inboundTag, sourcePort, targetPort, localPort,
targetDomain, network, protocol, user, vlessRoute, skipDNSResolve)
if inboundTag == "in" and network == router.NetworkTCP and matcher:AnyMatch(ctx:GetTargetIPs()) then
return "out", "rule"
end
end
`); err != nil {
b.Fatal(err)
}
L.SetContext(context.Background())
routeCtx := newLuaRouteTestContext()
for _, benchmark := range []struct {
name string
route func() (string, string, error)
}{
{"direct", func() (string, string, error) {
route, err := r.PickRoute(routeCtx)
if err != nil {
return "", "", err
}
return route.GetOutboundTag(), route.GetRuleTag(), nil
}},
{"lua_hook", func() (string, string, error) {
return r.callLuaHook(L, routeCtx)
}},
} {
b.Run(benchmark.name, func(b *testing.B) {
b.ReportAllocs()
b.ResetTimer()
var tag, rule string
var err error
for i := 0; i < b.N; i++ {
tag, rule, err = benchmark.route()
if err != nil {
b.Fatal(err)
}
}
b.StopTimer()
if tag != "out" || rule != "rule" {
b.Fatalf("route() = %q, %q; want out, rule", tag, rule)
}
})
}
}
var _ routing.Context = (*luaRouteTestContext)(nil)
+114 -76
View File
@@ -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,10 +17,8 @@ import (
// Router is an implementation of routing.Router.
type Router struct {
domainStrategy Config_DomainStrategy
rules atomic.Pointer[[]*Rule]
scriptPath string
script *scriptEngine
balancers atomic.Pointer[map[string]*Balancer]
rules []*Rule
balancers map[string]*Balancer
dns dns.Client
ctx context.Context
@@ -42,23 +38,61 @@ type Route struct {
// Init initializes the Router.
func (r *Router) Init(ctx context.Context, config *Config, d dns.Client, ohm outbound.Manager, dispatcher routing.Dispatcher) error {
r.domainStrategy = config.DomainStrategy
r.scriptPath = config.Script
r.dns = d
r.ctx = ctx
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.
func (r *Router) PickRoute(ctx routing.Context) (routing.Route, error) {
if r.script != nil {
return r.script.pickRoute(ctx)
}
originalCtx := ctx
rule, ctx, err := r.pickRouteInternal(ctx)
if err != nil {
@@ -90,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)
@@ -113,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{
@@ -126,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,
@@ -202,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
}
@@ -217,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
}
@@ -228,19 +276,12 @@ func (r *Router) pickRouteInternal(ctx routing.Context) (*Rule, routing.Context,
// Start implements common.Runnable.
func (r *Router) Start() error {
if r.scriptPath != "" {
engine, err := newScriptEngine(r.scriptPath, r)
if err != nil {
return errors.New("failed to initialize routing script").Base(err)
}
r.script = engine
}
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()
}
@@ -249,12 +290,9 @@ func closeWebhooks(rules []*Rule) {
// Close implements common.Closable.
func (r *Router) Close() error {
if r.script != nil {
r.script.close()
}
r.mu.Lock()
defer r.mu.Unlock()
closeWebhooks(*r.rules.Load())
r.closeWebhooks()
return nil
}
-68
View File
@@ -1,68 +0,0 @@
package router
import (
"time"
"github.com/xtls/xray-core/app/dns"
"github.com/xtls/xray-core/common"
"github.com/xtls/xray-core/common/errors"
"github.com/xtls/xray-core/common/geodata"
"github.com/xtls/xray-core/common/log"
xlua "github.com/xtls/xray-core/common/lua"
"github.com/xtls/xray-core/features/routing"
lua "github.com/yuin/gopher-lua"
)
const scriptExecutionTimeout = 6 * time.Second
type scriptEngine struct {
router *Router
pool *xlua.Pool
}
func newScriptEngine(path string, router *Router) (*scriptEngine, error) {
program, err := xlua.CompileFile(path)
if err != nil {
return nil, err
}
e := &scriptEngine{router: router}
e.pool, err = xlua.NewPool(router.ctx, scriptExecutionTimeout, program.NewStateFactory(
scriptExecutionTimeout*20,
func(L *lua.LState) {
geodata.RegisterLua(L)
log.RegisterLua(L)
router.RegisterLua(L)
dns.RegisterLua(L, router.dns)
},
func(L *lua.LState) error {
if L.GetGlobal("HandleRoute").Type() != lua.LTFunction {
return errors.New("routing script must define HandleRoute(...)")
}
return nil
}))
if err != nil {
return nil, err
}
errors.LogInfo(router.ctx, "routing script initialized from ", path)
return e, nil
}
func (e *scriptEngine) close() {
e.pool.Close()
}
func (e *scriptEngine) pickRoute(ctx routing.Context) (routing.Route, error) {
var tag, ruleTag string
err := e.pool.WithState(nil, 0, func(L *lua.LState) error {
var hookErr error
tag, ruleTag, hookErr = e.router.callLuaHook(L, ctx)
return hookErr
})
if err != nil {
return nil, err
}
if tag == "" {
return nil, common.ErrNoClue
}
return &Route{Context: ctx, outboundTag: tag, ruleTag: ruleTag}, nil
}
-372
View File
@@ -1,372 +0,0 @@
package router
import (
"context"
stdnet "net"
"os"
"path/filepath"
"sync"
"sync/atomic"
"testing"
"time"
wireDNS "github.com/miekg/dns"
"github.com/xtls/xray-core/app/dispatcher"
appdns "github.com/xtls/xray-core/app/dns"
"github.com/xtls/xray-core/app/proxyman"
_ "github.com/xtls/xray-core/app/proxyman/outbound"
"github.com/xtls/xray-core/common"
"github.com/xtls/xray-core/common/net"
"github.com/xtls/xray-core/common/serial"
"github.com/xtls/xray-core/core"
featureDNS "github.com/xtls/xray-core/features/dns"
"github.com/xtls/xray-core/features/outbound"
"github.com/xtls/xray-core/features/routing"
routing_session "github.com/xtls/xray-core/features/routing/session"
"github.com/xtls/xray-core/proxy/blackhole"
"github.com/xtls/xray-core/proxy/freedom"
)
type luaRouteDNSClient struct {
featureDNS.Client
lookup func(string, featureDNS.IPOption) ([]net.IP, uint32, error)
}
func (d *luaRouteDNSClient) LookupIP(domain string, option featureDNS.IPOption) ([]net.IP, uint32, error) {
return d.lookup(domain, option)
}
type luaRouteOutboundManager struct{ outbound.Manager }
func (*luaRouteOutboundManager) Select(selectors []string) []string { return selectors }
func writeRouteScript(t *testing.T, script string) string {
t.Helper()
path := filepath.Join(t.TempDir(), "route.lua")
if err := os.WriteFile(path, []byte(script), 0o600); err != nil {
t.Fatal(err)
}
return path
}
func startLuaRouter(t *testing.T, script string, d featureDNS.Client, config *Config) *Router {
t.Helper()
if config == nil {
config = &Config{}
}
config.Script = writeRouteScript(t, script)
r := new(Router)
if err := r.Init(context.Background(), config, d, &luaRouteOutboundManager{}, nil); err != nil {
t.Fatal(err)
}
if err := r.Start(); err != nil {
t.Fatal(err)
}
t.Cleanup(func() {
if err := r.Close(); err != nil {
t.Error(err)
}
})
return r
}
func TestRouterScriptStartup(t *testing.T) {
for _, tc := range []struct{ name, script string }{
{"syntax error", "function HandleRoute("},
{"missing hook", "value = 1"},
{"initialization error", `error("setup failed")`},
} {
t.Run(tc.name, func(t *testing.T) {
r := new(Router)
if err := r.Init(context.Background(), &Config{Script: writeRouteScript(t, tc.script)}, nil, nil, nil); err != nil {
t.Fatal(err)
}
defer r.Close()
if err := r.Start(); err == nil {
t.Fatal("Start accepted an invalid routing script")
}
})
}
}
func TestRouterScriptRouting(t *testing.T) {
var dnsCalls atomic.Int32
d := &luaRouteDNSClient{lookup: func(string, featureDNS.IPOption) ([]net.IP, uint32, error) {
dnsCalls.Add(1)
return []net.IP{{1, 2, 3, 4}}, 60, nil
}}
r := startLuaRouter(t, `
function HandleRoute(ctx, inbound)
if inbound == "miss" then return nil end
return "lua-out", "lua-rule"
end`, d, &Config{
DomainStrategy: Config_IpOnDemand,
Rule: []*RoutingRule{{
TargetTag: &RoutingRule_Tag{Tag: "json-out"},
Networks: []net.Network{net.Network_TCP},
}},
})
ctx := newLuaRouteTestContext()
ctx.Content.SkipDNSResolve = false
route, err := r.PickRoute(ctx)
if err != nil || route.GetOutboundTag() != "lua-out" || route.GetRuleTag() != "lua-rule" || route.(*Route).Context != ctx {
t.Fatalf("route = %v, %v", route, err)
}
ctx.Inbound.Tag = "miss"
if route, err := r.PickRoute(ctx); route != nil || err != common.ErrNoClue {
t.Fatalf("miss = %v, %v", route, err)
}
if dnsCalls.Load() != 0 {
t.Fatal("script routing implicitly resolved DNS")
}
}
func TestRouterScriptModules(t *testing.T) {
ips := []net.IP{{127, 0, 0, 7}}
calls := 0
d := &luaRouteDNSClient{lookup: func(domain string, option featureDNS.IPOption) ([]net.IP, uint32, error) {
calls++
if domain != "mixed.example." || !option.IPv4Enable || option.IPv6Enable || !option.FakeEnable {
t.Fatalf("dns.Query arguments = %q, %+v", domain, option)
}
return ips, 17, nil
}}
r := startLuaRouter(t, `
local dns = require("xray.dns")
local matcher = require("xray.geodata").BuildIPMatcher("127.0.0.0/8")
assert(dns.Servers == nil and type(dns.Query) == "function")
assert(type(require("xray.log").Info) == "function")
function HandleRoute(ctx, inbound, sourcePort, targetPort, localPort, domain)
local ips, ttl, err = dns.Query(domain, true, false, true)
assert(not err and ttl == 17)
assert(matcher:AnyMatch(ips) and matcher:AnyMatch(ctx:GetTargetIPs()))
return "out"
end`, d, nil)
if _, err := r.PickRoute(newLuaRouteTestContext()); err != nil {
t.Fatal(err)
}
if calls != 1 {
t.Fatalf("DNS calls = %d, want 1", calls)
}
}
func TestRouterScriptBalancerReload(t *testing.T) {
config := func(tag string) *Config {
return &Config{BalancingRule: []*BalancingRule{{
Tag: "balance", Strategy: "roundrobin", OutboundSelector: []string{tag},
}}}
}
r := startLuaRouter(t, `
local router = require("xray.router")
function HandleRoute()
local tag, err = router:PickOutbound("balance")
return tag, "balanced", err
end`, nil, config("old"))
pick := func(want string) {
t.Helper()
route, err := r.PickRoute(&routing_session.Context{})
if err != nil || route.GetOutboundTag() != want || route.GetRuleTag() != "balanced" {
t.Fatalf("route = %v, %v, want %q", route, err, want)
}
}
pick("old")
if err := r.SetOverrideTarget("balance", "override"); err != nil {
t.Fatal(err)
}
pick("override")
if err := r.SetOverrideTarget("balance", ""); err != nil {
t.Fatal(err)
}
if err := r.ReloadRules(config("new"), false); err != nil {
t.Fatal(err)
}
pick("new")
}
func TestRouterScriptConcurrentBalancerReload(t *testing.T) {
config := func(tag string) *Config {
return &Config{BalancingRule: []*BalancingRule{{
Tag: "balance", Strategy: "roundrobin", OutboundSelector: []string{tag},
}}}
}
r := startLuaRouter(t, `
local router = require("xray.router")
function HandleRoute()
local tag, err = router:PickOutbound("balance")
return tag, nil, err
end`, nil, config("a"))
var wg sync.WaitGroup
for range 4 {
wg.Go(func() {
for range 20 {
route, err := r.PickRoute(&routing_session.Context{})
if err != nil {
t.Errorf("PickRoute: %v", err)
return
}
if tag := route.GetOutboundTag(); tag != "a" && tag != "b" {
t.Errorf("unexpected tag %q", tag)
}
}
})
}
wg.Go(func() {
for range 20 {
for _, tag := range []string{"a", "b"} {
if err := r.ReloadRules(config(tag), false); err != nil {
t.Error(err)
return
}
}
}
})
wg.Wait()
}
func TestRouterScriptStateReuse(t *testing.T) {
r := startLuaRouter(t, `
local calls = 0
function HandleRoute(ctx, inbound)
calls = calls + 1
if inbound == "miss" then return nil end
if inbound == "fail" then error("failed") end
return tostring(calls)
end`, nil, nil)
ctx := newLuaRouteTestContext()
pick := func(want string) {
t.Helper()
route, err := r.PickRoute(ctx)
if err != nil || route.GetOutboundTag() != want {
t.Fatalf("route = %v, %v, want %q", route, err, want)
}
}
pick("1")
ctx.Inbound.Tag = "miss"
if _, err := r.PickRoute(ctx); err != common.ErrNoClue {
t.Fatalf("miss = %v", err)
}
ctx.Inbound.Tag = "in"
pick("3")
ctx.Inbound.Tag = "fail"
if _, err := r.PickRoute(ctx); err == nil {
t.Fatal("script error was ignored")
}
ctx.Inbound.Tag = "in"
pick("1")
}
func TestRouterScriptDNSDispatcherReentry(t *testing.T) {
conn, err := stdnet.ListenPacket("udp4", "127.0.0.1:0")
if err != nil {
t.Fatal(err)
}
port := conn.LocalAddr().(*stdnet.UDPAddr).Port
ready, stopped := make(chan struct{}), make(chan error, 1)
var queries atomic.Int32
server := &wireDNS.Server{
PacketConn: conn,
NotifyStartedFunc: func() {
close(ready)
},
Handler: wireDNS.HandlerFunc(func(w wireDNS.ResponseWriter, query *wireDNS.Msg) {
queries.Add(1)
response := new(wireDNS.Msg).SetReply(query)
for _, question := range query.Question {
if question.Name == "nested.example." && question.Qtype == wireDNS.TypeA {
response.Answer = append(response.Answer, &wireDNS.A{
Hdr: wireDNS.RR_Header{Name: question.Name, Rrtype: wireDNS.TypeA, Class: wireDNS.ClassINET, Ttl: 60},
A: stdnet.IP{127, 0, 0, 7},
})
}
}
if err := w.WriteMsg(response); err != nil {
t.Error(err)
}
}),
}
go func() { stopped <- server.ActivateAndServe() }()
defer func() {
server.Shutdown()
select {
case err := <-stopped:
if err != nil {
t.Error(err)
}
case <-time.After(3 * time.Second):
t.Error("DNS server did not stop")
}
}()
select {
case <-ready:
case err := <-stopped:
t.Fatalf("DNS server startup: %v", err)
case <-time.After(3 * time.Second):
t.Fatal("DNS server did not start")
}
dnsScript := writeRouteScript(t, `
local server = require("xray.dns").Servers[1]
function HandleDNSQuery(domain, ipv4, ipv6, fake)
return server:Query(domain, ipv4, ipv6, fake)
end`)
routerScript := writeRouteScript(t, `
local router = require("xray.router")
local dns = require("xray.dns")
local matcher = require("xray.geodata").BuildIPMatcher("127.0.0.7")
local active = false
function HandleRoute(ctx, inbound, sourcePort, targetPort, localPort, domain, network,
protocol, user, vlessRoute, skipDNSResolve)
assert(not active, "borrowed Router VM reentered")
if inbound == "dns" then
assert(network == router.NetworkUDP and skipDNSResolve == false)
return "direct", "dns-route"
end
active = true
local ips, ttl, err = dns.Query("nested.example", true, false, false)
assert(not err and matcher:AnyMatch(ips) and active)
active = false
return "direct", "outer-route"
end`)
instance, err := core.New(&core.Config{
App: []*serial.TypedMessage{
serial.ToTypedMessage(&appdns.Config{
Tag: "dns", Script: dnsScript, DisableCache: true,
NameServer: []*appdns.NameServer{{
Id: "upstream", TimeoutMs: 1000,
Address: &net.Endpoint{
Network: net.Network_UDP,
Address: &net.IPOrDomain{Address: &net.IPOrDomain_Ip{Ip: []byte{127, 0, 0, 1}}},
Port: uint32(port),
},
}},
}),
serial.ToTypedMessage(&Config{Script: routerScript}),
serial.ToTypedMessage(&dispatcher.Config{}),
serial.ToTypedMessage(&proxyman.OutboundConfig{}),
},
Outbound: []*core.OutboundHandlerConfig{
{Tag: "default", ProxySettings: serial.ToTypedMessage(&blackhole.Config{})},
{Tag: "direct", ProxySettings: serial.ToTypedMessage(&freedom.Config{
FinalRules: []*freedom.FinalRuleConfig{{Action: freedom.RuleAction_Allow}},
})},
},
})
if err != nil {
t.Fatal(err)
}
defer instance.Close()
if err := instance.Start(); err != nil {
t.Fatal(err)
}
r := instance.GetFeature(routing.RouterType()).(*Router)
route, err := r.PickRoute(newLuaRouteTestContext())
if err != nil || route.GetOutboundTag() != "direct" || route.GetRuleTag() != "outer-route" {
t.Fatalf("nested DNS routing = %v, %v", route, err)
}
if queries.Load() == 0 {
t.Fatal("DNS query did not pass through the dispatcher")
}
}
+23 -17
View File
@@ -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
View File
@@ -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
View File
@@ -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 {
+7 -19
View File
@@ -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
}
}
+1 -3
View File
@@ -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")
+1 -11
View File
@@ -5,8 +5,7 @@ import (
)
type windowsReader struct {
bufs []syscall.WSABuf
ready bool
bufs []syscall.WSABuf
}
func (r *windowsReader) Init(bs []*Buffer) {
@@ -16,7 +15,6 @@ func (r *windowsReader) Init(bs []*Buffer) {
for _, b := range bs {
r.bufs = append(r.bufs, syscall.WSABuf{Len: uint32(Size), Buf: &b.v[0]})
}
r.ready = false
}
func (r *windowsReader) Clear() {
@@ -27,14 +25,6 @@ func (r *windowsReader) Clear() {
}
func (r *windowsReader) Read(fd uintptr) int32 {
// On the first invocation, we return -1 to indicate "not ready"
// to make rawConn.Read wait for readability using the runtime's own mechanism
// because syscall.WSARecv() is a blocking call when used with nil OVERLAPPED
if !r.ready {
r.ready = true
return -1
}
var nBytes uint32
var flags uint32
err := syscall.WSARecv(syscall.Handle(fd), &r.bufs[0], uint32(len(r.bufs)), &nBytes, &flags, nil, nil)
+1 -3
View File
@@ -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() {
+3 -3
View File
@@ -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
View File
@@ -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
}
+15 -6
View File
@@ -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)
}
-49
View File
@@ -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() }
+50 -33
View File
@@ -82,10 +82,19 @@ func (f *MphDomainMatcherFactory) BuildMatcher(rules []*DomainRule) (DomainMatch
}
g.Add(m, uint32(i))
case *DomainRule_Geosite:
err := loadSiteMatchers(v.Geosite, func(m strmatcher.Matcher) { g.Add(m, uint32(i)) })
domains, err := loadSiteWithAttrs(v.Geosite.File, v.Geosite.Code, v.Geosite.Attrs)
if err != nil {
return nil, err
}
for j, d := range domains {
domains[j] = nil // peak mem
m, err := parseDomain(d)
if err != nil {
errors.LogError(context.Background(), "ignore invalid geosite entry in ", v.Geosite.File, ":", v.Geosite.Code, " at index ", j, ", ", err)
continue
}
g.Add(m, uint32(i))
}
default:
panic("unknown domain rule type")
}
@@ -99,12 +108,12 @@ func (f *MphDomainMatcherFactory) BuildMatcher(rules []*DomainRule) (DomainMatch
return g, nil
}
type CompactMphDomainMatcherFactory struct {
type CompactDomainMatcherFactory struct {
sync.Mutex
shared *utils.WeakCacheMap[string, strmatcher.MphValueMatcher]
shared *utils.WeakCacheMap[string, strmatcher.LinearAnyMatcher]
}
func (f *CompactMphDomainMatcherFactory) getOrCreateFrom(rule *GeoSiteRule) (*strmatcher.MphValueMatcher, error) {
func (f *CompactDomainMatcherFactory) getOrCreateFrom(rule *GeoSiteRule) (strmatcher.MatcherSet, error) {
key := rule.File + ":" + rule.Code + "@" + rule.Attrs
f.Lock()
@@ -116,23 +125,33 @@ func (f *CompactMphDomainMatcherFactory) getOrCreateFrom(rule *GeoSiteRule) (*st
}
errors.LogDebug(context.Background(), "geodata geosite matcher cache MISS ", key)
s := strmatcher.NewMphValueMatcher()
if err := loadSiteMatchers(rule, func(m strmatcher.Matcher) { s.Add(m, 0) }); err != nil {
s := strmatcher.NewLinearAnyMatcher()
domains, err := loadSiteWithAttrs(rule.File, rule.Code, rule.Attrs)
if err != nil {
return nil, err
}
if err := s.Build(); err != nil {
return nil, err
for i, d := range domains {
domains[i] = nil // peak mem
m, err := parseDomain(d)
if err != nil {
errors.LogError(context.Background(), "ignore invalid geosite entry in ", rule.File, ":", rule.Code, " at index ", i, ", ", err)
continue
}
s.Add(m)
}
f.shared.Store(key, s)
return s, nil
return s, err
}
// BuildMatcher implements DomainMatcherFactory.
func (f *CompactMphDomainMatcherFactory) BuildMatcher(rules []*DomainRule) (DomainMatcher, error) {
func (f *CompactDomainMatcherFactory) BuildMatcher(rules []*DomainRule) (DomainMatcher, error) {
if len(rules) == 0 {
return nil, errors.New("empty domain rule list")
}
compact := new(CompactMphDomainMatcher)
compact := &CompactDomainMatcher{
matchers: make([]strmatcher.MatcherSet, 0, len(rules)),
values: make([]uint32, 0, len(rules)),
}
for i, r := range rules {
switch v := r.Value.(type) {
case *DomainRule_Custom:
@@ -149,7 +168,8 @@ func (f *CompactMphDomainMatcherFactory) BuildMatcher(rules []*DomainRule) (Doma
if err != nil {
return nil, err
}
compact.combiner.Add(m, uint32(i))
compact.matchers = append(compact.matchers, m)
compact.values = append(compact.values, uint32(i))
default:
panic("unknown domain rule type")
}
@@ -157,40 +177,37 @@ func (f *CompactMphDomainMatcherFactory) BuildMatcher(rules []*DomainRule) (Doma
return compact, nil
}
type CompactMphDomainMatcher struct {
type CompactDomainMatcher struct {
custom strmatcher.ValueMatcher
combiner strmatcher.MphValueMatcherCombiner
matchers []strmatcher.MatcherSet
values []uint32
}
// Match implements DomainMatcher.
func (c *CompactMphDomainMatcher) Match(input string) []uint32 {
result := c.combiner.Match(input)
func (c *CompactDomainMatcher) Match(input string) []uint32 {
var result []uint32
if c.custom != nil {
result = append(c.custom.Match(input), result...)
result = append(result, c.custom.Match(input)...)
}
for i, m := range c.matchers {
if m.MatchAny(input) {
result = append(result, c.values[i])
}
}
return result
}
// MatchAny implements DomainMatcher.
func (c *CompactMphDomainMatcher) MatchAny(input string) bool {
func (c *CompactDomainMatcher) MatchAny(input string) bool {
if c.custom != nil && c.custom.MatchAny(input) {
return true
}
return c.combiner.MatchAny(input)
}
// loadSiteMatchers calls add with a matcher for every domain of the geosite rule and logs the invalid ones.
func loadSiteMatchers(rule *GeoSiteRule, add func(strmatcher.Matcher)) error {
i := 0
return loadSite(rule.File, rule.Code, rule.Attrs, func(t Domain_Type, value []byte) {
m, err := parseDomain(&Domain{Type: t, Value: string(value)})
if err != nil {
errors.LogError(context.Background(), "ignore invalid geosite entry in ", rule.File, ":", rule.Code, " at index ", i, ", ", err)
} else {
add(m)
for _, m := range c.matchers {
if m.MatchAny(input) {
return true
}
i++
})
}
return false
}
func parseDomain(d *Domain) (strmatcher.Matcher, error) {
@@ -214,7 +231,7 @@ func parseDomain(d *Domain) (strmatcher.Matcher, error) {
func newDomainMatcherFactory() DomainMatcherFactory {
switch runtime.GOOS {
case "ios", "android":
return &CompactMphDomainMatcherFactory{shared: utils.NewWeakCacheMap[string, strmatcher.MphValueMatcher]()}
return &CompactDomainMatcherFactory{shared: utils.NewWeakCacheMap[string, strmatcher.LinearAnyMatcher]()}
default:
return &MphDomainMatcherFactory{shared: utils.NewWeakCacheMap[string, strmatcher.MphValueMatcher]()}
}
+2 -76
View File
@@ -4,7 +4,6 @@ import (
"path/filepath"
"reflect"
"slices"
"sync"
"testing"
"github.com/xtls/xray-core/common/geodata/strmatcher"
@@ -12,7 +11,7 @@ import (
)
func TestCompactDomainMatcher_PreservesCustomRuleIndices(t *testing.T) {
factory := &CompactMphDomainMatcherFactory{shared: utils.NewWeakCacheMap[string, strmatcher.MphValueMatcher]()}
factory := &CompactDomainMatcherFactory{shared: utils.NewWeakCacheMap[string, strmatcher.LinearAnyMatcher]()}
matcher, err := factory.BuildMatcher([]*DomainRule{
{Value: &DomainRule_Custom{Custom: &Domain{Type: Domain_Full, Value: "example.com"}}},
{Value: &DomainRule_Custom{Custom: &Domain{Type: Domain_Domain, Value: "example.com"}}},
@@ -33,7 +32,7 @@ func TestCompactDomainMatcher_PreservesCustomRuleIndices(t *testing.T) {
func TestCompactDomainMatcher_PreservesMixedRuleIndices(t *testing.T) {
t.Setenv("xray.location.asset", filepath.Join("..", "..", "resources"))
factory := &CompactMphDomainMatcherFactory{shared: utils.NewWeakCacheMap[string, strmatcher.MphValueMatcher]()}
factory := &CompactDomainMatcherFactory{shared: utils.NewWeakCacheMap[string, strmatcher.LinearAnyMatcher]()}
matcher, err := factory.BuildMatcher([]*DomainRule{
{Value: &DomainRule_Geosite{Geosite: &GeoSiteRule{File: DefaultGeoSiteDat, Code: "CN"}}},
{Value: &DomainRule_Custom{Custom: &Domain{Type: Domain_Full, Value: "163.com"}}},
@@ -73,76 +72,3 @@ func TestMphDomainMatcher_MatchReturnsDetachedSlice(t *testing.T) {
t.Fatalf("Match() after caller mutation = %v, want %v", gotAgain, []uint32{0, 1})
}
}
// DNS sorts every Match result in place, so a matcher must never hand out a
// slice it keeps, also when only its keyword or regex part matches.
func TestDomainMatcher_MatchResultsCanBeSortedConcurrently(t *testing.T) {
t.Setenv("xray.location.asset", filepath.Join("..", "..", "resources"))
rules := []*DomainRule{
{Value: &DomainRule_Custom{Custom: &Domain{Type: Domain_Full, Value: "example.com"}}},
{Value: &DomainRule_Custom{Custom: &Domain{Type: Domain_Domain, Value: "example.com"}}},
{Value: &DomainRule_Custom{Custom: &Domain{Type: Domain_Substr, Value: "exam"}}},
{Value: &DomainRule_Custom{Custom: &Domain{Type: Domain_Regex, Value: `^ex.*\.org$`}}},
{Value: &DomainRule_Custom{Custom: &Domain{Type: Domain_Substr, Value: "exam"}}},
{Value: &DomainRule_Geosite{Geosite: &GeoSiteRule{File: DefaultGeoSiteDat, Code: "CN"}}},
{Value: &DomainRule_Custom{Custom: &Domain{Type: Domain_Full, Value: "only.full.test"}}},
}
cases := []struct {
input string
want []uint32
}{
{"example.com", []uint32{0, 1, 2, 4}},
{"www.example.com", []uint32{1, 2, 4}},
{"exam.net", []uint32{2, 4}}, // keyword part only
{"example.org", []uint32{2, 3, 4}},
{"163.com", []uint32{5}},
{"www.163.com", []uint32{5}},
{"only.full.test", []uint32{6}}, // full part only
{"nomatch.test", nil},
}
factories := map[string]DomainMatcherFactory{
"mph": &MphDomainMatcherFactory{shared: utils.NewWeakCacheMap[string, strmatcher.MphValueMatcher]()},
"compact": &CompactMphDomainMatcherFactory{shared: utils.NewWeakCacheMap[string, strmatcher.MphValueMatcher]()},
}
for name, factory := range factories {
t.Run(name, func(t *testing.T) {
matcher, err := factory.BuildMatcher(rules)
if err != nil {
t.Fatalf("BuildMatcher() failed: %v", err)
}
for _, c := range cases {
got := matcher.Match(c.input)
if sorted := slices.Sorted(slices.Values(got)); !slices.Equal(sorted, c.want) {
t.Fatalf("Match(%q) = %v, want %v", c.input, sorted, c.want)
}
got = got[:cap(got)]
for j := range got {
got[j] = ^uint32(0)
}
if again := slices.Sorted(slices.Values(matcher.Match(c.input))); !slices.Equal(again, c.want) {
t.Fatalf("Match(%q) after caller mutation = %v, want %v", c.input, again, c.want)
}
}
var wg sync.WaitGroup
for range 8 {
wg.Add(1)
go func() {
defer wg.Done()
for range 500 {
for _, c := range cases {
got := matcher.Match(c.input)
slices.Sort(got)
if !slices.Equal(got, c.want) {
t.Errorf("Match(%q) = %v, want %v", c.input, got, c.want)
return
}
}
}
}()
}
wg.Wait()
})
}
}
+62 -213
View File
@@ -5,14 +5,11 @@ import (
"bytes"
"io"
"runtime"
"slices"
"strings"
"unicode/utf8"
"github.com/xtls/xray-core/common/errors"
"github.com/xtls/xray-core/common/platform/filesystem"
"google.golang.org/protobuf/encoding/protowire"
"google.golang.org/protobuf/proto"
)
@@ -55,56 +52,17 @@ func loadIP(file, code string) ([]*CIDR, error) {
return geoip.Cidr, nil
}
// loadSite calls fn, in file order, with the type and value of every domain of the geosite code
// that has all the "@"-separated attrs. It decodes the entry while reading the file instead of
// unmarshalling it into a []*Domain, so value is only valid during fn.
func loadSite(file, code, attrs string, fn func(Domain_Type, []byte)) error {
runtime.GC() // peak mem
r, err := filesystem.OpenAsset(file)
func loadSite(file, code string) ([]*Domain, error) {
bs, err := loadFile(file, code)
if err != nil {
return errors.New("failed to open ", file).Base(err)
return nil, err
}
defer r.Close()
br := bufio.NewReaderSize(r, 64*1024)
n, err := seek(br, []byte(code))
if err != nil {
return errors.New("failed to load code ", code, " from ", file).Base(err)
defer runtime.GC() // peak mem
var geosite GeoSite
if err := proto.Unmarshal(bs, &geosite); err != nil {
return nil, errors.New("error unmarshal Site in ", file, ":", code).Base(err)
}
loadErr := func(err error) error {
if err == io.EOF {
err = io.ErrUnexpectedEOF
}
return errors.New("failed to load code ", code, " from ", file).Base(err)
}
unmarshalErr := func(err error) error {
return errors.New("error unmarshal Site in ", file, ":", code).Base(err)
}
d := newSiteDecoder(attrs, fn)
for n > 0 {
w, err := br.Peek(min(n, br.Size()))
if err != nil {
return loadErr(err)
}
used, err := d.decode(w, len(w) < n)
if err != nil {
return unmarshalErr(err)
}
if used == 0 {
break // a field longer than the buffer
}
br.Discard(used)
n -= used
}
if n > 0 {
w := make([]byte, n)
if _, err := io.ReadFull(br, w); err != nil {
return loadErr(err)
}
if _, err := d.decode(w, false); err != nil {
return unmarshalErr(err)
}
}
return nil
return geosite.Domain, nil
}
func decodeVarint(br *bufio.Reader) (uint64, error) {
@@ -124,63 +82,68 @@ func decodeVarint(br *bufio.Reader) (uint64, error) {
}
func find(r io.Reader, code []byte, readBody bool) ([]byte, error) {
br := bufio.NewReaderSize(r, 64*1024)
bodyL, err := seek(br, code)
if err != nil || !readBody {
return nil, err
}
out := make([]byte, bodyL)
if _, err := io.ReadFull(br, out); err != nil {
return nil, err
}
return out, nil
}
// seek advances br to the body of the entry for code and returns the body length.
func seek(br *bufio.Reader, code []byte) (int, error) {
codeL := len(code)
if codeL == 0 {
return 0, errors.New("empty code")
return nil, errors.New("empty code")
}
br := bufio.NewReaderSize(r, 64*1024)
need := 2 + codeL // TODO: if code too long
prefixBuf := make([]byte, need)
for {
if _, err := br.ReadByte(); err != nil {
return 0, err
return nil, err
}
x, err := decodeVarint(br)
if err != nil {
return 0, err
return nil, err
}
bodyL := int(x)
if bodyL <= 0 {
return 0, errors.New("invalid body length: ", bodyL)
return nil, errors.New("invalid body length: ", bodyL)
}
// Peek no more than the buffer holds: a code longer than the buffer cannot match a single
// length byte anyway, so a short peek only skips it, as base find (io.ReadFull) does.
prefix, err := br.Peek(min(bodyL, need, br.Size()))
if err != nil {
if err == io.EOF && len(prefix) > 0 {
err = io.ErrUnexpectedEOF // as io.ReadFull
prefixL := bodyL
if prefixL > need {
prefixL = need
}
prefix := prefixBuf[:prefixL]
if _, err := io.ReadFull(br, prefix); err != nil {
return nil, err
}
match := false
if bodyL >= need {
if int(prefix[1]) == codeL && bytes.Equal(prefix[2:need], code) {
if !readBody {
return nil, nil
}
match = true
}
return 0, err
}
if bodyL >= need && len(prefix) >= need && int(prefix[1]) == codeL && bytes.Equal(prefix[2:], code) {
return bodyL, nil
remain := bodyL - prefixL
if match {
out := make([]byte, bodyL)
copy(out, prefix)
if remain > 0 {
if _, err := io.ReadFull(br, out[prefixL:]); err != nil {
return nil, err
}
}
return out, nil
}
if _, err := br.Discard(bodyL); err != nil {
return 0, err
if remain > 0 {
if _, err := br.Discard(remain); err != nil {
return nil, err
}
}
}
}
// AttributeMatcher, HasAttrMatcher, AllAttrsMatcher and NewAllAttrsMatcher are the exported
// attribute helpers that have been part of this package's API since #5814. The streaming loader
// above filters attributes itself without building a *Domain, so it does not use them, but they
// are kept for external callers. Their behaviour is unchanged.
type AttributeMatcher interface {
Match(*Domain) bool
}
@@ -222,137 +185,23 @@ func NewAllAttrsMatcher(attrs string) AttributeMatcher {
return m
}
var errInvalidUTF8 = errors.New("string field contains invalid UTF-8")
func loadSiteWithAttrs(file, code, attrs string) ([]*Domain, error) {
domains, err := loadSite(file, code)
if err != nil {
return nil, err
}
type siteDecoder struct {
want []string
has []bool
fn func(Domain_Type, []byte)
}
matcher := NewAllAttrsMatcher(attrs)
if matcher == nil {
return domains, nil
}
func newSiteDecoder(attrs string, fn func(Domain_Type, []byte)) *siteDecoder {
d := &siteDecoder{fn: fn}
if attrs != "" {
d.want = strings.Split(attrs, "@")
d.has = make([]bool, len(d.want))
filtered := make([]*Domain, 0, len(domains))
for _, d := range domains {
if matcher.Match(d) {
filtered = append(filtered, d)
}
}
return d
}
// decode walks the whole fields at the start of b, a part of an encoded GeoSite (see geodat.proto),
// calls fn for every domain that has all attrs and returns how many bytes it used. A field cut off
// by the end of b is an error unless more is set. It accepts and rejects what proto.Unmarshal does.
func (d *siteDecoder) decode(b []byte, more bool) (int, error) {
used := 0
for used < len(b) {
f, n, err := consumeField(b[used:])
if err == io.ErrUnexpectedEOF && more {
break
}
if err != nil {
return used, err
}
used += n
if f.typ != protowire.BytesType {
continue
}
switch f.num {
case 1: // code
if !utf8.Valid(f.v) {
return used, errInvalidUTF8
}
case 2: // domain
t, value, err := decodeDomain(f.v, d.want, d.has)
if err != nil {
return used, err
}
if !slices.Contains(d.has, false) {
d.fn(t, value)
}
}
}
return used, nil
}
// decodeDomain decodes an encoded Domain and sets has[i] if one of its attributes has the key want[i].
func decodeDomain(b []byte, want []string, has []bool) (t Domain_Type, value []byte, err error) {
clear(has)
for len(b) > 0 {
f, n, err := consumeField(b)
if err != nil {
return 0, nil, err
}
b = b[n:]
switch {
case f.num == 1 && f.typ == protowire.VarintType: // type
t = Domain_Type(f.x)
case f.num == 2 && f.typ == protowire.BytesType: // value
if !utf8.Valid(f.v) {
return 0, nil, errInvalidUTF8
}
value = f.v
case f.num == 3 && f.typ == protowire.BytesType: // attribute
key, err := decodeAttributeKey(f.v)
if err != nil {
return 0, nil, err
}
for i, w := range want {
if string(key) == w {
has[i] = true
}
}
}
}
return t, value, nil
}
// decodeAttributeKey returns the key of an encoded Domain.Attribute.
func decodeAttributeKey(b []byte) ([]byte, error) {
var key []byte
for len(b) > 0 {
f, n, err := consumeField(b)
if err != nil {
return nil, err
}
b = b[n:]
if f.num == 1 && f.typ == protowire.BytesType {
if !utf8.Valid(f.v) {
return nil, errInvalidUTF8
}
key = f.v
}
}
return key, nil
}
type protoField struct {
num protowire.Number
typ protowire.Type
v []byte // payload of a length-delimited field
x uint64 // value of a varint field
}
// consumeField parses the first field of an encoded message and returns it with its length.
func consumeField(b []byte) (protoField, int, error) {
num, typ, n := protowire.ConsumeTag(b)
if n < 0 {
return protoField{}, 0, protowire.ParseError(n)
}
if num > protowire.MaxValidNumber {
return protoField{}, 0, errors.New("invalid field number ", num)
}
f := protoField{num: num, typ: typ}
var m int
switch typ {
case protowire.BytesType:
f.v, m = protowire.ConsumeBytes(b[n:])
case protowire.VarintType:
f.x, m = protowire.ConsumeVarint(b[n:])
default:
m = protowire.ConsumeFieldValue(num, typ, b[n:])
}
if m < 0 {
return protoField{}, 0, protowire.ParseError(m)
}
return f, n + m, nil
return filtered, nil
}
-283
View File
@@ -1,283 +0,0 @@
package geodata
import (
"fmt"
"os"
"path/filepath"
"slices"
"strings"
"testing"
"google.golang.org/protobuf/encoding/protowire"
"google.golang.org/protobuf/proto"
)
type siteEntry struct {
Type Domain_Type
Value string
}
// unmarshalSite is what loadSite used to do: proto.Unmarshal, then keep the domains that have all attrs.
func unmarshalSite(b []byte, attrs string) ([]siteEntry, error) {
var site GeoSite
if err := proto.Unmarshal(b, &site); err != nil {
return nil, err
}
var entries []siteEntry
for _, d := range site.Domain {
ok := true
for _, key := range strings.Split(attrs, "@") {
ok = ok && (attrs == "" || slices.ContainsFunc(d.Attribute, func(a *Domain_Attribute) bool { return a.Key == key }))
}
if ok {
entries = append(entries, siteEntry{d.Type, d.Value})
}
}
return entries, nil
}
func checkDecodeSite(t *testing.T, name string, b []byte, attrs string) {
t.Helper()
want, wantErr := unmarshalSite(b, attrs)
var got []siteEntry
_, err := newSiteDecoder(attrs, func(typ Domain_Type, value []byte) {
got = append(got, siteEntry{typ, string(value)})
}).decode(b, false)
if (err == nil) != (wantErr == nil) {
t.Fatalf("%s@%s: error %v, proto.Unmarshal: %v", name, attrs, err, wantErr)
}
if err == nil && !slices.Equal(got, want) {
t.Fatalf("%s@%s: got %v, want %v", name, attrs, got, want)
}
}
func TestDecodeSiteMatchesUnmarshal(t *testing.T) {
bs, err := os.ReadFile(filepath.Join("..", "..", "resources", DefaultGeoSiteDat))
if err != nil {
t.Fatal(err)
}
for len(bs) > 0 {
num, typ, n := protowire.ConsumeTag(bs)
if n < 0 || num != 1 || typ != protowire.BytesType {
t.Fatal("unexpected GeoSiteList field")
}
entry, m := protowire.ConsumeBytes(bs[n:])
if m < 0 {
t.Fatal(protowire.ParseError(m))
}
bs = bs[n+m:]
var site GeoSite
if err := proto.Unmarshal(entry, &site); err != nil {
t.Fatal(err)
}
queries := []string{"", "none"}
for _, d := range site.Domain {
for _, a := range d.Attribute {
if !slices.Contains(queries, a.Key) {
queries = append(queries, a.Key, a.Key+"@none")
}
}
}
for _, attrs := range queries {
checkDecodeSite(t, site.Code, entry, attrs)
}
}
}
func TestDecodeSiteUnusualEncodings(t *testing.T) {
field := func(num protowire.Number, v []byte) []byte {
return protowire.AppendBytes(protowire.AppendTag(nil, num, protowire.BytesType), v)
}
typ := func(v Domain_Type) []byte {
return protowire.AppendVarint(protowire.AppendTag(nil, 1, protowire.VarintType), uint64(v))
}
value := func(s string) []byte { return field(2, []byte(s)) }
attr := func(keys ...string) []byte {
var b []byte
for _, k := range keys {
b = append(b, field(1, []byte(k))...)
}
return field(3, b)
}
domain := func(fields ...[]byte) []byte { return field(2, slices.Concat(fields...)) }
unknown := protowire.AppendFixed32(protowire.AppendTag(nil, 9, protowire.Fixed32Type), 1)
for name, b := range map[string][]byte{
"unknown field": domain(typ(Domain_Full), unknown, value("example.com")),
"repeated value": domain(value("a.com"), typ(Domain_Full), value("b.com")),
"repeated type": domain(typ(Domain_Full), value("a.com"), typ(Domain_Regex)),
"repeated key": domain(value("a.com"), attr("cn", "ads")),
"type as bytes": domain(field(1, []byte("x")), value("a.com")),
"no value": domain(typ(Domain_Domain), attr("cn")),
"truncated": domain(typ(Domain_Full), value("example.com"))[:10],
"invalid utf8": domain(value("example.\xff")),
"invalid key": domain(value("a.com"), attr("\xff")),
"bad field": protowire.AppendVarint(protowire.AppendTag(nil, protowire.MaxValidNumber+1, protowire.VarintType), 1),
"stray end group": protowire.AppendTag(nil, 5, protowire.EndGroupType),
} {
for _, attrs := range []string{"", "cn", "ads", "cn@ads"} {
checkDecodeSite(t, name, b, attrs)
}
}
}
// TestLoadSiteReadsInPieces covers what real lists never do: an entry far longer than the read
// buffer, with a field longer than the buffer in the middle, and a file cut short.
func TestLoadSiteReadsInPieces(t *testing.T) {
site := &GeoSite{Code: "BIG"}
for i := range 5000 {
d := &Domain{Type: Domain_Domain, Value: strings.Repeat("x", i%40) + ".example.com"}
if i%3 == 0 {
d.Attribute = []*Domain_Attribute{{Key: "cn"}}
}
if i == 2500 {
d = &Domain{Type: Domain_Regex, Value: strings.Repeat("a", 100_000)}
}
site.Domain = append(site.Domain, d)
}
list := &GeoSiteList{Entry: []*GeoSite{{Code: "SMALL", Domain: []*Domain{{Type: Domain_Full, Value: "a.com"}}}, site}}
bs, err := proto.Marshal(list)
if err != nil {
t.Fatal(err)
}
entry, err := proto.Marshal(site)
if err != nil {
t.Fatal(err)
}
dir := t.TempDir()
t.Setenv("xray.location.asset", dir)
write := func(b []byte) {
if err := os.WriteFile(filepath.Join(dir, "big.dat"), b, 0o644); err != nil {
t.Fatal(err)
}
}
for _, attrs := range []string{"", "cn"} {
want, _ := unmarshalSite(entry, attrs)
var got []siteEntry
write(bs)
err := loadSite("big.dat", "BIG", attrs, func(typ Domain_Type, value []byte) {
got = append(got, siteEntry{typ, string(value)})
})
if err != nil || !slices.Equal(got, want) {
t.Fatalf("attrs %q: %d entries, want %d, error %v", attrs, len(got), len(want), err)
}
for _, cut := range []int{30_000, len(bs) - 150_000, len(bs) - 1} {
write(bs[:cut])
if err := loadSite("big.dat", "BIG", attrs, func(Domain_Type, []byte) {}); err == nil {
t.Fatalf("file cut at %d of %d: no error", cut, len(bs))
}
}
}
}
// oneEntryGeoSiteFile wraps an encoded GeoSite as a one-entry GeoSiteList, the file loadSite reads.
func oneEntryGeoSiteFile(entry []byte) []byte {
return protowire.AppendBytes(protowire.AppendTag(nil, 1, protowire.BytesType), entry)
}
// TestLoadSiteWindowedMatchesSingleShot checks that the windowed reader in loadSite (its Peek/Discard
// loop, the more-break when a field is cut by a window edge, the used==0 fallback for a field longer
// than the buffer, and the tail path) reaches exactly the same result as decoding the whole entry at
// once, for a category several 64 KiB windows long, valid and then mutated near a window edge and
// early in the file: same error-or-not, and the same emitted (type, value) sequence when both accept.
func TestLoadSiteWindowedMatchesSingleShot(t *testing.T) {
const window = 64 * 1024
site := &GeoSite{Code: "BIG"}
for i := range 12000 { // ~250 KiB, four windows
d := &Domain{Type: Domain_Domain, Value: fmt.Sprintf("host%d.%s.example.com", i, strings.Repeat("y", i%30))}
if i%3 == 0 {
d.Attribute = []*Domain_Attribute{{Key: "cn"}}
}
site.Domain = append(site.Domain, d)
}
// a field longer than the buffer, straddling the third window, to force the used==0 fallback
site.Domain = slices.Insert(site.Domain, 8000, &Domain{Type: Domain_Regex, Value: strings.Repeat("a", 90_000)})
entry, err := proto.Marshal(site)
if err != nil {
t.Fatal(err)
}
dir := t.TempDir()
t.Setenv("xray.location.asset", dir)
// mutations of the encoded entry: unchanged, a byte flipped at several offsets (early windows and
// either side of a window edge), and truncations at the same places.
type mut struct {
name string
make func([]byte) []byte
}
muts := []mut{{"valid", func(b []byte) []byte { return b }}}
for _, off := range []int{3, 40, 4000, window - 1, window, window + 1, 2*window - 2, 2 * window} {
if off < len(entry) {
off := off
muts = append(muts, mut{fmt.Sprintf("flip@%d", off), func(b []byte) []byte {
c := slices.Clone(b)
c[off] ^= 0xff
return c
}})
muts = append(muts, mut{fmt.Sprintf("cut@%d", off), func(b []byte) []byte { return slices.Clone(b[:off]) }})
}
}
for _, attrs := range []string{"", "cn"} {
for _, m := range muts {
e := m.make(entry)
// single-shot reference: decode the whole entry in one call
var want []siteEntry
_, wantErr := newSiteDecoder(attrs, func(typ Domain_Type, value []byte) {
want = append(want, siteEntry{typ, string(value)})
}).decode(e, false)
// windowed: loadSite reads the file 64 KiB at a time
if err := os.WriteFile(filepath.Join(dir, "w.dat"), oneEntryGeoSiteFile(e), 0o644); err != nil {
t.Fatal(err)
}
var got []siteEntry
gotErr := loadSite("w.dat", "BIG", attrs, func(typ Domain_Type, value []byte) {
got = append(got, siteEntry{typ, string(value)})
})
if (gotErr == nil) != (wantErr == nil) {
t.Fatalf("%s attrs=%q: windowed err %v, single-shot err %v", m.name, attrs, gotErr, wantErr)
}
if gotErr == nil && !slices.Equal(got, want) {
t.Fatalf("%s attrs=%q: windowed got %d entries, single-shot %d", m.name, attrs, len(got), len(want))
}
}
}
}
// TestLoadSiteLongCode covers a geosite entry whose code is longer than the 64 KiB read buffer. seek
// must skip it (find compares a single length byte, so it never matches such a code) and still find a
// later entry, and looking the long code up must fail cleanly, like a missing code, not panic.
func TestLoadSiteLongCode(t *testing.T) {
longCode := strings.Repeat("Z", 70000)
list := &GeoSiteList{Entry: []*GeoSite{
{Code: "FIRST", Domain: []*Domain{{Type: Domain_Full, Value: "first.com"}}},
{Code: longCode, Domain: []*Domain{{Type: Domain_Full, Value: "huge.com"}}},
{Code: "AFTER", Domain: []*Domain{{Type: Domain_Domain, Value: "after.com"}}},
}}
bs, err := proto.Marshal(list)
if err != nil {
t.Fatal(err)
}
dir := t.TempDir()
t.Setenv("xray.location.asset", dir)
if err := os.WriteFile(filepath.Join(dir, "lc.dat"), bs, 0o644); err != nil {
t.Fatal(err)
}
collect := func(code string) ([]siteEntry, error) {
var got []siteEntry
err := loadSite("lc.dat", code, "", func(typ Domain_Type, value []byte) {
got = append(got, siteEntry{typ, string(value)})
})
return got, err
}
if got, err := collect("FIRST"); err != nil || !slices.Equal(got, []siteEntry{{Domain_Full, "first.com"}}) {
t.Fatalf("FIRST: %v %v", got, err)
}
if got, err := collect("AFTER"); err != nil || !slices.Equal(got, []siteEntry{{Domain_Domain, "after.com"}}) {
t.Fatalf("AFTER (past the oversized entry): %v %v", got, err)
}
if _, err := collect(longCode); err == nil {
t.Fatal("oversized code: expected a not-found error, got nil")
}
}
-58
View File
@@ -1,58 +0,0 @@
package geodata
import (
lua "github.com/yuin/gopher-lua"
luar "layeh.com/gopher-luar"
)
// RegisterLua makes xray.geodata available to require in an LState.
func RegisterLua(L *lua.LState) {
L.PreloadModule("xray.geodata", func(L *lua.LState) int {
module := L.NewTable()
module.RawSetString("BuildDomainMatcher", L.NewFunction(func(L *lua.LState) int {
parsed, err := ParseDomainRules(luaRules(L), Domain_Domain)
if err != nil {
L.RaiseError("%v", err)
return 0
}
matcher, err := DomainReg.BuildDomainMatcher(parsed)
if err != nil {
L.RaiseError("%v", err)
return 0
}
L.Push(luar.New(L, matcher))
return 1
}))
module.RawSetString("BuildIPMatcher", L.NewFunction(func(L *lua.LState) int {
parsed, err := ParseIPRules(luaRules(L))
if err != nil {
L.RaiseError("%v", err)
return 0
}
matcher, err := IPReg.BuildIPMatcher(parsed)
if err != nil {
L.RaiseError("%v", err)
return 0
}
L.Push(luar.New(L, matcher))
return 1
}))
L.Push(module)
return 1
})
}
func luaRules(L *lua.LState) []string {
rules := make([]string, L.GetTop())
for i := range rules {
value, ok := L.Get(i + 1).(lua.LString)
if !ok {
L.RaiseError("geodata rules must be strings")
return nil
}
rules[i] = string(value)
}
return rules
}
-66
View File
@@ -1,66 +0,0 @@
package geodata
import (
"testing"
"github.com/xtls/xray-core/common/net"
lua "github.com/yuin/gopher-lua"
)
func TestLuaIPMatcher(t *testing.T) {
L := lua.NewState()
defer L.Close()
RegisterLua(L)
ip := L.NewUserData()
ip.Value = net.ParseIP("127.0.0.1")
L.SetGlobal("ip", ip)
ips := L.NewUserData()
ips.Value = []net.IP{ip.Value.(net.IP), net.ParseIP("8.8.8.8")}
L.SetGlobal("ips", ips)
if err := L.DoString(`
local matcher = require("xray.geodata").BuildIPMatcher("127.0.0.0/8", "::1")
assert(matcher:Match(ip))
assert(matcher:AnyMatch(ips))
assert(not matcher:Matches(ips))
local matched, unmatched = matcher:FilterIPs(ips)
assert(type(matched) == "userdata" and type(unmatched) == "userdata")
assert(#matched == 1 and #unmatched == 1)
`); err != nil {
t.Fatal(err)
}
}
func TestLuaDomainMatcher(t *testing.T) {
L := lua.NewState()
defer L.Close()
RegisterLua(L)
if err := L.DoString(`
local matcher = require("xray.geodata").BuildDomainMatcher("example.com", "full:other.com")
assert(matcher:MatchAny("example.com"))
assert(matcher:MatchAny("www.example.com"))
assert(matcher:MatchAny("other.com"))
assert(not matcher:MatchAny("www.other.com"))
assert(#(matcher:Match("www.example.com")) == 1)
`); err != nil {
t.Fatal(err)
}
}
func TestLuaMatchersRejectInvalidRules(t *testing.T) {
for _, tc := range []struct {
name string
script string
}{
{"IP rule", `require("xray.geodata").BuildIPMatcher("not-an-ip")`},
{"non-string domain rule", `require("xray.geodata").BuildDomainMatcher("example.com", true)`},
} {
t.Run(tc.name, func(t *testing.T) {
L := lua.NewState()
defer L.Close()
RegisterLua(L)
if err := L.DoString(tc.script); err == nil {
t.Fatal("invalid geodata rule was accepted")
}
})
}
}
@@ -1,7 +1,6 @@
package strmatcher_test
import (
"regexp"
"strconv"
"testing"
@@ -73,64 +72,6 @@ func BenchmarkSubstrMatcher(b *testing.B) {
})
}
func BenchmarkRegexMatcher(b *testing.B) {
patterns := []string{ // taken from geosite
`(^|\.)91porn\.(best|com|cool|fun|group|party|plus|site|tw|work)$`,
`(^|\.)91porn[0-9]{3}\.me$`,
`(^|\.)apiproxy-device-prod-nlb-.+\.amazonaws\.com$`,
`(^|\.)dualstack\.apiproxy-.+\.amazonaws\.com$`,
`(^|\.)aqdk[0-9]{3}\.com$`,
`(^|\.)bilibili3(0[1-9]|1[0-2])\.xyz$`,
`(^|\.)byyum([3589]|2[235689]|3[34]|4[1-9]|5[1-79]|6[0134679])?\.com$`,
`(^|\.)fiftymvapi\..+$`,
`(^|\.)gossipfuli[0-9]{3,4}\.xyz$`,
`(^|\.)kpkuang\.(bond|fun|info|one|us)$`,
`(^|\.)rule34\.(asia|us|world|xxx|xyz)$`,
`(^|\.)[a-z][1-9][0-9][a-z]\.com$`,
`.+\.awsdns-[0-9][0-9]\.(co\.uk|com|net|org)$`,
`.+\.dkr\.ecr\.[^\.]+\.amazonaws\.com$`,
`^(.+\.)*zh\.okaapps\.com$`,
`^cdn\d-epicgames-\d+\.file\.myqcloud\.com$`,
`^chatgpt-async-webps-prod-\S+-\d+\.webpubsub\.azure\.com$`,
`^r+[0-9]+(---|\.)sn-(2x3|ni5|j5o)\w{5}\.googlevideo\.com$`,
`^speed\.(coe|open)\.ad\.[a-z]{2,6}\.prod\.hosts\.ooklaserver\.net$`,
`javdb\d+\.com$`,
}
domains := []string{
"www.google.com", "rr3---sn-4g5edndy.googlevideo.com", "r1---sn-2x3abcde.googlevideo.com", "i.ytimg.com",
"graph.facebook.com", "api.twitter.com", "www.baidu.com", "github.com", "objects.githubusercontent.com",
"login.microsoftonline.com", "e1234.dscb.akamaiedge.net", "d1a2b3c4d5e6f7.cloudfront.net",
"s3.us-east-1.amazonaws.com", "123456789012.dkr.ecr.us-east-1.amazonaws.com", "www.wikipedia.org",
"discord.com", "telegram.org", "store.steampowered.com", "www.91porn.com", "ns-1234.awsdns-12.org",
}
bench := func(b *testing.B, ctor func(pattern string) func(string) bool) {
var matchers []func(string) bool
for _, p := range patterns {
matchers = append(matchers, ctor(p))
}
b.ResetTimer()
for i := 0; i < b.N; i++ {
for _, d := range domains {
for _, match := range matchers {
_ = match(d)
}
}
}
}
b.Run("regexp", func(b *testing.B) {
bench(b, func(pattern string) func(string) bool {
return regexp.MustCompile(pattern).MatchString
})
})
b.Run("prefilter", func(b *testing.B) {
bench(b, func(pattern string) func(string) bool {
m, err := Regex.New(pattern)
common.Must(err)
return m.Match
})
})
}
// Utility functions for benchmark
func benchmarkMatcherType(b *testing.B, t Type, ctor func() MatcherGroup) {
+12 -8
View File
@@ -52,9 +52,7 @@ func (g *MphIndexMatcher) Add(matcher Matcher) uint32 {
func (g *MphIndexMatcher) Build() error {
if g.mph != nil {
runtime.GC() // peak mem
if err := g.mph.Build(); err != nil {
return err
}
g.mph.Build()
}
runtime.GC() // peak mem
if g.ac != nil {
@@ -66,17 +64,23 @@ func (g *MphIndexMatcher) Build() error {
// Match implements IndexMatcher.Match.
func (g *MphIndexMatcher) Match(input string) []uint32 {
var result []uint32
result := make([][]uint32, 0, 5)
if g.mph != nil {
result = g.mph.Match(input) // a new slice, returned without another copy
if matches := g.mph.Match(input); len(matches) > 0 {
result = append(result, matches)
}
}
if g.ac != nil {
result = append(result, g.ac.Match(input)...)
if matches := g.ac.Match(input); len(matches) > 0 {
result = append(result, matches)
}
}
if g.regex != nil {
result = append(result, g.regex.Match(input)...)
if matches := g.regex.Match(input); len(matches) > 0 {
result = append(result, matches)
}
}
return result
return CompositeMatches(result)
}
// MatchAny implements IndexMatcher.MatchAny.
@@ -78,10 +78,6 @@ func TestMphIndexMatcher(t *testing.T) {
Input: "example.com",
Output: []uint32{10, 4},
},
{
Input: "apis.org",
Output: []uint32{2, 6},
},
}
matcherGroup := NewMphIndexMatcher()
for _, rule := range rules {
@@ -91,13 +87,8 @@ func TestMphIndexMatcher(t *testing.T) {
}
matcherGroup.Build()
for _, test := range cases {
m := matcherGroup.Match(test.Input)
if !reflect.DeepEqual(m, test.Output) {
t.Error("unexpected output: ", m, " for test case ", test)
}
clear(m) // the caller owns the result, so this must not change the next one
if m := matcherGroup.Match(test.Input); !reflect.DeepEqual(m, test.Output) {
t.Error("unexpected output after clearing the previous one: ", m, " for test case ", test)
t.Error("unexpected output: ", m, " for test case ", test)
}
}
}
+136 -378
View File
@@ -1,440 +1,198 @@
package strmatcher
import (
"bytes"
"cmp"
"encoding/binary"
"errors"
"math"
"slices"
"math/bits"
"runtime"
"sort"
"strings"
"unsafe"
)
// Flags of a level1 slot, stored above the record offset.
const (
mphDomain = 1 << 31 // matches the pattern and its subdomains
mphFull = 1 << 30 // matches the pattern only
mphParent = 1 << 29 // matches subdomains only, from a pattern with a leading dot
mphOffMask = mphParent - 1
)
// PrimeRK is the prime base used in Rabin-Karp algorithm.
const PrimeRK = 16777619
// Kinds of an added pattern, indexes of mphKinds.
const (
mphKindFull = iota
mphKindParent
mphKindDomain
)
// mphKinds are the slot flags in the order Match reports their values.
var mphKinds = [...]uint32{mphFull, mphParent, mphDomain}
// mphMultipliers are odd multipliers for the suffix hash. Build moves to the next one if two patterns collide.
var mphMultipliers = [...]uint64{0x9e3779b97f4a7c15, 0xc2b2ae3d27d4eb4f, 0x165667b19e3779f9, 0x27d4eb2f165667c5}
var (
errMphCollision = errors.New("strmatcher: suffix hash collision in MphMatcherGroup")
errMphBuilt = errors.New("strmatcher: MphMatcherGroup is already built")
)
type mphEntry struct {
off uint32 // pattern start in buf
value uint32
n uint32 // pattern length
kind uint8
// RollingHash calculates the rolling murmurHash of given string based on a provided suffix hash.
func RollingHash(hash uint32, input string) uint32 {
for i := len(input) - 1; i >= 0; i-- {
hash = hash*PrimeRK + uint32(input[i])
}
return hash
}
// MphMatcherGroup is an implementation of MatcherGroup for Full and Domain matchers.
// Each distinct pattern is stored once as a record in arena: its length (255 means a uvarint length follows),
// its bytes and, if the group holds more than one distinct value, its values. A minimal perfect hash table
// built with hash, displace and compress (http://cmph.sourceforge.net/papers/esa09.pdf) maps a pattern to its
// record. Patterns are hashed from the right, so one pass over the input hashes all its parent domains.
type MphMatcherGroup struct {
arena string
level0 []uint16 // bucket -> seed
level1 []uint32 // slot -> flags | record offset
fp []uint8 // slot -> low byte of its pattern's hash, rejects most misses without reading arena
n0, n1 uint32
mul uint64 // multiplier of the suffix hash
single uint32 // the only value if !multi
multi bool
// MemHash is the hash function used by go map, it utilizes available hardware instructions(behaves
// as aeshash if aes instruction is available).
// With different seed, each MemHash<seed> performs as distinct hash functions.
func MemHash(seed uint32, input string) uint32 {
return uint32(strhash(unsafe.Pointer(&input), uintptr(seed))) // nosemgrep
}
buf []byte // build only, patterns in Add order
entries []mphEntry
const (
mphMatchTypeCount = 2 // Full and Domain
)
type mphRuleInfo struct {
rollingHash uint32
matchers [mphMatchTypeCount][]uint32
}
// MphMatcherGroup is an implementation of MatcherGroup.
// It implements Rabin-Karp algorithm and minimal perfect hash table for Full and Domain matcher.
type MphMatcherGroup struct {
rules []string // RuleIdx -> pattern string, index 0 reserved for failed lookup
values [][]uint32 // RuleIdx -> registered matcher values for the pattern (Full Matcher takes precedence)
level0 []uint32 // RollingHash & Mask -> seed for Memhash
level0Mask uint32 // Mask restricting RollingHash to 0 ~ len(level0)
level1 []uint32 // Memhash<seed> & Mask -> stored index for rules
level1Mask uint32 // Mask for restricting Memhash<seed> to 0 ~ len(level1)
ruleInfos *map[string]mphRuleInfo
}
func NewMphMatcherGroup() *MphMatcherGroup {
return new(MphMatcherGroup)
return &MphMatcherGroup{
rules: []string{""},
values: [][]uint32{nil},
level0: nil,
level0Mask: 0,
level1: nil,
level1Mask: 0,
ruleInfos: &map[string]mphRuleInfo{}, // Only used for building, destroyed after build complete
}
}
// AddFullMatcher implements MatcherGroupForFull.
func (g *MphMatcherGroup) AddFullMatcher(matcher FullMatcher, value uint32) {
g.add(matcher.Pattern(), mphKindFull, value)
pattern := strings.ToLower(matcher.Pattern())
g.addPattern(0, "", pattern, matcher.Type(), value)
}
// AddDomainMatcher implements MatcherGroupForDomain.
func (g *MphMatcherGroup) AddDomainMatcher(matcher DomainMatcher, value uint32) {
g.add(matcher.Pattern(), mphKindDomain, value)
pattern := strings.ToLower(matcher.Pattern())
hash := g.addPattern(0, "", pattern, matcher.Type(), value) // For full domain match
g.addPattern(hash, pattern, ".", matcher.Type(), value) // For partial domain match
}
func (g *MphMatcherGroup) add(pattern string, kind uint8, value uint32) {
if g.arena != "" {
panic(errMphBuilt)
}
pattern = strings.ToLower(pattern)
off := uint32(len(g.buf))
g.buf = append(g.buf, pattern...)
g.entries = append(g.entries, mphEntry{off: off, value: value, n: uint32(len(pattern)), kind: kind})
if len(pattern) > 0 && pattern[0] == '.' {
// ".x" has always matched "*.x" as well, so it also gets a parent-only record for "x"
g.entries = append(g.entries, mphEntry{off: off + 1, value: value, n: uint32(len(pattern) - 1), kind: mphKindParent})
func (g *MphMatcherGroup) addPattern(suffixHash uint32, suffixPattern string, pattern string, matcherType Type, value uint32) uint32 {
fullPattern := pattern + suffixPattern
info, found := (*g.ruleInfos)[fullPattern]
if !found {
info = mphRuleInfo{rollingHash: RollingHash(suffixHash, pattern)}
g.rules = append(g.rules, fullPattern)
g.values = append(g.values, nil)
}
info.matchers[matcherType] = append(info.matchers[matcherType], value)
(*g.ruleInfos)[fullPattern] = info
return info.rollingHash
}
func (g *MphMatcherGroup) key(i uint32) []byte {
e := &g.entries[i]
return g.buf[e.off : e.off+e.n]
}
// Build builds the hash table. It must be called once, after the last Add.
// Build builds a minimal perfect hash table for insert rules.
// Algorithm used: Hash, displace, and compress. See http://cmph.sourceforge.net/papers/esa09.pdf
func (g *MphMatcherGroup) Build() error {
if g.arena != "" {
return errMphBuilt
}
if uint64(len(g.buf)) > math.MaxUint32 {
return errors.New("too many rules for MphMatcherGroup")
}
recs := g.writeRecords()
if len(g.arena) > mphOffMask {
return errors.New("too many rules for MphMatcherGroup")
}
hashes := make([]uint64, len(recs))
for _, mul := range mphMultipliers {
for i, rec := range recs {
hashes[i] = mphMix(mphHash(mul, g.recKey(rec)))
}
g.mul = mul
if err := g.place(recs, hashes); err != errMphCollision {
return err
}
}
return errMphCollision
}
ruleCount := len(*g.ruleInfos)
g.level0 = make([]uint32, nextPow2(ruleCount/4))
g.level0Mask = uint32(len(g.level0) - 1)
g.level1 = make([]uint32, nextPow2(ruleCount))
g.level1Mask = uint32(len(g.level1) - 1)
// writeRecords writes one record per distinct pattern to arena and returns flags | offset of each.
func (g *MphMatcherGroup) writeRecords() []uint32 {
g.multi = false
if len(g.entries) > 0 {
g.single = g.entries[0].value
for _, e := range g.entries {
if e.value != g.single {
g.multi = true
break
}
}
// Create buckets based on all rule's rolling hash
buckets := make([][]uint32, len(g.level0))
for ruleIdx := 1; ruleIdx < len(g.rules); ruleIdx++ { // Traverse rules starting from index 1 (0 reserved for failed lookup)
ruleInfo := (*g.ruleInfos)[g.rules[ruleIdx]]
bucketIdx := ruleInfo.rollingHash & g.level0Mask
buckets[bucketIdx] = append(buckets[bucketIdx], uint32(ruleIdx))
g.values[ruleIdx] = append(ruleInfo.matchers[Full], ruleInfo.matchers[Domain]...) // nolint:gocritic
}
// Equal patterns become neighbours in Add order, so their values keep their priority
order := make([]uint32, len(g.entries))
for i := range order {
order[i] = uint32(i)
}
slices.SortFunc(order, func(a, b uint32) int {
return cmp.Or(bytes.Compare(g.key(a), g.key(b)), cmp.Compare(a, b))
})
g.ruleInfos = nil // Set ruleInfos nil to release memory
runtime.GC() // peak mem
size := len(g.buf) + len(g.entries) + 2
if g.multi {
size += 3 * len(g.entries)
// Sort buckets in descending order with respect to each bucket's size
bucketIdxs := make([]int, len(buckets))
for bucketIdx := range buckets {
bucketIdxs[bucketIdx] = bucketIdx
}
arena := make([]byte, 0, size)
recs := make([]uint32, 0, len(order))
var vals [len(mphKinds)][]uint32
for i := 0; i < len(order); {
k := g.key(order[i])
for t := range vals {
vals[t] = vals[t][:0]
}
for ; i < len(order) && bytes.Equal(g.key(order[i]), k); i++ {
e := &g.entries[order[i]]
if !slices.Contains(vals[e.kind], e.value) {
vals[e.kind] = append(vals[e.kind], e.value)
}
}
rec := uint32(len(arena))
if len(k) < 255 {
arena = append(arena, byte(len(k)))
} else {
arena = binary.AppendUvarint(append(arena, 255), uint64(len(k)))
}
arena = append(arena, k...)
for t, v := range vals {
if len(v) == 0 {
continue
}
rec |= mphKinds[t]
if g.multi {
arena = binary.AppendUvarint(arena, uint64(len(v)))
for _, x := range v {
arena = binary.AppendUvarint(arena, uint64(x))
sort.Slice(bucketIdxs, func(i, j int) bool { return len(buckets[bucketIdxs[i]]) > len(buckets[bucketIdxs[j]]) })
// Exercise Hash, Displace, and Compress algorithm to construct minimal perfect hash table
occupied := make([]bool, len(g.level1)) // Whether a second-level hash has been already used
hashedBucket := make([]uint32, 0, 4) // Second-level hashes for each rule in a specific bucket
for _, bucketIdx := range bucketIdxs {
bucket := buckets[bucketIdx]
hashedBucket = hashedBucket[:0]
seed := uint32(0)
for len(hashedBucket) != len(bucket) {
for _, ruleIdx := range bucket {
memHash := MemHash(seed, g.rules[ruleIdx]) & g.level1Mask
if occupied[memHash] { // Collision occurred with this seed
for _, hash := range hashedBucket { // Revert all values in this hashed bucket
occupied[hash] = false
g.level1[hash] = 0
}
hashedBucket = hashedBucket[:0]
seed++ // Try next seed
break
}
occupied[memHash] = true
g.level1[memHash] = ruleIdx // The final value in the hash table
hashedBucket = append(hashedBucket, memHash)
}
}
recs = append(recs, rec)
}
// Lookups may point one byte past a pattern, and an empty group needs a record at offset 0 for empty slots
arena = append(arena, 0)
if len(recs) == 0 {
arena = append(arena, 0)
}
g.buf, g.entries = nil, nil
if cap(arena)-len(arena) > len(arena)/32 {
arena = slices.Clone(arena)
}
g.arena = unsafe.String(unsafe.SliceData(arena), len(arena)) // arena is not written after this
return recs
}
// place fills level0, level1 and fp: records are bucketed by hash, and each bucket, largest first, gets
// the first seed that puts all its records in free slots.
func (g *MphMatcherGroup) place(recs []uint32, hashes []uint64) error {
r := len(recs)
n0, n1 := max(1, r/3), max(1, r+r/99)
g.n0, g.n1 = uint32(n0), uint32(n1)
g.level0 = make([]uint16, n0)
g.level1 = make([]uint32, n1)
g.fp = make([]uint8, n1)
start := make([]uint32, n0+1)
for _, h := range hashes {
start[g.bucket(h)+1]++
}
for b := range n0 {
start[b+1] += start[b]
}
members := make([]uint32, r)
fill := slices.Clone(start[:n0])
for i, h := range hashes {
b := g.bucket(h)
members[fill[b]] = uint32(i)
fill[b]++
}
fill = nil
buckets := make([]uint32, n0)
for b := range buckets {
buckets[b] = uint32(b)
}
slices.SortStableFunc(buckets, func(a, b uint32) int {
return cmp.Compare(start[b+1]-start[b], start[a+1]-start[a])
})
occupied := make([]uint64, (n1+63)/64)
var slots []uint32
next:
for _, b := range buckets {
m := members[start[b]:start[b+1]]
if len(m) == 0 {
break
}
for i := range m {
for j := range i {
if hashes[m[i]] == hashes[m[j]] {
return errMphCollision // no seed can separate them
}
}
}
search:
for seed := range math.MaxUint16 + 1 {
slots = slots[:0]
for _, ri := range m {
s := g.slot(hashes[ri], uint16(seed))
if occupied[s/64]&(1<<(s%64)) != 0 || slices.Contains(slots, s) {
continue search
}
slots = append(slots, s)
}
for k, ri := range m {
s := slots[k]
occupied[s/64] |= 1 << (s % 64)
g.level1[s] = recs[ri]
g.fp[s] = uint8(hashes[ri])
}
g.level0[b] = uint16(seed)
continue next
}
return errors.New("strmatcher: no seed found for a bucket in MphMatcherGroup")
g.level0[bucketIdx] = seed // Displacement value for this bucket
}
return nil
}
// mphHash is the suffix hash of s, taken from the right: the hash of s[i:] is the state after reading s[i].
func mphHash(mul uint64, s string) uint64 {
h := uint64(0)
for i := len(s) - 1; i >= 0; i-- {
h = h*mul + uint64(s[i])
}
return h
}
// mphMix spreads the weak low bits of a suffix hash.
func mphMix(h uint64) uint64 {
h ^= h >> 32
h *= 0xd6e8feb86659fd93
return h ^ h>>32
}
func (g *MphMatcherGroup) bucket(f uint64) uint32 {
return uint32(((f >> 32) * uint64(g.n0)) >> 32)
}
func (g *MphMatcherGroup) slot(f uint64, seed uint16) uint32 {
x := ((f ^ uint64(seed)*0x9e3779b97f4a7c15) * 0xc4ceb9fe1a85ec53) >> 32
return uint32((x * uint64(g.n1)) >> 32)
}
func (g *MphMatcherGroup) uvarint(p uint32) (x, next uint32) {
for shift := 0; ; shift += 7 {
c := g.arena[p]
p++
x |= uint32(c&0x7f) << shift
if c < 0x80 {
return x, p
}
}
}
// recSpan returns where the pattern of the record at off starts and how long it is.
func (g *MphMatcherGroup) recSpan(off uint32) (p, n uint32) {
n, p = uint32(g.arena[off]), off+1
if n == 255 {
n, p = g.uvarint(p)
}
return p, n
}
func (g *MphMatcherGroup) recKey(rec uint32) string {
p, n := g.recSpan(rec & mphOffMask)
return g.arena[p : p+n]
}
// lookup returns the level1 entry of s, or 0 if s is not a pattern. h is the suffix hash of s.
func (g *MphMatcherGroup) lookup(h uint64, s string) uint32 {
f := mphMix(h)
// bucket < n0 == len(level0) and slot < n1 == len(level1) == len(fp), skip the bounds checks
seed := *(*uint16)(unsafe.Add(unsafe.Pointer(unsafe.SliceData(g.level0)), uintptr(g.bucket(f))*2))
slot := uintptr(g.slot(f, seed))
if *(*uint8)(unsafe.Add(unsafe.Pointer(unsafe.SliceData(g.fp)), slot)) != uint8(f) {
return 0
}
e := *(*uint32)(unsafe.Add(unsafe.Pointer(unsafe.SliceData(g.level1)), slot*4))
if len(s) < 255 {
// A record whose length byte is len(s) has len(s) pattern bytes after it
p := unsafe.Add(unsafe.Pointer(unsafe.StringData(g.arena)), e&mphOffMask)
if int(*(*byte)(p)) == len(s) && unsafe.String((*byte)(unsafe.Add(p, 1)), len(s)) == s {
return e
}
return 0
}
if g.recKey(e) == s {
return e
// Lookup searches for input in minimal perfect hash table and returns its index. 0 indicates not found.
func (g *MphMatcherGroup) Lookup(rollingHash uint32, input string) uint32 {
i0 := rollingHash & g.level0Mask
seed := g.level0[i0]
i1 := MemHash(seed, input) & g.level1Mask
if n := g.level1[i1]; g.rules[n] == input {
return n
}
return 0
}
// appendValues appends the values of record e for the flags in want, in mphKinds order.
func (g *MphMatcherGroup) appendValues(dst []uint32, e, want uint32) []uint32 {
if !g.multi {
for _, flag := range mphKinds {
if e&want&flag != 0 {
dst = append(dst, g.single)
}
}
return dst
}
if e&want == 0 {
return dst
}
p, n := g.recSpan(e & mphOffMask)
p += n
for _, flag := range mphKinds {
if e&flag == 0 {
continue
}
var count, v uint32
for count, p = g.uvarint(p); count > 0; count-- {
v, p = g.uvarint(p)
if want&flag != 0 {
dst = append(dst, v)
}
}
}
return dst
}
// Match implements MatcherGroup.Match. Values of an exact match come first (Full, then Domain), then those of
// the parent domains, nearest first.
// Match implements MatcherGroup.Match.
func (g *MphMatcherGroup) Match(input string) []uint32 {
var stack [8]uint32
parents := stack[:0] // TLD side first
h, mul := uint64(0), g.mul
matches := make([][]uint32, 0, 5)
hash := uint32(0)
for i := len(input) - 1; i >= 0; i-- {
hash = hash*PrimeRK + uint32(input[i])
if input[i] == '.' {
if e := g.lookup(h, input[i+1:]); e&(mphDomain|mphParent) != 0 {
parents = append(parents, e)
if mphIdx := g.Lookup(hash, input[i:]); mphIdx != 0 {
matches = append(matches, g.values[mphIdx])
}
}
h = h*mul + uint64(input[i])
}
exact := g.lookup(h, input)
if exact&(mphFull|mphDomain) == 0 && len(parents) == 0 {
return nil
if mphIdx := g.Lookup(hash, input); mphIdx != 0 {
matches = append(matches, g.values[mphIdx])
}
result := g.appendValues(make([]uint32, 0, len(parents)+1), exact, mphFull|mphDomain)
for k := len(parents) - 1; k >= 0; k-- {
result = g.appendValues(result, parents[k], mphParent|mphDomain)
}
return result
return CompositeMatchesReverse(matches)
}
// MatchAny implements MatcherGroup.MatchAny.
func (g *MphMatcherGroup) MatchAny(input string) bool {
h, mul := uint64(0), g.mul
for i := len(input) - 1; i >= 0; i-- {
if input[i] == '.' && g.lookup(h, input[i+1:])&(mphDomain|mphParent) != 0 {
return true
}
h = h*mul + uint64(input[i])
}
return g.lookup(h, input)&(mphFull|mphDomain) != 0
}
// mphSuffix is the suffix hash of input[off:], a parent domain of the input.
type mphSuffix struct {
h uint64
off int
}
// mphSuffixes appends the suffix hashes of the parent domains of input to dst, TLD side first, and returns them
// with the hash of input itself: what MatchAny computes, computed once for several groups.
func mphSuffixes(dst []mphSuffix, mul uint64, input string) ([]mphSuffix, uint64) {
h := uint64(0)
hash := uint32(0)
for i := len(input) - 1; i >= 0; i-- {
hash = hash*PrimeRK + uint32(input[i])
if input[i] == '.' {
dst = append(dst, mphSuffix{h, i + 1})
if g.Lookup(hash, input[i:]) != 0 {
return true
}
}
h = h*mul + uint64(input[i])
}
return dst, h
return g.Lookup(hash, input) != 0
}
// matchAnyHashed is MatchAny with parents and h from mphSuffixes(_, mul, input).
func (g *MphMatcherGroup) matchAnyHashed(input string, parents []mphSuffix, h, mul uint64) bool {
if g.mul != mul {
return g.MatchAny(input) // built with a later multiplier after a collision
func nextPow2(v int) int {
if v <= 1 {
return 1
}
for _, p := range parents {
if g.lookup(p.h, input[p.off:])&(mphDomain|mphParent) != 0 {
return true
}
}
return g.lookup(h, input)&(mphFull|mphDomain) != 0
const MaxUInt = ^uint(0)
n := (MaxUInt >> bits.LeadingZeros(uint(v))) + 1
return int(n)
}
//go:noescape
//go:linkname strhash runtime.strhash
func strhash(p unsafe.Pointer, h uintptr) uintptr
@@ -1,108 +0,0 @@
package strmatcher
import (
"slices"
"testing"
)
func TestMphMatcherGroupHashCollision(t *testing.T) {
saved := mphMultipliers
defer func() { mphMultipliers = saved }()
mphMultipliers[0] = 1 // anagrams collide
g := NewMphMatcherGroup()
g.AddFullMatcher(FullMatcher("ab.com"), 1)
g.AddDomainMatcher(DomainMatcher("ba.com"), 2)
g.AddDomainMatcher(DomainMatcher("com"), 3)
if err := g.Build(); err != nil {
t.Fatal(err)
}
if g.mul != saved[1] {
t.Errorf("multiplier %#x, want the second one %#x", g.mul, saved[1])
}
for input, want := range map[string][]uint32{"ab.com": {1, 3}, "x.ba.com": {2, 3}, "x.ab.com": {3}, "ba.com": {2, 3}} {
if m := g.Match(input); !slices.Equal(m, want) {
t.Errorf("Match(%q) = %v, want %v", input, m, want)
}
}
// Thue-Morse strings of 2048 bytes and their complements collide for every odd multiplier
mphMultipliers = saved
a, b := make([]byte, 2048), make([]byte, 2048)
for i := range a {
a[i], b[i] = "ab"[bitsOnes(i)%2], "ba"[bitsOnes(i)%2]
}
g = NewMphMatcherGroup()
g.AddFullMatcher(FullMatcher(a), 1)
g.AddFullMatcher(FullMatcher(b), 1)
if err := g.Build(); err != errMphCollision {
t.Errorf("Build() = %v, want %v", err, errMphCollision)
}
}
func bitsOnes(i int) int {
n := 0
for ; i > 0; i &= i - 1 {
n++
}
return n
}
func TestMphValueMatcherCombiner(t *testing.T) {
build := func(matchers ...Matcher) *MphValueMatcher {
m := NewMphValueMatcher()
for _, x := range matchers {
m.Add(x, 0)
}
if err := m.Build(); err != nil {
t.Fatal(err)
}
return m
}
regex, err := Regex.New(`^a\d+\.net$`)
if err != nil {
t.Fatal(err)
}
saved := mphMultipliers
t.Cleanup(func() { mphMultipliers = saved })
mphMultipliers[0] = 1 // anagrams collide, so this one falls back to its own hash pass
collided := build(FullMatcher("ab.com"), DomainMatcher("ba.com"))
mphMultipliers = saved
if collided.mph.mul == mphMultipliers[0] {
t.Fatal("collided matcher uses the first multiplier")
}
matchers := []*MphValueMatcher{
build(DomainMatcher("example.com"), FullMatcher("full.org"), DomainMatcher(".dot.io")),
collided,
build(regex, SubstrMatcher("keyword")),
build(),
build(DomainMatcher("com"), DomainMatcher("a.b.c.d.e.f.g.h.i.j.k.l.m.n.o.p.q.r.s")),
}
var s MphValueMatcherCombiner
for i, m := range matchers {
s.Add(m, uint32(10+i))
}
inputs := []string{
"", ".", "..", "com", "example.com", "www.example.com", "xexample.com", "example.com.", "full.org", "x.full.org",
"dot.io", "x.dot.io", ".dot.io", "ab.com", "x.ab.com", "ba.com", "x.ba.com", "a12.net", "a12.net.x", "my-keyword.org",
"a.b.c.d.e.f.g.h.i.j.k.l.m.n.o.p.q.r.s", "0.a.b.c.d.e.f.g.h.i.j.k.l.m.n.o.p.q.r.s", "b.c.d.e.f.g.h.i.j.k.l.m.n.o.p.q.r.s",
"x.y.z.1.2.3.4.5.6.7.8.9.10.11.12.13.14.15.16.17.ab.com", "x.y.z.1.2.3.4.5.6.7.8.9.10.11.12.13.14.15.16.17.org",
}
for _, input := range inputs {
var want []uint32
for i, m := range matchers {
if m.MatchAny(input) {
want = append(want, uint32(10+i))
}
}
if got := s.Match(input); !slices.Equal(got, want) {
t.Errorf("Match(%q) = %v, want %v", input, got, want)
}
if got := s.MatchAny(input); got != (len(want) > 0) {
t.Errorf("MatchAny(%q) = %v", input, got)
}
}
if n := testing.AllocsPerRun(100, func() { s.MatchAny("www.a.b.c.example.org") }); n != 0 {
t.Errorf("MatchAny allocates %v times", n)
}
}
@@ -1,10 +1,7 @@
package strmatcher_test
import (
"math/rand"
"reflect"
"slices"
"strings"
"testing"
"github.com/xtls/xray-core/common"
@@ -279,142 +276,3 @@ func TestEmptyMphMatcherGroup(t *testing.T) {
t.Error("Expect [], but ", r)
}
}
func TestMphMatcherGroupRandom(t *testing.T) {
inputs := []string{""} // All strings over "ab." up to 7 bytes
for i := 0; len(inputs[i]) < 7; i++ {
for _, c := range []string{"a", "b", "."} {
inputs = append(inputs, inputs[i]+c)
}
}
for seed := int64(0); seed < 300; seed++ {
r := rand.New(rand.NewSource(seed))
g := NewMphMatcherGroup()
full, domain := map[string][]uint32{}, map[string][]uint32{} // Stored pattern -> values
for value := uint32(r.Intn(200)); value > 0; value-- {
pattern := make([]byte, r.Intn(8))
for i := range pattern {
pattern[i] = "ab."[r.Intn(3)]
}
if p := string(pattern); r.Intn(2) == 0 {
g.AddFullMatcher(FullMatcher(p), value)
full[p] = append(full[p], value)
} else {
g.AddDomainMatcher(DomainMatcher(p), value)
domain[p] = append(domain[p], value)
domain["."+p] = append(domain["."+p], value)
}
}
common.Must(g.Build())
for _, input := range inputs {
keys := []string{input} // Whole input first, then "." suffixes from longest to shortest
for i := range len(input) {
if input[i] == '.' {
keys = append(keys, input[i:])
}
}
var want []uint32
for _, k := range keys {
want = append(append(want, full[k]...), domain[k]...)
}
// Compared as sets: Match reports a value once per matching pattern, and orders them differently
// from want for patterns and inputs with a leading dot
m := g.Match(input)
if !slices.Equal(sortedSet(m), sortedSet(want)) {
t.Fatalf("seed %d: Match(%q) = %v, want %v", seed, input, m, want)
}
if m := g.MatchAny(input); m != (len(want) > 0) {
t.Fatalf("seed %d: MatchAny(%q) = %v", seed, input, m)
}
}
}
}
func TestMphMatcherGroupAppend(t *testing.T) {
g := NewMphMatcherGroup()
g.AddFullMatcher(FullMatcher("a.com"), 1)
g.AddFullMatcher(FullMatcher("b.com"), 2)
g.Build()
if m := append(g.Match("a.com"), 3); !slices.Equal(m, []uint32{1, 3}) {
t.Error("expect [1 3], but ", m)
}
if m := g.Match("b.com"); !slices.Equal(m, []uint32{2}) {
t.Error("expect [2], but ", m)
}
}
func sortedSet(v []uint32) []uint32 {
v = slices.Clone(v)
slices.Sort(v)
return slices.Compact(v)
}
func TestMphMatcherGroupLongPattern(t *testing.T) {
long := strings.Repeat("a", 300) + ".com"
for _, values := range [][4]uint32{{1, 2, 3, 4}, {7, 7, 7, 7}} {
g := NewMphMatcherGroup()
g.AddDomainMatcher(DomainMatcher(long), values[0])
g.AddFullMatcher(FullMatcher("x."+long), values[1])
g.AddFullMatcher(FullMatcher(long[:255]), values[2]) // the shortest pattern stored with a long length
g.AddFullMatcher(FullMatcher(long[:254]), values[3])
common.Must(g.Build())
cases := []struct {
input string
want []uint32
}{
{long, []uint32{values[0]}},
{"www." + long, []uint32{values[0]}},
{"x." + long, []uint32{values[1], values[0]}},
{long[1:], nil},
{"a" + long, nil},
{long[:255], []uint32{values[2]}},
{long[:254], []uint32{values[3]}},
{long[:256], nil},
{long[:253], nil},
}
for _, c := range cases {
if m := g.Match(c.input); !slices.Equal(m, c.want) {
t.Errorf("Match(%d bytes) = %v, want %v", len(c.input), m, c.want)
}
if m := g.MatchAny(c.input); m != (c.want != nil) {
t.Errorf("MatchAny(%d bytes) = %v", len(c.input), m)
}
}
}
// A pattern longer than 65535 bytes builds and matches: a record's length is a uvarint,
// so the only cap was the build-time length field, now widened to uint32.
huge := strings.Repeat("a", 70000)
g := NewMphMatcherGroup()
g.AddFullMatcher(FullMatcher(strings.Repeat("a", 65535)), 1)
g.AddDomainMatcher(DomainMatcher(huge+".com"), 2)
g.AddFullMatcher(FullMatcher("a.com"), 3)
common.Must(g.Build())
if !g.MatchAny(strings.Repeat("a", 65535)) || g.MatchAny(strings.Repeat("a", 65534)) {
t.Error("wrong answer for a 65535-byte pattern")
}
if m := g.Match(huge + ".com"); !slices.Equal(m, []uint32{2}) {
t.Errorf("Match(%d-byte input) = %v, want [2]", len(huge)+4, m)
}
if m := g.Match("x." + huge + ".com"); !slices.Equal(m, []uint32{2}) {
t.Errorf("Match(subdomain of a %d-byte pattern) = %v, want [2]", len(huge)+4, m)
}
if g.MatchAny(huge) { // the 70000-byte label on its own is not a rule
t.Error("unexpected match for the bare 70000-byte label")
}
}
func TestMphMatcherGroupBuildOnce(t *testing.T) {
g := NewMphMatcherGroup()
g.AddFullMatcher(FullMatcher("a.com"), 1)
common.Must(g.Build())
if err := g.Build(); err == nil || !g.MatchAny("a.com") {
t.Errorf("second Build() = %v, MatchAny(a.com) = %v", err, g.MatchAny("a.com"))
}
defer func() {
if recover() == nil {
t.Error("Add after Build did not panic")
}
}()
g.AddDomainMatcher(DomainMatcher("b.com"), 2)
}
+11 -281
View File
@@ -2,12 +2,9 @@ package strmatcher
import (
"errors"
"math/bits"
"regexp"
"regexp/syntax"
"slices"
"strings"
"unicode"
"unicode/utf8"
"golang.org/x/net/idna"
@@ -76,274 +73,7 @@ func (m SubstrMatcher) Match(s string) bool {
// RegexMatcher is an implementation of Matcher.
type RegexMatcher struct {
pattern *regexp.Regexp
literals []string // every match contains all of them, longest first
tail []byteSet // tail[i] holds the bytes a matching input can have i bytes before its end
rest *byteSet // the bytes it can have further before, nil if any
}
func newRegexMatcher(pattern string) (Matcher, error) {
regex, err := regexp.Compile(pattern)
if err != nil {
return nil, err
}
m := &RegexMatcher{pattern: regex}
if re, err := syntax.Parse(pattern, syntax.Perl); err == nil { // same flags as regexp.Compile
m.literals = requiredLiterals(re, nil)
slices.SortStableFunc(m.literals, func(a, b string) int { return len(b) - len(a) })
m.tail, m.rest = tailGuard(re)
}
return m, nil
}
// byteSet is a set of bytes. The bytes >= 0x80 share one bit with 0x7f.
type byteSet [4]uint32
func (s *byteSet) add(c byte) { c = min(c, 0x7f); s[c>>5] |= 1 << (c & 31) }
func (s *byteSet) has(c byte) bool { c = min(c, 0x7f); return s[c>>5]&(1<<(c&31)) != 0 }
func (s *byteSet) or(t *byteSet) {
for i := range s {
s[i] |= t[i]
}
}
var allBytes = byteSet{^uint32(0), ^uint32(0), ^uint32(0), ^uint32(0)}
// tailLen is how many positions before the end of the input tailGuard tells apart.
const tailLen = 8
// tailBudget caps how many repetition steps tailGuard walks. Only nested repeats can make the
// walk explode, so only they are charged: a flat pattern, however long, is walked once and keeps
// its guard.
const tailBudget = 100000
// tailWalk is a set of positions in the input, counted in bytes before its end.
type tailWalk struct {
at uint32 // bit i: exactly i bytes before the end, for i < tailLen
far bool // tailLen or more bytes before the end
free bool // not tied to the end of the input yet
}
func (w tailWalk) union(v tailWalk) tailWalk {
return tailWalk{w.at | v.at, w.far || v.far, w.free || v.free}
}
type tailBuilder struct {
tail [tailLen]byteSet
rest byteSet
void bool
work int
}
// tailGuard walks re backwards from the end of the input and collects the bytes an input
// matching re can have at each position before its end. It returns nil, nil when a branch
// of re does not end with $ or when nested repeats push the walk past tailBudget.
func tailGuard(re *syntax.Regexp) ([]byteSet, *byteSet) {
var b tailBuilder
w := b.walk(re, tailWalk{free: true})
b.stop(w)
if b.void {
return nil, nil
}
if w.at != 0 { // a match can start here, so any bytes can come before
for i := bits.TrailingZeros32(w.at); i < tailLen; i++ {
b.tail[i] = allBytes
}
}
if w.at != 0 || w.far {
b.rest = allBytes
}
n := tailLen
for n > 0 && b.tail[n-1] == b.rest {
n--
}
var tail []byteSet
if n > 0 {
tail = slices.Clone(b.tail[:n])
}
if b.rest != allBytes {
rest := b.rest
return tail, &rest
}
return tail, nil
}
// stop ends the paths of w. One that never met $ lets its match be followed by anything.
func (b *tailBuilder) stop(w tailWalk) {
if w.free {
b.void = true
}
}
func (b *tailBuilder) walk(re *syntax.Regexp, w tailWalk) tailWalk {
if w == (tailWalk{}) || b.void {
return w
}
switch re.Op {
case syntax.OpNoMatch:
return tailWalk{}
case syntax.OpLiteral:
for i := len(re.Rune) - 1; i >= 0; i-- {
var set byteSet
set.add(byte(min(re.Rune[i], utf8.RuneSelf)))
if re.Flags&syntax.FoldCase != 0 {
for f := unicode.SimpleFold(re.Rune[i]); f != re.Rune[i]; f = unicode.SimpleFold(f) {
set.add(byte(min(f, utf8.RuneSelf)))
}
}
w = b.step(w, &set)
}
return w
case syntax.OpCharClass:
var set byteSet
for i := 0; i+1 < len(re.Rune); i += 2 {
for r := min(re.Rune[i], utf8.RuneSelf); r <= min(re.Rune[i+1], utf8.RuneSelf); r++ {
set.add(byte(r))
}
}
return b.step(w, &set)
case syntax.OpAnyChar, syntax.OpAnyCharNotNL: // a domain has no \n to reject
return b.step(w, &allBytes)
case syntax.OpBeginText: // nothing comes before
b.stop(w)
return tailWalk{}
case syntax.OpEndText:
out := tailWalk{at: w.at & 1}
if w.free {
out.at = 1
}
return out
case syntax.OpCapture:
return b.walk(re.Sub[0], w)
case syntax.OpConcat:
for i := len(re.Sub) - 1; i >= 0; i-- {
w = b.walk(re.Sub[i], w)
}
return w
case syntax.OpAlternate:
var out tailWalk
for _, sub := range re.Sub {
out = out.union(b.walk(sub, w))
}
return out
case syntax.OpQuest:
return b.repeat(re.Sub[0], w, 1)
case syntax.OpStar:
return b.repeat(re.Sub[0], w, -1)
case syntax.OpPlus:
return b.repeat(re.Sub[0], b.walk(re.Sub[0], w), -1)
case syntax.OpRepeat:
for i := 0; i < re.Min; i++ {
if b.charge() {
return w
}
w = b.walk(re.Sub[0], w)
}
if re.Max < 0 {
return b.repeat(re.Sub[0], w, -1)
}
return b.repeat(re.Sub[0], w, re.Max-re.Min)
}
return w // empty match, line and word boundaries: no constraint
}
// charge counts one repetition step and reports whether the walk has run out of budget. Only
// repeats re-walk their body, so charging them alone bounds the blow-up of nested repeats while
// leaving a single linear pass, of any length, free.
func (b *tailBuilder) charge() bool {
b.work++
if b.work > tailBudget {
b.void = true
}
return b.void
}
// repeat walks back over up to n more repetitions of re, any number if n < 0.
func (b *tailBuilder) repeat(re *syntax.Regexp, w tailWalk, n int) tailWalk {
for ; n != 0; n-- {
if b.charge() {
return w
}
next := w.union(b.walk(re, w))
if next == w {
break
}
w = next
}
return w
}
// step walks back over one character whose last byte is in set. A character that can be
// non-ASCII can take up to 4 bytes, all >= 0x80; regexp matches an invalid byte as U+FFFD.
func (b *tailBuilder) step(w tailWalk, set *byteSet) tailWalk {
out := tailWalk{far: w.far, free: w.free}
if w.far {
b.rest.or(set)
}
width := 1
if set.has(0x80) {
width = utf8.UTFMax
}
for i := 0; i < tailLen; i++ {
if w.at&(1<<i) == 0 {
continue
}
b.tail[i].or(set)
for n := 1; n <= width; n++ {
if j := i + n; j < tailLen {
out.at |= 1 << j
if n < width {
b.tail[j].add(0x80)
}
} else {
out.far = true
if n < width {
b.rest.add(0x80)
}
}
}
}
return out
}
// mayMatch reports whether s passes the tail guard.
func (m *RegexMatcher) mayMatch(s string) bool {
n := len(s)
if m.rest == nil {
n = min(n, len(m.tail))
}
for i := 0; i < n; i++ {
set := m.rest
if i < len(m.tail) {
set = &m.tail[i]
}
if !set.has(s[len(s)-1-i]) {
return false
}
}
return true
}
// requiredLiterals appends to dst the case-sensitive strings that every match of re contains.
func requiredLiterals(re *syntax.Regexp, dst []string) []string {
switch re.Op {
case syntax.OpLiteral:
// regexp matches U+FFFD against invalid UTF-8 bytes, strings.Contains does not
if re.Flags&syntax.FoldCase == 0 && !slices.Contains(re.Rune, utf8.RuneError) {
dst = append(dst, string(re.Rune))
}
case syntax.OpCapture, syntax.OpPlus:
dst = requiredLiterals(re.Sub[0], dst)
case syntax.OpRepeat:
if re.Min > 0 {
dst = requiredLiterals(re.Sub[0], dst)
}
case syntax.OpConcat:
for _, sub := range re.Sub {
dst = requiredLiterals(sub, dst)
}
}
return dst
pattern *regexp.Regexp
}
func (*RegexMatcher) Type() Type {
@@ -359,14 +89,6 @@ func (m *RegexMatcher) String() string {
}
func (m *RegexMatcher) Match(s string) bool {
if !m.mayMatch(s) {
return false
}
for _, l := range m.literals {
if !strings.Contains(s, l) {
return false
}
}
return m.pattern.MatchString(s)
}
@@ -380,7 +102,11 @@ func (t Type) New(pattern string) (Matcher, error) {
case Domain:
return DomainMatcher(pattern), nil
case Regex: // 1. regex matching is case-sensitive
return newRegexMatcher(pattern)
regex, err := regexp.Compile(pattern)
if err != nil {
return nil, err
}
return &RegexMatcher{pattern: regex}, nil
default:
return nil, errors.New("unknown matcher type")
}
@@ -409,7 +135,11 @@ func (t Type) NewDomainPattern(pattern string) (Matcher, error) {
}
return DomainMatcher(pattern), nil
case Regex: // Regex's charset not in LDH subset
return newRegexMatcher(pattern)
regex, err := regexp.Compile(pattern)
if err != nil {
return nil, err
}
return &RegexMatcher{pattern: regex}, nil
default:
return nil, errors.New("unknown matcher type")
}
@@ -1,233 +0,0 @@
package strmatcher
import (
"hash/fnv"
"math/rand/v2"
"regexp"
"regexp/syntax"
"slices"
"strconv"
"strings"
"testing"
"unicode"
"unicode/utf8"
)
var regexLiteralCases = []struct {
pattern string
literals []string
}{
{`(^|\.)91porn\.(best|com)$`, []string{"91porn."}},
{`.+\.awsdns-cn-[0-9][0-9]\.(biz|com|net|top)$`, []string{".awsdns-cn-", "."}},
{`^r+[0-9]+(---|\.)sn-(2x3|ni5|j5o)\w{5}\.googlevideo\.com$`, []string{".googlevideo.com", "sn-", "r"}},
{`(?i)abc`, nil},
{`ab(?i:CD)ef`, []string{"ab", "ef"}},
{`(abc)?x`, []string{"x"}},
{`(abc)*x`, []string{"x"}},
{`x{0,3}yy`, []string{"yy"}},
{`(ab)+c{2}`, []string{"ab", "c"}},
{`abc|abd`, []string{"ab"}},
{`\Qa.b\E`, []string{"a.b"}},
{`a\x{FFFD}b`, nil},
{`^[^.]+$`, nil},
}
func TestRegexRequiredLiterals(t *testing.T) {
for _, test := range regexLiteralCases {
m, err := newRegexMatcher(test.pattern)
if err != nil {
t.Fatal(err)
}
if got := m.(*RegexMatcher).literals; !slices.Equal(got, test.literals) {
t.Errorf("%s: got %q, want %q", test.pattern, got, test.literals)
}
}
}
var regexTailCases = []struct {
pattern string
guard bool
match []string // inputs the pattern matches
reject []string // inputs the tail guard alone rejects
}{
{`^[a-z]([a-z0-9-]{0,61}[a-z0-9])?$`, true, []string{"a", "localhost", "x-1"}, []string{"www.example.com", "localhost.", "LOCALHOST", "a b"}},
{`(^|\.)[a-z][1-9][0-9][a-z]\.com$`, true, []string{"a12b.com", "x.q10z.com"}, []string{"google.com", "a12b.co", "a12b.com.", "ab12.com"}},
{`^hses[1-7]?\.akamaized\.net$`, true, []string{"hses.akamaized.net", "hses3.akamaized.net"}, []string{"xhses.akamaized.net", "www.hses.akamaized.net"}},
{`(?i)k\.net$`, true, []string{"k.net", "K.NET", "\u212a.net"}, []string{"x.net", "k.nex"}},
{`[^.]+\.cn$`, true, []string{"a.cn", "\xff.cn", "\u4e2d.cn"}, []string{"a.cnn", "a.c"}},
{`\x{FFFD}$`, true, []string{"\xff", "a\xc3", "\uFFFD"}, []string{"a", "\xff."}},
{`^.\.cn$`, true, []string{"a.cn", "\u4E2D.cn", "\xff.cn"}, []string{"ab.cn"}},
{`^$`, true, []string{""}, []string{"a"}},
{`(^|\.)youyuapi\..+$`, false, []string{"youyuapi.com"}, nil},
{`abc`, false, []string{"abc", "xabcx"}, nil},
{`^ab`, false, []string{"ab", "abc"}, nil},
{`a$|b`, false, []string{"a", "bx"}, nil},
{`(?m)a$`, false, []string{"a", "a\nb"}, nil},
{strings.Repeat(`(?:abcdefgh(?:a`, 20) + strings.Repeat(`)*)*`, 20) + `\.com$`, false, []string{".com", "abcdefgha.com"}, nil}, // over tailBudget
}
func TestRegexTailGuard(t *testing.T) {
for _, test := range regexTailCases {
m, err := newRegexMatcher(test.pattern)
if err != nil {
t.Fatal(err)
}
rm := m.(*RegexMatcher)
if guard := rm.tail != nil || rm.rest != nil; guard != test.guard {
t.Errorf("%s: guard %v, want %v", test.pattern, guard, test.guard)
}
for _, s := range test.match {
if !rm.pattern.MatchString(s) || !rm.Match(s) {
t.Errorf("%s: %q does not match", test.pattern, s)
}
}
for _, s := range test.reject {
if rm.pattern.MatchString(s) || rm.mayMatch(s) {
t.Errorf("%s: %q passes the guard", test.pattern, s)
}
}
}
}
// TestRegexTailGuardFlatAlternation checks that a long but non-recursive pattern keeps its
// guard. Only nested repeats are charged against tailBudget, so a flat alternation of many
// names, however large, is walked once and guarded; its guard is checked against regexp.
func TestRegexTailGuardFlatAlternation(t *testing.T) {
var sb strings.Builder
sb.WriteString("(?:")
for i := 0; i < 20000; i++ {
if i > 0 {
sb.WriteByte('|')
}
sb.WriteString("name")
sb.WriteString(strconv.Itoa(i))
}
sb.WriteString(`)\.example\.com$`)
m, err := newRegexMatcher(sb.String())
if err != nil {
t.Fatal(err)
}
rm := m.(*RegexMatcher)
if rm.tail == nil && rm.rest == nil {
t.Fatal("flat alternation of 20000 names lost its guard")
}
for _, s := range []string{"name0.example.com", "name19999.example.com", "x.name12345.example.com"} {
if !rm.pattern.MatchString(s) || !rm.Match(s) {
t.Errorf("%q should match", s)
}
}
for _, s := range []string{"name0.example.org", "name0.example.com.", "name0.example.con", "google.com"} {
if rm.pattern.MatchString(s) {
t.Fatalf("test bug: %q matches the pattern", s)
}
if rm.mayMatch(s) {
t.Errorf("%q should be rejected by the guard", s)
}
}
}
// sampleMatch appends a string that re matches, assertions aside, unless it runs out of
// budget, which it spends one per call so that nested repeats stay cheap.
func sampleMatch(sb *strings.Builder, re *syntax.Regexp, rnd *rand.Rand, budget *int) {
if *budget <= 0 {
return
}
*budget--
switch re.Op {
case syntax.OpLiteral:
for _, r := range re.Rune {
if re.Flags&syntax.FoldCase != 0 {
for n := rnd.IntN(4); n > 0; n-- {
r = unicode.SimpleFold(r)
}
}
sampleRune(sb, r, rnd)
}
case syntax.OpCharClass:
if len(re.Rune) > 0 {
i := rnd.IntN(len(re.Rune)/2) * 2
sampleRune(sb, re.Rune[i]+rnd.Int32N(min(re.Rune[i+1]-re.Rune[i]+1, 300)), rnd)
}
case syntax.OpAnyChar, syntax.OpAnyCharNotNL:
sampleRune(sb, []rune{'a', '.', '\n', 0xe9, 0x212a, utf8.RuneError}[rnd.IntN(6)], rnd)
case syntax.OpCapture:
sampleMatch(sb, re.Sub[0], rnd, budget)
case syntax.OpConcat:
for _, sub := range re.Sub {
sampleMatch(sb, sub, rnd, budget)
}
case syntax.OpAlternate:
sampleMatch(sb, re.Sub[rnd.IntN(len(re.Sub))], rnd, budget)
case syntax.OpQuest, syntax.OpStar, syntax.OpPlus, syntax.OpRepeat:
lo, hi := 0, 3
switch re.Op {
case syntax.OpQuest:
hi = 1
case syntax.OpPlus:
lo = 1
case syntax.OpRepeat:
lo, hi = re.Min, re.Min+3
if re.Max >= 0 {
hi = min(hi, re.Max)
}
}
for n := lo + rnd.IntN(hi-lo+1); n > 0; n-- {
sampleMatch(sb, re.Sub[0], rnd, budget)
}
}
}
func sampleRune(sb *strings.Builder, r rune, rnd *rand.Rand) {
if r == utf8.RuneError && rnd.IntN(2) == 0 {
sb.WriteByte(0x80 | byte(rnd.IntN(0x80))) // regexp matches an invalid byte as U+FFFD
return
}
sb.WriteRune(r)
}
func FuzzRegexMatcher(f *testing.F) {
inputs := []string{
"", "x", "yy", "abd", "ccc", "ABC", "abCDef", "abcdef", "abababcc", "a.b", "a\xffb", "a\uFFFDb",
"www.91porn.com", "ns1.awsdns-cn-01.top", "r1---sn-2x3abcde.googlevideo.com",
}
for _, test := range regexLiteralCases {
for _, s := range inputs {
f.Add(test.pattern, s)
}
}
for _, test := range regexTailCases {
for _, s := range append(test.match, test.reject...) {
f.Add(test.pattern, s)
}
}
f.Fuzz(func(t *testing.T, pattern, s string) {
re, err := regexp.Compile(pattern)
if err != nil {
return
}
m, _ := newRegexMatcher(pattern)
check := func(s string) {
if got, want := m.Match(s), re.MatchString(s); got != want {
t.Errorf("pattern %q, input %q: got %v, want %v", pattern, s, got, want)
}
}
check(s)
// random inputs seldom match, so also try strings built from the pattern
parsed, _ := syntax.Parse(pattern, syntax.Perl)
h := fnv.New64a()
h.Write([]byte(s))
rnd := rand.New(rand.NewPCG(h.Sum64(), 1))
for range 8 {
var sb strings.Builder
budget := 256
sampleMatch(&sb, parsed, rnd, &budget)
sample := sb.String()
check(sample)
check(s + sample)
if len(sample) > 0 && len(s) > 0 {
i := rnd.IntN(len(sample))
check(sample[:i] + s[:1] + sample[i+1:])
}
}
})
}
+12 -67
View File
@@ -46,9 +46,7 @@ func (g *MphValueMatcher) Add(matcher Matcher, value uint32) {
func (g *MphValueMatcher) Build() error {
if g.mph != nil {
runtime.GC() // peak mem
if err := g.mph.Build(); err != nil {
return err
}
g.mph.Build()
}
runtime.GC() // peak mem
if g.ac != nil {
@@ -60,17 +58,23 @@ func (g *MphValueMatcher) Build() error {
// Match implements ValueMatcher.Match.
func (g *MphValueMatcher) Match(input string) []uint32 {
var result []uint32
result := make([][]uint32, 0, 5)
if g.mph != nil {
result = g.mph.Match(input) // a new slice, returned without another copy
if matches := g.mph.Match(input); len(matches) > 0 {
result = append(result, matches)
}
}
if g.ac != nil {
result = append(result, g.ac.Match(input)...)
if matches := g.ac.Match(input); len(matches) > 0 {
result = append(result, matches)
}
}
if g.regex != nil {
result = append(result, g.regex.Match(input)...)
if matches := g.regex.Match(input); len(matches) > 0 {
result = append(result, matches)
}
}
return result
return CompositeMatches(result)
}
// MatchAny implements ValueMatcher.MatchAny.
@@ -83,62 +87,3 @@ func (g *MphValueMatcher) MatchAny(input string) bool {
}
return g.regex != nil && g.regex.MatchAny(input)
}
func (g *MphValueMatcher) matchAnyHashed(input string, parents []mphSuffix, h, mul uint64) bool {
if g.mph != nil && g.mph.matchAnyHashed(input, parents, h, mul) {
return true
}
if g.ac != nil && g.ac.MatchAny(input) {
return true
}
return g.regex != nil && g.regex.MatchAny(input)
}
// MphValueMatcherCombiner combines several built MphValueMatchers, each bound to one value, and matches an input
// against them as their MatchAny would, hashing the input once for all of them.
type MphValueMatcherCombiner struct {
matchers []*MphValueMatcher
values []uint32
}
// Add adds a built matcher that stands for value.
func (s *MphValueMatcherCombiner) Add(m *MphValueMatcher, value uint32) {
s.matchers = append(s.matchers, m)
s.values = append(s.values, value)
}
// Match returns the values of the matchers that match input, in Add order.
func (s *MphValueMatcherCombiner) Match(input string) []uint32 {
if len(s.matchers) == 0 {
return nil
}
var stack [16]mphSuffix
mul := mphMultipliers[0]
parents, h := mphSuffixes(stack[:0], mul, input)
var result []uint32
for i, m := range s.matchers {
if m.matchAnyHashed(input, parents, h, mul) {
result = append(result, s.values[i])
}
}
return result
}
// MatchAny returns true as soon as one matcher matches input.
func (s *MphValueMatcherCombiner) MatchAny(input string) bool {
switch len(s.matchers) {
case 0:
return false
case 1:
return s.matchers[0].MatchAny(input) // nothing to share, and it stops at the first matching suffix
}
var stack [16]mphSuffix
mul := mphMultipliers[0]
parents, h := mphSuffixes(stack[:0], mul, input)
for _, m := range s.matchers {
if m.matchAnyHashed(input, parents, h, mul) {
return true
}
}
return false
}
+25 -21
View File
@@ -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
}
-4
View File
@@ -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()
-61
View File
@@ -1,61 +0,0 @@
package log
import (
"path/filepath"
"strings"
lua "github.com/yuin/gopher-lua"
)
// RegisterLua makes xray.log available to require in an LState.
func RegisterLua(L *lua.LState) {
L.PreloadModule("xray.log", func(L *lua.LState) int {
module := L.NewTable()
var source, prefix string // cache
for name, severity := range map[string]Severity{
"Debug": Severity_Debug,
"Info": Severity_Info,
"Warning": Severity_Warning,
"Error": Severity_Error,
} {
module.RawSetString(name, L.NewFunction(func(L *lua.LState) int {
if GetSeverity() < severity {
return 0
}
var content strings.Builder
// Prefix with the calling script's filename.
if caller, ok := L.GetStack(1); ok {
if _, err := L.GetInfo("S", caller, lua.LNil); err == nil && caller.Source != "" {
if caller.Source != source {
source = caller.Source
prefix = filepath.Base(strings.TrimPrefix(source, "@")) + ": "
}
content.WriteString(prefix)
}
}
for i := 1; i <= L.GetTop(); i++ {
content.WriteString(luaLogString(L, L.Get(i)))
}
Record(&GeneralMessage{
Severity: severity,
Content: content.String(),
})
return 0
}))
}
L.Push(module)
return 1
})
}
func luaLogString(L *lua.LState, value lua.LValue) string {
if ud, ok := value.(*lua.LUserData); ok {
if err, ok := ud.Value.(error); ok {
return err.Error()
}
}
if _, ok := L.GetMetaField(value, "__tostring").(*lua.LFunction); ok {
return L.ToStringMeta(value).String()
}
return value.String()
}
-213
View File
@@ -1,213 +0,0 @@
package log
import (
"errors"
"fmt"
"os"
"path/filepath"
"testing"
lua "github.com/yuin/gopher-lua"
)
type luaLogHandler struct {
messages []Message
}
func (h *luaLogHandler) Handle(msg Message) {
h.messages = append(h.messages, msg)
}
func TestLuaLog(t *testing.T) {
previous := logHandler.Load()
t.Cleanup(func() { logHandler.Store(previous) })
handler := &luaLogHandler{}
RegisterHandler(handler)
L := lua.NewState()
defer L.Close()
RegisterLua(L)
nativeError := L.NewUserData()
nativeError.Value = fmt.Errorf("lookup failed: %w", errors.New("upstream timeout"))
L.SetGlobal("nativeError", nativeError)
path := filepath.Join(t.TempDir(), "logging.lua")
if err := os.WriteFile(path, []byte(`
local log = require("xray.log")
assert(log == require("xray.log"))
log.Debug("query: ", "example.com")
log.Info("count=", 42, ", enabled=", true, ", value=", nil)
log.Warning(setmetatable({}, {
__tostring = function() return "fallback" end
}))
assert(select("#", log.Error("failed")) == 0)
log.Error("DNS failed: ", nativeError)
log.Warning(nativeError)
local ok, err = pcall(function() error("Lua failure", 0) end)
assert(not ok)
log.Error(err)
local calls = 0
local custom = setmetatable({}, {
__tostring = function() calls = calls + 1; return "custom" end
})
log.Info(custom, custom)
assert(calls == 2)
log.Info("a", "b", "c", "d", "e", "f", "g", "h", "i", "j", "k", "l")
log.Info()
function logHook()
log.Info("hook")
end
`), 0o600); err != nil {
t.Fatal(err)
}
if err := L.DoFile(path); err != nil {
t.Fatal(err)
}
if err := L.DoString(`
logHook()
require("xray.log").Info("anonymous")
`); err != nil {
t.Fatal(err)
}
other := filepath.Join(t.TempDir(), "other.lua")
if err := os.WriteFile(other, []byte(`
local log = require("xray.log")
log.Info("other")
logHook()
log.Info("other again")
`), 0o600); err != nil {
t.Fatal(err)
}
if err := L.DoFile(other); err != nil {
t.Fatal(err)
}
want := []struct {
severity Severity
message string
}{
{Severity_Debug, "[Debug] logging.lua: query: example.com"},
{Severity_Info, "[Info] logging.lua: count=42, enabled=true, value=nil"},
{Severity_Warning, "[Warning] logging.lua: fallback"},
{Severity_Error, "[Error] logging.lua: failed"},
{Severity_Error, "[Error] logging.lua: DNS failed: lookup failed: upstream timeout"},
{Severity_Warning, "[Warning] logging.lua: lookup failed: upstream timeout"},
{Severity_Error, "[Error] logging.lua: Lua failure"},
{Severity_Info, "[Info] logging.lua: customcustom"},
{Severity_Info, "[Info] logging.lua: abcdefghijkl"},
{Severity_Info, "[Info] logging.lua: "},
{Severity_Info, "[Info] logging.lua: hook"},
{Severity_Info, "[Info] <string>: anonymous"},
{Severity_Info, "[Info] other.lua: other"},
{Severity_Info, "[Info] logging.lua: hook"},
{Severity_Info, "[Info] other.lua: other again"},
}
if len(handler.messages) != len(want) {
t.Fatalf("logged %d messages, want %d", len(handler.messages), len(want))
}
for i, expected := range want {
msg, ok := handler.messages[i].(*GeneralMessage)
if !ok {
t.Fatalf("message %d has type %T, want *GeneralMessage", i, handler.messages[i])
}
if msg.Severity != expected.severity || msg.String() != expected.message {
t.Errorf("message %d = %q with severity %v, want %q with severity %v", i, msg.String(), msg.Severity, expected.message, expected.severity)
}
}
}
type luaSeverityLogHandler struct {
luaLogHandler
level Severity
}
func (h *luaSeverityLogHandler) Severity() Severity { return h.level }
func TestLuaLogSeverity(t *testing.T) {
previous := logHandler.Load()
t.Cleanup(func() { logHandler.Store(previous) })
L := lua.NewState()
defer L.Close()
RegisterLua(L)
for _, level := range []Severity{Severity_Unknown, Severity_Error, Severity_Warning, Severity_Info, Severity_Debug, Severity_Warning} {
t.Run(level.String(), func(t *testing.T) {
handler := &luaSeverityLogHandler{level: level}
RegisterHandler(handler)
want := []Severity{}
for _, severity := range []Severity{Severity_Error, Severity_Warning, Severity_Info, Severity_Debug} {
if severity <= level {
want = append(want, severity)
}
}
if err := L.DoString(fmt.Sprintf(`
local log = require("xray.log")
local calls = 0
local value = setmetatable({}, {
__tostring = function() calls = calls + 1; return "message" end
})
for _, write in ipairs({log.Error, log.Warning, log.Info, log.Debug}) do
assert(select("#", write(value)) == 0)
end
assert(calls == %d)
`, len(want))); err != nil {
t.Fatal(err)
}
if len(handler.messages) != len(want) {
t.Fatalf("logged %d messages, want %d", len(handler.messages), len(want))
}
for i, severity := range want {
msg := handler.messages[i].(*GeneralMessage)
if msg.Severity != severity || msg.Content != "<string>: message" {
t.Errorf("message %d = %v, want severity %v and content %q", i, msg, severity, "<string>: message")
}
}
})
}
}
type luaDiscardLogHandler struct{ level Severity }
func (luaDiscardLogHandler) Handle(Message) {}
func (h luaDiscardLogHandler) Severity() Severity { return h.level }
func BenchmarkLuaLog(b *testing.B) {
benchmarkLuaLog(b, Severity_Debug)
}
func BenchmarkLuaLogFiltered(b *testing.B) {
benchmarkLuaLog(b, Severity_Warning)
}
func benchmarkLuaLog(b *testing.B, level Severity) {
previous := logHandler.Load()
b.Cleanup(func() { logHandler.Store(previous) })
RegisterHandler(luaDiscardLogHandler{level: level})
L := lua.NewState()
defer L.Close()
RegisterLua(L)
if err := L.DoString(`custom = setmetatable({}, {__tostring = function() return "custom" end})`); err != nil {
b.Fatal(err)
}
for _, benchmark := range []struct {
name, arguments string
}{
{"strings", `"query: ", "example.com"`},
{"mixed", `"count=", 42, ", enabled=", true, ", value=", nil`},
{"many_arguments", `"a", "b", "c", "d", "e", "f", "g", "h", "i", "j", "k", "l"`},
{"tostring", "custom"},
} {
b.Run(benchmark.name, func(b *testing.B) {
if err := L.DoString(fmt.Sprintf(`local log = require("xray.log")
function benchmarkLog() log.Info(%s) end`, benchmark.arguments)); err != nil {
b.Fatal(err)
}
fn := L.GetGlobal("benchmarkLog")
b.ReportAllocs()
b.ResetTimer()
for i := 0; i < b.N; i++ {
if err := L.CallByParam(lua.P{Fn: fn, NRet: 0, Protect: true}); err != nil {
b.Fatal(err)
}
}
})
}
}
-3
View File
@@ -1,3 +0,0 @@
// Package lua provides shared GopherLua programs, state management, and value
// conversion and validation helpers for Xray scripts.
package lua
-150
View File
@@ -1,150 +0,0 @@
package lua
import (
"context"
"errors"
"sync"
"time"
glua "github.com/yuin/gopher-lua"
)
const maxIdleStates = 16
// Pool lends each state to one caller at a time. It grows on contention and
// keeps up to maxIdleStates idle states until Close. Acquire/Release callers
// decide reusability; WithState uses its callback's error.
type Pool struct {
ctx context.Context
cancel context.CancelFunc
timeout time.Duration
factory LStateFactory
idle []*glua.LState
top int
mu sync.Mutex
active sync.WaitGroup
closed bool
}
// NewPool tests the factory by creating one state during initialization.
func NewPool(ctx context.Context, timeout time.Duration, factory LStateFactory) (*Pool, error) {
if timeout <= 0 {
return nil, errors.New("Lua pool timeout must be positive")
}
poolCtx, cancel := context.WithCancel(ctx)
state, err := factory(poolCtx)
if err != nil {
cancel()
return nil, err
}
return &Pool{ctx: poolCtx, cancel: cancel, timeout: timeout, factory: factory, idle: []*glua.LState{state}, top: state.GetTop()}, nil
}
// Acquire returns an initialized exclusive state, growing the pool if necessary.
// ctx is passed to the factory for state creation; nil uses the pool context.
func (p *Pool) Acquire(ctx context.Context) (*glua.LState, error) {
p.mu.Lock()
if p.closed {
p.mu.Unlock()
return nil, errors.New("Lua pool is closed")
}
if err := p.ctx.Err(); err != nil {
p.mu.Unlock()
return nil, err
}
if ctx == nil {
ctx = p.ctx
} else if err := ctx.Err(); err != nil {
p.mu.Unlock()
return nil, err
}
p.active.Add(1)
n := len(p.idle)
if n != 0 {
state := p.idle[n-1]
p.idle = p.idle[:n-1]
p.mu.Unlock()
return state, nil
}
p.mu.Unlock()
// TODO: Limit the total number of states. When the limit is reached, wait
// for a Release instead of creating another state; allow the wait to be
// cancelled by the caller or by Close.
state, err := p.factory(ctx)
if err != nil {
p.active.Done()
return nil, err
}
return state, nil
}
// WithState runs work on an exclusive state and releases it afterward.
// Nil ctx and zero timeout use pool defaults. The timeout starts after acquisition.
func (p *Pool) WithState(ctx context.Context, timeout time.Duration, work func(*glua.LState) error) error {
state, err := p.Acquire(ctx)
if err != nil {
return err
}
if ctx == nil {
ctx = p.ctx
}
if timeout == 0 {
timeout = p.timeout
}
ctx, cancel := context.WithTimeout(ctx, timeout)
state.SetContext(ctx)
reusable := false
defer func() {
cancel()
p.Release(state, reusable)
}()
err = work(state)
reusable = err == nil
return err
}
// Release resets a state for reuse or closes it.
func (p *Pool) Release(state *glua.LState, reusable bool) {
if reusable {
state.RemoveContext()
state.SetTop(p.top)
p.mu.Lock()
if !p.closed && p.ctx.Err() == nil && len(p.idle) < maxIdleStates {
p.idle = append(p.idle, state)
} else {
reusable = false
}
p.mu.Unlock()
}
if !reusable {
state.Close()
}
p.active.Done()
}
// Close cancels the pool context, closes idle states, and waits for borrowed states.
func (p *Pool) Close() {
p.mu.Lock()
if !p.closed {
p.closed = true
p.cancel()
for _, state := range p.idle {
state.Close()
}
p.idle = nil
}
p.mu.Unlock()
p.active.Wait()
}
-466
View File
@@ -1,466 +0,0 @@
package lua
import (
"context"
"errors"
"testing"
"time"
glua "github.com/yuin/gopher-lua"
)
func newTestPool(t testing.TB, ctx context.Context, timeout time.Duration, factory LStateFactory) *Pool {
t.Helper()
pool, err := NewPool(ctx, timeout, factory)
if err != nil {
t.Fatal(err)
}
t.Cleanup(pool.Close)
return pool
}
func assertPoolCloseBlocked(t *testing.T, done <-chan struct{}) {
t.Helper()
select {
case <-done:
t.Fatal("Close returned while work was still active")
case <-time.After(20 * time.Millisecond):
}
}
func TestPoolTimeoutValidation(t *testing.T) {
for _, tc := range []struct {
name string
timeout time.Duration
wantErr bool
}{
{"zero", 0, true},
{"negative", -time.Nanosecond, true},
{"positive", time.Nanosecond, false},
} {
t.Run(tc.name, func(t *testing.T) {
called := false
pool, err := NewPool(context.Background(), tc.timeout, func(context.Context) (*glua.LState, error) {
called = true
return glua.NewState(), nil
})
if pool != nil {
t.Cleanup(pool.Close)
}
if (err != nil) != tc.wantErr {
t.Fatalf("NewPool error = %v, want error %t", err, tc.wantErr)
}
if tc.wantErr && (pool != nil || called) {
t.Fatal("invalid timeout created a pool or called the factory")
}
})
}
}
func TestPoolFactoryFailure(t *testing.T) {
failure := errors.New("factory failed")
_, err := NewPool(context.Background(), time.Second, func(context.Context) (*glua.LState, error) {
return nil, failure
})
if !errors.Is(err, failure) {
t.Fatalf("NewPool error = %v, want original factory error", err)
}
calls := 0
pool := newTestPool(t, context.Background(), time.Second, func(context.Context) (*glua.LState, error) {
calls++
if calls == 1 {
return glua.NewState(), nil
}
return nil, failure
})
state, err := pool.Acquire(nil)
if err != nil {
t.Fatal(err)
}
defer pool.Release(state, true)
err = pool.WithState(nil, 0, func(*glua.LState) error {
t.Error("work ran after factory failure")
return nil
})
if !errors.Is(err, failure) {
t.Fatalf("WithState error = %v, want original factory error", err)
}
}
func TestPoolReusesStatesAndLimitsIdle(t *testing.T) {
created := 0
pool := newTestPool(t, context.Background(), time.Second, func(context.Context) (*glua.LState, error) {
created++
return glua.NewState(), nil
})
var borrowed []*glua.LState
defer func() {
for _, state := range borrowed {
pool.Release(state, false)
}
}()
for range maxIdleStates + 3 {
state, err := pool.Acquire(nil)
if err != nil {
t.Fatal(err)
}
borrowed = append(borrowed, state)
state.SetContext(context.Background())
}
states := borrowed
for _, state := range states {
pool.Release(state, true)
}
borrowed = nil
open := 0
for _, state := range states {
if !state.IsClosed() {
if state.Context() != nil {
t.Fatal("Release left a context on a reusable state")
}
open++
}
}
if open != maxIdleStates {
t.Fatalf("retained %d states, want %d", open, maxIdleStates)
}
if err := pool.WithState(nil, 0, func(*glua.LState) error { return nil }); err != nil {
t.Fatal(err)
}
if created != len(states) {
t.Fatalf("created %d states, want %d", created, len(states))
}
pool.Close()
for _, state := range states {
if !state.IsClosed() {
t.Fatal("Close left an idle state open")
}
}
}
func TestPoolWithStateOptions(t *testing.T) {
key := struct{}{}
parent := context.WithValue(context.Background(), key, "pool")
caller := context.WithValue(context.Background(), key, "caller")
pool := newTestPool(t, parent, time.Second, func(context.Context) (*glua.LState, error) {
return glua.NewState(), nil
})
for _, tc := range []struct {
name string
ctx context.Context
timeout time.Duration
wantValue string
wantTimeout time.Duration
}{
{"defaults", nil, 0, "pool", time.Second},
{"context", caller, 0, "caller", time.Second},
{"timeout", nil, 2 * time.Second, "pool", 2 * time.Second},
{"both", caller, 2 * time.Second, "caller", 2 * time.Second},
} {
t.Run(tc.name, func(t *testing.T) {
started := time.Now()
err := pool.WithState(tc.ctx, tc.timeout, func(L *glua.LState) error {
ctx := L.Context()
if ctx.Value(key) != tc.wantValue {
t.Errorf("context value = %v, want %q", ctx.Value(key), tc.wantValue)
}
deadline, ok := ctx.Deadline()
if !ok || deadline.Before(started.Add(tc.wantTimeout)) || deadline.After(time.Now().Add(tc.wantTimeout)) {
t.Errorf("deadline = %v, want timeout %v", deadline, tc.wantTimeout)
}
return nil
})
if err != nil {
t.Fatal(err)
}
})
}
}
func TestPoolFactoryContext(t *testing.T) {
caller, cancel := context.WithTimeout(context.Background(), time.Minute)
defer cancel()
for _, tc := range []struct {
name string
ctx context.Context
}{
{"default", nil},
{"caller", caller},
} {
t.Run(tc.name, func(t *testing.T) {
var contexts []context.Context
pool := newTestPool(t, context.Background(), time.Second, func(ctx context.Context) (*glua.LState, error) {
contexts = append(contexts, ctx)
return glua.NewState(), nil
})
state, err := pool.Acquire(nil)
if err != nil {
t.Fatal(err)
}
defer pool.Release(state, true)
if err := pool.WithState(tc.ctx, 2*time.Second, func(*glua.LState) error { return nil }); err != nil {
t.Fatal(err)
}
want := tc.ctx
if want == nil {
want = pool.ctx
}
if len(contexts) != 2 || contexts[0] != pool.ctx || contexts[1] != want {
t.Fatal("factory did not receive the initialization and acquisition contexts unchanged")
}
})
}
}
func TestPoolWithStateLifecycle(t *testing.T) {
failure := errors.New("work failed")
for _, tc := range []struct {
name string
work func(*glua.LState, context.CancelFunc) error
reusable bool
wantPanic bool
wantErr error
}{
{"success", func(*glua.LState, context.CancelFunc) error { return nil }, true, false, nil},
{"canceled success", func(_ *glua.LState, cancel context.CancelFunc) error {
cancel()
return nil
}, true, false, nil},
{"error", func(*glua.LState, context.CancelFunc) error { return failure }, false, false, failure},
{"timeout", func(L *glua.LState, _ context.CancelFunc) error { return L.DoString("while true do end") }, false, false, nil},
{"panic", func(*glua.LState, context.CancelFunc) error { panic(failure) }, false, true, nil},
} {
t.Run(tc.name, func(t *testing.T) {
pool := newTestPool(t, context.Background(), 10*time.Millisecond, func(context.Context) (*glua.LState, error) {
state := glua.NewState()
state.Push(glua.LTrue)
return state, nil
})
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
var state *glua.LState
var workCtx context.Context
var recovered any
err := func() (err error) {
defer func() { recovered = recover() }()
return pool.WithState(ctx, 0, func(L *glua.LState) error {
state, workCtx = L, L.Context()
L.Push(glua.LFalse)
return tc.work(L, cancel)
})
}()
if tc.wantPanic {
if recovered != failure {
t.Fatalf("panic = %v, want original panic", recovered)
}
} else {
if recovered != nil || (err == nil) != tc.reusable {
t.Fatalf("WithState error = %v, panic = %v", err, recovered)
}
if tc.wantErr != nil && !errors.Is(err, tc.wantErr) {
t.Fatalf("WithState error = %v, want %v", err, tc.wantErr)
}
}
if workCtx.Err() == nil {
t.Fatal("WithState did not cancel the execution context")
}
if closed := state.IsClosed(); closed == tc.reusable {
t.Fatalf("state closed = %t, want %t", closed, !tc.reusable)
}
if tc.reusable && (state.Context() != nil || state.GetTop() != 1 || state.Get(1) != glua.LTrue) {
t.Fatal("WithState did not reset the state for reuse")
}
if err := pool.WithState(nil, 0, func(L *glua.LState) error {
if (L == state) != tc.reusable {
t.Error("unexpected state reuse")
}
return nil
}); err != nil {
t.Fatal(err)
}
})
}
}
func TestPoolClose(t *testing.T) {
pool := newTestPool(t, context.Background(), time.Second, func(context.Context) (*glua.LState, error) {
return glua.NewState(), nil
})
finishCtx, finish := context.WithCancel(context.Background())
t.Cleanup(finish)
started, done := make(chan *glua.LState, 1), make(chan error, 1)
var workCtx context.Context
go func() {
done <- pool.WithState(nil, 0, func(L *glua.LState) error {
workCtx = L.Context()
started <- L
<-finishCtx.Done()
return nil
})
}()
var state *glua.LState
select {
case state = <-started:
case <-time.After(time.Second):
t.Fatal("WithState did not start")
}
closed := make(chan struct{})
go func() {
pool.Close()
close(closed)
}()
select {
case <-workCtx.Done():
case <-time.After(time.Second):
t.Fatal("Close did not cancel work using the pool context")
}
if !errors.Is(workCtx.Err(), context.Canceled) {
t.Fatalf("work context error = %v, want context.Canceled", workCtx.Err())
}
assertPoolCloseBlocked(t, closed)
finish()
select {
case err := <-done:
if err != nil {
t.Fatalf("successful work returned an error: %v", err)
}
case <-time.After(time.Second):
t.Fatal("WithState did not finish")
}
select {
case <-closed:
case <-time.After(time.Second):
t.Fatal("Close did not finish after WithState")
}
if !state.IsClosed() {
t.Fatal("Release returned a state to a closed pool")
}
if state, err := pool.Acquire(nil); state != nil || err == nil || errors.Is(err, context.Canceled) {
t.Fatalf("Acquire after Close = %v, %v; want closed pool error", state, err)
}
pool.Close()
}
func TestPoolCloseWaitsForFactory(t *testing.T) {
finishCtx, finish := context.WithCancel(context.Background())
started, canceled := make(chan struct{}), make(chan struct{})
first := true
pool := newTestPool(t, context.Background(), time.Second, func(ctx context.Context) (*glua.LState, error) {
if first {
first = false
return glua.NewState(), nil
}
close(started)
<-ctx.Done()
close(canceled)
<-finishCtx.Done()
return nil, ctx.Err()
})
t.Cleanup(finish)
state, err := pool.Acquire(nil)
if err != nil {
t.Fatal(err)
}
pool.Release(state, false)
acquireDone := make(chan error, 1)
go func() {
_, err := pool.Acquire(nil)
acquireDone <- err
}()
select {
case <-started:
case <-time.After(time.Second):
t.Fatal("state creation did not start")
}
closed := make(chan struct{})
go func() {
pool.Close()
close(closed)
}()
select {
case <-canceled:
case <-time.After(time.Second):
t.Fatal("Close did not cancel state creation")
}
assertPoolCloseBlocked(t, closed)
finish()
select {
case err := <-acquireDone:
if !errors.Is(err, context.Canceled) {
t.Fatalf("Acquire error = %v, want context.Canceled", err)
}
case <-time.After(time.Second):
t.Fatal("state creation did not finish")
}
select {
case <-closed:
case <-time.After(time.Second):
t.Fatal("Close did not finish after state creation")
}
}
func TestPoolCloseWaitsForCallerContext(t *testing.T) {
pool := newTestPool(t, context.Background(), time.Minute, func(context.Context) (*glua.LState, error) {
return glua.NewState(), nil
})
ctx, cancel := context.WithCancel(context.Background())
t.Cleanup(cancel)
started, done := make(chan context.Context, 1), make(chan error, 1)
go func() {
done <- pool.WithState(ctx, 0, func(L *glua.LState) error {
started <- L.Context()
<-L.Context().Done()
return L.Context().Err()
})
}()
var workCtx context.Context
select {
case workCtx = <-started:
case <-time.After(time.Second):
t.Fatal("WithState did not start")
}
closed := make(chan struct{})
go func() {
pool.Close()
close(closed)
}()
select {
case <-pool.ctx.Done():
case <-time.After(time.Second):
t.Fatal("Close did not cancel the pool context")
}
assertPoolCloseBlocked(t, closed)
if workCtx.Err() != nil || ctx.Err() != nil {
t.Fatal("Close canceled the caller's execution context")
}
cancel()
select {
case err := <-done:
if !errors.Is(err, context.Canceled) {
t.Fatalf("WithState error = %v, want context.Canceled", err)
}
case <-time.After(time.Second):
t.Fatal("WithState did not stop after caller cancellation")
}
select {
case <-closed:
case <-time.After(time.Second):
t.Fatal("Close did not finish after WithState")
}
}
func BenchmarkPoolAcquireRelease(b *testing.B) {
pool := newTestPool(b, context.Background(), time.Second, func(context.Context) (*glua.LState, error) {
return glua.NewState(), nil
})
b.ReportAllocs()
b.ResetTimer()
for i := 0; i < b.N; i++ {
state, err := pool.Acquire(nil)
if err != nil {
b.Fatal(err)
}
pool.Release(state, true)
}
}
-77
View File
@@ -1,77 +0,0 @@
package lua
import (
"bufio"
"context"
"os"
"time"
glua "github.com/yuin/gopher-lua"
"github.com/yuin/gopher-lua/parse"
)
// Program holds immutable bytecode that can be run by independent LStates.
type Program struct {
proto *glua.FunctionProto
}
// LStateFactory returns a fully initialized state or nil and an error.
// Implementations must close partial states on failure; callers own successful states.
type LStateFactory func(context.Context) (*glua.LState, error)
// CompileFile reads and compiles a Lua file once.
func CompileFile(path string) (*Program, error) {
f, err := os.Open(path)
if err != nil {
return nil, err
}
defer f.Close()
chunk, err := parse.Parse(bufio.NewReader(f), path)
if err != nil {
return nil, err
}
proto, err := glua.Compile(chunk, path)
if err != nil {
return nil, err
}
return &Program{proto: proto}, nil
}
// NewState creates a state, runs register, executes the program under ctx, and
// runs validate. It removes the initialization context before returning a state
// owned by the caller.
func (p *Program) NewState(ctx context.Context, register func(*glua.LState), validate func(*glua.LState) error) (*glua.LState, error) {
L := glua.NewState()
valid := false
defer func() {
if !valid {
L.Close()
}
}()
L.SetContext(ctx)
defer L.RemoveContext()
if register != nil {
register(L)
}
L.Push(L.NewFunctionFromProto(p.proto))
// Execute the Lua script's top level.
if err := L.PCall(0, 0, nil); err != nil {
return nil, err
}
if validate != nil {
if err := validate(L); err != nil {
return nil, err
}
}
valid = true
return L, nil
}
// NewStateFactory returns a factory that gives each state an initialization timeout.
func (p *Program) NewStateFactory(initTimeout time.Duration, register func(*glua.LState), validate func(*glua.LState) error) LStateFactory {
return func(ctx context.Context) (*glua.LState, error) {
initCtx, cancel := context.WithTimeout(ctx, initTimeout)
defer cancel()
return p.NewState(initCtx, register, validate)
}
}
-76
View File
@@ -1,76 +0,0 @@
package lua
import (
"context"
"errors"
"os"
"path/filepath"
"testing"
glua "github.com/yuin/gopher-lua"
)
func TestProgramStatesAreIndependent(t *testing.T) {
path := filepath.Join(t.TempDir(), "state.lua")
if err := os.WriteFile(path, []byte("value = (value or 0) + 1"), 0o600); err != nil {
t.Fatal(err)
}
program, err := CompileFile(path)
if err != nil {
t.Fatal(err)
}
first, err := program.NewState(context.Background(), nil, nil)
if err != nil {
t.Fatal(err)
}
defer first.Close()
first.SetGlobal("value", glua.LNumber(42))
second, err := program.NewState(context.Background(), nil, nil)
if err != nil {
t.Fatal(err)
}
defer second.Close()
if got := second.GetGlobal("value"); got != glua.LNumber(1) {
t.Fatalf("second state value = %v, want 1", got)
}
}
func TestProgramInitializationObservesCancellation(t *testing.T) {
path := filepath.Join(t.TempDir(), "loop.lua")
if err := os.WriteFile(path, []byte("while true do end"), 0o600); err != nil {
t.Fatal(err)
}
program, err := CompileFile(path)
if err != nil {
t.Fatal(err)
}
ctx, cancel := context.WithCancel(context.Background())
cancel()
state, err := program.NewState(ctx, nil, nil)
if err == nil || state != nil {
if state != nil {
state.Close()
}
t.Fatalf("NewState with canceled context = %v, %v; want nil state and error", state, err)
}
}
func TestNewStateClosesFailedValidation(t *testing.T) {
path := filepath.Join(t.TempDir(), "state.lua")
if err := os.WriteFile(path, []byte("value = 1"), 0o600); err != nil {
t.Fatal(err)
}
program, err := CompileFile(path)
if err != nil {
t.Fatal(err)
}
wantErr := errors.New("invalid script")
var checked *glua.LState
L, err := program.NewState(context.Background(), nil, func(L *glua.LState) error {
checked = L
return wantErr
})
if L != nil || !errors.Is(err, wantErr) || checked == nil || !checked.IsClosed() {
t.Fatalf("state = %v, error = %v, checked state closed = %t", L, err, checked != nil && checked.IsClosed())
}
}
-95
View File
@@ -1,95 +0,0 @@
package lua
import (
"math"
"github.com/xtls/xray-core/common/errors"
glua "github.com/yuin/gopher-lua"
)
type number interface {
~int | ~int8 | ~int16 | ~int32 | ~int64 |
~uint | ~uint8 | ~uint16 | ~uint32 | ~uint64 | ~uintptr |
~float32 | ~float64
}
// PushNumber converts a Go number to a Lua number and pushes it.
func PushNumber[T number](L *glua.LState, value T) {
L.Push(glua.LNumber(value))
}
// PushString converts a Go string to a Lua string and pushes it.
func PushString(L *glua.LState, value string) {
L.Push(glua.LString(value))
}
// PushNil pushes Lua nil.
func PushNil(L *glua.LState) {
L.Push(glua.LNil)
}
// PushUserData pushes a native Go value without copying it.
func PushUserData(L *glua.LState, value any) {
ud := L.NewUserData()
ud.Value = value
L.Push(ud)
}
// PushError pushes nil or the original Go error as userdata.
func PushError(L *glua.LState, err error) {
if err == nil {
L.Push(glua.LNil)
return
}
PushUserData(L, err)
}
// ReadUserData reads a native Go value of type T without copying it.
// Other Lua values or userdata containing a different type return invalidMessage.
func ReadUserData[T any](value glua.LValue, invalidMessage string) (T, error) {
if ud, ok := value.(*glua.LUserData); ok {
if result, ok := ud.Value.(T); ok {
return result, nil
}
}
var zero T
return zero, errors.New(invalidMessage)
}
// ReadError accepts nil, a native Go error, or a Lua string.
// Native errors retain their identity; other values return invalidMessage.
func ReadError(value glua.LValue, invalidMessage string) error {
if value == glua.LNil {
return nil
}
if ud, ok := value.(*glua.LUserData); ok {
if err, ok := ud.Value.(error); ok {
return err
}
}
if message, ok := value.(glua.LString); ok {
return errors.New(string(message))
}
return errors.New(invalidMessage)
}
// ReadUint32 accepts only integral Lua numbers in the uint32 range.
func ReadUint32(value glua.LValue, invalidMessage string) (uint32, error) {
number, ok := value.(glua.LNumber)
if !ok || number < 0 || number > math.MaxUint32 || math.Trunc(float64(number)) != float64(number) {
return 0, errors.New(invalidMessage)
}
return uint32(number), nil
}
// ReadOptionalString accepts a Lua string or nil, which becomes an empty string.
// It does not coerce other values to strings.
func ReadOptionalString(value glua.LValue, invalidMessage string) (string, error) {
if value == glua.LNil {
return "", nil
}
if result, ok := value.(glua.LString); ok {
return string(result), nil
}
return "", errors.New(invalidMessage)
}
-121
View File
@@ -1,121 +0,0 @@
package lua
import (
"errors"
"math"
"strings"
"testing"
glua "github.com/yuin/gopher-lua"
)
func TestReadUint32(t *testing.T) {
for _, tc := range []struct {
name string
value glua.LValue
want uint32
wantErr bool
}{
{name: "zero", value: glua.LNumber(0)},
{name: "integer", value: glua.LNumber(45), want: 45},
{name: "maximum", value: glua.LNumber(math.MaxUint32), want: math.MaxUint32},
{name: "fraction", value: glua.LNumber(1.5), wantErr: true},
{name: "negative", value: glua.LNumber(-1), wantErr: true},
{name: "overflow", value: glua.LNumber(math.MaxUint32 + 1), wantErr: true},
{name: "NaN", value: glua.LNumber(math.NaN()), wantErr: true},
{name: "positive infinity", value: glua.LNumber(math.Inf(1)), wantErr: true},
{name: "negative infinity", value: glua.LNumber(math.Inf(-1)), wantErr: true},
{name: "nil", value: glua.LNil, wantErr: true},
{name: "numeric string", value: glua.LString("45"), wantErr: true},
{name: "boolean", value: glua.LTrue, wantErr: true},
} {
t.Run(tc.name, func(t *testing.T) {
got, err := ReadUint32(tc.value, "invalid number")
if got != tc.want || (err != nil) != tc.wantErr {
t.Fatalf("ReadUint32() = %d, %v; want %d, error %t", got, err, tc.want, tc.wantErr)
}
if err != nil && !strings.Contains(err.Error(), "invalid number") {
t.Fatalf("error = %v, want invalid number", err)
}
})
}
}
func TestReadOptionalString(t *testing.T) {
for _, tc := range []struct {
name string
value glua.LValue
want string
wantErr bool
}{
{name: "nil", value: glua.LNil},
{name: "empty", value: glua.LString("")},
{name: "string", value: glua.LString("out"), want: "out"},
{name: "number", value: glua.LNumber(1), wantErr: true},
{name: "boolean", value: glua.LFalse, wantErr: true},
} {
t.Run(tc.name, func(t *testing.T) {
got, err := ReadOptionalString(tc.value, "invalid string")
if got != tc.want || (err != nil) != tc.wantErr {
t.Fatalf("ReadOptionalString() = %q, %v; want %q, error %t", got, err, tc.want, tc.wantErr)
}
if err != nil && !strings.Contains(err.Error(), "invalid string") {
t.Fatalf("error = %v, want invalid string", err)
}
})
}
}
func TestUserDataRoundTrip(t *testing.T) {
L := glua.NewState()
defer L.Close()
want := []int{1, 2}
PushUserData(L, want)
if L.GetTop() != 1 {
t.Fatalf("stack top = %d, want 1", L.GetTop())
}
got, err := ReadUserData[[]int](L.Get(-1), "invalid userdata")
if err != nil || len(got) != len(want) || &got[0] != &want[0] {
t.Fatalf("userdata = %v, %v; want original slice", got, err)
}
PushUserData(L, []int(nil))
if got, err := ReadUserData[[]int](L.Get(-1), "invalid userdata"); err != nil || got != nil {
t.Fatalf("nil slice userdata = %v, %v", got, err)
}
for _, value := range []glua.LValue{glua.LNil, glua.LString("1"), L.NewTable(), L.Get(1)} {
if got, err := ReadUserData[int](value, "invalid userdata"); got != 0 || err == nil || !strings.Contains(err.Error(), "invalid userdata") {
t.Fatalf("ReadUserData(%v) = %d, %v; want invalid userdata", value, got, err)
}
}
}
func TestErrorRoundTrip(t *testing.T) {
L := glua.NewState()
defer L.Close()
want := errors.New("upstream failed")
for _, err := range []error{nil, want} {
PushError(L, err)
if L.GetTop() != 1 {
t.Fatalf("stack top = %d, want 1", L.GetTop())
}
if err == nil && L.Get(-1) != glua.LNil {
t.Fatalf("nil error pushed as %v", L.Get(-1))
}
if got := ReadError(L.Get(-1), "invalid error"); got != err {
t.Fatalf("ReadError() = %v, want original error %v", got, err)
}
L.Pop(1)
}
for _, message := range []string{"script failed", ""} {
if err := ReadError(glua.LString(message), "invalid error"); err == nil || !strings.Contains(err.Error(), message) {
t.Fatalf("string error = %v, want %q", err, message)
}
}
wrong := L.NewUserData()
wrong.Value = "not a native error"
for _, value := range []glua.LValue{glua.LTrue, glua.LNumber(1), L.NewTable(), wrong, L.NewUserData()} {
if err := ReadError(value, "invalid error"); err == nil || !strings.Contains(err.Error(), "invalid error") {
t.Fatalf("ReadError(%v) = %v, want invalid error", value, err)
}
}
}
+1 -1
View File
@@ -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
View File
@@ -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()
+1 -1
View File
@@ -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 {
-350
View File
@@ -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
-18
View File
@@ -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)
-356
View File
@@ -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[:])
}
-11
View File
@@ -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 -1
View File
@@ -1,4 +1,4 @@
//go:build !windows && !linux && !android && !darwin
//go:build !windows && !linux && !android
package net
-20
View File
@@ -1,20 +0,0 @@
package net
// PacketConnWrapper wraps a PacketConn into a Conn with a fixed destination address.
type PacketConnWrapper struct {
PacketConn
Dest Addr
}
func (c *PacketConnWrapper) Read(p []byte) (int, error) {
n, _, err := c.PacketConn.ReadFrom(p)
return n, err
}
func (c *PacketConnWrapper) Write(p []byte) (int, error) {
return c.PacketConn.WriteTo(p, c.Dest)
}
func (c *PacketConnWrapper) RemoteAddr() Addr {
return c.Dest
}
-40
View File
@@ -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...)
}
-48
View File
@@ -1,8 +1,6 @@
package platform // import "github.com/xtls/xray-core/common/platform"
import (
"errors"
"fmt"
"os"
"path/filepath"
"strconv"
@@ -92,49 +90,3 @@ func GetConfDirPath() string {
configPath := NewEnvFlag(ConfdirLocation).GetValue(func() string { return "" })
return configPath
}
// ResolveLuaFile finds a local Lua script and returns its absolute path.
// Relative paths: XRAY_LOCATION_CONFDIR > XRAY_LOCATION_CONFIG > working dir > executable dir.
func ResolveLuaFile(path string) (string, error) {
if path == "" {
return "", errors.New("Lua file path is empty")
}
paths := []string{path}
if !filepath.IsAbs(path) {
paths = nil
for _, dir := range []string{
GetConfDirPath(),
NewEnvFlag(ConfigLocation).GetValue(func() string { return "" }),
".",
getExecutableDir(),
} {
if dir != "" {
paths = append(paths, filepath.Join(dir, path))
}
}
}
return resolveFile(paths)
}
func resolveFile(paths []string) (string, error) {
var tried []string
for _, path := range paths {
path, err := filepath.Abs(path)
if err != nil {
return "", fmt.Errorf("failed to resolve file path: %w", err)
}
tried = append(tried, path)
info, err := os.Stat(path)
if errors.Is(err, os.ErrNotExist) {
continue
}
if err != nil {
return "", fmt.Errorf("failed to inspect file %q: %w", path, err)
}
if !info.Mode().IsRegular() {
return "", fmt.Errorf("file is not a regular file: %s", path)
}
return path, nil
}
return "", fmt.Errorf("file not found; tried %q: %w", tried, os.ErrNotExist)
}
-51
View File
@@ -1,7 +1,6 @@
package platform_test
import (
"errors"
"os"
"path/filepath"
"runtime"
@@ -65,53 +64,3 @@ func TestGetAssetLocation(t *testing.T) {
}
}
}
func TestResolveLuaFile(t *testing.T) {
workingDir := t.TempDir()
t.Chdir(workingDir)
executable, err := os.Executable()
common.Must(err)
file, err := os.CreateTemp(filepath.Dir(executable), "lua-*.lua")
common.Must(err)
common.Must(file.Close())
defer os.Remove(file.Name())
name := filepath.Base(file.Name())
paths := []string{
filepath.Join(t.TempDir(), name),
filepath.Join(t.TempDir(), name),
filepath.Join(workingDir, name),
file.Name(),
}
t.Setenv(ConfdirLocation, filepath.Dir(paths[0]))
t.Setenv(ConfigLocation, filepath.Dir(paths[1]))
for _, path := range paths[:3] {
common.Must(os.WriteFile(path, nil, 0o600))
}
if got, err := ResolveLuaFile(paths[2]); err != nil || got != paths[2] {
t.Fatalf("absolute path = %q, %v; want %q", got, err, paths[2])
}
for i, want := range paths {
if i == 2 {
t.Setenv(ConfdirLocation, "")
t.Setenv(ConfigLocation, "")
}
if got, err := ResolveLuaFile(name); err != nil || got != want {
t.Fatalf("resolved path = %q, %v; want %q", got, err, want)
}
common.Must(os.Remove(want))
}
if _, err := ResolveLuaFile(name); !errors.Is(err, os.ErrNotExist) {
t.Fatalf("missing file error = %v", err)
}
t.Setenv(ConfdirLocation, filepath.Dir(paths[0]))
t.Setenv(ConfigLocation, filepath.Dir(paths[1]))
common.Must(os.Mkdir(paths[0], 0o700))
common.Must(os.WriteFile(paths[1], nil, 0o600))
for _, path := range []string{"", name, filepath.Join(t.TempDir(), name)} {
if _, err := ResolveLuaFile(path); err == nil {
t.Fatalf("accepted invalid path %q", path)
}
}
}
+34 -25
View File
@@ -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, &timestamp); 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")
}
})
}
}
+10 -2
View File
@@ -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 (
+2
View File
@@ -11,6 +11,8 @@ enum SecurityType {
AUTO = 2;
AES128_GCM = 3;
CHACHA20_POLY1305 = 4;
NONE = 5; // [DEPRECATED 2023-06]
ZERO = 6;
}
message SecurityConfig {

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