Compare commits

..
Author SHA1 Message Date
patterniha 813f3810ef Direct/Freedom outbound: Skip domain resolution and finalRules with dialerProxy
https://github.com/XTLS/Xray-core/pull/6058#issuecomment-5578800461

When using dialerProxy, the target domain should not be resolved to an IP. Also, with dialerProxy freedom is no longer the final outbound, which makes finalRules invalid, so finalRules should not be applied at all.
2026-09-08 20:28:22 +03:30
304 changed files with 6287 additions and 32003 deletions
+3 -1
View File
@@ -67,7 +67,9 @@ jobs:
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
-3
View File
@@ -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")
+1 -1
View File
@@ -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
+2 -2
View File
@@ -87,7 +87,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:
@@ -217,7 +217,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
+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)
+2 -2
View File
@@ -115,7 +115,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 +145,7 @@ func (br *BalancingRule) Build(ohm outbound.Manager, dispatcher routing.Dispatch
}
s, ok := i.(*StrategyLeastLoadConfig)
if !ok {
return nil, errors.New("not a StrategyLeastLoadConfig")
return nil, errors.New("not a StrategyLeastLoadConfig").AtError()
}
leastLoadStrategy := NewLeastLoadStrategy(s)
return &Balancer{
+4 -14
View File
@@ -587,10 +587,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,13 +642,6 @@ 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 = "" +
@@ -708,12 +699,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" +
-2
View File
@@ -110,6 +110,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)
-17
View File
@@ -20,8 +20,6 @@ import (
type Router struct {
domainStrategy Config_DomainStrategy
rules atomic.Pointer[[]*Rule]
scriptPath string
script *scriptEngine
balancers atomic.Pointer[map[string]*Balancer]
dns dns.Client
@@ -42,7 +40,6 @@ 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
@@ -55,10 +52,6 @@ func (r *Router) Init(ctx context.Context, config *Config, d dns.Client, ohm out
// 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 {
@@ -228,13 +221,6 @@ 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
}
@@ -249,9 +235,6 @@ 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())
-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")
}
}
+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)
+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)
}
+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 {
-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
}
-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)
}
}
}
+1 -1
View File
@@ -7,7 +7,7 @@ import (
func (u *User) GetTypedAccount() (Account, error) {
if u.GetAccount() == nil {
return nil, errors.New("Account is missing")
return nil, errors.New("Account is missing").AtWarning()
}
rawAccount, err := u.Account.GetInstance()
+53
View File
@@ -0,0 +1,53 @@
package singbridge
import (
M "github.com/sagernet/sing/common/metadata"
N "github.com/sagernet/sing/common/network"
"github.com/xtls/xray-core/common/errors"
"github.com/xtls/xray-core/common/net"
)
func ToNetwork(network string) net.Network {
switch N.NetworkName(network) {
case N.NetworkTCP:
return net.Network_TCP
case N.NetworkUDP:
return net.Network_UDP
default:
return net.Network_Unknown
}
}
func ToDestination(socksaddr M.Socksaddr, network net.Network) (net.Destination, error) {
// IsFqdn() implicitly checks if the domain name is valid
if socksaddr.IsFqdn() {
return net.Destination{
Network: network,
Address: net.DomainAddress(socksaddr.Fqdn),
Port: net.Port(socksaddr.Port),
}, nil
}
// IsIP() implicitly checks if the IP address is valid
if socksaddr.IsIP() {
return net.Destination{
Network: network,
Address: net.IPAddress(socksaddr.Addr.AsSlice()),
Port: net.Port(socksaddr.Port),
}, nil
}
return net.Destination{}, errors.New("invalid socks address: ", socksaddr)
}
func ToSocksaddr(destination net.Destination) M.Socksaddr {
var addr M.Socksaddr
switch destination.Address.Family() {
case net.AddressFamilyDomain:
addr.Fqdn = destination.Address.Domain()
default:
addr.Addr = M.AddrFromIP(destination.Address.IP())
}
addr.Port = uint16(destination.Port)
return addr
}
+72
View File
@@ -0,0 +1,72 @@
package singbridge
import (
"context"
"os"
M "github.com/sagernet/sing/common/metadata"
N "github.com/sagernet/sing/common/network"
"github.com/xtls/xray-core/common/net"
"github.com/xtls/xray-core/common/net/cnc"
"github.com/xtls/xray-core/common/session"
"github.com/xtls/xray-core/proxy"
"github.com/xtls/xray-core/transport"
"github.com/xtls/xray-core/transport/internet"
"github.com/xtls/xray-core/transport/pipe"
)
var _ N.Dialer = (*XrayDialer)(nil)
type XrayDialer struct {
internet.Dialer
}
func NewDialer(dialer internet.Dialer) *XrayDialer {
return &XrayDialer{dialer}
}
func (d *XrayDialer) DialContext(ctx context.Context, network string, destination M.Socksaddr) (net.Conn, error) {
dest, err := ToDestination(destination, ToNetwork(network))
if err != nil {
return nil, err
}
return d.Dialer.Dial(ctx, dest)
}
func (d *XrayDialer) ListenPacket(ctx context.Context, destination M.Socksaddr) (net.PacketConn, error) {
return nil, os.ErrInvalid
}
type XrayOutboundDialer struct {
outbound proxy.Outbound
dialer internet.Dialer
}
func NewOutboundDialer(outbound proxy.Outbound, dialer internet.Dialer) *XrayOutboundDialer {
return &XrayOutboundDialer{outbound, dialer}
}
func (d *XrayOutboundDialer) DialContext(ctx context.Context, network string, destination M.Socksaddr) (net.Conn, error) {
dest, err := ToDestination(destination, ToNetwork(network))
if err != nil {
return nil, err
}
outbounds := session.OutboundsFromContext(ctx)
if len(outbounds) == 0 {
outbounds = []*session.Outbound{{}}
ctx = session.ContextWithOutbounds(ctx, outbounds)
}
ob := outbounds[len(outbounds)-1]
ob.Target = dest
opts := []pipe.Option{pipe.WithSizeLimit(64 * 1024)}
uplinkReader, uplinkWriter := pipe.New(opts...)
downlinkReader, downlinkWriter := pipe.New(opts...)
conn := cnc.NewConnection(cnc.ConnectionInputMulti(downlinkWriter), cnc.ConnectionOutputMulti(uplinkReader))
go d.outbound.Process(ctx, &transport.Link{Reader: downlinkReader, Writer: uplinkWriter}, d.dialer)
return conn, nil
}
func (d *XrayOutboundDialer) ListenPacket(ctx context.Context, destination M.Socksaddr) (net.PacketConn, error) {
return nil, os.ErrInvalid
}
+10
View File
@@ -0,0 +1,10 @@
package singbridge
import E "github.com/sagernet/sing/common/exceptions"
func ReturnError(err error) error {
if E.IsClosedOrCanceled(err) {
return nil
}
return err
}
+58
View File
@@ -0,0 +1,58 @@
package singbridge
import (
"context"
"io"
M "github.com/sagernet/sing/common/metadata"
N "github.com/sagernet/sing/common/network"
"github.com/xtls/xray-core/common/buf"
"github.com/xtls/xray-core/common/errors"
"github.com/xtls/xray-core/common/net"
"github.com/xtls/xray-core/features/routing"
"github.com/xtls/xray-core/transport"
)
var (
_ N.TCPConnectionHandler = (*Dispatcher)(nil)
_ N.UDPConnectionHandler = (*Dispatcher)(nil)
)
type Dispatcher struct {
upstream routing.Dispatcher
newErrorFunc func(values ...any) *errors.Error
}
func NewDispatcher(dispatcher routing.Dispatcher, newErrorFunc func(values ...any) *errors.Error) *Dispatcher {
return &Dispatcher{
upstream: dispatcher,
newErrorFunc: newErrorFunc,
}
}
func (d *Dispatcher) NewConnection(ctx context.Context, conn net.Conn, metadata M.Metadata) error {
dest, err := ToDestination(metadata.Destination, net.Network_TCP)
if err != nil {
return err
}
xConn := NewConn(conn)
return d.upstream.DispatchLink(ctx, dest, &transport.Link{
Reader: xConn,
Writer: xConn,
})
}
func (d *Dispatcher) NewPacketConnection(ctx context.Context, conn N.PacketConn, metadata M.Metadata) error {
dest, err := ToDestination(metadata.Destination, net.Network_UDP)
if err != nil {
return err
}
return d.upstream.DispatchLink(ctx, dest, &transport.Link{
Reader: buf.NewPacketReader(conn.(io.Reader)),
Writer: buf.NewWriter(conn.(io.Writer)),
})
}
func (d *Dispatcher) NewError(ctx context.Context, err error) {
errors.LogInfo(ctx, err.Error())
}
+70
View File
@@ -0,0 +1,70 @@
package singbridge
import (
"context"
"github.com/sagernet/sing/common/logger"
"github.com/xtls/xray-core/common/errors"
)
var _ logger.ContextLogger = (*XrayLogger)(nil)
type XrayLogger struct {
newError func(values ...any) *errors.Error
}
func NewLogger(newErrorFunc func(values ...any) *errors.Error) *XrayLogger {
return &XrayLogger{
newErrorFunc,
}
}
func (l *XrayLogger) Trace(args ...any) {
}
func (l *XrayLogger) Debug(args ...any) {
errors.LogDebug(context.Background(), args...)
}
func (l *XrayLogger) Info(args ...any) {
errors.LogInfo(context.Background(), args...)
}
func (l *XrayLogger) Warn(args ...any) {
errors.LogWarning(context.Background(), args...)
}
func (l *XrayLogger) Error(args ...any) {
errors.LogError(context.Background(), args...)
}
func (l *XrayLogger) Fatal(args ...any) {
}
func (l *XrayLogger) Panic(args ...any) {
}
func (l *XrayLogger) TraceContext(ctx context.Context, args ...any) {
}
func (l *XrayLogger) DebugContext(ctx context.Context, args ...any) {
errors.LogDebug(ctx, args...)
}
func (l *XrayLogger) InfoContext(ctx context.Context, args ...any) {
errors.LogInfo(ctx, args...)
}
func (l *XrayLogger) WarnContext(ctx context.Context, args ...any) {
errors.LogWarning(ctx, args...)
}
func (l *XrayLogger) ErrorContext(ctx context.Context, args ...any) {
errors.LogError(ctx, args...)
}
func (l *XrayLogger) FatalContext(ctx context.Context, args ...any) {
}
func (l *XrayLogger) PanicContext(ctx context.Context, args ...any) {
}
+107
View File
@@ -0,0 +1,107 @@
package singbridge
import (
"context"
"time"
B "github.com/sagernet/sing/common/buf"
"github.com/sagernet/sing/common/bufio"
M "github.com/sagernet/sing/common/metadata"
"github.com/xtls/xray-core/common"
"github.com/xtls/xray-core/common/buf"
"github.com/xtls/xray-core/common/net"
"github.com/xtls/xray-core/common/signal"
"github.com/xtls/xray-core/transport"
)
func CopyPacketConn(ctx context.Context, inboundConn net.Conn, link *transport.Link, destination net.Destination, serverConn net.PacketConn) error {
cancel := func() {
common.Interrupt(link.Reader)
common.Interrupt(serverConn)
}
conn := &PacketConnWrapper{
Reader: link.Reader,
Writer: link.Writer,
Dest: destination,
Conn: inboundConn,
T: signal.CancelAfterInactivity(ctx, cancel, 300*time.Second),
}
return ReturnError(bufio.CopyPacketConn(ctx, conn, bufio.NewPacketConn(serverConn)))
}
type PacketConnWrapper struct {
buf.Reader
buf.Writer
net.Conn
Dest net.Destination
cached buf.MultiBuffer
// A simple patch to avoid goroutine leak since sing infra cannot awake read block by write err
T *signal.ActivityTimer
}
func (w *PacketConnWrapper) ReadPacket(buffer *B.Buffer) (addr M.Socksaddr, err error) {
w.T.Update()
defer func() {
if err != nil {
// uplinkonly
w.T.SetTimeout(2 * time.Second)
}
}()
if w.cached != nil {
mb, bb := buf.SplitFirst(w.cached)
if bb == nil {
w.cached = nil
} else {
buffer.Write(bb.Bytes())
w.cached = mb
var destination net.Destination
if bb.UDP != nil {
destination = *bb.UDP
} else {
destination = w.Dest
}
bb.Release()
return ToSocksaddr(destination), nil
}
}
mb, err := w.ReadMultiBuffer()
nb, bb := buf.SplitFirst(mb)
if bb == nil {
return M.Socksaddr{}, nil
} else {
buffer.Write(bb.Bytes())
w.cached = nb
var destination net.Destination
if bb.UDP != nil {
destination = *bb.UDP
} else {
destination = w.Dest
}
bb.Release()
return ToSocksaddr(destination), nil
}
}
func (w *PacketConnWrapper) WritePacket(buffer *B.Buffer, destination M.Socksaddr) (err error) {
w.T.Update()
defer func() {
if err != nil {
// downlinkonly
w.T.SetTimeout(5 * time.Second)
}
}()
endpoint, err := ToDestination(destination, net.Network_UDP)
if err != nil {
return err
}
vBuf := buf.New()
vBuf.Write(buffer.Bytes())
vBuf.UDP = &endpoint
return w.WriteMultiBuffer(buf.MultiBuffer{vBuf})
}
func (w *PacketConnWrapper) Close() error {
buf.ReleaseMulti(w.cached)
return nil
}
+81
View File
@@ -0,0 +1,81 @@
package singbridge
import (
"context"
"io"
"net"
"time"
"github.com/sagernet/sing/common/bufio"
"github.com/xtls/xray-core/common"
"github.com/xtls/xray-core/common/buf"
"github.com/xtls/xray-core/common/signal"
"github.com/xtls/xray-core/transport"
)
func CopyConn(ctx context.Context, inboundConn net.Conn, link *transport.Link, serverConn net.Conn) error {
conn := &PipeConnWrapper{
W: link.Writer,
Conn: inboundConn,
}
if ir, ok := link.Reader.(io.Reader); ok {
conn.R = ir
} else {
conn.R = &buf.BufferedReader{Reader: link.Reader}
}
cancel := func() {
common.Interrupt(link.Reader)
common.Interrupt(serverConn)
}
conn.T = signal.CancelAfterInactivity(ctx, cancel, 300*time.Second)
return ReturnError(bufio.CopyConn(ctx, conn, serverConn))
}
type PipeConnWrapper struct {
R io.Reader
W buf.Writer
net.Conn
// A simple patch to avoid goroutine leak since sing infra cannot awake read block by write err
T *signal.ActivityTimer
}
func (w *PipeConnWrapper) Close() error {
return nil
}
func (w *PipeConnWrapper) Read(b []byte) (n int, err error) {
w.T.Update()
n, err = w.R.Read(b)
if err != nil {
// uplinkonly
w.T.SetTimeout(2 * time.Second)
}
return
}
func (w *PipeConnWrapper) Write(p []byte) (n int, err error) {
w.T.Update()
n = len(p)
var mb buf.MultiBuffer
pLen := len(p)
for pLen > 0 {
buffer := buf.New()
if pLen > buf.Size {
_, err = buffer.Write(p[:buf.Size])
p = p[buf.Size:]
} else {
buffer.Write(p)
}
pLen -= int(buffer.Len())
mb = append(mb, buffer)
}
err = w.W.WriteMultiBuffer(mb)
if err != nil {
n = 0
buf.ReleaseMulti(mb)
// downlinkonly
w.T.SetTimeout(5 * time.Second)
}
return
}
+66
View File
@@ -0,0 +1,66 @@
package singbridge
import (
"time"
"github.com/sagernet/sing/common"
"github.com/sagernet/sing/common/bufio"
N "github.com/sagernet/sing/common/network"
"github.com/xtls/xray-core/common/buf"
"github.com/xtls/xray-core/common/net"
)
var (
_ buf.Reader = (*Conn)(nil)
_ buf.TimeoutReader = (*Conn)(nil)
_ buf.Writer = (*Conn)(nil)
)
type Conn struct {
net.Conn
writer N.VectorisedWriter
}
func NewConn(conn net.Conn) *Conn {
writer, _ := bufio.CreateVectorisedWriter(conn)
return &Conn{
Conn: conn,
writer: writer,
}
}
func (c *Conn) ReadMultiBuffer() (buf.MultiBuffer, error) {
buffer, err := buf.ReadBuffer(c.Conn)
if err != nil {
return nil, err
}
return buf.MultiBuffer{buffer}, nil
}
func (c *Conn) ReadMultiBufferTimeout(duration time.Duration) (buf.MultiBuffer, error) {
err := c.SetReadDeadline(time.Now().Add(duration))
if err != nil {
return nil, err
}
defer c.SetReadDeadline(time.Time{})
return c.ReadMultiBuffer()
}
func (c *Conn) WriteMultiBuffer(bufferList buf.MultiBuffer) error {
defer buf.ReleaseMulti(bufferList)
if c.writer != nil {
bytesList := make([][]byte, len(bufferList))
for i, buffer := range bufferList {
bytesList[i] = buffer.Bytes()
}
return common.Error(bufio.WriteVectorised(c.writer, bytesList))
}
// Since this conn is only used by tun, we don't force buffer writes to merge.
for _, buffer := range bufferList {
_, err := c.Conn.Write(buffer.Bytes())
if err != nil {
return err
}
}
return nil
}
+2 -2
View File
@@ -16,7 +16,7 @@ var typeCreatorRegistry = make(map[reflect.Type]ConfigCreator)
func RegisterConfig(config interface{}, configCreator ConfigCreator) error {
configType := reflect.TypeOf(config)
if _, found := typeCreatorRegistry[configType]; found {
return errors.New(configType.Name() + " is already registered")
return errors.New(configType.Name() + " is already registered").AtError()
}
typeCreatorRegistry[configType] = configCreator
return nil
@@ -27,7 +27,7 @@ func CreateObject(ctx context.Context, config interface{}) (interface{}, error)
configType := reflect.TypeOf(config)
creator, found := typeCreatorRegistry[configType]
if !found {
return nil, errors.New(configType.String() + " is not registered")
return nil, errors.New(configType.String() + " is not registered").AtError()
}
return creator(ctx, config)
}
+4 -4
View File
@@ -125,7 +125,7 @@ func LoadConfig(formatName string, input interface{}) (*Config, error) {
}
if f == "" {
return nil, errors.New("Failed to get format of ", file)
return nil, errors.New("Failed to get format of ", file).AtWarning()
}
if f == "protobuf" {
@@ -142,7 +142,7 @@ func LoadConfig(formatName string, input interface{}) (*Config, error) {
if len(v) == 1 {
return configLoaderByName["protobuf"].Loader(v)
} else {
return nil, errors.New("Only one protobuf config file is allowed")
return nil, errors.New("Only one protobuf config file is allowed").AtWarning()
}
}
@@ -152,11 +152,11 @@ func LoadConfig(formatName string, input interface{}) (*Config, error) {
if f, found := configLoaderByName[formatName]; found {
return f.Loader(v)
} else {
return nil, errors.New("Unable to load config in", formatName)
return nil, errors.New("Unable to load config in", formatName).AtWarning()
}
}
return nil, errors.New("Unable to load config")
return nil, errors.New("Unable to load config").AtWarning()
}
func loadProtobufConfig(data []byte) (*Config, error) {
+1 -1
View File
@@ -20,7 +20,7 @@ import (
var (
Version_x byte = 26
Version_y byte = 9
Version_z byte = 30
Version_z byte = 8
)
var (
+1 -1
View File
@@ -13,7 +13,7 @@ type FakeDNSEngine interface {
var (
FakeIPv4Pool = "198.18.0.0/15"
FakeIPv6Pool = "2001:2::/48"
FakeIPv6Pool = "fc00::/18"
)
type FakeDNSEngineRev0 interface {
-3
View File
@@ -97,9 +97,6 @@ func New() *Client {
r := &net.Resolver{
PreferGo: true,
Dial: func(ctx context.Context, network, address string) (net.Conn, error) {
if internet.IsSkippedDNSServer(address) {
return nil, errors.New("skipped DNS server ", address)
}
return d.DialContext(ctx, network, address)
},
}
-23
View File
@@ -1,23 +0,0 @@
package localdns
import (
"context"
"net/netip"
"testing"
"github.com/xtls/xray-core/transport/internet"
)
func TestSkippedDNSServers(t *testing.T) {
internet.SkipDNSServers([]netip.Addr{netip.MustParseAddr("203.0.113.53")})
t.Cleanup(func() { internet.SkipDNSServers(nil) })
c := New()
if _, err := c.r.Dial(context.Background(), "udp", "203.0.113.53:53"); err == nil {
t.Error("a skipped DNS server was dialed")
}
conn, err := c.r.Dial(context.Background(), "udp", "127.0.0.1:53")
if err != nil {
t.Fatal(err)
}
conn.Close()
}
+11 -12
View File
@@ -18,26 +18,25 @@ require (
github.com/pires/go-proxyproto v0.15.0
github.com/refraction-networking/utls v1.8.3-0.20260301010127-aa6edf4b11af
github.com/robfig/cron/v3 v3.0.1
github.com/sagernet/sing v0.5.1
github.com/sagernet/sing-shadowsocks v0.2.7
github.com/stretchr/testify v1.12.1
github.com/vishvananda/netlink v1.3.1
github.com/xtls/reality v0.0.0-20260908062103-8cdf7bf9c7f0
github.com/yuin/gopher-lua v1.1.2
go4.org/netipx v0.0.0-20231129151722-fdeea329fbba
golang.org/x/crypto v0.57.0
golang.org/x/crypto v0.55.0
golang.org/x/exp v0.0.0-20240506185415-9bf2ced13842
golang.org/x/net v0.59.0
golang.org/x/sync v0.23.0
golang.org/x/sys v0.48.0
golang.org/x/net v0.58.0
golang.org/x/sync v0.22.0
golang.org/x/sys v0.47.0
golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2
golang.zx2c4.com/wireguard v0.0.0-20250521234502-f333402bd9cb
golang.zx2c4.com/wireguard/windows v1.1.1
google.golang.org/grpc v1.84.0
golang.zx2c4.com/wireguard/windows v1.0.1
google.golang.org/grpc v1.83.2
google.golang.org/protobuf v1.36.12
gvisor.dev/gvisor v0.0.0-20260122175437-89a5d21be8f0
h12.io/socks v1.0.3
layeh.com/gopher-luar v1.0.11
lukechampine.com/blake3 v1.4.1
mvdan.cc/gofumpt v0.12.0
)
require (
@@ -49,6 +48,7 @@ require (
github.com/juju/ratelimit v1.0.2 // indirect
github.com/klauspost/compress v1.17.4 // indirect
github.com/koron/go-ssdp v0.0.4 // indirect
github.com/kr/text v0.2.0 // indirect
github.com/libp2p/go-netroute v0.2.1 // indirect
github.com/pion/dtls/v3 v3.1.5 // indirect
github.com/pion/logging v0.2.4 // indirect
@@ -57,9 +57,8 @@ require (
github.com/vishvananda/netns v0.0.5 // indirect
github.com/wlynxg/anet v0.0.5 // indirect
go.yaml.in/yaml/v3 v3.0.5 // indirect
golang.org/x/text v0.42.0 // indirect
golang.org/x/text v0.41.0 // indirect
golang.org/x/time v0.14.0 // indirect
golang.org/x/tools v0.49.0 // indirect
google.golang.org/genproto/googleapis/rpc v0.0.0-20260706201446-f0a921348800 // indirect
google.golang.org/genproto/googleapis/rpc v0.0.0-20260526163538-3dc84a4a5aaa // indirect
gopkg.in/yaml.v2 v2.4.0 // indirect
)
+41 -33
View File
@@ -2,15 +2,17 @@ github.com/andybalholm/brotli v1.0.6 h1:Yf9fFpf49Zrxb9NlQaluyE92/+X7UVHlhMNJN2sx
github.com/andybalholm/brotli v1.0.6/go.mod h1:fO7iG3H7G2nSZ7m0zPUDn85XEX2GTukHGRSepvi9Eig=
github.com/apernet/quic-go v0.61.1-0.20260806010916-184d081eef3e h1:5mgtR5gwIgBKMiGI1QdXldZZ+SNor06Nbu1wCBulQBg=
github.com/apernet/quic-go v0.61.1-0.20260806010916-184d081eef3e/go.mod h1:x7qxEvX6MCVtDuBKHj3E+88+BtrbEMuAL5qGUKItjW8=
github.com/chzyer/logex v1.1.10/go.mod h1:+Ywpsq7O8HXn0nuIou7OrIPyXbp3wmkHB+jjWRnGsAI=
github.com/chzyer/readline v0.0.0-20180603132655-2972be24d48e/go.mod h1:nSuG5e5PlCu98SY8svDHJxuZscDgtXS6KTTbou5AhLI=
github.com/chzyer/test v0.0.0-20180213035817-a1ea475d72b1/go.mod h1:Q3SI9o4m/ZMnBNeIyt5eFwwo7qiLfzFZmjNmxjkiQlU=
github.com/cespare/xxhash/v2 v2.3.0 h1:UL815xU9SqsFlibzuggzjXhog7bL6oX9BbNZnL2UFvs=
github.com/cespare/xxhash/v2 v2.3.0/go.mod h1:VGX0DQ3Q6kWi7AoAeZDth3/j3BFtOZR5XLFGgcrjCOs=
github.com/cloudflare/circl v1.6.5 h1:O64F26HEqNhznd/hrC5KZXVKYuKM2rx4deZDTc4ihQA=
github.com/cloudflare/circl v1.6.5/go.mod h1:h5LNyxAc5nTue9DS5jT+48en2PSDYt3zdGnz5OstK6c=
github.com/creack/pty v1.1.9/go.mod h1:oKZEueFk5CKHvIhNR5MUki03XCEU+Q6VDXinZuGJ33E=
github.com/ghodss/yaml v1.0.1-0.20220118164431-d8423dcdf344 h1:Arcl6UOIS/kgO2nW3A65HN+7CMjSDP/gofXL4CZt1V4=
github.com/ghodss/yaml v1.0.1-0.20220118164431-d8423dcdf344/go.mod h1:GIjDIg/heH5DOkXY3YJ/wNhfHsQHoXGjl8G8amsYQ1I=
github.com/go-quicktest/qt v1.102.0 h1:HSQxCeh5YZH3EL3W39ixjtyaEhcWSXQHtHnMBzSs474=
github.com/go-quicktest/qt v1.102.0/go.mod h1:p4lGIVX+8Wa6ZPNDvqcxq36XpUDLh42FLetFU7odllI=
github.com/go-logr/logr v1.4.3 h1:CjnDlHq8ikf6E492q6eKboGOC0T8CDaOvkHCIg8idEI=
github.com/go-logr/logr v1.4.3/go.mod h1:9T104GzyrTigFIr8wt5mBrctHMim0Nb2HLGrmQ40KvY=
github.com/go-logr/stdr v1.2.2 h1:hSWxHoqTgW2S2qGc0LTAI563KZ5YKYRhT3MFKZMbjag=
github.com/go-logr/stdr v1.2.2/go.mod h1:mMo/vtBO5dYbehREoey6XUKy/eSumjCCveDpRre4VKE=
github.com/golang/mock v1.7.0-rc.1 h1:YojYx61/OLFsiv6Rw1Z96LpldJIy31o+UHmwAUMJ6/U=
github.com/golang/mock v1.7.0-rc.1/go.mod h1:s42URUywIqd+OcERslBJvOjepvNymP31m3q8d/GkuRs=
github.com/golang/protobuf v1.5.4 h1:i7eJL8qZTpSEXOPTxNKhASYpMn+8e5Q6AdndVa1dWek=
@@ -71,8 +73,12 @@ github.com/refraction-networking/utls v1.8.3-0.20260301010127-aa6edf4b11af h1:er
github.com/refraction-networking/utls v1.8.3-0.20260301010127-aa6edf4b11af/go.mod h1:jkSOEkLqn+S/jtpEHPOsVv/4V4EVnelwbMQl4vCWXAM=
github.com/robfig/cron/v3 v3.0.1 h1:WdRxkvbJztn8LMz/QEvLN5sBU+xKpSqwwUO1Pjr4qDs=
github.com/robfig/cron/v3 v3.0.1/go.mod h1:eQICP3HwyT7UooqI/z+Ov+PtYAWygg1TEWWzGIFLtro=
github.com/rogpeppe/go-internal v1.16.0 h1:O9DK+vNMDVGLr2BeZqmpLeMjiMNkuXfcqntWbZV6S5g=
github.com/rogpeppe/go-internal v1.16.0/go.mod h1:DrUVZyrJU+txYW5/1kwtXQSMFio52ZOxX7yM1VHvnxs=
github.com/rogpeppe/go-internal v1.10.0 h1:TMyTOH3F/DB16zRVcYyreMH6GnZZrwQVAoYjRBZyWFQ=
github.com/rogpeppe/go-internal v1.10.0/go.mod h1:UQnix2H7Ngw/k4C5ijL5+65zddjncjaFoBhdsK/akog=
github.com/sagernet/sing v0.5.1 h1:mhL/MZVq0TjuvHcpYcFtmSD1BFOxZ/+8ofbNZcg1k1Y=
github.com/sagernet/sing v0.5.1/go.mod h1:ARkL0gM13/Iv5VCZmci/NuoOlePoIsW0m7BWfln/Hak=
github.com/sagernet/sing-shadowsocks v0.2.7 h1:zaopR1tbHEw5Nk6FAkM05wCslV6ahVegEZaKMv9ipx8=
github.com/sagernet/sing-shadowsocks v0.2.7/go.mod h1:0rIKJZBR65Qi0zwdKezt4s57y/Tl1ofkaq6NlkzVuyE=
github.com/stretchr/testify v1.12.1 h1:EuwCh5fleGS7H32xRwO3wRGT7DxrDhLAT6FF8MpWDWE=
github.com/stretchr/testify v1.12.1/go.mod h1:MDEgiDPPsNp5cuIrHPPCyornHKgEVbtFUmoNlxoYthg=
github.com/vishvananda/netlink v1.3.1 h1:3AEMt62VKqz90r0tmNhog0r/PpWKmrEShJU0wJW6bV0=
@@ -84,9 +90,18 @@ github.com/wlynxg/anet v0.0.5/go.mod h1:eay5PRQr7fIVAMbTbchTnO9gG65Hg/uYGdc7mguH
github.com/xtls/reality v0.0.0-20260908062103-8cdf7bf9c7f0 h1:rb+fKQFhz+5I2PPuQsNYxI5mUU840XWYtRF0ZBjvkws=
github.com/xtls/reality v0.0.0-20260908062103-8cdf7bf9c7f0/go.mod h1:DsJblcWDGt76+FVqBVwbwRhxyyNJsGV48gJLch0OOWI=
github.com/yuin/goldmark v1.4.1/go.mod h1:mwnBkeHKe2W/ZEtQ+71ViKU8L12m81fl3OWwC1Zlc8k=
github.com/yuin/gopher-lua v0.0.0-20190206043414-8bfc7677f583/go.mod h1:gqRgreBUhTSL0GeU64rtZ3Uq3wtjOa/TB2YfrtkCbVQ=
github.com/yuin/gopher-lua v1.1.2 h1:yF/FjE3hD65tBbt0VXLE13HWS9h34fdzJmrWRXwobGA=
github.com/yuin/gopher-lua v1.1.2/go.mod h1:7aRmXIWl37SqRf0koeyylBEzJ+aPt8A+mmkQ4f1ntR8=
go.opentelemetry.io/auto/sdk v1.2.1 h1:jXsnJ4Lmnqd11kwkBV2LgLoFMZKizbCi5fNZ/ipaZ64=
go.opentelemetry.io/auto/sdk v1.2.1/go.mod h1:KRTj+aOaElaLi+wW1kO/DZRXwkF4C5xPbEe3ZiIhN7Y=
go.opentelemetry.io/otel v1.44.0 h1:JjwHmHpA4iZ3wBxluu2fbbE7j4kqlE8jXyAyPXH7HqU=
go.opentelemetry.io/otel v1.44.0/go.mod h1:BMgjTHL9WPRlRjL2oZCBTL4whCGtXch2H4BhOPIAyYc=
go.opentelemetry.io/otel/metric v1.44.0 h1:1w0gILTcHdr3YI+ixLyjemwrVnsMURbTZFrSYCdDdmc=
go.opentelemetry.io/otel/metric v1.44.0/go.mod h1:8O7hanEPBNgEMmybD3s2VBKcgWOCsA6tzHBPODAiquo=
go.opentelemetry.io/otel/sdk v1.44.0 h1:nHYwb9lK+fJPU/dnT6s7W7Z8itMWyqrnVfbheVYrZ58=
go.opentelemetry.io/otel/sdk v1.44.0/go.mod h1:Osuydd3Se74nqjAKxid74N5eC+jfEqfTegHRnq58oK0=
go.opentelemetry.io/otel/sdk/metric v1.44.0 h1:3LlKgI+VjbVsjNRFZJZAJ30WjXC5VkNRks6si09iEfI=
go.opentelemetry.io/otel/sdk/metric v1.44.0/go.mod h1:5B5pMARnXxKhltooO4xUuCBorl65a4EpnTalObqOigA=
go.opentelemetry.io/otel/trace v1.44.0 h1:jxF5CsGYCe74MCRx2X4g7WsY/VBKRqqpNvXlX/6gtIk=
go.opentelemetry.io/otel/trace v1.44.0/go.mod h1:oLl1jrMQAVo6v3GAggN+1VH9VIz9iUSvW53sW1Q8PIE=
go.uber.org/mock v0.5.2 h1:LbtPTcP8A5k9WPXj54PPPbjcI4Y6lhyOZXn+VS7wNko=
go.uber.org/mock v0.5.2/go.mod h1:wLlUxC2vVTPTaE3UD51E0BGOAElKrILxhVSDYQLld5o=
go.yaml.in/yaml/v3 v3.0.5 h1:N6y/pJk8buWs9NY5ERU2HSMfm+IuD/OtfdAnq6kESPw=
@@ -95,8 +110,8 @@ go4.org/netipx v0.0.0-20231129151722-fdeea329fbba h1:0b9z3AuHCjxk0x/opv64kcgZLBs
go4.org/netipx v0.0.0-20231129151722-fdeea329fbba/go.mod h1:PLyyIXexvUFg3Owu6p/WfdlivPbZJsZdgWZlrGope/Y=
golang.org/x/crypto v0.0.0-20190308221718-c2843e01d9a2/go.mod h1:djNgcEr1/C05ACkg1iLfiJU5Ep61QUkGW8qpdssI0+w=
golang.org/x/crypto v0.0.0-20191011191535-87dc89f01550/go.mod h1:yigFU9vqHzYiE8UmvKecakEJjdnWj3jj499lnFckfCI=
golang.org/x/crypto v0.57.0 h1:3ZVCjf8Ggz7zneR/EHRVx68Ctf+2pmIMP2UFhh9cC6M=
golang.org/x/crypto v0.57.0/go.mod h1:Fdz0i5U6CoizGwLda9DttjSk6qlZo25zYNtR+ycvuZA=
golang.org/x/crypto v0.55.0 h1:+KWHjbgOaAQ66dh/YlkZKHlz9ZUlq61AFirAR9ntP8M=
golang.org/x/crypto v0.55.0/go.mod h1:uq0V9dE/fzQuJtbnL+2EhWOE63vo164FY8xqEnV9xis=
golang.org/x/exp v0.0.0-20240506185415-9bf2ced13842 h1:vr/HnozRka3pE4EsMEg1lgkXJkTFJCVUX+S/ZT6wYzM=
golang.org/x/exp v0.0.0-20240506185415-9bf2ced13842/go.mod h1:XtvwrStGgqGPLc4cjQfWqZHG1YFdYs6swckp8vpsjnc=
golang.org/x/lint v0.0.0-20200302205851-738671d3881b/go.mod h1:3xt1FjdF8hUf6vQPIChWIBhFzV8gjjsPE/fR3IyQdNY=
@@ -105,13 +120,12 @@ golang.org/x/mod v0.5.1/go.mod h1:5OXOZSfqPIIbmVBIIKWRFfZjPR0E5r58TLhUjH0a2Ro=
golang.org/x/net v0.0.0-20190404232315-eb5bcb51f2a3/go.mod h1:t9HGtf8HONx5eT2rtn7q6eTqICYqUVnKs3thJo3Qplg=
golang.org/x/net v0.0.0-20190620200207-3b0461eec859/go.mod h1:z5CRVTTTmAJ677TzLLGU+0bjPO0LkuOLi4/5GtJWs/s=
golang.org/x/net v0.0.0-20211015210444-4f30a5c0130f/go.mod h1:9nx3DQGgdP8bBQD5qxJ1jj9UTztislL4KSBs9R2vV5Y=
golang.org/x/net v0.59.0 h1:5zfYln+w5XCxwrnMMJPufRgNoXEaGxl0wo5GqPXyues=
golang.org/x/net v0.59.0/go.mod h1:2DA/G1UfVbCpQPeWTmMPGY7Cs2PkBkwu743bVX5PIVg=
golang.org/x/net v0.58.0 h1:ynWG7rqYi4ccpTEuPZ2QGWHktVEM9DMCj9yzDE0Q7To=
golang.org/x/net v0.58.0/go.mod h1:YwCddHnFlT7eLQqVprV19OnhLGtc5xOKgE0RyqgfWAU=
golang.org/x/sync v0.0.0-20190423024810-112230192c58/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
golang.org/x/sync v0.0.0-20210220032951-036812b2e83c/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
golang.org/x/sync v0.23.0 h1:KameEIfc1IkluZyXWLn39Wd4tURc6GbCiISGiZm2bQk=
golang.org/x/sync v0.23.0/go.mod h1:sUUOizhqBxiL6pEWpqNLUiaJn1ShEbZ6BBqskPbjZm0=
golang.org/x/sys v0.0.0-20190204203706-41f3e6584952/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY=
golang.org/x/sync v0.22.0 h1:SZjpbeLmrCk4xhRSZFNZW5gFUeCeFgjekvI/+gfScek=
golang.org/x/sync v0.22.0/go.mod h1:9xrNwdLfx4jkKbNva9FpL6vEN7evnE43NNNJQ2LF3+0=
golang.org/x/sys v0.0.0-20190215142949-d0b11bdaac8a/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY=
golang.org/x/sys v0.0.0-20190412213103-97732733099d/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs=
golang.org/x/sys v0.0.0-20201119102817-f84b799fce68/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs=
@@ -119,22 +133,20 @@ golang.org/x/sys v0.0.0-20210423082822-04245dca01da/go.mod h1:h1NjWce9XRLGQEsW7w
golang.org/x/sys v0.0.0-20211019181941-9d821ace8654/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
golang.org/x/sys v0.2.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
golang.org/x/sys v0.10.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
golang.org/x/sys v0.48.0 h1:bbX/i/6MgT9BVLM9RT1thmxL04yeTAhbEz4SyadbXoo=
golang.org/x/sys v0.48.0/go.mod h1:hNLxWAXmnKAxqDtdwIYC4bM9oQPEecfsnNMuSxOs3og=
golang.org/x/sys v0.47.0 h1:o7XGOvZQCADBQQ4Y7VNq2dRWQR7JmOUW8Kxx4ZsNgWs=
golang.org/x/sys v0.47.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw=
golang.org/x/term v0.0.0-20201126162022-7de9c90e9dd1/go.mod h1:bj7SfCRtBDWHUb9snDiAeCFNEtKQo2Wmx5Cou7ajbmo=
golang.org/x/text v0.3.0/go.mod h1:NqM8EUOU14njkJ3fqMW+pc6Ldnwhi/IjpwHt7yyuwOQ=
golang.org/x/text v0.3.6/go.mod h1:5Zoc/QRtKVWzQhOtBMvqHzDpF6irO9z98xDceosuGiQ=
golang.org/x/text v0.3.7/go.mod h1:u+2+/6zg+i71rQMx5EYifcz6MCKuco9NR6JIITiCfzQ=
golang.org/x/text v0.42.0 h1:JbOZXgfeCPU9gacVtYliJqOhD+zhrEqK4LfdpmlUZqI=
golang.org/x/text v0.42.0/go.mod h1:ojzP1Z+2QtioaF8DTtO8K5q7JWVVYwZKenzujK0Zd0E=
golang.org/x/text v0.41.0 h1:vz/seA0lnX87Othu2f/0L24RcgrXD9/YFTSuGjj3rH8=
golang.org/x/text v0.41.0/go.mod h1:jvf1O8ajNzZqhSrQBPbutR/EB83Cc0CFrezNQIwbb5M=
golang.org/x/time v0.14.0 h1:MRx4UaLrDotUKUdCIqzPC48t1Y9hANFKIRpNx+Te8PI=
golang.org/x/time v0.14.0/go.mod h1:eL/Oa2bBBK0TkX57Fyni+NgnyQQN4LitPmob2Hjnqw4=
golang.org/x/tools v0.0.0-20180917221912-90fa682c2a6e/go.mod h1:n7NCudcB/nEzxVGmLbDWY5pfWTLqBcC2KZ6jyYvM4mQ=
golang.org/x/tools v0.0.0-20191119224855-298f0cb1881e/go.mod h1:b+2E5dAYhXwXZwtnZ6UAqBI28+e2cm9otk0dWdXHAEo=
golang.org/x/tools v0.0.0-20200130002326-2f3ba24bd6e7/go.mod h1:TB2adYChydJhpapKDTa4BR/hXlZSLoq2Wpct/0txZ28=
golang.org/x/tools v0.1.8/go.mod h1:nABZi5QlRsZVlzPpHl034qft6wpY4eDcsTt5AaioBiU=
golang.org/x/tools v0.49.0 h1:3NI7VXzL9+1WZD52Dx2ttoPwD5DWrFGpl9mFZDlmisI=
golang.org/x/tools v0.49.0/go.mod h1:SJNXV9DBKT0UbdttsQjbfJlAE/q+y36++zo3uL3N0Oo=
golang.org/x/xerrors v0.0.0-20190717185122-a985d3407aa7/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0=
golang.org/x/xerrors v0.0.0-20191011141410-1b5146add898/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0=
golang.org/x/xerrors v0.0.0-20200804184101-5ec99f83aff1/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0=
@@ -142,14 +154,14 @@ golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2 h1:B82qJJgjvYKsXS9jeu
golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2/go.mod h1:deeaetjYA+DHMHg+sMSMI58GrEteJUUzzw7en6TJQcI=
golang.zx2c4.com/wireguard v0.0.0-20250521234502-f333402bd9cb h1:whnFRlWMcXI9d+ZbWg+4sHnLp52d5yiIPUxMBSt4X9A=
golang.zx2c4.com/wireguard v0.0.0-20250521234502-f333402bd9cb/go.mod h1:rpwXGsirqLqN2L0JDJQlwOboGHmptD5ZD6T2VmcqhTw=
golang.zx2c4.com/wireguard/windows v1.1.1 h1:8/H97U1v1PNDNcBsMZgU3KFuND9MQdTsU2NOwmCXArE=
golang.zx2c4.com/wireguard/windows v1.1.1/go.mod h1:+fbT3FFdX4zzYDLwJh5+HPEcNN/3HyNdzhNSVsQM+zs=
golang.zx2c4.com/wireguard/windows v1.0.1 h1:eOxiDVbywPC+ZQqvdCK7x+ZwWXKbYv50TtH8ysFIbw8=
golang.zx2c4.com/wireguard/windows v1.0.1/go.mod h1:+fbT3FFdX4zzYDLwJh5+HPEcNN/3HyNdzhNSVsQM+zs=
gonum.org/v1/gonum v0.17.0 h1:VbpOemQlsSMrYmn7T2OUvQ4dqxQXU+ouZFQsZOx50z4=
gonum.org/v1/gonum v0.17.0/go.mod h1:El3tOrEuMpv2UdMrbNlKEh9vd86bmQ6vqIcDwxEOc1E=
google.golang.org/genproto/googleapis/rpc v0.0.0-20260706201446-f0a921348800 h1:qEHAMpSaUhtD0p3NbEEI83HwNGFxEwaSJ1G9PLnCBZE=
google.golang.org/genproto/googleapis/rpc v0.0.0-20260706201446-f0a921348800/go.mod h1:4Hqkh8ycfw05ld/3BWL7rJOSfebL2Q+DVDeRgYgxUU8=
google.golang.org/grpc v1.84.0 h1:soMyaPJ8pAak5PIQ0DGBUir0XRo2fRoMqhNWMLlLxO0=
google.golang.org/grpc v1.84.0/go.mod h1:ljCht0DrxQrXBDRTZp52Qxh3Ffk8CdYm2sj4O2QN2C0=
google.golang.org/genproto/googleapis/rpc v0.0.0-20260526163538-3dc84a4a5aaa h1:mZHHdPZl0dbGHCflZgAq/Q468DWVFcU2whhB2KAo8fk=
google.golang.org/genproto/googleapis/rpc v0.0.0-20260526163538-3dc84a4a5aaa/go.mod h1:4Hqkh8ycfw05ld/3BWL7rJOSfebL2Q+DVDeRgYgxUU8=
google.golang.org/grpc v1.83.2 h1:EManeRomTObA0BU7I8vXgg/78uE5MJ9M8B39EX2WscU=
google.golang.org/grpc v1.83.2/go.mod h1:YPI1hK3kDked6iHvgX3tR0y+nX/qpMFKhPgFsokw1S8=
google.golang.org/protobuf v1.36.12 h1:pJOKDDOyeXErUroCihFAd5LQuwXBSpVnKGrj5o/fwxc=
google.golang.org/protobuf v1.36.12/go.mod h1:HTf+CrKn2C3g5S8VImy6tdcUvCska2kB7j23XfzDpco=
gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0=
@@ -162,9 +174,5 @@ gvisor.dev/gvisor v0.0.0-20260122175437-89a5d21be8f0 h1:Lk6hARj5UPY47dBep70OD/TI
gvisor.dev/gvisor v0.0.0-20260122175437-89a5d21be8f0/go.mod h1:QkHjoMIBaYtpVufgwv3keYAbln78mBoCuShZrPrer1Q=
h12.io/socks v1.0.3 h1:Ka3qaQewws4j4/eDQnOdpr4wXsC//dXtWvftlIcCQUo=
h12.io/socks v1.0.3/go.mod h1:AIhxy1jOId/XCz9BO+EIgNL2rQiPTBNnOfnVnQ+3Eck=
layeh.com/gopher-luar v1.0.11 h1:8zJudpKI6HWkoh9eyyNFaTM79PY6CAPcIr6X/KTiliw=
layeh.com/gopher-luar v1.0.11/go.mod h1:TPnIVCZ2RJBndm7ohXyaqfhzjlZ+OA2SZR/YwL8tECk=
lukechampine.com/blake3 v1.4.1 h1:I3Smz7gso8w4/TunLKec6K2fn+kyKtDxr/xcQEN84Wg=
lukechampine.com/blake3 v1.4.1/go.mod h1:QFosUxmjB8mnrWFSNwKmvxHpfY72bmD2tQ0kBMM3kwo=
mvdan.cc/gofumpt v0.12.0 h1:1Lbudkz2kpM9Cjz2pL4M19u7q+GaEhCTNf7N9mfpcho=
mvdan.cc/gofumpt v0.12.0/go.mod h1:SmBHHrljiZu/uoypeKup3rFzP6eoC9UwCp2iH5E3jZA=
-14
View File
@@ -14,11 +14,9 @@ import (
"github.com/xtls/xray-core/common/errors"
"github.com/xtls/xray-core/common/geodata"
"github.com/xtls/xray-core/common/net"
"github.com/xtls/xray-core/common/platform"
)
type NameServerConfig struct {
ID string `json:"id"`
Address *Address `json:"address"`
ClientIP *Address `json:"clientIp"`
Port uint16 `json:"port"`
@@ -45,7 +43,6 @@ func (c *NameServerConfig) UnmarshalJSON(data []byte) error {
}
var advanced struct {
ID string `json:"id"`
Address *Address `json:"address"`
ClientIP *Address `json:"clientIp"`
Port uint16 `json:"port"`
@@ -63,7 +60,6 @@ func (c *NameServerConfig) UnmarshalJSON(data []byte) error {
UnexpectedIPs StringList `json:"unexpectedIPs"`
}
if err := json.Unmarshal(data, &advanced); err == nil {
c.ID = advanced.ID
c.Address = advanced.Address
c.ClientIP = advanced.ClientIP
c.Port = advanced.Port
@@ -138,7 +134,6 @@ func (c *NameServerConfig) Build() (*dns.NameServer, error) {
}
return &dns.NameServer{
Id: c.ID,
Address: &net.Endpoint{
Network: net.Network_UDP,
Address: c.Address.Build(),
@@ -164,7 +159,6 @@ func (c *NameServerConfig) Build() (*dns.NameServer, error) {
// DNSConfig is a JSON serializable object for dns.Config
type DNSConfig struct {
Servers []*NameServerConfig `json:"servers"`
Script string `json:"script"`
Hosts *HostsWrapper `json:"hosts"`
ClientIP *Address `json:"clientIp"`
Tag string `json:"tag"`
@@ -284,14 +278,6 @@ func (c *DNSConfig) Build() (*dns.Config, error) {
QueryStrategy: resolveQueryStrategy(c.QueryStrategy),
}
if c.Script != "" {
path, err := platform.ResolveLuaFile(c.Script)
if err != nil {
return nil, errors.New("failed to resolve DNS script: ", c.Script).Base(err)
}
config.Script = path
}
if c.ClientIP != nil {
if !c.ClientIP.Family().IsIP() {
return nil, errors.New("not an IP address:", c.ClientIP.String())
-50
View File
@@ -2,8 +2,6 @@ package conf_test
import (
"encoding/json"
"os"
"path/filepath"
"testing"
"github.com/google/go-cmp/cmp"
@@ -124,51 +122,3 @@ func TestDNSConfigParsing(t *testing.T) {
}
}
}
func TestDNSScriptConfig(t *testing.T) {
dir := t.TempDir()
t.Setenv("xray.location.confdir", dir)
path := filepath.Join(dir, "lookup.lua")
if err := os.WriteFile(path, []byte("function HandleDNSQuery(domain, ipv4, ipv6, fake) end"), 0o600); err != nil {
t.Fatal(err)
}
for _, tc := range []struct {
name string
script string
wantError bool
}{
{"relative", "lookup.lua", false},
{"absolute", path, false},
{"missing", "missing.lua", true},
{"directory", dir, true},
} {
t.Run(tc.name, func(t *testing.T) {
built, err := (&DNSConfig{Script: tc.script}).Build()
if tc.wantError {
if err == nil {
t.Fatal("Build accepted an invalid script path")
}
return
}
if err != nil {
t.Fatal(err)
}
if built.Script != path {
t.Fatalf("script path = %q, want %q", built.Script, path)
}
})
}
var parsed DNSConfig
if err := json.Unmarshal([]byte(`{"servers":[{"id":"primary","address":"1.1.1.1"}]}`), &parsed); err != nil {
t.Fatal(err)
}
built, err := parsed.Build()
if err != nil {
t.Fatal(err)
}
if len(built.NameServer) != 1 || built.NameServer[0].Id != "primary" {
t.Fatalf("nameserver IDs = %v, want primary", built.NameServer)
}
}
+2 -2
View File
@@ -97,7 +97,7 @@ func (v *HTTPClientConfig) Build() (proto.Message, error) {
user.Email = v.Email
} else {
if err := json.Unmarshal(rawUser, user); err != nil {
return nil, errors.New("failed to parse HTTP user").Base(err)
return nil, errors.New("failed to parse HTTP user").Base(err).AtError()
}
}
account := new(HTTPAccount)
@@ -106,7 +106,7 @@ func (v *HTTPClientConfig) Build() (proto.Message, error) {
account.Password = v.Password
} else {
if err := json.Unmarshal(rawUser, account); err != nil {
return nil, errors.New("failed to parse HTTP account").Base(err)
return nil, errors.New("failed to parse HTTP account").Base(err).AtError()
}
}
user.Account = serial.ToTypedMessage(account.Build())
+1 -1
View File
@@ -18,7 +18,7 @@ func RegisterConfigureFilePostProcessingStage(name string, stage ConfigureFilePo
func PostProcessConfigureFile(conf *Config) error {
for k, v := range configureFilePostProcessingStages {
if err := v.Process(conf); err != nil {
return errors.New("Rejected by Postprocessing Stage ", k).Base(err)
return errors.New("Rejected by Postprocessing Stage ", k).AtError().Base(err)
}
}
return nil
+2 -2
View File
@@ -13,7 +13,7 @@ type ConfigCreatorCache map[string]ConfigCreator
func (v ConfigCreatorCache) RegisterCreator(id string, creator ConfigCreator) error {
if _, found := v[id]; found {
return errors.New(id, " already registered.")
return errors.New(id, " already registered.").AtError()
}
v[id] = creator
@@ -61,7 +61,7 @@ func (v *JSONConfigLoader) Load(raw []byte) (interface{}, string, error) {
}
rawID, found := obj[v.idKey]
if !found {
return nil, "", errors.New(v.idKey, " not found in JSON context")
return nil, "", errors.New(v.idKey, " not found in JSON context").AtError()
}
var id string
if err := json.Unmarshal(rawID, &id); err != nil {
-103
View File
@@ -1,103 +0,0 @@
package conf
import (
"net/netip"
"strings"
"github.com/xtls/xray-core/common/errors"
"github.com/xtls/xray-core/common/protocol"
"github.com/xtls/xray-core/common/serial"
"github.com/xtls/xray-core/proxy/masque"
"google.golang.org/protobuf/proto"
)
type MasqueClientConfig struct {
Address *Address `json:"address"`
Port uint16 `json:"port"`
RemoteDNS []string `json:"remoteDNS"`
}
func (c *MasqueClientConfig) Build() (proto.Message, error) {
if c.Address == nil {
return nil, errors.New(`MASQUE: "address" is not set`)
}
if c.Port == 0 {
return nil, errors.New(`MASQUE: "port" is not set`)
}
for _, s := range c.RemoteDNS {
if _, err := netip.ParseAddr(s); err != nil {
return nil, errors.New(`MASQUE: invalid "remoteDNS" `, s).Base(err)
}
}
return &masque.ClientConfig{
Server: &protocol.ServerEndpoint{
Address: c.Address.Build(),
Port: uint32(c.Port),
},
RemoteDns: c.RemoteDNS,
}, nil
}
type MasqueUserConfig struct {
Pass string `json:"pass"`
Level uint32 `json:"level"`
Email string `json:"email"`
}
type MasqueServerConfig struct {
Users []*MasqueUserConfig `json:"users"`
Clients []*MasqueUserConfig `json:"clients"`
Address []string `json:"address"`
MTU uint32 `json:"mtu"`
}
func (c *MasqueServerConfig) Build() (proto.Message, error) {
if c.Clients != nil {
c.Users = c.Clients
}
config := &masque.ServerConfig{
Address: c.Address,
Mtu: c.MTU,
}
emails := make(map[string]bool)
for _, user := range c.Users {
if user.Email == "" {
return nil, errors.New(`MASQUE: "email" is empty`)
}
if strings.Contains(user.Email, ":") {
return nil, errors.New(`MASQUE: invalid "email" `, user.Email)
}
if user.Pass == "" {
return nil, errors.New(`MASQUE: "pass" of `, user.Email, ` is empty`)
}
email := strings.ToLower(user.Email)
if emails[email] {
return nil, errors.New(`MASQUE: duplicate "email" `, user.Email)
}
emails[email] = true
config.Users = append(config.Users, &protocol.User{
Email: user.Email,
Level: user.Level,
Account: serial.ToTypedMessage(&masque.Account{Password: user.Pass}),
})
}
if len(c.Address) == 0 {
return nil, errors.New(`MASQUE: "address" is not set`)
}
var v4, v6 bool
for _, s := range c.Address {
prefix, err := netip.ParsePrefix(s)
if err != nil {
return nil, errors.New(`MASQUE: invalid "address" `, s).Base(err)
}
if prefix.Addr().Is4() && v4 || prefix.Addr().Is6() && v6 {
return nil, errors.New(`MASQUE: "address" takes at most one IPv4 and one IPv6 prefix`)
}
v4 = v4 || prefix.Addr().Is4()
v6 = v6 || prefix.Addr().Is6()
}
if c.MTU != 0 && (c.MTU < 1280 || c.MTU > 65535) {
return nil, errors.New(`MASQUE: "mtu" must be between 1280 and 65535`)
}
return config, nil
}
-188
View File
@@ -1,188 +0,0 @@
package conf_test
import (
"encoding/json"
"testing"
"github.com/xtls/xray-core/common/protocol"
"github.com/xtls/xray-core/common/serial"
. "github.com/xtls/xray-core/infra/conf"
masqueproxy "github.com/xtls/xray-core/proxy/masque"
"github.com/xtls/xray-core/transport/internet/masque"
)
func TestMasqueConfig(t *testing.T) {
creator := func() Buildable {
return new(MasqueConfig)
}
runMultiTestCase(t, []TestCase{
{
Input: `{}`,
Parser: loadJSON(creator),
Output: &masque.Config{Path: "/.well-known/masque/ip/*/*/"},
},
{
Input: `{
"host": "example.com:8443",
"path": "/.well-known/masque/ip/{target}/{ipproto}/",
"headers": {"Authorization": "Basic dTpw"}
}`,
Parser: loadJSON(creator),
Output: &masque.Config{
Host: "example.com:8443",
Path: "/.well-known/masque/ip/*/*/",
Headers: map[string]string{"Authorization": "Basic dTpw"},
},
},
{
Input: `{"path": "/masque/ip{?target,ipproto}"}`,
Parser: loadJSON(creator),
Output: &masque.Config{Path: "/masque/ip?target=*&ipproto=*"},
},
{
Input: `{"user": "u", "pass": "p:q", "headers": {"X-Token": "a"}}`,
Parser: loadJSON(creator),
Output: &masque.Config{
Path: "/.well-known/masque/ip/*/*/",
Headers: map[string]string{"Authorization": "Basic dTpwOnE=", "X-Token": "a"},
},
},
})
for _, input := range []string{
`{"path": "/masque/{target}/{ipproto}/{dns}"}`,
`{"path": "masque"}`,
`{"host": "example.com/path"}`,
`{"headers": {"host": "example.com"}}`,
`{"headers": {"Capsule-Protocol": "?0"}}`,
`{"headers": {"X Token": "a"}}`,
`{"headers": {"X-Token": "a\r\nb"}}`,
`{"user": "u:v", "pass": "p"}`,
`{"user": "u", "pass": "p", "headers": {"authorization": "Basic dTpw"}}`,
} {
if _, err := loadJSON(creator)(input); err == nil {
t.Errorf("expected an error for %s", input)
}
}
}
func TestMasqueOutboundConfig(t *testing.T) {
build := func(s string) error {
c := new(OutboundDetourConfig)
if err := json.Unmarshal([]byte(s), c); err != nil {
return err
}
_, err := c.Build()
return err
}
if err := build(`{
"protocol": "masque",
"settings": {"address": "example.com", "port": 443},
"streamSettings": {"network": "masque", "security": "tls"},
"mux": {"enabled": false, "concurrency": -1}
}`); err != nil {
t.Error(err)
}
for _, input := range []string{
`{"protocol": "masque", "settings": {"address": "example.com"}, "streamSettings": {"network": "masque", "security": "tls"}}`,
`{"protocol": "masque", "settings": {"address": "example.com", "port": 443}, "streamSettings": {"network": "masque", "security": "tls"}, "mux": {"enabled": true}}`,
`{"protocol": "masque", "settings": {"address": "example.com", "port": 443}, "streamSettings": {"network": "masque", "security": "tls"}, "mux": {"enabled": true, "concurrency": -1}}`,
`{"protocol": "freedom", "streamSettings": {"network": "masque", "security": "tls"}}`,
} {
if err := build(input); err == nil {
t.Errorf("expected an error for %s", input)
}
}
}
func TestMasqueServerConfig(t *testing.T) {
creator := func() Buildable {
return new(MasqueServerConfig)
}
runMultiTestCase(t, []TestCase{
{
Input: `{
"users": [{"email": "u@example.com", "pass": "p", "level": 1}],
"address": ["10.13.0.1/24", "fd13::1/64"],
"mtu": 1400
}`,
Parser: loadJSON(creator),
Output: &masqueproxy.ServerConfig{
Users: []*protocol.User{{
Email: "u@example.com",
Level: 1,
Account: serial.ToTypedMessage(&masqueproxy.Account{Password: "p"}),
}},
Address: []string{"10.13.0.1/24", "fd13::1/64"},
Mtu: 1400,
},
},
{
Input: `{"clients": [{"email": "u", "pass": "p:q"}], "address": ["10.13.0.1/24"]}`,
Parser: loadJSON(creator),
Output: &masqueproxy.ServerConfig{
Users: []*protocol.User{{
Email: "u",
Account: serial.ToTypedMessage(&masqueproxy.Account{Password: "p:q"}),
}},
Address: []string{"10.13.0.1/24"},
},
},
{
Input: `{"address": ["10.13.0.1/24"]}`,
Parser: loadJSON(creator),
Output: &masqueproxy.ServerConfig{
Address: []string{"10.13.0.1/24"},
},
},
})
for _, input := range []string{
`{"users": [{"email": "u:v", "pass": "p"}], "address": ["10.13.0.1/24"]}`,
`{"users": [{"email": "", "pass": "p"}], "address": ["10.13.0.1/24"]}`,
`{"users": [{"pass": "p"}], "address": ["10.13.0.1/24"]}`,
`{"users": [{"email": "u", "pass": ""}], "address": ["10.13.0.1/24"]}`,
`{"users": [{"email": "u", "pass": "p"}, {"email": "U", "pass": "q"}], "address": ["10.13.0.1/24"]}`,
`{"users": [{"email": "u", "pass": "p"}]}`,
`{"users": [{"email": "u", "pass": "p"}], "address": ["10.13.0.1"]}`,
`{"users": [{"email": "u", "pass": "p"}], "address": ["10.13.0.1/24", "10.14.0.1/24"]}`,
`{"users": [{"email": "u", "pass": "p"}], "address": ["fd13::1/64", "fd14::1/64"]}`,
`{"users": [{"email": "u", "pass": "p"}], "address": ["10.13.0.1/24"], "mtu": 1000}`,
`{"users": [{"email": "u", "pass": "p"}], "address": ["10.13.0.1/24"], "mtu": 70000}`,
} {
if _, err := loadJSON(creator)(input); err == nil {
t.Errorf("expected an error for %s", input)
}
}
}
func TestMasqueInboundConfig(t *testing.T) {
build := func(s string) error {
c := new(InboundDetourConfig)
if err := json.Unmarshal([]byte(s), c); err != nil {
return err
}
_, err := c.Build()
return err
}
if err := build(`{
"protocol": "masque",
"port": 443,
"settings": {"users": [{"email": "u@example.com", "pass": "p"}], "address": ["10.13.0.1/24"]},
"streamSettings": {"network": "masque", "security": "tls"}
}`); err != nil {
t.Error(err)
}
if err := build(`{
"protocol": "vless",
"port": 443,
"settings": {"users": [{"id": "27848739-7e62-4138-9fd3-098a63964b6b"}], "decryption": "none"},
"streamSettings": {"network": "masque", "security": "tls"}
}`); err == nil {
t.Error("expected an error for the masque transport on a vless inbound")
}
}
-11
View File
@@ -7,7 +7,6 @@ import (
"github.com/xtls/xray-core/app/router"
"github.com/xtls/xray-core/common/errors"
"github.com/xtls/xray-core/common/geodata"
"github.com/xtls/xray-core/common/platform"
"github.com/xtls/xray-core/common/serial"
"google.golang.org/protobuf/proto"
@@ -73,7 +72,6 @@ type RouterConfig struct {
RuleList []json.RawMessage `json:"rules"`
DomainStrategy *string `json:"domainStrategy"`
Balancers []*BalancingRule `json:"balancers"`
Script string `json:"script"`
}
func (c *RouterConfig) getDomainStrategy() router.Config_DomainStrategy {
@@ -94,15 +92,6 @@ func (c *RouterConfig) getDomainStrategy() router.Config_DomainStrategy {
func (c *RouterConfig) Build() (*router.Config, error) {
config := new(router.Config)
if c.Script != "" {
path, err := platform.ResolveLuaFile(c.Script)
if err != nil {
return nil, errors.New("failed to resolve routing script").Base(err)
}
config.Script = path
}
config.DomainStrategy = c.getDomainStrategy()
var rawRuleList []json.RawMessage
-38
View File
@@ -2,8 +2,6 @@ package conf_test
import (
"encoding/json"
"os"
"path/filepath"
"testing"
"time"
_ "unsafe"
@@ -238,39 +236,3 @@ func TestRouterConfig(t *testing.T) {
},
})
}
func TestRouterScriptConfig(t *testing.T) {
dir := t.TempDir()
t.Setenv("xray.location.confdir", dir)
path := filepath.Join(dir, "route.lua")
if err := os.WriteFile(path, []byte("function HandleRoute() end"), 0o600); err != nil {
t.Fatal(err)
}
for _, tc := range []struct {
name string
script string
wantError bool
}{
{"relative", "route.lua", false},
{"absolute", path, false},
{"missing", "missing.lua", true},
{"directory", dir, true},
} {
t.Run(tc.name, func(t *testing.T) {
built, err := (&RouterConfig{Script: tc.script}).Build()
if tc.wantError {
if err == nil {
t.Fatal("Build accepted invalid script path")
}
return
}
if err != nil {
t.Fatal(err)
}
if built.Script != path {
t.Fatalf("script path = %q, want %q", built.Script, path)
}
})
}
}
+1 -1
View File
@@ -30,7 +30,7 @@ func MergeConfigFromFiles(files []*core.ConfigSource) (string, error) {
if j, ok := creflect.MarshalToJson(c, true); ok {
return j, nil
}
return "", errors.New("marshal to json failed.")
return "", errors.New("marshal to json failed.").AtError()
}
func mergeConfigs(files []*core.ConfigSource) (*conf.Config, error) {
+56 -37
View File
@@ -3,6 +3,8 @@ package conf
import (
"strings"
"github.com/sagernet/sing-shadowsocks/shadowaead_2022"
C "github.com/sagernet/sing/common"
"github.com/xtls/xray-core/common/errors"
"github.com/xtls/xray-core/common/protocol"
"github.com/xtls/xray-core/common/serial"
@@ -53,7 +55,7 @@ func (v *ShadowsocksServerConfig) Build() (proto.Message, error) {
v.Users = v.Clients
}
if _, err := shadowsocks_2022.GetCipherMethod(v.Cipher); err == nil {
if C.Contains(shadowaead_2022.List, v.Cipher) {
return buildShadowsocks2022(v)
}
@@ -109,14 +111,12 @@ func (v *ShadowsocksServerConfig) Build() (proto.Message, error) {
}
func buildShadowsocks2022(v *ShadowsocksServerConfig) (proto.Message, error) {
v.Cipher = strings.ToLower(v.Cipher)
if len(v.Users) == 0 {
config := new(shadowsocks_2022.ServerConfig)
config.Method = v.Cipher
config.Key = v.Password
config.Network = v.NetworkList.Build()
config.Email = v.Email
config.Level = int32(v.Level)
return config, nil
}
@@ -171,7 +171,6 @@ func buildShadowsocks2022(v *ShadowsocksServerConfig) (proto.Message, error) {
Email: user.Email,
Address: user.Address.Build(),
Port: uint32(user.Port),
Level: int32(user.Level),
})
}
return config, nil
@@ -215,43 +214,63 @@ func (v *ShadowsocksClientConfig) Build() (proto.Message, error) {
return nil, errors.New(`Shadowsocks settings: "servers" should have one and only one member. Multiple endpoints in "servers" should use multiple Shadowsocks outbounds and routing balancer instead`)
}
server := v.Servers[0]
if server.Address == nil {
return nil, errors.New("Shadowsocks server address is not set.")
}
if server.Port == 0 {
return nil, errors.New("Invalid Shadowsocks port.")
}
if server.Password == "" {
return nil, errors.New("Shadowsocks password is not specified.")
if len(v.Servers) == 1 {
server := v.Servers[0]
if C.Contains(shadowaead_2022.List, server.Cipher) {
if server.Address == nil {
return nil, errors.New("Shadowsocks server address is not set.")
}
if server.Port == 0 {
return nil, errors.New("Invalid Shadowsocks port.")
}
if server.Password == "" {
return nil, errors.New("Shadowsocks password is not specified.")
}
config := new(shadowsocks_2022.ClientConfig)
config.Address = server.Address.Build()
config.Port = uint32(server.Port)
config.Method = server.Cipher
config.Key = server.Password
return config, nil
}
}
if _, err := shadowsocks_2022.GetCipherMethod(server.Cipher); err == nil {
config := new(shadowsocks_2022.ClientConfig)
config.Address = server.Address.Build()
config.Port = uint32(server.Port)
config.Method = server.Cipher
config.Key = server.Password
return config, nil
}
config := new(shadowsocks.ClientConfig)
account := &shadowsocks.Account{
Password: server.Password,
for _, server := range v.Servers {
if C.Contains(shadowaead_2022.List, server.Cipher) {
return nil, errors.New("Shadowsocks 2022 accept no multi servers")
}
if server.Address == nil {
return nil, errors.New("Shadowsocks server address is not set.")
}
if server.Port == 0 {
return nil, errors.New("Invalid Shadowsocks port.")
}
if server.Password == "" {
return nil, errors.New("Shadowsocks password is not specified.")
}
account := &shadowsocks.Account{
Password: server.Password,
}
account.CipherType = cipherFromString(server.Cipher)
if account.CipherType == shadowsocks.CipherType_UNKNOWN {
return nil, errors.New("unknown cipher method: ", server.Cipher)
}
ss := &protocol.ServerEndpoint{
Address: server.Address.Build(),
Port: uint32(server.Port),
User: &protocol.User{
Level: uint32(server.Level),
Email: server.Email,
Account: serial.ToTypedMessage(account),
},
}
config.Server = ss
break
}
account.CipherType = cipherFromString(server.Cipher)
if account.CipherType == shadowsocks.CipherType_UNKNOWN {
return nil, errors.New("unknown cipher method: ", server.Cipher)
}
ss := &protocol.ServerEndpoint{
Address: server.Address.Build(),
Port: uint32(server.Port),
User: &protocol.User{
Level: uint32(server.Level),
Email: server.Email,
Account: serial.ToTypedMessage(account),
},
}
config.Server = ss
return config, nil
}
+3 -2
View File
@@ -44,6 +44,7 @@ func (v *SocksServerConfig) Build() (proto.Message, error) {
case AuthMethodUserPass:
config.AuthType = socks.AuthType_PASSWORD
default:
// errors.New("unknown socks auth method: ", v.AuthMethod, ". Default to noauth.").AtWarning().WriteToLog()
config.AuthType = socks.AuthType_NO_AUTH
}
@@ -114,7 +115,7 @@ func (v *SocksClientConfig) Build() (proto.Message, error) {
user.Email = v.Email
} else {
if err := json.Unmarshal(rawUser, user); err != nil {
return nil, errors.New("failed to parse Socks user").Base(err)
return nil, errors.New("failed to parse Socks user").Base(err).AtError()
}
}
account := new(SocksAccount)
@@ -123,7 +124,7 @@ func (v *SocksClientConfig) Build() (proto.Message, error) {
account.Password = v.Password
} else {
if err := json.Unmarshal(rawUser, account); err != nil {
return nil, errors.New("failed to parse socks account").Base(err)
return nil, errors.New("failed to parse socks account").Base(err).AtError()
}
}
user.Account = serial.ToTypedMessage(account.Build())
+27 -234
View File
@@ -1,7 +1,6 @@
package conf
import (
"context"
"crypto/x509"
"encoding/base64"
"encoding/hex"
@@ -15,7 +14,6 @@ import (
googleuuid "github.com/google/uuid"
"github.com/xtls/xray-core/common/errors"
"github.com/xtls/xray-core/common/net"
"github.com/xtls/xray-core/common/serial"
"github.com/xtls/xray-core/transport/internet/finalmask/fragment"
"github.com/xtls/xray-core/transport/internet/finalmask/header/custom"
"github.com/xtls/xray-core/transport/internet/finalmask/mkcp/aes128gcm"
@@ -25,7 +23,6 @@ import (
"github.com/xtls/xray-core/transport/internet/finalmask/realm"
"github.com/xtls/xray-core/transport/internet/finalmask/salamander"
"github.com/xtls/xray-core/transport/internet/finalmask/sudoku"
"github.com/xtls/xray-core/transport/internet/finalmask/udphop"
"github.com/xtls/xray-core/transport/internet/finalmask/xdns"
"github.com/xtls/xray-core/transport/internet/finalmask/xicmp"
"github.com/xtls/xray-core/transport/internet/finalmask/xmc"
@@ -83,10 +80,9 @@ var (
"noise": func() interface{} { return new(NoiseMask) },
"salamander": func() interface{} { return new(Salamander) },
"sudoku": func() interface{} { return new(Sudoku) },
"xdns": func() interface{} { return new(XDNS) },
"xdns": func() interface{} { return new(Xdns) },
"xicmp": func() interface{} { return new(Xicmp) },
"realm": func() interface{} { return new(Realm) },
"udphop": func() interface{} { return new(UDPHop) },
}, "type", "settings")
)
@@ -310,27 +306,14 @@ type NoiseMask struct {
}
func (c *NoiseMask) Build() (proto.Message, error) {
noiseSlice := make([]*noise.Item, 0, len(c.Noise))
for _, item := range c.Noise {
if len(item.Packet) > 0 && item.Rand.To > 0 {
return nil, errors.New("len(item.Packet) > 0 && item.Rand.To > 0")
}
if strings.ToLower(item.Type) == "exp" {
var exp string
if err := json.Unmarshal(item.Packet, &exp); err != nil {
return nil, errors.New(`"packet" of noise "type": "exp" must be a string`).Base(err)
}
segments, err := parseNoiseExp(exp)
if err != nil {
return nil, err
}
noiseSlice = append(noiseSlice, &noise.Item{
Segments: segments,
DelayMin: int64(item.Delay.From),
DelayMax: int64(item.Delay.To),
})
continue
}
}
noiseSlice := make([]*noise.Item, 0, len(c.Noise))
for _, item := range c.Noise {
if item.RandRange == nil {
item.RandRange = &Int32Range{From: 0, To: 255}
}
@@ -359,88 +342,6 @@ func (c *NoiseMask) Build() (proto.Message, error) {
}, nil
}
var noiseExpPattern = regexp.MustCompile(`<\s*([a-z]+)(?:\s+([^>]*?))?\s*>`)
func parseNoiseExp(exp string) ([]*noise.Segment, error) {
var segments []*noise.Segment
matches := noiseExpPattern.FindAllStringSubmatchIndex(exp, -1)
last := 0
for _, m := range matches {
if strings.TrimSpace(exp[last:m[0]]) != "" {
return nil, errors.New("invalid noise exp near ", exp[last:m[0]])
}
last = m[1]
key := exp[m[2]:m[3]]
arg := ""
if m[4] >= 0 {
arg = exp[m[4]:m[5]]
}
segment, err := buildNoiseSegment(key, arg)
if err != nil {
return nil, err
}
segments = append(segments, segment)
}
if strings.TrimSpace(exp[last:]) != "" {
return nil, errors.New("invalid noise exp near ", exp[last:])
}
if len(segments) == 0 {
return nil, errors.New("empty noise exp: ", exp)
}
return segments, nil
}
func buildNoiseSegment(key, arg string) (*noise.Segment, error) {
sizeSegment := func(kind noise.Segment_Kind) (*noise.Segment, error) {
if arg == "" {
return nil, errors.New("<", key, "> in noise exp needs a size")
}
lo, hi, err := ParseRangeString(arg)
if err != nil {
return nil, err
}
if lo < 0 || hi < lo || hi > 65535 {
return nil, errors.New("invalid size in noise exp: ", arg)
}
return &noise.Segment{Kind: kind, MinSize: int64(lo), MaxSize: int64(hi)}, nil
}
switch key {
case "b":
hexStr := strings.TrimPrefix(strings.TrimPrefix(strings.Join(strings.Fields(arg), ""), "0x"), "0X")
if len(hexStr) == 0 {
return nil, errors.New("empty bytes in noise exp")
}
raw, err := hex.DecodeString(hexStr)
if err != nil {
return nil, errors.New("invalid hex in noise exp: ", arg).Base(err)
}
return &noise.Segment{Kind: noise.Segment_BYTES, Bytes: raw}, nil
case "r":
return sizeSegment(noise.Segment_RANDOM)
case "rc":
return sizeSegment(noise.Segment_RANDOM_ASCII)
case "rd":
return sizeSegment(noise.Segment_RANDOM_DIGIT)
case "t":
if arg != "" {
return nil, errors.New("<t> in noise exp takes no argument")
}
return &noise.Segment{Kind: noise.Segment_TIMESTAMP}, nil
case "c":
if arg != "" {
return nil, errors.New("<c> in noise exp takes no argument")
}
return &noise.Segment{Kind: noise.Segment_COUNTER}, nil
case "n":
if arg != "" {
return nil, errors.New("<n> in noise exp takes no argument")
}
return &noise.Segment{Kind: noise.Segment_NONCE}, nil
default:
return nil, errors.New("unknown <", key, "> in noise exp")
}
}
type UDPItem struct {
Rand int32 `json:"rand"`
RandRange *Int32Range `json:"randRange"`
@@ -791,88 +692,32 @@ func (c *Sudoku) Build() (proto.Message, error) {
}, nil
}
type XDNSDomain struct {
Name string `json:"name"`
LenLimit int32 `json:"lenLimit"`
LabelLimit int32 `json:"labelLimit"`
Types []int32 `json:"types"`
Edns0 int32 `json:"edns0"`
type Xdns struct {
Domain json.RawMessage `json:"domain"`
Domains []string `json:"domains"`
Resolvers []string `json:"resolvers"`
}
type XDNSResolverTCP struct {
Addr string `json:"addr"`
}
func (c *XDNSResolverTCP) Build() (proto.Message, error) {
return &xdns.TCPResolverProto{Addr: c.Addr}, nil
}
type XDNSResolverUDP struct {
Addr string `json:"addr"`
}
func (c *XDNSResolverUDP) Build() (proto.Message, error) {
return &xdns.UDPResolverProto{Addr: c.Addr}, nil
}
var xdnsLoader = NewJSONConfigLoader(ConfigCreatorCache{
"tcp": func() interface{} { return new(XDNSResolverTCP) },
"udp": func() interface{} { return new(XDNSResolverUDP) },
}, "type", "settings")
type XDNSResolver struct {
Type string `json:"type"`
Settings json.RawMessage `json:"settings"`
}
type XDNS struct {
Domains []XDNSDomain `json:"domains"`
Resolvers []XDNSResolver `json:"resolvers"`
ExtraPoll int32 `json:"extraPoll"`
}
func (c *XDNS) Build() (proto.Message, error) {
var domains []*xdns.DomainProto
var resolvers []*serial.TypedMessage
for i := range c.Domains {
if c.Domains[i].LenLimit == 0 {
c.Domains[i].LenLimit = 255
}
if c.Domains[i].LabelLimit == 0 {
c.Domains[i].LabelLimit = 63
}
types := make([]uint16, 0, len(c.Domains[i].Types))
for j := range c.Domains[i].Types {
types = append(types, uint16(c.Domains[i].Types[j]))
}
domain, err := xdns.NewDomain(c.Domains[i].Name, int(c.Domains[i].LenLimit), int(c.Domains[i].LabelLimit), types, uint16(c.Domains[i].Edns0))
if err != nil {
return nil, err
}
errors.LogInfo(context.Background(), domain.Show())
domains = append(domains, &xdns.DomainProto{
Name: c.Domains[i].Name,
LenLimit: c.Domains[i].LenLimit,
LabelLimit: c.Domains[i].LabelLimit,
Types: c.Domains[i].Types,
Edns0: c.Domains[i].Edns0,
})
func (c *Xdns) Build() (proto.Message, error) {
if c.Domain != nil {
return nil, errors.PrintRemovedFeatureError("domain", "domains(server) & resolvers(client)")
}
for i := range c.Resolvers {
config, err := xdnsLoader.LoadWithID(c.Resolvers[i].Settings, c.Resolvers[i].Type)
if err != nil {
return nil, err
}
pm, err := config.(interface{ Build() (proto.Message, error) }).Build()
if err != nil {
return nil, err
}
resolvers = append(resolvers, serial.ToTypedMessage(pm))
if len(c.Domains) == 0 && len(c.Resolvers) == 0 {
return nil, errors.New("empty domains & empty resolvers")
}
if c.ExtraPoll < 0 || c.ExtraPoll > 3 {
return nil, errors.New("c.ExtraPoll < 0 || c.ExtraPoll > 3")
for _, r := range c.Resolvers {
if !strings.Contains(r, "+udp://") {
return nil, errors.New("invalid resolver ", r)
}
}
return &xdns.Config{Domains: domains, Resolvers: resolvers, ExtraPoll: c.ExtraPoll}, nil
return &xdns.Config{
Domains: c.Domains,
Resolvers: c.Resolvers,
}, nil
}
type XMC struct {
@@ -1060,59 +905,6 @@ func (c *Realm) Build() (proto.Message, error) {
}, nil
}
type UDPHop struct {
Mode string `json:"mode"`
Interval Int32Range `json:"interval"`
RemoteIPs []string `json:"remoteIPs"`
RemotePorts PortList `json:"remotePorts"`
}
func (c *UDPHop) Build() (proto.Message, error) {
var local, remote, remoteOnce bool
for _, mode := range strings.Split(c.Mode, ",") {
switch strings.ToLower(mode) {
case "intervallocal":
local = true
case "intervalremote":
remote = true
case "perconnremote":
remoteOnce = true
default:
return nil, errors.New("invalid mode ", mode)
}
}
var remoteIPs []string
for _, ip := range c.RemoteIPs {
prefix, err := netip.ParsePrefix(ip)
if err == nil {
remoteIPs = append(remoteIPs, prefix.String())
continue
}
addr, err := netip.ParseAddr(ip)
if err == nil {
remoteIPs = append(remoteIPs, netip.PrefixFrom(addr, addr.BitLen()).String())
continue
}
return nil, errors.New("invalid ip ", ip)
}
interval := c.Interval
if interval.From == 0 && interval.To == 0 {
interval.From, interval.To = 30, 30
}
if interval.From < 5 {
return nil, errors.New("interval must be at least 5")
}
return &udphop.Config{
Local: local,
Remote: remote,
RemoteOnce: remoteOnce,
IntervalMin: int64(interval.From),
IntervalMax: int64(interval.To),
RemoteIPs: remoteIPs,
RemotePorts: c.RemotePorts.Build().Ports(),
}, nil
}
type Mask struct {
Type string `json:"type"`
Settings *json.RawMessage `json:"settings"`
@@ -1146,6 +938,7 @@ type QuicParamsConfig struct {
BrutalUp Bandwidth `json:"brutalUp"`
BrutalDown Bandwidth `json:"brutalDown"`
BrutalDisableLossCompensation bool `json:"brutalDisableLossCompensation"`
UdpHop UdpHop `json:"udpHop"`
InitStreamReceiveWindow uint64 `json:"initStreamReceiveWindow"`
MaxStreamReceiveWindow uint64 `json:"maxStreamReceiveWindow"`
InitConnectionReceiveWindow uint64 `json:"initConnectionReceiveWindow"`
@@ -1,136 +0,0 @@
package conf
import (
"encoding/json"
"testing"
"github.com/xtls/xray-core/transport/internet/finalmask/noise"
)
func expPacket(exp string) json.RawMessage {
b, _ := json.Marshal(exp)
return b
}
func buildNoiseExp(exp string) (*noise.Config, error) {
msg, err := (&NoiseMask{Noise: []NoiseItem{{Type: "exp", Packet: expPacket(exp)}}}).Build()
if err != nil {
return nil, err
}
return msg.(*noise.Config), nil
}
func TestNoiseExp(t *testing.T) {
cfg, err := buildNoiseExp("<b 0d0a0d0a><t><r 24><rc 20-40><rd 8><c><n>")
if err != nil {
t.Fatal(err)
}
segments := cfg.Items[0].Segments
if len(segments) != 7 {
t.Fatalf("got %d segments, want 7", len(segments))
}
want := []struct {
kind noise.Segment_Kind
bytes []byte
min, max int64
}{
{noise.Segment_BYTES, []byte{0x0d, 0x0a, 0x0d, 0x0a}, 0, 0},
{noise.Segment_TIMESTAMP, nil, 0, 0},
{noise.Segment_RANDOM, nil, 24, 24},
{noise.Segment_RANDOM_ASCII, nil, 20, 40},
{noise.Segment_RANDOM_DIGIT, nil, 8, 8},
{noise.Segment_COUNTER, nil, 0, 0},
{noise.Segment_NONCE, nil, 0, 0},
}
for i, w := range want {
s := segments[i]
if s.Kind != w.kind || s.MinSize != w.min || s.MaxSize != w.max || string(s.Bytes) != string(w.bytes) {
t.Errorf("segment %d = %+v, want %+v", i, s, w)
}
}
}
func TestNoiseExpStripsHexPrefix(t *testing.T) {
cfg, err := buildNoiseExp("<b 0x16030100>")
if err != nil {
t.Fatal(err)
}
if got := cfg.Items[0].Segments[0].Bytes; string(got) != string([]byte{0x16, 0x03, 0x01, 0x00}) {
t.Errorf("got %x", got)
}
}
func TestNoiseExpWhitespace(t *testing.T) {
if _, err := buildNoiseExp(" <b 00> <t> "); err != nil {
t.Errorf("surrounding whitespace should be allowed: %v", err)
}
cfg, err := buildNoiseExp("<b 0d 0a 0d 0a>")
if err != nil {
t.Fatal(err)
}
if got := cfg.Items[0].Segments[0].Bytes; string(got) != "\r\n\r\n" {
t.Errorf("got %x", got)
}
}
func TestNoiseExpRejects(t *testing.T) {
for _, exp := range []string{
"<x 1>",
"<b>",
"<b zz>",
"<b 0d0>",
"<r>",
"<r -1>",
"<r 40-20>",
"<r 70000>",
"<t 5>",
"<n 5>",
"garbage<t>",
"<t> tail",
"<t><b>",
} {
if _, err := buildNoiseExp(exp); err == nil {
t.Errorf("expected an error for %q", exp)
}
}
}
func TestNoiseExpConflicts(t *testing.T) {
if _, err := (&NoiseMask{Noise: []NoiseItem{{Type: "exp", Packet: expPacket("<t>"), Rand: Int32Range{From: 10, To: 20}}}}).Build(); err == nil {
t.Error("exp with rand should be rejected")
}
for _, packet := range []string{``, `[1, 2]`, `5`} {
if _, err := (&NoiseMask{Noise: []NoiseItem{{Type: "exp", Packet: json.RawMessage(packet)}}}).Build(); err == nil {
t.Errorf("expected an error for packet %q", packet)
}
}
}
func TestNoiseExpFromJSON(t *testing.T) {
var mask NoiseMask
if err := json.Unmarshal([]byte(`{"noise": [
{"type": "exp", "packet": "<b 504f5354><rd 10-20>", "delay": "1-3"},
{"type": "EXP", "packet": "<t>"},
{"type": "str", "packet": "<t>"},
{"rand": "10-20"}
]}`), &mask); err != nil {
t.Fatal(err)
}
msg, err := mask.Build()
if err != nil {
t.Fatal(err)
}
items := msg.(*noise.Config).Items
if len(items[0].Segments) != 2 || items[0].DelayMin != 1 || items[0].DelayMax != 3 {
t.Errorf("item 0 = %+v", items[0])
}
if len(items[1].Segments) != 1 || items[1].Segments[0].Kind != noise.Segment_TIMESTAMP {
t.Errorf("item 1 = %+v", items[1])
}
if len(items[2].Segments) != 0 || string(items[2].Packet) != "<t>" {
t.Errorf("item 2 = %+v", items[2])
}
if len(items[3].Segments) != 0 || items[3].RandMin != 10 || items[3].RandMax != 20 {
t.Errorf("item 3 = %+v", items[3])
}
}
+20 -37
View File
@@ -36,10 +36,6 @@ func (p TransportProtocol) Build() (string, error) {
return "", errors.PrintRemovedFeatureError("QUIC transport (without web service, etc.)", "XHTTP stream-one H3")
case "hysteria":
return "hysteria", nil
case "masque":
return "masque", nil
case "xdrive":
return "xdrive", nil
default:
return "", errors.New("Config: unknown transport protocol: ", p)
}
@@ -63,8 +59,6 @@ type StreamConfig struct {
WSSettings *WebSocketConfig `json:"wsSettings"`
HTTPUPGRADESettings *HttpUpgradeConfig `json:"httpupgradeSettings"`
HysteriaSettings *HysteriaConfig `json:"hysteriaSettings"`
MASQUESettings *MasqueConfig `json:"masqueSettings"`
XDRIVESettings *XDriveConfig `json:"xdriveSettings"`
SocketSettings *SocketConfig `json:"sockopt"`
}
@@ -198,26 +192,6 @@ func (c *StreamConfig) Build() (*internet.StreamConfig, error) {
Settings: serial.ToTypedMessage(hs),
})
}
if c.MASQUESettings != nil {
ms, err := c.MASQUESettings.Build()
if err != nil {
return nil, errors.New("Failed to build MASQUE config.").Base(err)
}
config.TransportSettings = append(config.TransportSettings, &internet.TransportConfig{
ProtocolName: "masque",
Settings: serial.ToTypedMessage(ms),
})
}
if c.XDRIVESettings != nil {
xs, err := c.XDRIVESettings.Build()
if err != nil {
return nil, errors.New("Failed to build XDRIVE config.").Base(err)
}
config.TransportSettings = append(config.TransportSettings, &internet.TransportConfig{
ProtocolName: "xdrive",
Settings: serial.ToTypedMessage(xs),
})
}
if c.SocketSettings != nil {
ss, err := c.SocketSettings.Build()
if err != nil {
@@ -279,6 +253,10 @@ func (c *StreamConfig) Build() (*internet.StreamConfig, error) {
return nil, errors.New("unknown congestion control: ", c.FinalMask.QuicParams.Congestion, ", valid values: reno, bbr, brutal, force-brutal")
}
if (c.FinalMask.QuicParams.UdpHop.Interval.From != 0 && c.FinalMask.QuicParams.UdpHop.Interval.From < 5) || (c.FinalMask.QuicParams.UdpHop.Interval.To != 0 && c.FinalMask.QuicParams.UdpHop.Interval.To < 5) {
return nil, errors.New("Interval must be at least 5")
}
if c.FinalMask.QuicParams.InitStreamReceiveWindow > 0 && c.FinalMask.QuicParams.InitStreamReceiveWindow < 16384 {
return nil, errors.New("InitStreamReceiveWindow must be at least 16384")
}
@@ -312,17 +290,22 @@ func (c *StreamConfig) Build() (*internet.StreamConfig, error) {
BrutalUp: up,
BrutalDown: down,
BrutalDisableLossCompensation: c.FinalMask.QuicParams.BrutalDisableLossCompensation,
InitStreamReceiveWindow: c.FinalMask.QuicParams.InitStreamReceiveWindow,
MaxStreamReceiveWindow: c.FinalMask.QuicParams.MaxStreamReceiveWindow,
InitConnReceiveWindow: c.FinalMask.QuicParams.InitConnectionReceiveWindow,
MaxConnReceiveWindow: c.FinalMask.QuicParams.MaxConnectionReceiveWindow,
MaxIdleTimeout: c.FinalMask.QuicParams.MaxIdleTimeout,
KeepAlivePeriod: c.FinalMask.QuicParams.KeepAlivePeriod,
DisablePathMtuDiscovery: c.FinalMask.QuicParams.DisablePathMTUDiscovery,
DisableChromeParrot: c.FinalMask.QuicParams.DisableChromeParrot,
DisableGSO: c.FinalMask.QuicParams.DisableGSO,
MaxIncomingStreams: c.FinalMask.QuicParams.MaxIncomingStreams,
DisableStatelessReset: c.FinalMask.QuicParams.DisableStatelessReset,
UdpHop: &internet.UdpHop{
Ports: c.FinalMask.QuicParams.UdpHop.PortList.Build().Ports(),
IntervalMin: int64(c.FinalMask.QuicParams.UdpHop.Interval.From),
IntervalMax: int64(c.FinalMask.QuicParams.UdpHop.Interval.To),
},
InitStreamReceiveWindow: c.FinalMask.QuicParams.InitStreamReceiveWindow,
MaxStreamReceiveWindow: c.FinalMask.QuicParams.MaxStreamReceiveWindow,
InitConnReceiveWindow: c.FinalMask.QuicParams.InitConnectionReceiveWindow,
MaxConnReceiveWindow: c.FinalMask.QuicParams.MaxConnectionReceiveWindow,
MaxIdleTimeout: c.FinalMask.QuicParams.MaxIdleTimeout,
KeepAlivePeriod: c.FinalMask.QuicParams.KeepAlivePeriod,
DisablePathMtuDiscovery: c.FinalMask.QuicParams.DisablePathMTUDiscovery,
DisableChromeParrot: c.FinalMask.QuicParams.DisableChromeParrot,
DisableGSO: c.FinalMask.QuicParams.DisableGSO,
MaxIncomingStreams: c.FinalMask.QuicParams.MaxIncomingStreams,
DisableStatelessReset: c.FinalMask.QuicParams.DisableStatelessReset,
}
}
}
+30 -119
View File
@@ -1,9 +1,8 @@
package conf
import (
"encoding/base64"
"context"
"encoding/json"
"maps"
"math/big"
"net/url"
"sort"
@@ -22,12 +21,9 @@ import (
"github.com/xtls/xray-core/transport/internet/httpupgrade"
"github.com/xtls/xray-core/transport/internet/hysteria"
"github.com/xtls/xray-core/transport/internet/kcp"
"github.com/xtls/xray-core/transport/internet/masque"
"github.com/xtls/xray-core/transport/internet/splithttp"
"github.com/xtls/xray-core/transport/internet/tcp"
"github.com/xtls/xray-core/transport/internet/websocket"
"github.com/xtls/xray-core/transport/internet/xdrive"
"golang.org/x/net/http/httpguts"
"google.golang.org/protobuf/proto"
)
@@ -126,7 +122,7 @@ func (v *AuthenticatorRequest) Build() (*http.RequestConfig, error) {
for _, key := range headerNames {
value := v.Headers[key]
if value == nil {
return nil, errors.New("empty HTTP header value: " + key)
return nil, errors.New("empty HTTP header value: " + key).AtError()
}
config.Header = append(config.Header, &http.Header{
Name: key,
@@ -194,7 +190,7 @@ func (v *AuthenticatorResponse) Build() (*http.ResponseConfig, error) {
for _, key := range headerNames {
value := v.Headers[key]
if value == nil {
return nil, errors.New("empty HTTP header value: " + key)
return nil, errors.New("empty HTTP header value: " + key).AtError()
}
config.Header = append(config.Header, &http.Header{
Name: key,
@@ -244,11 +240,11 @@ func (c *TCPConfig) Build() (proto.Message, error) {
if len(c.HeaderConfig) > 0 {
headerConfig, _, err := tcpHeaderLoader.Load(c.HeaderConfig)
if err != nil {
return nil, errors.New("invalid TCP header config").Base(err)
return nil, errors.New("invalid TCP header config").Base(err).AtError()
}
ts, err := headerConfig.(Buildable).Build()
if err != nil {
return nil, errors.New("invalid TCP header config").Base(err)
return nil, errors.New("invalid TCP header config").Base(err).AtError()
}
config.HeaderSettings = serial.ToTypedMessage(ts)
}
@@ -538,6 +534,10 @@ type KCPConfig struct {
// Build implements Buildable.
func (c *KCPConfig) Build() (proto.Message, error) {
if c.HeaderConfig != nil || c.Seed != nil {
return nil, errors.PrintRemovedFeatureError("mkcp header & seed", "finalmask/udp header-* & mkcp-original & mkcp-aes128gcm")
}
config := common.Must2(internet.CreateTransportConfig(kcp.ProtocolName)).(*kcp.Config)
if c.Mtu != nil {
@@ -560,16 +560,16 @@ func (c *KCPConfig) Build() (proto.Message, error) {
}
if config.Mtu < 21 {
return nil, errors.New("MTU must be at least 21")
return nil, errors.New("Mtu must be at least 21").AtError()
}
if config.Tti < 10 || config.Tti > 1000 {
return nil, errors.New("TTI must be between 10 and 1000")
return nil, errors.New("invalid mKCP TTI: ", c.Tti).AtError()
}
if config.CwndMultiplier < 1 {
return nil, errors.New("CwndMultiplier must be at least 1")
return nil, errors.New("CwndMultiplier must be at least 1").AtError()
}
if config.GetSendingBufferSize() == 0 {
return nil, errors.New("MaxSendingWindow must be at least ", config.Mtu)
return nil, errors.New("MaxSendingWindow must be >= Mtu").AtError()
}
return config, nil
@@ -739,6 +739,11 @@ func (b Bandwidth) Bps() (uint64, error) {
return uint64(val*float64(mul)) / 8, nil
}
type UdpHop struct {
PortList PortList `json:"ports"`
Interval Int32Range `json:"interval"`
}
type Masquerade struct {
Type string `json:"type"`
@@ -755,8 +760,14 @@ type Masquerade struct {
}
type HysteriaConfig struct {
Version int32 `json:"version"`
Auth string `json:"auth"`
Version int32 `json:"version"`
Auth string `json:"auth"`
Congestion *string `json:"congestion"`
Up *Bandwidth `json:"up"`
Down *Bandwidth `json:"down"`
UdpHop *UdpHop `json:"udphop"`
UdpIdleTimeout int64 `json:"udpIdleTimeout"`
Masquerade Masquerade `json:"masquerade"`
}
@@ -766,6 +777,10 @@ func (c *HysteriaConfig) Build() (proto.Message, error) {
return nil, errors.New("version != 2")
}
if c.Congestion != nil || c.Up != nil || c.Down != nil || c.UdpHop != nil {
errors.LogWarning(context.Background(), "congestion & up & down & udphop move to finalmask/quicParams")
}
if c.UdpIdleTimeout != 0 && (c.UdpIdleTimeout < 2 || c.UdpIdleTimeout > 600) {
return nil, errors.New("UdpIdleTimeout must be between 2 and 600")
}
@@ -790,63 +805,6 @@ func (c *HysteriaConfig) Build() (proto.Message, error) {
return config, nil
}
type MasqueConfig struct {
Host string `json:"host"`
Path string `json:"path"`
User string `json:"user"`
Pass string `json:"pass"`
Headers map[string]string `json:"headers"`
}
func (c *MasqueConfig) Build() (proto.Message, error) {
path := c.Path
if path == "" {
path = masque.DefaultPath
}
path = strings.NewReplacer(
"{target}", "*", "{ipproto}", "*",
"{?target,ipproto}", "?target=*&ipproto=*", "{?ipproto,target}", "?ipproto=*&target=*",
"{&target,ipproto}", "&target=*&ipproto=*", "{&ipproto,target}", "&ipproto=*&target=*",
).Replace(path)
if !strings.HasPrefix(path, "/") || strings.ContainsAny(path, "{}") {
return nil, errors.New(`invalid "path": `, path, `, only the variables {target} and {ipproto} are supported`)
}
if c.Host != "" {
if u, err := url.Parse("https://" + c.Host); err != nil || u.Host != c.Host {
return nil, errors.New(`invalid "host": `, c.Host)
}
}
for k, v := range c.Headers {
if !httpguts.ValidHeaderFieldName(k) || !httpguts.ValidHeaderFieldValue(v) {
return nil, errors.New(`invalid header in "headers": `, strconv.Quote(k))
}
switch strings.ToLower(k) {
case "host", "capsule-protocol":
return nil, errors.New(`"headers" can't contain "`, k, `"`)
case "authorization":
if c.User != "" || c.Pass != "" {
return nil, errors.New(`"headers" can't contain "`, k, `" when "user" or "pass" is set`)
}
}
}
headers := c.Headers
if c.User != "" || c.Pass != "" {
if strings.Contains(c.User, ":") {
return nil, errors.New(`invalid "user": `, c.User)
}
headers = maps.Clone(c.Headers)
if headers == nil {
headers = make(map[string]string)
}
headers["Authorization"] = "Basic " + base64.StdEncoding.EncodeToString([]byte(c.User+":"+c.Pass))
}
return &masque.Config{
Host: c.Host,
Path: path,
Headers: headers,
}, nil
}
func readFileOrString(f string, s []string) ([]byte, error) {
if len(f) > 0 {
return filesystem.ReadCert(f)
@@ -856,50 +814,3 @@ func readFileOrString(f string, s []string) ([]byte, error) {
}
return nil, errors.New("both file and bytes are empty.")
}
type XDriveConfig struct {
RemoteFolder string `json:"remoteFolder"`
Service string `json:"service"`
Secrets []string `json:"secrets"`
SegmentBytes uint32 `json:"segmentBytes"`
FlushIntervalMs uint32 `json:"flushIntervalMs"`
PollIntervalMs uint32 `json:"pollIntervalMs"`
MaxPollIntervalMs uint32 `json:"maxPollIntervalMs"`
SessionTTLSeconds uint32 `json:"sessionTtlSeconds"`
Concurrency uint32 `json:"concurrency"`
EagerWindowMs uint32 `json:"eagerWindowMs"`
HoleTimeoutMs uint32 `json:"holeTimeoutMs"`
Template json.RawMessage `json:"template"`
}
// Build implements Buildable.
func (c *XDriveConfig) Build() (proto.Message, error) {
switch c.Service {
case "local":
case "Google Drive":
if len(c.Secrets) != 3 {
return nil, errors.New("Google Drive needs 3 secrets in order of ClientID, ClientSecret, RefreshToken")
}
case "template":
if len(c.Template) == 0 {
return nil, errors.New(`service "template" needs a "template" object`)
}
default:
return nil, errors.New("unsupported service")
}
config := &xdrive.Config{
RemoteFolder: c.RemoteFolder,
Service: c.Service,
Secrets: c.Secrets,
SegmentBytes: c.SegmentBytes,
FlushIntervalMs: c.FlushIntervalMs,
PollIntervalMs: c.PollIntervalMs,
MaxPollIntervalMs: c.MaxPollIntervalMs,
SessionTtlSeconds: c.SessionTTLSeconds,
Concurrency: c.Concurrency,
EagerWindowMs: c.EagerWindowMs,
HoleTimeoutMs: c.HoleTimeoutMs,
Template: string(c.Template),
}
return config, nil
}

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