mirror of
https://github.com/XTLS/Xray-core.git
synced 2026-10-10 09:35:41 +00:00
Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
adff74795e |
@@ -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
|
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 {
|
} else {
|
||||||
if err != common.ErrNoClue {
|
|
||||||
errors.LogErrorInner(ctx, err, "failed to pick route for ", destination)
|
|
||||||
}
|
|
||||||
errors.LogInfo(ctx, "default route for ", destination)
|
errors.LogInfo(ctx, "default route for ", destination)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
+4
-23
@@ -93,7 +93,6 @@ type NameServer struct {
|
|||||||
UnexpectedIp []*geodata.IPRule `protobuf:"bytes,13,rep,name=unexpected_ip,json=unexpectedIp,proto3" json:"unexpected_ip,omitempty"`
|
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"`
|
ActUnprior bool `protobuf:"varint,14,opt,name=actUnprior,proto3" json:"actUnprior,omitempty"`
|
||||||
PolicyID uint32 `protobuf:"varint,17,opt,name=policyID,proto3" json:"policyID,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
|
unknownFields protoimpl.UnknownFields
|
||||||
sizeCache protoimpl.SizeCache
|
sizeCache protoimpl.SizeCache
|
||||||
}
|
}
|
||||||
@@ -240,13 +239,6 @@ func (x *NameServer) GetPolicyID() uint32 {
|
|||||||
return 0
|
return 0
|
||||||
}
|
}
|
||||||
|
|
||||||
func (x *NameServer) GetId() string {
|
|
||||||
if x != nil {
|
|
||||||
return x.Id
|
|
||||||
}
|
|
||||||
return ""
|
|
||||||
}
|
|
||||||
|
|
||||||
type Config struct {
|
type Config struct {
|
||||||
state protoimpl.MessageState `protogen:"open.v1"`
|
state protoimpl.MessageState `protogen:"open.v1"`
|
||||||
// NameServer list used by this DNS client.
|
// NameServer list used by this DNS client.
|
||||||
@@ -266,8 +258,6 @@ type Config struct {
|
|||||||
DisableFallback bool `protobuf:"varint,10,opt,name=disableFallback,proto3" json:"disableFallback,omitempty"`
|
DisableFallback bool `protobuf:"varint,10,opt,name=disableFallback,proto3" json:"disableFallback,omitempty"`
|
||||||
DisableFallbackIfMatch bool `protobuf:"varint,11,opt,name=disableFallbackIfMatch,proto3" json:"disableFallbackIfMatch,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"`
|
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
|
unknownFields protoimpl.UnknownFields
|
||||||
sizeCache protoimpl.SizeCache
|
sizeCache protoimpl.SizeCache
|
||||||
}
|
}
|
||||||
@@ -379,13 +369,6 @@ func (x *Config) GetEnableParallelQuery() bool {
|
|||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
|
|
||||||
func (x *Config) GetScript() string {
|
|
||||||
if x != nil {
|
|
||||||
return x.Script
|
|
||||||
}
|
|
||||||
return ""
|
|
||||||
}
|
|
||||||
|
|
||||||
type Config_HostMapping struct {
|
type Config_HostMapping struct {
|
||||||
state protoimpl.MessageState `protogen:"open.v1"`
|
state protoimpl.MessageState `protogen:"open.v1"`
|
||||||
Domain *geodata.DomainRule `protobuf:"bytes,2,opt,name=domain,proto3" json:"domain,omitempty"`
|
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 = "" +
|
const file_app_dns_config_proto_rawDesc = "" +
|
||||||
"\n" +
|
"\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" +
|
"\n" +
|
||||||
"NameServer\x123\n" +
|
"NameServer\x123\n" +
|
||||||
"\aaddress\x18\x01 \x01(\v2\x19.xray.common.net.EndpointR\aaddress\x12\x1b\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" +
|
"\n" +
|
||||||
"actUnprior\x18\x0e \x01(\bR\n" +
|
"actUnprior\x18\x0e \x01(\bR\n" +
|
||||||
"actUnprior\x12\x1a\n" +
|
"actUnprior\x12\x1a\n" +
|
||||||
"\bpolicyID\x18\x11 \x01(\rR\bpolicyID\x12\x0e\n" +
|
"\bpolicyID\x18\x11 \x01(\rR\bpolicyIDB\x0f\n" +
|
||||||
"\x02id\x18\x12 \x01(\tR\x02idB\x0f\n" +
|
|
||||||
"\r_disableCacheB\r\n" +
|
"\r_disableCacheB\r\n" +
|
||||||
"\v_serveStaleB\x12\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" +
|
"\x06Config\x129\n" +
|
||||||
"\vname_server\x18\x05 \x03(\v2\x18.xray.app.dns.NameServerR\n" +
|
"\vname_server\x18\x05 \x03(\v2\x18.xray.app.dns.NameServerR\n" +
|
||||||
"nameServer\x12\x1b\n" +
|
"nameServer\x12\x1b\n" +
|
||||||
@@ -498,8 +480,7 @@ const file_app_dns_config_proto_rawDesc = "" +
|
|||||||
"\x0fdisableFallback\x18\n" +
|
"\x0fdisableFallback\x18\n" +
|
||||||
" \x01(\bR\x0fdisableFallback\x126\n" +
|
" \x01(\bR\x0fdisableFallback\x126\n" +
|
||||||
"\x16disableFallbackIfMatch\x18\v \x01(\bR\x16disableFallbackIfMatch\x120\n" +
|
"\x16disableFallbackIfMatch\x18\v \x01(\bR\x16disableFallbackIfMatch\x120\n" +
|
||||||
"\x13enableParallelQuery\x18\x0e \x01(\bR\x13enableParallelQuery\x12\x16\n" +
|
"\x13enableParallelQuery\x18\x0e \x01(\bR\x13enableParallelQuery\x1a}\n" +
|
||||||
"\x06script\x18\x0f \x01(\tR\x06script\x1a}\n" +
|
|
||||||
"\vHostMapping\x127\n" +
|
"\vHostMapping\x127\n" +
|
||||||
"\x06domain\x18\x02 \x01(\v2\x1f.xray.common.geodata.DomainRuleR\x06domain\x12\x0e\n" +
|
"\x06domain\x18\x02 \x01(\v2\x1f.xray.common.geodata.DomainRuleR\x06domain\x12\x0e\n" +
|
||||||
"\x02ip\x18\x03 \x03(\fR\x02ip\x12%\n" +
|
"\x02ip\x18\x03 \x03(\fR\x02ip\x12%\n" +
|
||||||
|
|||||||
@@ -27,7 +27,6 @@ message NameServer {
|
|||||||
repeated xray.common.geodata.IPRule unexpected_ip = 13;
|
repeated xray.common.geodata.IPRule unexpected_ip = 13;
|
||||||
bool actUnprior = 14;
|
bool actUnprior = 14;
|
||||||
uint32 policyID = 17;
|
uint32 policyID = 17;
|
||||||
string id = 18;
|
|
||||||
}
|
}
|
||||||
|
|
||||||
enum QueryStrategy {
|
enum QueryStrategy {
|
||||||
@@ -74,7 +73,4 @@ message Config {
|
|||||||
bool disableFallbackIfMatch = 11;
|
bool disableFallbackIfMatch = 11;
|
||||||
|
|
||||||
bool enableParallelQuery = 14;
|
bool enableParallelQuery = 14;
|
||||||
|
|
||||||
// Absolute path to the Lua DNS query script.
|
|
||||||
string script = 15;
|
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -31,8 +31,6 @@ type DNS struct {
|
|||||||
domainMatcher geodata.DomainMatcher
|
domainMatcher geodata.DomainMatcher
|
||||||
matcherInfos []*DomainMatcherInfo
|
matcherInfos []*DomainMatcherInfo
|
||||||
checkSystem bool
|
checkSystem bool
|
||||||
script *scriptEngine
|
|
||||||
scriptPath string
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// DomainMatcherInfo contains information attached to index returned by Server.domainMatcher.
|
// 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,
|
disableFallbackIfMatch: config.DisableFallbackIfMatch,
|
||||||
enableParallelQuery: config.EnableParallelQuery,
|
enableParallelQuery: config.EnableParallelQuery,
|
||||||
checkSystem: checkSystem,
|
checkSystem: checkSystem,
|
||||||
scriptPath: config.Script,
|
|
||||||
}, nil
|
}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -193,21 +190,11 @@ func (*DNS) Type() interface{} {
|
|||||||
|
|
||||||
// Start implements common.Runnable.
|
// Start implements common.Runnable.
|
||||||
func (s *DNS) Start() error {
|
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
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// Close implements common.Closable.
|
// Close implements common.Closable.
|
||||||
func (s *DNS) Close() error {
|
func (s *DNS) Close() error {
|
||||||
if s.script != nil {
|
|
||||||
s.script.close()
|
|
||||||
}
|
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -292,9 +279,6 @@ func (s *DNS) LookupIP(domain string, option dns.IPOption) ([]net.IP, uint32, er
|
|||||||
}
|
}
|
||||||
|
|
||||||
// Name servers lookup
|
// Name servers lookup
|
||||||
if s.script != nil {
|
|
||||||
return s.script.query(domain, option)
|
|
||||||
}
|
|
||||||
if s.enableParallelQuery {
|
if s.enableParallelQuery {
|
||||||
return s.parallelQuery(domain, option)
|
return s.parallelQuery(domain, option)
|
||||||
} else {
|
} else {
|
||||||
|
|||||||
-169
@@ -1,169 +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 {
|
|
||||||
pushIPs := xlua.NewSlicePusher[net.IP](L)
|
|
||||||
|
|
||||||
serverList := L.CreateTable(len(servers), 0)
|
|
||||||
for i, client := range servers {
|
|
||||||
server := L.CreateTable(0, 2)
|
|
||||||
|
|
||||||
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)
|
|
||||||
}
|
|
||||||
pushIPs(L, ips)
|
|
||||||
xlua.PushNumber(L, ttl)
|
|
||||||
xlua.PushError(L, err)
|
|
||||||
return 3
|
|
||||||
}))
|
|
||||||
serverList.RawSetInt(i+1, server)
|
|
||||||
}
|
|
||||||
|
|
||||||
module := L.CreateTable(0, 2)
|
|
||||||
if servers != nil {
|
|
||||||
module.RawSetString("Servers", serverList)
|
|
||||||
}
|
|
||||||
if client != nil {
|
|
||||||
module.RawSetString("Query", newLuaClientQuery(L, client, pushIPs))
|
|
||||||
}
|
|
||||||
L.Push(module)
|
|
||||||
return 1
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
func newLuaClientQuery(L *lua.LState, client featureDNS.Client, pushIPs func(*lua.LState, []net.IP)) *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)
|
|
||||||
pushIPs(L, ips)
|
|
||||||
xlua.PushNumber(L, ttl)
|
|
||||||
xlua.PushError(L, err)
|
|
||||||
return 3
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
// callLuaQuery runs HandleDNSQuery and leaves (ips, ttl, err) on the stack.
|
|
||||||
func callLuaQuery(L *lua.LState, domain string, option featureDNS.IPOption) error {
|
|
||||||
fn := L.GetGlobal("HandleDNSQuery")
|
|
||||||
if fn.Type() != lua.LTFunction {
|
|
||||||
return errors.New("DNS script must define HandleDNSQuery(...)")
|
|
||||||
}
|
|
||||||
|
|
||||||
return 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))
|
|
||||||
}
|
|
||||||
|
|
||||||
// readLuaQueryResult reads (ips, ttl, err) from the stack without copying the IPs.
|
|
||||||
func readLuaQueryResult(L *lua.LState) ([]net.IP, uint32, error) {
|
|
||||||
if err := xlua.ReadError(L.Get(-1), "DNS script error must be an error or string"); err != nil {
|
|
||||||
return nil, 0, err
|
|
||||||
}
|
|
||||||
|
|
||||||
ttl, err := xlua.ReadUint32(L.Get(-2), "DNS script returned invalid TTL")
|
|
||||||
if err != nil {
|
|
||||||
return nil, 0, err
|
|
||||||
}
|
|
||||||
|
|
||||||
addresses := L.Get(-3)
|
|
||||||
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
|
|
||||||
}
|
|
||||||
@@ -1,118 +0,0 @@
|
|||||||
package dns
|
|
||||||
|
|
||||||
import (
|
|
||||||
"context"
|
|
||||||
"os"
|
|
||||||
"path/filepath"
|
|
||||||
"testing"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
"github.com/xtls/xray-core/common/net"
|
|
||||||
featureDNS "github.com/xtls/xray-core/features/dns"
|
|
||||||
lua "github.com/yuin/gopher-lua"
|
|
||||||
)
|
|
||||||
|
|
||||||
// BenchmarkLuaDNSHook isolates scalar argument bridging and a fixed return.
|
|
||||||
// It excludes upstream queries, result decoding, and state pool management.
|
|
||||||
func BenchmarkLuaDNSHook(b *testing.B) {
|
|
||||||
L := lua.NewState()
|
|
||||||
b.Cleanup(L.Close)
|
|
||||||
if err := L.DoString(`
|
|
||||||
function HandleDNSQuery(domain, ipv4, ipv6, fake)
|
|
||||||
return true
|
|
||||||
end
|
|
||||||
`); err != nil {
|
|
||||||
b.Fatal(err)
|
|
||||||
}
|
|
||||||
L.SetContext(context.Background())
|
|
||||||
option := featureDNS.IPOption{IPv4Enable: true}
|
|
||||||
if err := callLuaQuery(L, "example.com", option); err != nil {
|
|
||||||
b.Fatal(err)
|
|
||||||
}
|
|
||||||
if L.Get(-3) != lua.LTrue {
|
|
||||||
b.Fatal("hook did not return true")
|
|
||||||
}
|
|
||||||
L.Pop(3)
|
|
||||||
b.ReportAllocs()
|
|
||||||
b.ResetTimer()
|
|
||||||
for i := 0; i < b.N; i++ {
|
|
||||||
if err := callLuaQuery(L, "example.com", option); err != nil {
|
|
||||||
b.Fatal(err)
|
|
||||||
}
|
|
||||||
L.Pop(3)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// BenchmarkLuaDNSQuery queries the same preselected, in-memory upstream.
|
|
||||||
// client_query compares Client.QueryIP to a preloaded server:Query hook.
|
|
||||||
// script_query additionally measures production pool and timeout management.
|
|
||||||
// These cases do not measure DNS.LookupIP server selection or network latency.
|
|
||||||
func BenchmarkLuaDNSQuery(b *testing.B) {
|
|
||||||
ctx := context.Background()
|
|
||||||
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{ctx: ctx, clients: []*Client{client}}
|
|
||||||
const script = `
|
|
||||||
local server = require("xray.dns").Servers[1]
|
|
||||||
function HandleDNSQuery(domain, ipv4, ipv6, fake)
|
|
||||||
return server:Query(domain, ipv4, ipv6, fake)
|
|
||||||
end
|
|
||||||
`
|
|
||||||
L := lua.NewState()
|
|
||||||
b.Cleanup(L.Close)
|
|
||||||
server.registerLua(L)
|
|
||||||
if err := L.DoString(script); err != nil {
|
|
||||||
b.Fatal(err)
|
|
||||||
}
|
|
||||||
L.SetContext(ctx)
|
|
||||||
|
|
||||||
path := filepath.Join(b.TempDir(), "query.lua")
|
|
||||||
if err := os.WriteFile(path, []byte(script), 0o600); err != nil {
|
|
||||||
b.Fatal(err)
|
|
||||||
}
|
|
||||||
engine, err := newScriptEngine(path, server)
|
|
||||||
if err != nil {
|
|
||||||
b.Fatal(err)
|
|
||||||
}
|
|
||||||
b.Cleanup(engine.close)
|
|
||||||
for _, bench := range []struct {
|
|
||||||
name string
|
|
||||||
query func() ([]net.IP, uint32, error)
|
|
||||||
}{
|
|
||||||
{"client_query/native", func() ([]net.IP, uint32, error) {
|
|
||||||
return client.QueryIP(ctx, "example.com", option)
|
|
||||||
}},
|
|
||||||
{"client_query/lua", func() ([]net.IP, uint32, error) {
|
|
||||||
if err := callLuaQuery(L, "example.com", option); err != nil {
|
|
||||||
return nil, 0, err
|
|
||||||
}
|
|
||||||
ips, ttl, err := readLuaQueryResult(L)
|
|
||||||
L.Pop(3)
|
|
||||||
return ips, ttl, err
|
|
||||||
}},
|
|
||||||
{"script_query/lua", func() ([]net.IP, uint32, error) {
|
|
||||||
return engine.query("example.com", option)
|
|
||||||
}},
|
|
||||||
} {
|
|
||||||
b.Run(bench.name, func(b *testing.B) {
|
|
||||||
ips, ttl, err := bench.query()
|
|
||||||
if err != nil || ttl != 60 || len(ips) != 1 || !ips[0].Equal(ip) {
|
|
||||||
b.Fatalf("query() = %v, TTL %d, %v; want %v, TTL 60", ips, ttl, err, ip)
|
|
||||||
}
|
|
||||||
b.ReportAllocs()
|
|
||||||
b.ResetTimer()
|
|
||||||
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)
|
|
||||||
}
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@@ -1,272 +0,0 @@
|
|||||||
package dns
|
|
||||||
|
|
||||||
import (
|
|
||||||
"context"
|
|
||||||
go_errors "errors"
|
|
||||||
"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 TestReadLuaQueryResult(t *testing.T) {
|
|
||||||
wantIPs := []net.IP{net.ParseIP("8.8.8.8"), {127, 0, 0, 1}, net.ParseIP("::1")}
|
|
||||||
nativeErr := go_errors.New("upstream failed")
|
|
||||||
for _, tc := range []struct {
|
|
||||||
name, values string
|
|
||||||
wantIPs []net.IP
|
|
||||||
wantTTL uint32
|
|
||||||
wantErr error
|
|
||||||
wantMessage string
|
|
||||||
}{
|
|
||||||
{name: "IPs", values: `ips, 45`, wantIPs: wantIPs, wantTTL: 45},
|
|
||||||
{name: "nil IPs", values: `nil, 0`, wantErr: featureDNS.ErrEmptyResponse},
|
|
||||||
{name: "empty IPs", values: `emptyIPs, 0`, wantErr: featureDNS.ErrEmptyResponse},
|
|
||||||
{name: "native error", values: `nil, nil, nativeError`, wantErr: nativeErr},
|
|
||||||
{name: "string error", values: `nil, nil, "blocked"`, wantMessage: "blocked"},
|
|
||||||
{name: "fractional TTL", values: `ips, 1.5`, wantMessage: "invalid TTL"},
|
|
||||||
{name: "oversized TTL", values: `ips, 4294967296`, wantMessage: "invalid TTL"},
|
|
||||||
{name: "negative TTL", values: `ips, -1`, wantMessage: "invalid TTL"},
|
|
||||||
{name: "NaN TTL", values: `ips, 0/0`, wantMessage: "invalid TTL"},
|
|
||||||
{name: "missing TTL", values: `ips`, wantMessage: "invalid TTL"},
|
|
||||||
{name: "string IPs", values: `"127.0.0.1", 60`, wantMessage: "native IP slice"},
|
|
||||||
{name: "wrong userdata", values: `ip, 60`, wantMessage: "native IP slice"},
|
|
||||||
{name: "invalid error", values: `ips, 60, false`, wantMessage: "error or string"},
|
|
||||||
} {
|
|
||||||
t.Run(tc.name, func(t *testing.T) {
|
|
||||||
L := lua.NewState()
|
|
||||||
defer L.Close()
|
|
||||||
for name, value := range map[string]any{"ips": wantIPs, "ip": wantIPs[0], "emptyIPs": []net.IP(nil), "nativeError": nativeErr} {
|
|
||||||
ud := L.NewUserData()
|
|
||||||
ud.Value = value
|
|
||||||
L.SetGlobal(name, ud)
|
|
||||||
}
|
|
||||||
fn, err := L.LoadString("return " + tc.values)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := L.CallByParam(lua.P{Fn: fn, NRet: 3, Protect: true}); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
ips, ttl, err := readLuaQueryResult(L)
|
|
||||||
switch {
|
|
||||||
case tc.wantErr != nil:
|
|
||||||
if err != tc.wantErr {
|
|
||||||
t.Fatalf("error = %v, want original error %v", err, tc.wantErr)
|
|
||||||
}
|
|
||||||
case tc.wantMessage != "":
|
|
||||||
if err == nil || !strings.Contains(err.Error(), tc.wantMessage) {
|
|
||||||
t.Fatalf("error = %v, want %q", err, tc.wantMessage)
|
|
||||||
}
|
|
||||||
case err != nil:
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if ttl != tc.wantTTL || len(ips) != len(tc.wantIPs) {
|
|
||||||
t.Fatalf("result = %v, TTL %d; want %v, TTL %d", ips, ttl, tc.wantIPs, tc.wantTTL)
|
|
||||||
}
|
|
||||||
for i := range ips {
|
|
||||||
if !ips[i].Equal(tc.wantIPs[i]) {
|
|
||||||
t.Fatalf("IP %d = %v, want %v", i, ips[i], tc.wantIPs[i])
|
|
||||||
}
|
|
||||||
}
|
|
||||||
if len(ips) != 0 && &ips[0] != &tc.wantIPs[0] {
|
|
||||||
t.Fatal("result copied the IP slice")
|
|
||||||
}
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestCallLuaQueryCancellation(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 := callLuaQuery(L, "example.com", featureDNS.IPOption{IPv4Enable: true})
|
|
||||||
if err == nil {
|
|
||||||
t.Fatal("callLuaQuery did not stop after context cancellation")
|
|
||||||
}
|
|
||||||
if L.Context() != ctx {
|
|
||||||
t.Fatal("callLuaQuery changed the Lua state's context")
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestCallLuaQuery(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)
|
|
||||||
}
|
|
||||||
if err := callLuaQuery(L, "ExAmPlE.CoM", featureDNS.IPOption{IPv4Enable: true}); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if L.GetTop() != 3 || L.Get(1) != addresses || L.Get(2) != lua.LNumber(60) || L.Get(3) != lua.LNil {
|
|
||||||
t.Fatal("callLuaQuery did not leave the three query results on 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(#ips == 2 and ips[1]:String() == "127.0.0.1" and ips[2]:String() == "8.8.8.8")
|
|
||||||
assert(matcher:Match(ips[1]) and not matcher:Match(ips[2]))
|
|
||||||
assert(matcher:AnyMatch(ips))
|
|
||||||
local matched, unmatched = matcher:FilterIPs(ips)
|
|
||||||
assert(#matched == 1 and #unmatched == 1)
|
|
||||||
assert(matched[1]:Equal(ips[1]) and unmatched[1]:Equal(ips[2]))
|
|
||||||
return matched, ttl, err
|
|
||||||
end
|
|
||||||
`); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
L.SetContext(context.Background())
|
|
||||||
if err := callLuaQuery(L, "example.com", option); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
got, ttl, err := readLuaQueryResult(L)
|
|
||||||
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}, net.ParseIP("::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))
|
|
||||||
assert(#ips == 2 and ips[1]:String() == "127.0.0.1" and ips[2]:String() == "::1")
|
|
||||||
assert(matcher:Match(ips[1]) and not matcher:Match(ips[2]))
|
|
||||||
`); 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)
|
|
||||||
assert(#serverIPs == 1 and #clientIPs == 1)
|
|
||||||
assert(serverIPs[1]:String() == "127.0.0.1" and serverIPs[1]:Equal(clientIPs[1]))
|
|
||||||
`); 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)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestLuaDNSQueryEmptyIPs(t *testing.T) {
|
|
||||||
for _, tc := range []struct {
|
|
||||||
name string
|
|
||||||
ips []net.IP
|
|
||||||
}{
|
|
||||||
{"nil", nil},
|
|
||||||
{"empty", []net.IP{}},
|
|
||||||
} {
|
|
||||||
t.Run(tc.name, func(t *testing.T) {
|
|
||||||
L := lua.NewState()
|
|
||||||
defer L.Close()
|
|
||||||
L.SetContext(context.Background())
|
|
||||||
L.SetGlobal("expectNil", lua.LBool(tc.ips == nil))
|
|
||||||
client := &luaDNSClient{lookup: func(string, featureDNS.IPOption) ([]net.IP, uint32, error) {
|
|
||||||
return tc.ips, 0, featureDNS.ErrEmptyResponse
|
|
||||||
}}
|
|
||||||
registerLua(L, []luaDNSServer{{query: func(_ context.Context, domain string, option featureDNS.IPOption) ([]net.IP, uint32, error) {
|
|
||||||
return client.LookupIP(domain, option)
|
|
||||||
}}}, client)
|
|
||||||
if err := L.DoString(`
|
|
||||||
local dns = require("xray.dns")
|
|
||||||
for _, query in ipairs({
|
|
||||||
function() return dns.Servers[1]:Query("empty.example", true, false, false) end,
|
|
||||||
function() return dns.Query("empty.example", true, false, false) end,
|
|
||||||
}) do
|
|
||||||
local ips, ttl, err = query()
|
|
||||||
assert(ttl == 0 and err)
|
|
||||||
if expectNil then
|
|
||||||
assert(ips == nil)
|
|
||||||
else
|
|
||||||
assert(type(ips) == "userdata" and #ips == 0)
|
|
||||||
assert(not pcall(function() return ips[1] end))
|
|
||||||
end
|
|
||||||
end
|
|
||||||
`); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
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
|
|
||||||
}
|
|
||||||
@@ -29,7 +29,6 @@ type Server interface {
|
|||||||
|
|
||||||
// Client is the interface for DNS client.
|
// Client is the interface for DNS client.
|
||||||
type Client struct {
|
type Client struct {
|
||||||
id string
|
|
||||||
server Server
|
server Server
|
||||||
skipFallback bool
|
skipFallback bool
|
||||||
expectedIPs geodata.IPMatcher
|
expectedIPs geodata.IPMatcher
|
||||||
@@ -98,7 +97,7 @@ func NewClient(
|
|||||||
ipOption dns.IPOption,
|
ipOption dns.IPOption,
|
||||||
updateRules func(bool),
|
updateRules func(bool),
|
||||||
) (*Client, error) {
|
) (*Client, error) {
|
||||||
client := &Client{id: ns.Id}
|
client := &Client{}
|
||||||
err := core.RequireFeatures(ctx, func(dispatcher routing.Dispatcher) error {
|
err := core.RequireFeatures(ctx, func(dispatcher routing.Dispatcher) error {
|
||||||
// Create a new server for each client for now
|
// Create a new server for each client for now
|
||||||
server, err := NewServer(ctx, ns.Address.AsDestination(), dispatcher, disableCache, serveStale, serveExpiredTTL, clientIP)
|
server, err := NewServer(ctx, ns.Address.AsDestination(), dispatcher, disableCache, serveStale, serveExpiredTTL, clientIP)
|
||||||
|
|||||||
@@ -49,5 +49,5 @@ func NewLocalNameServer() *LocalNameServer {
|
|||||||
|
|
||||||
// NewLocalDNSClient creates localdns client object for directly lookup in system DNS.
|
// NewLocalDNSClient creates localdns client object for directly lookup in system DNS.
|
||||||
func NewLocalDNSClient(ipOption dns.IPOption) *Client {
|
func NewLocalDNSClient(ipOption dns.IPOption) *Client {
|
||||||
return &Client{id: "localhost", server: NewLocalNameServer(), ipOption: &ipOption}
|
return &Client{server: NewLocalNameServer(), ipOption: &ipOption}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1,63 +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 {
|
|
||||||
pool *xlua.Pool
|
|
||||||
}
|
|
||||||
|
|
||||||
func newScriptEngine(path string, server *DNS) (*scriptEngine, error) {
|
|
||||||
program, err := xlua.CompileFile(path)
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
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 &scriptEngine{pool: pool}, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (e *scriptEngine) close() {
|
|
||||||
e.pool.Close()
|
|
||||||
}
|
|
||||||
|
|
||||||
func (e *scriptEngine) query(domain string, option dns.IPOption) (ips []net.IP, ttl uint32, queryErr error) {
|
|
||||||
if err := e.pool.WithState(nil, 0, func(L *lua.LState) error {
|
|
||||||
if err := callLuaQuery(L, domain, option); err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
ips, ttl, queryErr = readLuaQueryResult(L)
|
|
||||||
return nil
|
|
||||||
}); err != nil {
|
|
||||||
return nil, 0, err
|
|
||||||
}
|
|
||||||
return ips, ttl, queryErr
|
|
||||||
}
|
|
||||||
@@ -1,280 +0,0 @@
|
|||||||
package dns
|
|
||||||
|
|
||||||
import (
|
|
||||||
"context"
|
|
||||||
go_errors "errors"
|
|
||||||
"os"
|
|
||||||
"path/filepath"
|
|
||||||
"strings"
|
|
||||||
"testing"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
"github.com/xtls/xray-core/common/net"
|
|
||||||
featureDNS "github.com/xtls/xray-core/features/dns"
|
|
||||||
)
|
|
||||||
|
|
||||||
type scriptNameServer struct {
|
|
||||||
name string
|
|
||||||
answers map[string]net.IP
|
|
||||||
errors map[string]error
|
|
||||||
ttl uint32
|
|
||||||
calls int
|
|
||||||
}
|
|
||||||
|
|
||||||
func (s *scriptNameServer) Name() string { return s.name }
|
|
||||||
func (s *scriptNameServer) IsDisableCache() bool { return true }
|
|
||||||
|
|
||||||
func (s *scriptNameServer) QueryIP(ctx context.Context, domain string, _ featureDNS.IPOption) ([]net.IP, uint32, error) {
|
|
||||||
if err := ctx.Err(); err != nil {
|
|
||||||
return nil, 0, err
|
|
||||||
}
|
|
||||||
s.calls++
|
|
||||||
if err := s.errors[domain]; err != nil {
|
|
||||||
return nil, 0, err
|
|
||||||
}
|
|
||||||
ip, ok := s.answers[domain]
|
|
||||||
if !ok {
|
|
||||||
return nil, 0, featureDNS.ErrEmptyResponse
|
|
||||||
}
|
|
||||||
return []net.IP{ip}, s.ttl, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestDNSScriptQuery(t *testing.T) {
|
|
||||||
wantIP := net.ParseIP("127.0.0.1")
|
|
||||||
upstreamErr := go_errors.New("upstream failed")
|
|
||||||
for _, tc := range []struct {
|
|
||||||
name, body string
|
|
||||||
wantIPs []net.IP
|
|
||||||
wantTTL uint32
|
|
||||||
wantErr error
|
|
||||||
wantMessage string
|
|
||||||
wantCalls uint32
|
|
||||||
}{
|
|
||||||
{name: "IPs", body: `return server:Query(domain, ipv4, ipv6, fake)`, wantIPs: []net.IP{wantIP}, wantTTL: 60, wantCalls: 2},
|
|
||||||
{name: "empty result", body: `return nil, 0`, wantErr: featureDNS.ErrEmptyResponse, wantCalls: 2},
|
|
||||||
{name: "upstream error", body: `return server:Query("failed.example", ipv4, ipv6, fake)`, wantErr: upstreamErr, wantCalls: 2},
|
|
||||||
{name: "string error", body: `return nil, nil, "blocked"`, wantMessage: "blocked", wantCalls: 2},
|
|
||||||
{name: "invalid result", body: `return false, 0`, wantMessage: "native IP slice", wantCalls: 2},
|
|
||||||
{name: "execution error", body: `error("execution failed")`, wantMessage: "execution failed", wantCalls: 1},
|
|
||||||
} {
|
|
||||||
t.Run(tc.name, func(t *testing.T) {
|
|
||||||
script := `
|
|
||||||
local server = require("xray.dns").Servers[1]
|
|
||||||
local calls = 0
|
|
||||||
function HandleDNSQuery(domain, ipv4, ipv6, fake)
|
|
||||||
calls = calls + 1
|
|
||||||
if domain == "count.example" then
|
|
||||||
local ips, _, err = server:Query("good.example", ipv4, ipv6, fake)
|
|
||||||
return ips, calls, err
|
|
||||||
end
|
|
||||||
` + tc.body + `
|
|
||||||
end
|
|
||||||
`
|
|
||||||
path := filepath.Join(t.TempDir(), "query.lua")
|
|
||||||
if err := os.WriteFile(path, []byte(script), 0o600); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
option := featureDNS.IPOption{IPv4Enable: true}
|
|
||||||
upstream := &scriptNameServer{
|
|
||||||
name: "test",
|
|
||||||
answers: map[string]net.IP{"good.example": wantIP},
|
|
||||||
errors: map[string]error{"failed.example": upstreamErr},
|
|
||||||
ttl: 60,
|
|
||||||
}
|
|
||||||
server := &DNS{
|
|
||||||
ctx: context.Background(),
|
|
||||||
clients: []*Client{{server: upstream, ipOption: &option, timeoutMs: time.Second}},
|
|
||||||
}
|
|
||||||
engine, err := newScriptEngine(path, server)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
defer engine.close()
|
|
||||||
|
|
||||||
ips, ttl, err := engine.query("good.example", option)
|
|
||||||
switch {
|
|
||||||
case tc.wantErr != nil:
|
|
||||||
if err != tc.wantErr {
|
|
||||||
t.Fatalf("query error = %v, want original error %v", err, tc.wantErr)
|
|
||||||
}
|
|
||||||
case tc.wantMessage != "":
|
|
||||||
if err == nil || !strings.Contains(err.Error(), tc.wantMessage) {
|
|
||||||
t.Fatalf("query error = %v, want %q", err, tc.wantMessage)
|
|
||||||
}
|
|
||||||
case err != nil:
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if ttl != tc.wantTTL || len(ips) != len(tc.wantIPs) {
|
|
||||||
t.Fatalf("query = %v, TTL %d; want %v, TTL %d", ips, ttl, tc.wantIPs, tc.wantTTL)
|
|
||||||
}
|
|
||||||
for i := range ips {
|
|
||||||
if !ips[i].Equal(tc.wantIPs[i]) {
|
|
||||||
t.Fatalf("IP %d = %v, want %v", i, ips[i], tc.wantIPs[i])
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
ips, calls, err := engine.query("count.example", option)
|
|
||||||
if err != nil || calls != tc.wantCalls || len(ips) != 1 || !ips[0].Equal(wantIP) {
|
|
||||||
t.Fatalf("next query = %v, calls %d, %v; want %v, calls %d", ips, calls, err, wantIP, tc.wantCalls)
|
|
||||||
}
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
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 := &scriptNameServer{
|
|
||||||
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 := &scriptNameServer{
|
|
||||||
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 TestDNSScriptFakeDNSOption(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)
|
|
||||||
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 := &scriptNameServer{
|
|
||||||
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("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
-12
@@ -587,8 +587,6 @@ 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"`
|
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"`
|
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"`
|
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
|
unknownFields protoimpl.UnknownFields
|
||||||
sizeCache protoimpl.SizeCache
|
sizeCache protoimpl.SizeCache
|
||||||
}
|
}
|
||||||
@@ -644,13 +642,6 @@ func (x *Config) GetBalancingRule() []*BalancingRule {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (x *Config) GetScript() string {
|
|
||||||
if x != nil {
|
|
||||||
return x.Script
|
|
||||||
}
|
|
||||||
return ""
|
|
||||||
}
|
|
||||||
|
|
||||||
var File_app_router_config_proto protoreflect.FileDescriptor
|
var File_app_router_config_proto protoreflect.FileDescriptor
|
||||||
|
|
||||||
const file_app_router_config_proto_rawDesc = "" +
|
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" +
|
"\tbaselines\x18\x03 \x03(\x03R\tbaselines\x12\x1a\n" +
|
||||||
"\bexpected\x18\x04 \x01(\x05R\bexpected\x12\x16\n" +
|
"\bexpected\x18\x04 \x01(\x05R\bexpected\x12\x16\n" +
|
||||||
"\x06maxRTT\x18\x05 \x01(\x03R\x06maxRTT\x12\x1c\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" +
|
"\x06Config\x12O\n" +
|
||||||
"\x0fdomain_strategy\x18\x01 \x01(\x0e2&.xray.app.router.Config.DomainStrategyR\x0edomainStrategy\x120\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" +
|
"\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" +
|
"\x0ebalancing_rule\x18\x03 \x03(\v2\x1e.xray.app.router.BalancingRuleR\rbalancingRule\"B\n" +
|
||||||
"\x06script\x18\x04 \x01(\tR\x06script\"B\n" +
|
|
||||||
"\x0eDomainStrategy\x12\b\n" +
|
"\x0eDomainStrategy\x12\b\n" +
|
||||||
"\x04AsIs\x10\x00\x12\x10\n" +
|
"\x04AsIs\x10\x00\x12\x10\n" +
|
||||||
"\fIpIfNonMatch\x10\x02\x12\x0e\n" +
|
"\fIpIfNonMatch\x10\x02\x12\x0e\n" +
|
||||||
|
|||||||
@@ -110,6 +110,4 @@ message Config {
|
|||||||
DomainStrategy domain_strategy = 1;
|
DomainStrategy domain_strategy = 1;
|
||||||
repeated RoutingRule rule = 2;
|
repeated RoutingRule rule = 2;
|
||||||
repeated BalancingRule balancing_rule = 3;
|
repeated BalancingRule balancing_rule = 3;
|
||||||
// Absolute path to the Lua routing script.
|
|
||||||
string script = 4;
|
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1,175 +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.CreateTable(0, 7)
|
|
||||||
|
|
||||||
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) {
|
|
||||||
pushIPs := xlua.NewSlicePusher[net.IP](L)
|
|
||||||
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.CreateTable(0, 4)
|
|
||||||
L.SetFuncs(methods, map[string]lua.LGFunction{
|
|
||||||
"GetSourceIPs": func(L *lua.LState) int {
|
|
||||||
pushIPs(L, checkLuaContext(L).GetSourceIPs())
|
|
||||||
return 1
|
|
||||||
},
|
|
||||||
"GetTargetIPs": func(L *lua.LState) int {
|
|
||||||
pushIPs(L, checkLuaContext(L).GetTargetIPs())
|
|
||||||
return 1
|
|
||||||
},
|
|
||||||
"GetLocalIPs": func(L *lua.LState) int {
|
|
||||||
pushIPs(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
|
|
||||||
}
|
|
||||||
|
|
||||||
// callLuaRoute runs HandleRoute and leaves (outboundTag, ruleTag, err) on the stack.
|
|
||||||
func callLuaRoute(L *lua.LState, ctx routing.Context) error {
|
|
||||||
fn := L.GetGlobal("HandleRoute")
|
|
||||||
if fn.Type() != lua.LTFunction {
|
|
||||||
return errors.New("routing script must define HandleRoute(...)")
|
|
||||||
}
|
|
||||||
|
|
||||||
value := L.NewUserData()
|
|
||||||
value.Value = ctx
|
|
||||||
L.SetMetatable(value, L.GetTypeMetatable(luaContextType))
|
|
||||||
|
|
||||||
return L.CallByParam(lua.P{Fn: fn, NRet: 3, Protect: true},
|
|
||||||
value,
|
|
||||||
lua.LString(ctx.GetInboundTag()),
|
|
||||||
lua.LNumber(ctx.GetSourcePort()),
|
|
||||||
lua.LNumber(ctx.GetTargetPort()),
|
|
||||||
lua.LNumber(ctx.GetLocalPort()),
|
|
||||||
lua.LString(strings.ToLower(ctx.GetTargetDomain())),
|
|
||||||
lua.LNumber(ctx.GetNetwork()),
|
|
||||||
lua.LString(ctx.GetProtocol()),
|
|
||||||
lua.LString(ctx.GetUser()),
|
|
||||||
lua.LNumber(ctx.GetVlessRoute()),
|
|
||||||
lua.LBool(ctx.GetSkipDNSResolve()))
|
|
||||||
}
|
|
||||||
|
|
||||||
// readLuaRouteResult reads (outboundTag, ruleTag, err) from the stack.
|
|
||||||
func readLuaRouteResult(L *lua.LState) (string, string, error) {
|
|
||||||
if err := xlua.ReadError(L.Get(-1), "routing script error must be an error or string"); err != nil {
|
|
||||||
return "", "", err
|
|
||||||
}
|
|
||||||
|
|
||||||
outboundTag, err := xlua.ReadOptionalString(L.Get(-3), "routing script outboundTag must be a string or nil")
|
|
||||||
if err != nil || outboundTag == "" {
|
|
||||||
return "", "", err
|
|
||||||
}
|
|
||||||
|
|
||||||
ruleTag, err := xlua.ReadOptionalString(L.Get(-2), "routing script ruleTag must be a string")
|
|
||||||
if err != nil {
|
|
||||||
return "", "", err
|
|
||||||
}
|
|
||||||
|
|
||||||
return outboundTag, 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)
|
|
||||||
}
|
|
||||||
@@ -1,223 +0,0 @@
|
|||||||
package router
|
|
||||||
|
|
||||||
import (
|
|
||||||
"context"
|
|
||||||
"fmt"
|
|
||||||
"os"
|
|
||||||
"path/filepath"
|
|
||||||
"strings"
|
|
||||||
"testing"
|
|
||||||
|
|
||||||
"github.com/xtls/xray-core/common/geodata"
|
|
||||||
"github.com/xtls/xray-core/common/net"
|
|
||||||
"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"
|
|
||||||
)
|
|
||||||
|
|
||||||
func benchmarkRouteContext(target net.Destination) *routing_session.Context {
|
|
||||||
// Use the production context: its IP getters construct a slice per call.
|
|
||||||
// The cached IP slices in luaRouteTestContext would undercount this cost.
|
|
||||||
return &routing_session.Context{
|
|
||||||
Inbound: &session.Inbound{
|
|
||||||
Tag: "in",
|
|
||||||
Source: net.TCPDestination(net.LocalHostIP, 1234),
|
|
||||||
Local: net.TCPDestination(net.LocalHostIP, 5678),
|
|
||||||
},
|
|
||||||
Outbound: &session.Outbound{Target: target},
|
|
||||||
Content: &session.Content{Protocol: "tls"},
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func benchmarkRouteState(b *testing.B, r *Router, script string) *lua.LState {
|
|
||||||
b.Helper()
|
|
||||||
L := lua.NewState()
|
|
||||||
b.Cleanup(L.Close)
|
|
||||||
r.RegisterLua(L)
|
|
||||||
geodata.RegisterLua(L)
|
|
||||||
if err := L.DoString(script); err != nil {
|
|
||||||
b.Fatal(err)
|
|
||||||
}
|
|
||||||
L.SetContext(context.Background())
|
|
||||||
return L
|
|
||||||
}
|
|
||||||
|
|
||||||
// BenchmarkLuaRouteHook isolates argument bridging and a fixed-return hook.
|
|
||||||
// It excludes rules, result decoding, the state pool, and Route construction.
|
|
||||||
func BenchmarkLuaRouteHook(b *testing.B) {
|
|
||||||
L := benchmarkRouteState(b, new(Router), `
|
|
||||||
function HandleRoute(ctx, inboundTag, sourcePort, targetPort, localPort,
|
|
||||||
targetDomain, network, protocol, user, vlessRoute, skipDNSResolve)
|
|
||||||
return "out", "rule"
|
|
||||||
end
|
|
||||||
`)
|
|
||||||
ctx := benchmarkRouteContext(net.TCPDestination(net.LocalHostIP, 443))
|
|
||||||
if err := callLuaRoute(L, ctx); err != nil {
|
|
||||||
b.Fatal(err)
|
|
||||||
}
|
|
||||||
outboundTag, ruleTag, err := readLuaRouteResult(L)
|
|
||||||
L.Pop(3)
|
|
||||||
if err != nil || outboundTag != "out" || ruleTag != "rule" {
|
|
||||||
b.Fatalf("hook() = %q, %q, %v", outboundTag, ruleTag, err)
|
|
||||||
}
|
|
||||||
b.ReportAllocs()
|
|
||||||
b.ResetTimer()
|
|
||||||
for i := 0; i < b.N; i++ {
|
|
||||||
if err := callLuaRoute(L, ctx); err != nil {
|
|
||||||
b.Fatal(err)
|
|
||||||
}
|
|
||||||
L.Pop(3)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// BenchmarkLuaRoute compares equivalent ordered rules on the same session.
|
|
||||||
// rules returns tags only; pick_route uses Router.PickRoute on both sides.
|
|
||||||
// All compilation, matcher construction, and pool startup are outside timing.
|
|
||||||
func BenchmarkLuaRoute(b *testing.B) {
|
|
||||||
for _, name := range []string{"scalar", "ip", "domain", "domain_32_last"} {
|
|
||||||
b.Run(name, func(b *testing.B) {
|
|
||||||
config, script, ctx, wantTag, wantRule := benchmarkRouteFixture(b, name)
|
|
||||||
native := new(Router)
|
|
||||||
if err := native.Init(context.Background(), config, nil, nil, nil); err != nil {
|
|
||||||
b.Fatal(err)
|
|
||||||
}
|
|
||||||
L := benchmarkRouteState(b, native, script)
|
|
||||||
|
|
||||||
path := filepath.Join(b.TempDir(), "route.lua")
|
|
||||||
if err := os.WriteFile(path, []byte(script), 0o600); err != nil {
|
|
||||||
b.Fatal(err)
|
|
||||||
}
|
|
||||||
scripted := new(Router)
|
|
||||||
if err := scripted.Init(context.Background(), &Config{Script: path}, nil, nil, nil); err != nil {
|
|
||||||
b.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := scripted.Start(); err != nil {
|
|
||||||
b.Fatal(err)
|
|
||||||
}
|
|
||||||
b.Cleanup(func() {
|
|
||||||
if err := scripted.Close(); err != nil {
|
|
||||||
b.Error(err)
|
|
||||||
}
|
|
||||||
})
|
|
||||||
|
|
||||||
for _, bench := range []struct {
|
|
||||||
name string
|
|
||||||
route func() (string, string, error)
|
|
||||||
}{
|
|
||||||
{"rules/native", func() (string, string, error) {
|
|
||||||
rule, _, err := native.pickRouteInternal(ctx)
|
|
||||||
if err != nil {
|
|
||||||
return "", "", err
|
|
||||||
}
|
|
||||||
tag, err := rule.GetTag()
|
|
||||||
return tag, rule.RuleTag, err
|
|
||||||
}},
|
|
||||||
{"rules/lua", func() (string, string, error) {
|
|
||||||
if err := callLuaRoute(L, ctx); err != nil {
|
|
||||||
return "", "", err
|
|
||||||
}
|
|
||||||
tag, ruleTag, err := readLuaRouteResult(L)
|
|
||||||
L.Pop(3)
|
|
||||||
return tag, ruleTag, err
|
|
||||||
}},
|
|
||||||
{"pick_route/native", func() (string, string, error) {
|
|
||||||
return benchmarkPickRoute(native, ctx)
|
|
||||||
}},
|
|
||||||
{"pick_route/lua", func() (string, string, error) {
|
|
||||||
return benchmarkPickRoute(scripted, ctx)
|
|
||||||
}},
|
|
||||||
} {
|
|
||||||
b.Run(bench.name, func(b *testing.B) {
|
|
||||||
// Validate and warm both paths before measuring steady state.
|
|
||||||
tag, ruleTag, err := bench.route()
|
|
||||||
if err != nil || tag != wantTag || ruleTag != wantRule {
|
|
||||||
b.Fatalf("route() = %q, %q, %v; want %q, %q", tag, ruleTag, err, wantTag, wantRule)
|
|
||||||
}
|
|
||||||
b.ReportAllocs()
|
|
||||||
b.ResetTimer()
|
|
||||||
for i := 0; i < b.N; i++ {
|
|
||||||
tag, ruleTag, err = bench.route()
|
|
||||||
if err != nil {
|
|
||||||
b.Fatal(err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
b.StopTimer()
|
|
||||||
if tag != wantTag || ruleTag != wantRule {
|
|
||||||
b.Fatalf("route() = %q, %q; want %q, %q", tag, ruleTag, wantTag, wantRule)
|
|
||||||
}
|
|
||||||
})
|
|
||||||
}
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func benchmarkPickRoute(r *Router, ctx routing.Context) (string, string, error) {
|
|
||||||
route, err := r.PickRoute(ctx)
|
|
||||||
if err != nil {
|
|
||||||
return "", "", err
|
|
||||||
}
|
|
||||||
return route.GetOutboundTag(), route.GetRuleTag(), nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func benchmarkRouteFixture(b *testing.B, name string) (*Config, string, routing.Context, string, string) {
|
|
||||||
b.Helper()
|
|
||||||
config := new(Config)
|
|
||||||
ctx := benchmarkRouteContext(net.TCPDestination(net.LocalHostIP, 443))
|
|
||||||
prelude := `local router = require("xray.router")
|
|
||||||
local geodata = require("xray.geodata")
|
|
||||||
`
|
|
||||||
body := `if inboundTag == "in" and network == router.NetworkTCP then return "out", "rule" end`
|
|
||||||
wantTag, wantRule := "out", "rule"
|
|
||||||
if name == "scalar" || name == "ip" {
|
|
||||||
rule := &RoutingRule{
|
|
||||||
TargetTag: &RoutingRule_Tag{Tag: wantTag},
|
|
||||||
RuleTag: wantRule,
|
|
||||||
InboundTag: []string{"in"},
|
|
||||||
Networks: []net.Network{net.Network_TCP},
|
|
||||||
}
|
|
||||||
if name == "ip" {
|
|
||||||
var err error
|
|
||||||
rule.Ip, err = geodata.ParseIPRules([]string{"127.0.0.0/8"})
|
|
||||||
if err != nil {
|
|
||||||
b.Fatal(err)
|
|
||||||
}
|
|
||||||
prelude += `local matcher = geodata.BuildIPMatcher("127.0.0.0/8")` + "\n"
|
|
||||||
body = `if inboundTag == "in" and network == router.NetworkTCP and matcher:AnyMatch(ctx:GetTargetIPs()) then return "out", "rule" end`
|
|
||||||
}
|
|
||||||
config.Rule = []*RoutingRule{rule}
|
|
||||||
} else {
|
|
||||||
count := 1
|
|
||||||
if name == "domain_32_last" {
|
|
||||||
count = 32
|
|
||||||
}
|
|
||||||
var rules strings.Builder
|
|
||||||
rules.WriteString("local rules = {\n")
|
|
||||||
for i := 0; i < count; i++ {
|
|
||||||
domain := fmt.Sprintf("route-%d.example.com", i)
|
|
||||||
tag, ruleTag := fmt.Sprintf("out-%d", i), fmt.Sprintf("rule-%d", i)
|
|
||||||
domains, err := geodata.ParseDomainRules([]string{"full:" + domain}, geodata.Domain_Domain)
|
|
||||||
if err != nil {
|
|
||||||
b.Fatal(err)
|
|
||||||
}
|
|
||||||
config.Rule = append(config.Rule, &RoutingRule{
|
|
||||||
TargetTag: &RoutingRule_Tag{Tag: tag}, RuleTag: ruleTag, Domain: domains,
|
|
||||||
})
|
|
||||||
fmt.Fprintf(&rules, "{geodata.BuildDomainMatcher(%q), %q, %q},\n", "full:"+domain, tag, ruleTag)
|
|
||||||
if i == count-1 {
|
|
||||||
ctx.Outbound.Target = net.TCPDestination(net.DomainAddress(domain), 443)
|
|
||||||
wantTag, wantRule = tag, ruleTag
|
|
||||||
}
|
|
||||||
}
|
|
||||||
rules.WriteString("}\n")
|
|
||||||
prelude += rules.String()
|
|
||||||
body = `for i = 1, #rules do
|
|
||||||
local rule = rules[i]
|
|
||||||
if rule[1]:MatchAny(targetDomain) then return rule[2], rule[3] end
|
|
||||||
end`
|
|
||||||
}
|
|
||||||
script := prelude + `function HandleRoute(ctx, inboundTag, sourcePort, targetPort, localPort,
|
|
||||||
targetDomain, network, protocol, user, vlessRoute, skipDNSResolve)
|
|
||||||
` + body + "\nend\n"
|
|
||||||
return config, script, ctx, wantTag, wantRule
|
|
||||||
}
|
|
||||||
@@ -1,274 +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) *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 L
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestLuaRouteBinding(t *testing.T) {
|
|
||||||
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(#sourceIPs == 1 and #targetIPs == 1 and #localIPs == 1)
|
|
||||||
assert(sourceIPs[1]:String() == "127.0.0.2" and targetIPs[1]:String() == "127.0.0.3")
|
|
||||||
assert(localIPs[1]:String() == "127.0.0.1")
|
|
||||||
assert(matcher:Match(sourceIPs[1]) and matcher:Match(targetIPs[1]) and matcher:Match(localIPs[1]))
|
|
||||||
assert(matcher:AnyMatch(sourceIPs) and matcher:AnyMatch(targetIPs) and matcher:AnyMatch(localIPs))
|
|
||||||
local matched = matcher:FilterIPs(targetIPs)
|
|
||||||
assert(#matched == 1 and matched[1]:Equal(targetIPs[1]))
|
|
||||||
assert(attributes.key == "value" and attributes.missing == nil)
|
|
||||||
assert(not pcall(function() attributes.key = "changed" end))
|
|
||||||
return "out", "rule"
|
|
||||||
end`)
|
|
||||||
|
|
||||||
ctx := newLuaRouteTestContext()
|
|
||||||
if err := callLuaRoute(L, ctx); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if L.GetTop() != 3 || L.Get(1) != lua.LString("out") || L.Get(2) != lua.LString("rule") || L.Get(3) != lua.LNil {
|
|
||||||
t.Fatal("callLuaRoute did not leave the three route results on the stack")
|
|
||||||
}
|
|
||||||
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 TestLuaRouteEmptyIPs(t *testing.T) {
|
|
||||||
for _, tc := range []struct {
|
|
||||||
name string
|
|
||||||
ips []net.IP
|
|
||||||
}{
|
|
||||||
{"nil", nil},
|
|
||||||
{"empty", []net.IP{}},
|
|
||||||
} {
|
|
||||||
t.Run(tc.name, func(t *testing.T) {
|
|
||||||
L := newLuaRouterState(t, `
|
|
||||||
function HandleRoute(ctx)
|
|
||||||
for _, name in ipairs({"GetSourceIPs", "GetTargetIPs", "GetLocalIPs"}) do
|
|
||||||
local ips = ctx[name](ctx)
|
|
||||||
if expectNil then
|
|
||||||
assert(ips == nil)
|
|
||||||
else
|
|
||||||
assert(type(ips) == "userdata" and #ips == 0)
|
|
||||||
assert(not pcall(function() return ips[1] end))
|
|
||||||
end
|
|
||||||
end
|
|
||||||
return "out"
|
|
||||||
end
|
|
||||||
`)
|
|
||||||
L.SetGlobal("expectNil", lua.LBool(tc.ips == nil))
|
|
||||||
ctx := newLuaRouteTestContext()
|
|
||||||
ctx.sourceIPs, ctx.targetIPs, ctx.localIPs = tc.ips, tc.ips, tc.ips
|
|
||||||
if err := callLuaRoute(L, ctx); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestReadLuaRouteResult(t *testing.T) {
|
|
||||||
nativeErr := go_errors.New("native failure")
|
|
||||||
for _, tc := range []struct {
|
|
||||||
name, values string
|
|
||||||
wantTag, wantRule string
|
|
||||||
wantErr error
|
|
||||||
wantMessage string
|
|
||||||
}{
|
|
||||||
{name: "route", values: `"out", "rule"`, wantTag: "out", wantRule: "rule"},
|
|
||||||
{name: "no match", values: `nil`},
|
|
||||||
{name: "empty tag", values: `""`},
|
|
||||||
{name: "no match ignores rule", values: `nil, false`},
|
|
||||||
{name: "empty tag ignores rule", values: `"", false`},
|
|
||||||
{name: "missing rule", values: `"out"`, wantTag: "out"},
|
|
||||||
{name: "invalid tag", values: `1`, wantMessage: "outboundTag"},
|
|
||||||
{name: "invalid rule", values: `"out", false`, wantMessage: "ruleTag"},
|
|
||||||
{name: "string error", values: `nil, nil, "script failure"`, wantMessage: "script failure"},
|
|
||||||
{name: "native error", values: `nil, nil, nativeError`, wantErr: nativeErr},
|
|
||||||
{name: "error overrides invalid tags", values: `false, false, nativeError`, wantErr: nativeErr},
|
|
||||||
{name: "invalid error", values: `"out", "rule", false`, wantMessage: "error or string"},
|
|
||||||
{name: "wrong error userdata", values: `"out", "rule", wrongError`, wantMessage: "error or string"},
|
|
||||||
} {
|
|
||||||
t.Run(tc.name, func(t *testing.T) {
|
|
||||||
L := lua.NewState()
|
|
||||||
defer L.Close()
|
|
||||||
for name, value := range map[string]any{"nativeError": nativeErr, "wrongError": "not a native error"} {
|
|
||||||
ud := L.NewUserData()
|
|
||||||
ud.Value = value
|
|
||||||
L.SetGlobal(name, ud)
|
|
||||||
}
|
|
||||||
fn, err := L.LoadString("return " + tc.values)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if err := L.CallByParam(lua.P{Fn: fn, NRet: 3, Protect: true}); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
outboundTag, ruleTag, err := readLuaRouteResult(L)
|
|
||||||
if outboundTag != tc.wantTag || ruleTag != tc.wantRule {
|
|
||||||
t.Fatalf("result = %q, %q, %v; want %q, %q", outboundTag, ruleTag, err, tc.wantTag, tc.wantRule)
|
|
||||||
}
|
|
||||||
switch {
|
|
||||||
case tc.wantErr != nil:
|
|
||||||
if err != tc.wantErr {
|
|
||||||
t.Fatalf("error = %v, want original error", err)
|
|
||||||
}
|
|
||||||
case tc.wantMessage != "":
|
|
||||||
if err == nil || !strings.Contains(err.Error(), tc.wantMessage) {
|
|
||||||
t.Fatalf("error = %v, want %q", err, tc.wantMessage)
|
|
||||||
}
|
|
||||||
case err != nil:
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestCallLuaRouteCancellation(t *testing.T) {
|
|
||||||
L := newLuaRouterState(t, `function HandleRoute() while true do end end`)
|
|
||||||
ctx, cancel := context.WithCancel(context.Background())
|
|
||||||
cancel()
|
|
||||||
L.SetContext(ctx)
|
|
||||||
if err := callLuaRoute(L, &routing_session.Context{}); err == nil {
|
|
||||||
t.Fatal("callLuaRoute did not stop after context cancellation")
|
|
||||||
}
|
|
||||||
if L.Context() != ctx {
|
|
||||||
t.Fatal("callLuaRoute changed the Lua state's context")
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
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)
|
|
||||||
}
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
var _ routing.Context = (*luaRouteTestContext)(nil)
|
|
||||||
@@ -20,8 +20,6 @@ import (
|
|||||||
type Router struct {
|
type Router struct {
|
||||||
domainStrategy Config_DomainStrategy
|
domainStrategy Config_DomainStrategy
|
||||||
rules atomic.Pointer[[]*Rule]
|
rules atomic.Pointer[[]*Rule]
|
||||||
scriptPath string
|
|
||||||
script *scriptEngine
|
|
||||||
balancers atomic.Pointer[map[string]*Balancer]
|
balancers atomic.Pointer[map[string]*Balancer]
|
||||||
dns dns.Client
|
dns dns.Client
|
||||||
|
|
||||||
@@ -42,7 +40,6 @@ type Route struct {
|
|||||||
// Init initializes the Router.
|
// Init initializes the Router.
|
||||||
func (r *Router) Init(ctx context.Context, config *Config, d dns.Client, ohm outbound.Manager, dispatcher routing.Dispatcher) error {
|
func (r *Router) Init(ctx context.Context, config *Config, d dns.Client, ohm outbound.Manager, dispatcher routing.Dispatcher) error {
|
||||||
r.domainStrategy = config.DomainStrategy
|
r.domainStrategy = config.DomainStrategy
|
||||||
r.scriptPath = config.Script
|
|
||||||
r.dns = d
|
r.dns = d
|
||||||
r.ctx = ctx
|
r.ctx = ctx
|
||||||
r.ohm = ohm
|
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.
|
// PickRoute implements routing.Router.
|
||||||
func (r *Router) PickRoute(ctx routing.Context) (routing.Route, error) {
|
func (r *Router) PickRoute(ctx routing.Context) (routing.Route, error) {
|
||||||
if r.script != nil {
|
|
||||||
return r.script.pickRoute(ctx)
|
|
||||||
}
|
|
||||||
|
|
||||||
originalCtx := ctx
|
originalCtx := ctx
|
||||||
rule, ctx, err := r.pickRouteInternal(ctx)
|
rule, ctx, err := r.pickRouteInternal(ctx)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -228,13 +221,6 @@ func (r *Router) pickRouteInternal(ctx routing.Context) (*Rule, routing.Context,
|
|||||||
|
|
||||||
// Start implements common.Runnable.
|
// Start implements common.Runnable.
|
||||||
func (r *Router) Start() error {
|
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
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -249,9 +235,6 @@ func closeWebhooks(rules []*Rule) {
|
|||||||
|
|
||||||
// Close implements common.Closable.
|
// Close implements common.Closable.
|
||||||
func (r *Router) Close() error {
|
func (r *Router) Close() error {
|
||||||
if r.script != nil {
|
|
||||||
r.script.close()
|
|
||||||
}
|
|
||||||
r.mu.Lock()
|
r.mu.Lock()
|
||||||
defer r.mu.Unlock()
|
defer r.mu.Unlock()
|
||||||
closeWebhooks(*r.rules.Load())
|
closeWebhooks(*r.rules.Load())
|
||||||
|
|||||||
@@ -1,76 +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 {
|
|
||||||
pool *xlua.Pool
|
|
||||||
}
|
|
||||||
|
|
||||||
func newScriptEngine(path string, router *Router) (*scriptEngine, error) {
|
|
||||||
program, err := xlua.CompileFile(path)
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
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 &scriptEngine{pool: pool}, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (e *scriptEngine) close() {
|
|
||||||
e.pool.Close()
|
|
||||||
}
|
|
||||||
|
|
||||||
func (e *scriptEngine) pickRoute(ctx routing.Context) (routing.Route, error) {
|
|
||||||
var outboundTag, ruleTag string
|
|
||||||
var routeErr error
|
|
||||||
|
|
||||||
if err := e.pool.WithState(nil, 0, func(L *lua.LState) error {
|
|
||||||
if err := callLuaRoute(L, ctx); err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
outboundTag, ruleTag, routeErr = readLuaRouteResult(L)
|
|
||||||
return nil
|
|
||||||
}); err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
if routeErr != nil {
|
|
||||||
return nil, routeErr
|
|
||||||
}
|
|
||||||
if outboundTag == "" {
|
|
||||||
return nil, common.ErrNoClue
|
|
||||||
}
|
|
||||||
|
|
||||||
return &Route{Context: ctx, outboundTag: outboundTag, ruleTag: ruleTag}, nil
|
|
||||||
}
|
|
||||||
@@ -1,382 +0,0 @@
|
|||||||
package router
|
|
||||||
|
|
||||||
import (
|
|
||||||
"context"
|
|
||||||
stdnet "net"
|
|
||||||
"os"
|
|
||||||
"path/filepath"
|
|
||||||
"strings"
|
|
||||||
"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) {
|
|
||||||
for _, tc := range []struct {
|
|
||||||
name, body string
|
|
||||||
wantTag, wantRule string
|
|
||||||
wantErr error
|
|
||||||
wantMessage string
|
|
||||||
wantCalls string
|
|
||||||
}{
|
|
||||||
{name: "route", body: `return "lua-out", "lua-rule"`, wantTag: "lua-out", wantRule: "lua-rule", wantCalls: "2"},
|
|
||||||
{name: "no match", body: `return nil`, wantErr: common.ErrNoClue, wantCalls: "2"},
|
|
||||||
{name: "empty tag", body: `return ""`, wantErr: common.ErrNoClue, wantCalls: "2"},
|
|
||||||
{name: "balancer error", body: `local tag, err = router:PickOutbound("missing"); return tag, nil, err`, wantMessage: "not found", wantCalls: "2"},
|
|
||||||
{name: "string error", body: `return nil, nil, "blocked"`, wantMessage: "blocked", wantCalls: "2"},
|
|
||||||
{name: "invalid tag", body: `return false`, wantMessage: "outboundTag", wantCalls: "2"},
|
|
||||||
{name: "invalid rule", body: `return "lua-out", false`, wantMessage: "ruleTag", wantCalls: "2"},
|
|
||||||
{name: "execution error", body: `error("execution failed")`, wantMessage: "execution failed", wantCalls: "1"},
|
|
||||||
} {
|
|
||||||
t.Run(tc.name, func(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
|
|
||||||
}}
|
|
||||||
script := `
|
|
||||||
local router = require("xray.router")
|
|
||||||
local calls = 0
|
|
||||||
function HandleRoute(ctx, inbound)
|
|
||||||
calls = calls + 1
|
|
||||||
if inbound == "count" then return "lua-out", tostring(calls) end
|
|
||||||
` + tc.body + `
|
|
||||||
end
|
|
||||||
`
|
|
||||||
r := startLuaRouter(t, script, 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)
|
|
||||||
switch {
|
|
||||||
case tc.wantErr != nil:
|
|
||||||
if err != tc.wantErr {
|
|
||||||
t.Fatalf("route error = %v, want %v", err, tc.wantErr)
|
|
||||||
}
|
|
||||||
case tc.wantMessage != "":
|
|
||||||
if err == nil || !strings.Contains(err.Error(), tc.wantMessage) {
|
|
||||||
t.Fatalf("route error = %v, want %q", err, tc.wantMessage)
|
|
||||||
}
|
|
||||||
case err != nil:
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if tc.wantTag == "" {
|
|
||||||
if route != nil {
|
|
||||||
t.Fatalf("route = %v, want nil", route)
|
|
||||||
}
|
|
||||||
} else if route == nil || route.GetOutboundTag() != tc.wantTag || route.GetRuleTag() != tc.wantRule || route.(*Route).Context != ctx {
|
|
||||||
t.Fatalf("route = %v; want %q, %q and original context", route, tc.wantTag, tc.wantRule)
|
|
||||||
}
|
|
||||||
|
|
||||||
ctx.Inbound.Tag = "count"
|
|
||||||
route, err = r.PickRoute(ctx)
|
|
||||||
if err != nil || route == nil || route.GetOutboundTag() != "lua-out" || route.GetRuleTag() != tc.wantCalls {
|
|
||||||
t.Fatalf("next route = %v, %v; want lua-out, calls %s", route, err, tc.wantCalls)
|
|
||||||
}
|
|
||||||
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 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,163 +0,0 @@
|
|||||||
package geodata
|
|
||||||
|
|
||||||
import (
|
|
||||||
xlua "github.com/xtls/xray-core/common/lua"
|
|
||||||
"github.com/xtls/xray-core/common/net"
|
|
||||||
lua "github.com/yuin/gopher-lua"
|
|
||||||
)
|
|
||||||
|
|
||||||
// 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.CreateTable(0, 2)
|
|
||||||
|
|
||||||
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
|
|
||||||
}
|
|
||||||
xlua.PushWithDirectMethods(L, matcher, map[string]xlua.DirectMethod{
|
|
||||||
"Match": newLuaDomainMatch(xlua.NewSlicePusher[uint32](L)),
|
|
||||||
"MatchAny": luaDomainMatchAny,
|
|
||||||
})
|
|
||||||
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
|
|
||||||
}
|
|
||||||
xlua.PushWithDirectMethods(L, matcher, map[string]xlua.DirectMethod{
|
|
||||||
"Match": luaIPMatch,
|
|
||||||
"AnyMatch": luaIPAnyMatch,
|
|
||||||
"Matches": luaIPMatches,
|
|
||||||
"FilterIPs": newLuaIPFilterIPs(xlua.NewSlicePusher[net.IP](L)),
|
|
||||||
})
|
|
||||||
return 1
|
|
||||||
}))
|
|
||||||
|
|
||||||
L.Push(module)
|
|
||||||
return 1
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
// Read native Go values by type assertion; slices keep their original storage.
|
|
||||||
func readLuaIPMatcherArgs[T any](L *lua.LState) (IPMatcher, T, bool) {
|
|
||||||
var input T
|
|
||||||
if L.GetTop() != 2 {
|
|
||||||
return nil, input, false
|
|
||||||
}
|
|
||||||
value, ok := L.Get(1).(*lua.LUserData)
|
|
||||||
if !ok {
|
|
||||||
return nil, input, false
|
|
||||||
}
|
|
||||||
matcher, ok := value.Value.(IPMatcher)
|
|
||||||
if !ok {
|
|
||||||
return nil, input, false
|
|
||||||
}
|
|
||||||
if L.Get(2) == lua.LNil {
|
|
||||||
return matcher, input, true
|
|
||||||
}
|
|
||||||
value, ok = L.Get(2).(*lua.LUserData)
|
|
||||||
if !ok {
|
|
||||||
return nil, input, false
|
|
||||||
}
|
|
||||||
input, ok = value.Value.(T)
|
|
||||||
return matcher, input, ok
|
|
||||||
}
|
|
||||||
|
|
||||||
func luaIPMatch(L *lua.LState) (int, bool) {
|
|
||||||
matcher, ip, ok := readLuaIPMatcherArgs[net.IP](L)
|
|
||||||
if !ok {
|
|
||||||
return 0, false
|
|
||||||
}
|
|
||||||
L.Push(lua.LBool(matcher.Match(ip)))
|
|
||||||
return 1, true
|
|
||||||
}
|
|
||||||
|
|
||||||
func luaIPAnyMatch(L *lua.LState) (int, bool) {
|
|
||||||
matcher, ips, ok := readLuaIPMatcherArgs[[]net.IP](L)
|
|
||||||
if !ok {
|
|
||||||
return 0, false
|
|
||||||
}
|
|
||||||
L.Push(lua.LBool(matcher.AnyMatch(ips)))
|
|
||||||
return 1, true
|
|
||||||
}
|
|
||||||
|
|
||||||
func luaIPMatches(L *lua.LState) (int, bool) {
|
|
||||||
matcher, ips, ok := readLuaIPMatcherArgs[[]net.IP](L)
|
|
||||||
if !ok {
|
|
||||||
return 0, false
|
|
||||||
}
|
|
||||||
L.Push(lua.LBool(matcher.Matches(ips)))
|
|
||||||
return 1, true
|
|
||||||
}
|
|
||||||
|
|
||||||
func newLuaIPFilterIPs(pushIPs func(*lua.LState, []net.IP)) xlua.DirectMethod {
|
|
||||||
return func(L *lua.LState) (int, bool) {
|
|
||||||
matcher, ips, ok := readLuaIPMatcherArgs[[]net.IP](L)
|
|
||||||
if !ok {
|
|
||||||
return 0, false
|
|
||||||
}
|
|
||||||
matched, unmatched := matcher.FilterIPs(ips)
|
|
||||||
pushIPs(L, matched)
|
|
||||||
pushIPs(L, unmatched)
|
|
||||||
return 2, true
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func newLuaDomainMatch(pushMatches func(*lua.LState, []uint32)) xlua.DirectMethod {
|
|
||||||
return func(L *lua.LState) (int, bool) {
|
|
||||||
if L.GetTop() == 2 {
|
|
||||||
if value, ok := L.Get(1).(*lua.LUserData); ok {
|
|
||||||
matcher, validMatcher := value.Value.(DomainMatcher)
|
|
||||||
domain, validDomain := L.Get(2).(lua.LString)
|
|
||||||
if validMatcher && validDomain {
|
|
||||||
pushMatches(L, matcher.Match(string(domain)))
|
|
||||||
return 1, true
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return 0, false
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func luaDomainMatchAny(L *lua.LState) (int, bool) {
|
|
||||||
if L.GetTop() == 2 {
|
|
||||||
if value, ok := L.Get(1).(*lua.LUserData); ok {
|
|
||||||
matcher, validMatcher := value.Value.(DomainMatcher)
|
|
||||||
domain, validDomain := L.Get(2).(lua.LString)
|
|
||||||
if validMatcher && validDomain {
|
|
||||||
L.Push(lua.LBool(matcher.MatchAny(string(domain))))
|
|
||||||
return 1, true
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return 0, false
|
|
||||||
}
|
|
||||||
|
|
||||||
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
|
|
||||||
}
|
|
||||||
@@ -1,172 +0,0 @@
|
|||||||
package geodata
|
|
||||||
|
|
||||||
import (
|
|
||||||
"fmt"
|
|
||||||
"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")
|
|
||||||
}
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestLuaMatcherArgumentsAndAliases(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)
|
|
||||||
if err := L.DoString(`
|
|
||||||
local geodata = require("xray.geodata")
|
|
||||||
local matcher = geodata.BuildIPMatcher("127.0.0.0/8")
|
|
||||||
assert(matcher.Match == matcher.match and matcher.AnyMatch == matcher.anyMatch)
|
|
||||||
assert(matcher.Matches == matcher.matches and matcher.FilterIPs == matcher.filterIPs)
|
|
||||||
assert(matcher:match(ip))
|
|
||||||
assert(matcher:anyMatch({ip}) and matcher:matches({ip}))
|
|
||||||
assert(not matcher:AnyMatch(nil))
|
|
||||||
assert(matcher:Matches(nil) == matcher:Matches({}))
|
|
||||||
local matched, unmatched = matcher:FilterIPs({ip})
|
|
||||||
assert(#matched == 1 and matched[1]:Equal(ip))
|
|
||||||
assert(matcher:AnyMatch(matched) and matcher:Matches(matched))
|
|
||||||
local filtered, excluded = matcher:filterIPs(matched)
|
|
||||||
assert(#filtered == 1 and #excluded == 0 and filtered[1]:Equal(ip))
|
|
||||||
local emptyMatched, emptyUnmatched = matcher:FilterIPs(nil)
|
|
||||||
assert(#emptyMatched == 0 and #emptyUnmatched == 0)
|
|
||||||
matcher:SetReverse(true)
|
|
||||||
assert(not matcher:Match(ip) and not matcher:AnyMatch(matched))
|
|
||||||
matcher:ToggleReverse()
|
|
||||||
assert(matcher:Match(ip) and matcher:AnyMatch(matched))
|
|
||||||
assert(matcher.missing == nil)
|
|
||||||
|
|
||||||
local domain = geodata.BuildDomainMatcher("full:example.com")
|
|
||||||
assert(domain.Match == domain.match and domain.MatchAny == domain.matchAny)
|
|
||||||
assert(domain:matchAny("example.com"))
|
|
||||||
assert(#domain:Match("example.com") == 1)
|
|
||||||
assert(domain:match("example.com")[1] == 0)
|
|
||||||
assert(not pcall(function() matcher:AnyMatch() end))
|
|
||||||
assert(not pcall(function() matcher:AnyMatch(matched, true) end))
|
|
||||||
assert(not pcall(function() matcher.AnyMatch(ip, matched) end))
|
|
||||||
assert(not pcall(function() matcher:Match(true) end))
|
|
||||||
assert(not pcall(function() domain:MatchAny(123) end))
|
|
||||||
assert(not pcall(function() domain:MatchAny("example.com", true) end))
|
|
||||||
assert(not pcall(function() matcher:FilterIPs(true) end))
|
|
||||||
assert(not pcall(function() matcher:FilterIPs(matched, true) end))
|
|
||||||
assert(not pcall(function() domain:Match(123) end))
|
|
||||||
assert(not pcall(function() domain:Match("example.com", true) end))
|
|
||||||
`); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// BenchmarkLuaMatcherCall measures repeated calls with prebuilt matchers and inputs.
|
|
||||||
func BenchmarkLuaMatcherCall(b *testing.B) {
|
|
||||||
L := lua.NewState()
|
|
||||||
defer L.Close()
|
|
||||||
RegisterLua(L)
|
|
||||||
ip := net.ParseIP("127.0.0.1")
|
|
||||||
for name, value := range map[string]any{"ip": ip, "ips": []net.IP{ip}} {
|
|
||||||
ud := L.NewUserData()
|
|
||||||
ud.Value = value
|
|
||||||
L.SetGlobal(name, ud)
|
|
||||||
}
|
|
||||||
if err := L.DoString(`
|
|
||||||
local geodata = require("xray.geodata")
|
|
||||||
ipMatcher = geodata.BuildIPMatcher("127.0.0.0/8")
|
|
||||||
domainMatcher = geodata.BuildDomainMatcher("full:example.com")
|
|
||||||
`); err != nil {
|
|
||||||
b.Fatal(err)
|
|
||||||
}
|
|
||||||
for _, benchmark := range []struct {
|
|
||||||
name, expression string
|
|
||||||
}{
|
|
||||||
{"ip_match", "ipMatcher:Match(ip)"},
|
|
||||||
{"ip_match_lower", "ipMatcher:match(ip)"},
|
|
||||||
{"ip_any_match", "ipMatcher:AnyMatch(ips)"},
|
|
||||||
{"ip_any_match_lower", "ipMatcher:anyMatch(ips)"},
|
|
||||||
{"ip_matches", "ipMatcher:Matches(ips)"},
|
|
||||||
{"ip_matches_lower", "ipMatcher:matches(ips)"},
|
|
||||||
{"domain_match_any", `domainMatcher:MatchAny("example.com")`},
|
|
||||||
{"domain_match_any_lower", `domainMatcher:matchAny("example.com")`},
|
|
||||||
{"ip_filter", "select(1, ipMatcher:FilterIPs(ips)) ~= nil"},
|
|
||||||
{"ip_filter_lower", "select(1, ipMatcher:filterIPs(ips)) ~= nil"},
|
|
||||||
{"domain_match", `#domainMatcher:Match("example.com") == 1`},
|
|
||||||
{"domain_match_lower", `#domainMatcher:match("example.com") == 1`},
|
|
||||||
{"ip_lua_table", "ipMatcher:AnyMatch({ip})"},
|
|
||||||
{"ip_lua_table_lower", "ipMatcher:anyMatch({ip})"},
|
|
||||||
} {
|
|
||||||
b.Run(benchmark.name, func(b *testing.B) {
|
|
||||||
if err := L.DoString(fmt.Sprintf("function benchmarkMatch() return %s end", benchmark.expression)); err != nil {
|
|
||||||
b.Fatal(err)
|
|
||||||
}
|
|
||||||
fn := L.GetGlobal("benchmarkMatch")
|
|
||||||
b.ReportAllocs()
|
|
||||||
b.ResetTimer()
|
|
||||||
for i := 0; i < b.N; i++ {
|
|
||||||
if err := L.CallByParam(lua.P{Fn: fn, NRet: 1, Protect: true}); err != nil {
|
|
||||||
b.Fatal(err)
|
|
||||||
}
|
|
||||||
if L.Get(-1) != lua.LTrue {
|
|
||||||
b.Fatal("matcher returned false")
|
|
||||||
}
|
|
||||||
L.Pop(1)
|
|
||||||
}
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@@ -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.CreateTable(0, 4)
|
|
||||||
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()
|
|
||||||
}
|
|
||||||
@@ -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)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@@ -1,3 +0,0 @@
|
|||||||
// Package lua provides shared GopherLua programs, state management, and value
|
|
||||||
// conversion and validation helpers for Xray scripts.
|
|
||||||
package lua
|
|
||||||
@@ -1,65 +0,0 @@
|
|||||||
package lua
|
|
||||||
|
|
||||||
import (
|
|
||||||
glua "github.com/yuin/gopher-lua"
|
|
||||||
luar "layeh.com/gopher-luar"
|
|
||||||
)
|
|
||||||
|
|
||||||
// NewSlicePusher captures luar's slice metatable during state initialization.
|
|
||||||
// The returned function wraps slices without reflection or metatable lookup,
|
|
||||||
// and pushes nil for nil slices. Use it with this state or its coroutines.
|
|
||||||
func NewSlicePusher[T any](L *glua.LState) func(*glua.LState, []T) {
|
|
||||||
metatable := luar.New(L, []T{}).(*glua.LUserData).Metatable
|
|
||||||
return func(L *glua.LState, values []T) {
|
|
||||||
if values == nil {
|
|
||||||
L.Push(glua.LNil)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
userdata := L.NewUserData()
|
|
||||||
userdata.Value = values
|
|
||||||
userdata.Metatable = metatable
|
|
||||||
L.Push(userdata)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// DirectMethod handles a Lua call without luar's reflected method invocation.
|
|
||||||
// It returns the result count and whether it handled the arguments. On false,
|
|
||||||
// it must leave the stack unchanged for the original luar wrapper.
|
|
||||||
type DirectMethod func(L *glua.LState) (nresults int, handled bool)
|
|
||||||
|
|
||||||
// PushWithDirectMethods pushes a luar userdata with typed Go method bindings.
|
|
||||||
// Handled calls bypass luar's argument conversion and reflect.Call; method lookup
|
|
||||||
// uses the methods table directly instead of luar's reflected __index handler.
|
|
||||||
// value must expose methods only. Bindings and their closures are installed once
|
|
||||||
// per Go type per LState, outside the method-call hot path.
|
|
||||||
func PushWithDirectMethods(L *glua.LState, value any, directMethods map[string]DirectMethod) {
|
|
||||||
userdata := luar.New(L, value).(*glua.LUserData)
|
|
||||||
metatable := userdata.Metatable.(*glua.LTable)
|
|
||||||
methods := metatable.RawGetString("methods").(*glua.LTable)
|
|
||||||
if metatable.RawGetString("__index") != methods {
|
|
||||||
for name, direct := range directMethods {
|
|
||||||
original := methods.RawGetString(name)
|
|
||||||
fn := L.NewFunction(func(L *glua.LState) int {
|
|
||||||
if nresults, handled := direct(L); handled {
|
|
||||||
return nresults
|
|
||||||
}
|
|
||||||
return callLuarMethod(L, original)
|
|
||||||
})
|
|
||||||
// Keep luar's method aliases on the same direct binding.
|
|
||||||
for key, method := methods.Next(glua.LNil); key != glua.LNil; key, method = methods.Next(key) {
|
|
||||||
if method == original {
|
|
||||||
methods.RawSet(key, fn)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
metatable.RawSetString("__index", methods)
|
|
||||||
}
|
|
||||||
L.Push(userdata)
|
|
||||||
}
|
|
||||||
|
|
||||||
func callLuarMethod(L *glua.LState, method glua.LValue) int {
|
|
||||||
nargs := L.GetTop()
|
|
||||||
L.Insert(method, 1)
|
|
||||||
L.Call(nargs, glua.MultRet)
|
|
||||||
return L.GetTop()
|
|
||||||
}
|
|
||||||
@@ -1,84 +0,0 @@
|
|||||||
package lua
|
|
||||||
|
|
||||||
import (
|
|
||||||
"net"
|
|
||||||
"testing"
|
|
||||||
|
|
||||||
glua "github.com/yuin/gopher-lua"
|
|
||||||
luar "layeh.com/gopher-luar"
|
|
||||||
)
|
|
||||||
|
|
||||||
func TestSlicePusher(t *testing.T) {
|
|
||||||
L := glua.NewState()
|
|
||||||
defer L.Close()
|
|
||||||
push := NewSlicePusher[int](L)
|
|
||||||
values := []int{3, 5}
|
|
||||||
L.SetGlobal("getValues", L.NewFunction(func(L *glua.LState) int {
|
|
||||||
push(L, values)
|
|
||||||
return 1
|
|
||||||
}))
|
|
||||||
if err := L.DoString(`
|
|
||||||
local values = getValues()
|
|
||||||
assert(#values == 2 and values[1] == 3 and values[2] == 5)
|
|
||||||
values[2] = 7
|
|
||||||
local co = coroutine.create(function()
|
|
||||||
local values = getValues()
|
|
||||||
assert(#values == 2 and values[1] == 3 and values[2] == 7)
|
|
||||||
return true
|
|
||||||
end)
|
|
||||||
local ok, result = coroutine.resume(co)
|
|
||||||
assert(ok and result == true)
|
|
||||||
`); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
if values[1] != 7 {
|
|
||||||
t.Fatal("slice storage was copied")
|
|
||||||
}
|
|
||||||
push(L, nil)
|
|
||||||
if L.Get(-1) != glua.LNil {
|
|
||||||
t.Fatal("nil slice must push Lua nil")
|
|
||||||
}
|
|
||||||
L.Pop(1)
|
|
||||||
push(L, []int{})
|
|
||||||
L.SetGlobal("empty", L.Get(-1))
|
|
||||||
L.Pop(1)
|
|
||||||
if err := L.DoString(`assert(type(empty) == "userdata" and #empty == 0)`); err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestSlicePusherMetatablePerState(t *testing.T) {
|
|
||||||
first := glua.NewState()
|
|
||||||
defer first.Close()
|
|
||||||
second := glua.NewState()
|
|
||||||
defer second.Close()
|
|
||||||
NewSlicePusher[int](first)(first, []int{1})
|
|
||||||
NewSlicePusher[int](second)(second, []int{1})
|
|
||||||
if first.Get(-1).(*glua.LUserData).Metatable == second.Get(-1).(*glua.LUserData).Metatable {
|
|
||||||
t.Fatal("independent states share a slice metatable")
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func BenchmarkSlicePusher(b *testing.B) {
|
|
||||||
L := glua.NewState()
|
|
||||||
defer L.Close()
|
|
||||||
ips := []net.IP{net.ParseIP("127.0.0.1")}
|
|
||||||
pushIPs := NewSlicePusher[net.IP](L)
|
|
||||||
for _, benchmark := range []struct {
|
|
||||||
name string
|
|
||||||
push func(*glua.LState, []net.IP)
|
|
||||||
}{
|
|
||||||
{"bare", func(L *glua.LState, ips []net.IP) { PushUserData(L, ips) }},
|
|
||||||
{"luar", func(L *glua.LState, ips []net.IP) { L.Push(luar.New(L, ips)) }},
|
|
||||||
{"cached", pushIPs},
|
|
||||||
} {
|
|
||||||
b.Run(benchmark.name, func(b *testing.B) {
|
|
||||||
b.ReportAllocs()
|
|
||||||
b.ResetTimer()
|
|
||||||
for i := 0; i < b.N; i++ {
|
|
||||||
benchmark.push(L, ips)
|
|
||||||
L.Pop(1)
|
|
||||||
}
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@@ -1,151 +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[n-1] = nil
|
|
||||||
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()
|
|
||||||
}
|
|
||||||
@@ -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)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@@ -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)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@@ -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())
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@@ -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)
|
|
||||||
}
|
|
||||||
@@ -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)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@@ -7,11 +7,9 @@ import (
|
|||||||
|
|
||||||
"github.com/golang/mock/gomock"
|
"github.com/golang/mock/gomock"
|
||||||
"github.com/xtls/xray-core/common"
|
"github.com/xtls/xray-core/common"
|
||||||
"github.com/xtls/xray-core/common/buf"
|
|
||||||
"github.com/xtls/xray-core/common/errors"
|
"github.com/xtls/xray-core/common/errors"
|
||||||
"github.com/xtls/xray-core/common/mux"
|
"github.com/xtls/xray-core/common/mux"
|
||||||
"github.com/xtls/xray-core/common/net"
|
"github.com/xtls/xray-core/common/net"
|
||||||
"github.com/xtls/xray-core/common/protocol"
|
|
||||||
"github.com/xtls/xray-core/common/session"
|
"github.com/xtls/xray-core/common/session"
|
||||||
"github.com/xtls/xray-core/testing/mocks"
|
"github.com/xtls/xray-core/testing/mocks"
|
||||||
"github.com/xtls/xray-core/transport"
|
"github.com/xtls/xray-core/transport"
|
||||||
@@ -116,48 +114,3 @@ func TestClientWorkerClose(t *testing.T) {
|
|||||||
|
|
||||||
common.Must(w2.Close())
|
common.Must(w2.Close())
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestClientWorkerUDPSource(t *testing.T) {
|
|
||||||
downR, downW := pipe.New(pipe.WithoutSizeLimit())
|
|
||||||
upR, upW := pipe.New(pipe.WithoutSizeLimit())
|
|
||||||
worker, err := mux.NewClientWorker(transport.Link{Reader: downR, Writer: upW}, mux.ClientStrategy{})
|
|
||||||
common.Must(err)
|
|
||||||
|
|
||||||
inR, inW := pipe.New(pipe.WithoutSizeLimit())
|
|
||||||
outR, outW := pipe.New(pipe.WithoutSizeLimit())
|
|
||||||
ctx := session.ContextWithOutbounds(context.Background(), []*session.Outbound{{
|
|
||||||
Target: net.UDPDestination(net.ParseAddress("8.8.8.8"), 53),
|
|
||||||
}})
|
|
||||||
if !worker.Dispatch(ctx, &transport.Link{Reader: inR, Writer: outW}) {
|
|
||||||
t.Fatal("failed to dispatch")
|
|
||||||
}
|
|
||||||
b := buf.New()
|
|
||||||
b.WriteString("query")
|
|
||||||
common.Must(inW.WriteMultiBuffer(buf.MultiBuffer{b}))
|
|
||||||
mb, err := upR.ReadMultiBuffer() // New frame, the session is UDP from now on
|
|
||||||
common.Must(err)
|
|
||||||
buf.ReleaseMulti(mb)
|
|
||||||
|
|
||||||
srcs := []net.Destination{
|
|
||||||
net.UDPDestination(net.ParseAddress("1.1.1.1"), 1111),
|
|
||||||
net.UDPDestination(net.DomainAddress("example.com"), 2222),
|
|
||||||
net.UDPDestination(net.ParseAddress("3.3.3.3"), 3333),
|
|
||||||
}
|
|
||||||
w := mux.NewResponseWriter(1, downW, protocol.TransferTypePacket)
|
|
||||||
var got buf.MultiBuffer
|
|
||||||
for i := range srcs {
|
|
||||||
b := buf.New()
|
|
||||||
b.WriteString("reply")
|
|
||||||
b.UDP = &srcs[i]
|
|
||||||
common.Must(w.WriteMultiBuffer(buf.MultiBuffer{b}))
|
|
||||||
// keep earlier replies around while the next frame is parsed
|
|
||||||
mb, err := outR.ReadMultiBuffer()
|
|
||||||
common.Must(err)
|
|
||||||
got = append(got, mb...)
|
|
||||||
}
|
|
||||||
for i, b := range got {
|
|
||||||
if b.UDP == nil || *b.UDP != srcs[i] {
|
|
||||||
t.Errorf("reply %d: source = %v, want %v", i, b.UDP, srcs[i])
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|||||||
@@ -14,16 +14,15 @@ import (
|
|||||||
type PacketReader struct {
|
type PacketReader struct {
|
||||||
reader io.Reader
|
reader io.Reader
|
||||||
eof bool
|
eof bool
|
||||||
dest net.Destination
|
dest *net.Destination
|
||||||
}
|
}
|
||||||
|
|
||||||
// NewPacketReader creates a new PacketReader.
|
// NewPacketReader creates a new PacketReader.
|
||||||
// dest is copied because the caller reuses it for the next frame.
|
|
||||||
func NewPacketReader(reader io.Reader, dest *net.Destination) *PacketReader {
|
func NewPacketReader(reader io.Reader, dest *net.Destination) *PacketReader {
|
||||||
return &PacketReader{
|
return &PacketReader{
|
||||||
reader: reader,
|
reader: reader,
|
||||||
eof: false,
|
eof: false,
|
||||||
dest: *dest,
|
dest: dest,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -48,8 +47,8 @@ func (r *PacketReader) ReadMultiBuffer() (buf.MultiBuffer, error) {
|
|||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
r.eof = true
|
r.eof = true
|
||||||
if r.dest.Network == net.Network_UDP {
|
if r.dest != nil && r.dest.Network == net.Network_UDP {
|
||||||
b.UDP = &r.dest // only one packet is read, so b owns r.dest
|
b.UDP = r.dest
|
||||||
}
|
}
|
||||||
return buf.MultiBuffer{b}, nil
|
return buf.MultiBuffer{b}, nil
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1,8 +1,6 @@
|
|||||||
package platform // import "github.com/xtls/xray-core/common/platform"
|
package platform // import "github.com/xtls/xray-core/common/platform"
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"errors"
|
|
||||||
"fmt"
|
|
||||||
"os"
|
"os"
|
||||||
"path/filepath"
|
"path/filepath"
|
||||||
"strconv"
|
"strconv"
|
||||||
@@ -92,49 +90,3 @@ func GetConfDirPath() string {
|
|||||||
configPath := NewEnvFlag(ConfdirLocation).GetValue(func() string { return "" })
|
configPath := NewEnvFlag(ConfdirLocation).GetValue(func() string { return "" })
|
||||||
return configPath
|
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)
|
|
||||||
}
|
|
||||||
|
|||||||
@@ -1,7 +1,6 @@
|
|||||||
package platform_test
|
package platform_test
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"errors"
|
|
||||||
"os"
|
"os"
|
||||||
"path/filepath"
|
"path/filepath"
|
||||||
"runtime"
|
"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)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|||||||
+2
-2
@@ -19,8 +19,8 @@ import (
|
|||||||
|
|
||||||
var (
|
var (
|
||||||
Version_x byte = 26
|
Version_x byte = 26
|
||||||
Version_y byte = 10
|
Version_y byte = 9
|
||||||
Version_z byte = 10
|
Version_z byte = 30
|
||||||
)
|
)
|
||||||
|
|
||||||
var (
|
var (
|
||||||
|
|||||||
+1
-3
@@ -92,8 +92,6 @@ type Instance struct {
|
|||||||
|
|
||||||
// Instance state
|
// Instance state
|
||||||
func (server *Instance) IsRunning() bool {
|
func (server *Instance) IsRunning() bool {
|
||||||
server.statusLock.Lock()
|
|
||||||
defer server.statusLock.Unlock()
|
|
||||||
return server.running
|
return server.running
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -322,7 +320,7 @@ func (s *Instance) RequireFeatures(callback interface{}, optional bool) error {
|
|||||||
|
|
||||||
// AddFeature registers a feature into current Instance.
|
// AddFeature registers a feature into current Instance.
|
||||||
func (s *Instance) AddFeature(feature features.Feature) error {
|
func (s *Instance) AddFeature(feature features.Feature) error {
|
||||||
if s.IsRunning() {
|
if s.running {
|
||||||
if err := feature.Start(); err != nil {
|
if err := feature.Start(); err != nil {
|
||||||
errors.LogInfoInner(s.ctx, err, "failed to start feature")
|
errors.LogInfoInner(s.ctx, err, "failed to start feature")
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -21,7 +21,6 @@ require (
|
|||||||
github.com/stretchr/testify v1.12.1
|
github.com/stretchr/testify v1.12.1
|
||||||
github.com/vishvananda/netlink v1.3.1
|
github.com/vishvananda/netlink v1.3.1
|
||||||
github.com/xtls/reality v0.0.0-20260908062103-8cdf7bf9c7f0
|
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
|
go4.org/netipx v0.0.0-20231129151722-fdeea329fbba
|
||||||
golang.org/x/crypto v0.57.0
|
golang.org/x/crypto v0.57.0
|
||||||
golang.org/x/exp v0.0.0-20240506185415-9bf2ced13842
|
golang.org/x/exp v0.0.0-20240506185415-9bf2ced13842
|
||||||
@@ -35,7 +34,6 @@ require (
|
|||||||
google.golang.org/protobuf v1.36.12
|
google.golang.org/protobuf v1.36.12
|
||||||
gvisor.dev/gvisor v0.0.0-20260122175437-89a5d21be8f0
|
gvisor.dev/gvisor v0.0.0-20260122175437-89a5d21be8f0
|
||||||
h12.io/socks v1.0.3
|
h12.io/socks v1.0.3
|
||||||
layeh.com/gopher-luar v1.0.11
|
|
||||||
lukechampine.com/blake3 v1.4.1
|
lukechampine.com/blake3 v1.4.1
|
||||||
mvdan.cc/gofumpt v0.12.0
|
mvdan.cc/gofumpt v0.12.0
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -2,9 +2,6 @@ github.com/andybalholm/brotli v1.0.6 h1:Yf9fFpf49Zrxb9NlQaluyE92/+X7UVHlhMNJN2sx
|
|||||||
github.com/andybalholm/brotli v1.0.6/go.mod h1:fO7iG3H7G2nSZ7m0zPUDn85XEX2GTukHGRSepvi9Eig=
|
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 h1:5mgtR5gwIgBKMiGI1QdXldZZ+SNor06Nbu1wCBulQBg=
|
||||||
github.com/apernet/quic-go v0.61.1-0.20260806010916-184d081eef3e/go.mod h1:x7qxEvX6MCVtDuBKHj3E+88+BtrbEMuAL5qGUKItjW8=
|
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/cloudflare/circl v1.6.5 h1:O64F26HEqNhznd/hrC5KZXVKYuKM2rx4deZDTc4ihQA=
|
github.com/cloudflare/circl v1.6.5 h1:O64F26HEqNhznd/hrC5KZXVKYuKM2rx4deZDTc4ihQA=
|
||||||
github.com/cloudflare/circl v1.6.5/go.mod h1:h5LNyxAc5nTue9DS5jT+48en2PSDYt3zdGnz5OstK6c=
|
github.com/cloudflare/circl v1.6.5/go.mod h1:h5LNyxAc5nTue9DS5jT+48en2PSDYt3zdGnz5OstK6c=
|
||||||
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 h1:Arcl6UOIS/kgO2nW3A65HN+7CMjSDP/gofXL4CZt1V4=
|
||||||
@@ -84,9 +81,6 @@ 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 h1:rb+fKQFhz+5I2PPuQsNYxI5mUU840XWYtRF0ZBjvkws=
|
||||||
github.com/xtls/reality v0.0.0-20260908062103-8cdf7bf9c7f0/go.mod h1:DsJblcWDGt76+FVqBVwbwRhxyyNJsGV48gJLch0OOWI=
|
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/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.uber.org/mock v0.5.2 h1:LbtPTcP8A5k9WPXj54PPPbjcI4Y6lhyOZXn+VS7wNko=
|
go.uber.org/mock v0.5.2 h1:LbtPTcP8A5k9WPXj54PPPbjcI4Y6lhyOZXn+VS7wNko=
|
||||||
go.uber.org/mock v0.5.2/go.mod h1:wLlUxC2vVTPTaE3UD51E0BGOAElKrILxhVSDYQLld5o=
|
go.uber.org/mock v0.5.2/go.mod h1:wLlUxC2vVTPTaE3UD51E0BGOAElKrILxhVSDYQLld5o=
|
||||||
go.yaml.in/yaml/v3 v3.0.5 h1:N6y/pJk8buWs9NY5ERU2HSMfm+IuD/OtfdAnq6kESPw=
|
go.yaml.in/yaml/v3 v3.0.5 h1:N6y/pJk8buWs9NY5ERU2HSMfm+IuD/OtfdAnq6kESPw=
|
||||||
@@ -113,7 +107,6 @@ golang.org/x/sync v0.0.0-20190423024810-112230192c58/go.mod h1:RxMgew5VJxzue5/jJ
|
|||||||
golang.org/x/sync v0.0.0-20210220032951-036812b2e83c/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 h1:KameEIfc1IkluZyXWLn39Wd4tURc6GbCiISGiZm2bQk=
|
||||||
golang.org/x/sync v0.23.0/go.mod h1:sUUOizhqBxiL6pEWpqNLUiaJn1ShEbZ6BBqskPbjZm0=
|
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/sys v0.0.0-20190215142949-d0b11bdaac8a/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY=
|
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-20190412213103-97732733099d/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs=
|
||||||
golang.org/x/sys v0.0.0-20201119102817-f84b799fce68/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs=
|
golang.org/x/sys v0.0.0-20201119102817-f84b799fce68/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs=
|
||||||
@@ -164,8 +157,6 @@ gvisor.dev/gvisor v0.0.0-20260122175437-89a5d21be8f0 h1:Lk6hARj5UPY47dBep70OD/TI
|
|||||||
gvisor.dev/gvisor v0.0.0-20260122175437-89a5d21be8f0/go.mod h1:QkHjoMIBaYtpVufgwv3keYAbln78mBoCuShZrPrer1Q=
|
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 h1:Ka3qaQewws4j4/eDQnOdpr4wXsC//dXtWvftlIcCQUo=
|
||||||
h12.io/socks v1.0.3/go.mod h1:AIhxy1jOId/XCz9BO+EIgNL2rQiPTBNnOfnVnQ+3Eck=
|
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 h1:I3Smz7gso8w4/TunLKec6K2fn+kyKtDxr/xcQEN84Wg=
|
||||||
lukechampine.com/blake3 v1.4.1/go.mod h1:QFosUxmjB8mnrWFSNwKmvxHpfY72bmD2tQ0kBMM3kwo=
|
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 h1:1Lbudkz2kpM9Cjz2pL4M19u7q+GaEhCTNf7N9mfpcho=
|
||||||
|
|||||||
@@ -14,11 +14,9 @@ import (
|
|||||||
"github.com/xtls/xray-core/common/errors"
|
"github.com/xtls/xray-core/common/errors"
|
||||||
"github.com/xtls/xray-core/common/geodata"
|
"github.com/xtls/xray-core/common/geodata"
|
||||||
"github.com/xtls/xray-core/common/net"
|
"github.com/xtls/xray-core/common/net"
|
||||||
"github.com/xtls/xray-core/common/platform"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
type NameServerConfig struct {
|
type NameServerConfig struct {
|
||||||
ID string `json:"id"`
|
|
||||||
Address *Address `json:"address"`
|
Address *Address `json:"address"`
|
||||||
ClientIP *Address `json:"clientIp"`
|
ClientIP *Address `json:"clientIp"`
|
||||||
Port uint16 `json:"port"`
|
Port uint16 `json:"port"`
|
||||||
@@ -45,7 +43,6 @@ func (c *NameServerConfig) UnmarshalJSON(data []byte) error {
|
|||||||
}
|
}
|
||||||
|
|
||||||
var advanced struct {
|
var advanced struct {
|
||||||
ID string `json:"id"`
|
|
||||||
Address *Address `json:"address"`
|
Address *Address `json:"address"`
|
||||||
ClientIP *Address `json:"clientIp"`
|
ClientIP *Address `json:"clientIp"`
|
||||||
Port uint16 `json:"port"`
|
Port uint16 `json:"port"`
|
||||||
@@ -63,7 +60,6 @@ func (c *NameServerConfig) UnmarshalJSON(data []byte) error {
|
|||||||
UnexpectedIPs StringList `json:"unexpectedIPs"`
|
UnexpectedIPs StringList `json:"unexpectedIPs"`
|
||||||
}
|
}
|
||||||
if err := json.Unmarshal(data, &advanced); err == nil {
|
if err := json.Unmarshal(data, &advanced); err == nil {
|
||||||
c.ID = advanced.ID
|
|
||||||
c.Address = advanced.Address
|
c.Address = advanced.Address
|
||||||
c.ClientIP = advanced.ClientIP
|
c.ClientIP = advanced.ClientIP
|
||||||
c.Port = advanced.Port
|
c.Port = advanced.Port
|
||||||
@@ -138,7 +134,6 @@ func (c *NameServerConfig) Build() (*dns.NameServer, error) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
return &dns.NameServer{
|
return &dns.NameServer{
|
||||||
Id: c.ID,
|
|
||||||
Address: &net.Endpoint{
|
Address: &net.Endpoint{
|
||||||
Network: net.Network_UDP,
|
Network: net.Network_UDP,
|
||||||
Address: c.Address.Build(),
|
Address: c.Address.Build(),
|
||||||
@@ -164,7 +159,6 @@ func (c *NameServerConfig) Build() (*dns.NameServer, error) {
|
|||||||
// DNSConfig is a JSON serializable object for dns.Config
|
// DNSConfig is a JSON serializable object for dns.Config
|
||||||
type DNSConfig struct {
|
type DNSConfig struct {
|
||||||
Servers []*NameServerConfig `json:"servers"`
|
Servers []*NameServerConfig `json:"servers"`
|
||||||
Script string `json:"script"`
|
|
||||||
Hosts *HostsWrapper `json:"hosts"`
|
Hosts *HostsWrapper `json:"hosts"`
|
||||||
ClientIP *Address `json:"clientIp"`
|
ClientIP *Address `json:"clientIp"`
|
||||||
Tag string `json:"tag"`
|
Tag string `json:"tag"`
|
||||||
@@ -284,14 +278,6 @@ func (c *DNSConfig) Build() (*dns.Config, error) {
|
|||||||
QueryStrategy: resolveQueryStrategy(c.QueryStrategy),
|
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 != nil {
|
||||||
if !c.ClientIP.Family().IsIP() {
|
if !c.ClientIP.Family().IsIP() {
|
||||||
return nil, errors.New("not an IP address:", c.ClientIP.String())
|
return nil, errors.New("not an IP address:", c.ClientIP.String())
|
||||||
|
|||||||
@@ -2,8 +2,6 @@ package conf_test
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"encoding/json"
|
"encoding/json"
|
||||||
"os"
|
|
||||||
"path/filepath"
|
|
||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
"github.com/google/go-cmp/cmp"
|
"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)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|||||||
@@ -7,7 +7,6 @@ import (
|
|||||||
"github.com/xtls/xray-core/app/router"
|
"github.com/xtls/xray-core/app/router"
|
||||||
"github.com/xtls/xray-core/common/errors"
|
"github.com/xtls/xray-core/common/errors"
|
||||||
"github.com/xtls/xray-core/common/geodata"
|
"github.com/xtls/xray-core/common/geodata"
|
||||||
"github.com/xtls/xray-core/common/platform"
|
|
||||||
"github.com/xtls/xray-core/common/serial"
|
"github.com/xtls/xray-core/common/serial"
|
||||||
|
|
||||||
"google.golang.org/protobuf/proto"
|
"google.golang.org/protobuf/proto"
|
||||||
@@ -73,7 +72,6 @@ type RouterConfig struct {
|
|||||||
RuleList []json.RawMessage `json:"rules"`
|
RuleList []json.RawMessage `json:"rules"`
|
||||||
DomainStrategy *string `json:"domainStrategy"`
|
DomainStrategy *string `json:"domainStrategy"`
|
||||||
Balancers []*BalancingRule `json:"balancers"`
|
Balancers []*BalancingRule `json:"balancers"`
|
||||||
Script string `json:"script"`
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *RouterConfig) getDomainStrategy() router.Config_DomainStrategy {
|
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) {
|
func (c *RouterConfig) Build() (*router.Config, error) {
|
||||||
config := new(router.Config)
|
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()
|
config.DomainStrategy = c.getDomainStrategy()
|
||||||
|
|
||||||
var rawRuleList []json.RawMessage
|
var rawRuleList []json.RawMessage
|
||||||
|
|||||||
@@ -2,8 +2,6 @@ package conf_test
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"encoding/json"
|
"encoding/json"
|
||||||
"os"
|
|
||||||
"path/filepath"
|
|
||||||
"testing"
|
"testing"
|
||||||
"time"
|
"time"
|
||||||
_ "unsafe"
|
_ "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)
|
|
||||||
}
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|||||||
@@ -170,17 +170,9 @@ func (s *Server) processTCP(ctx context.Context, conn stat.Connection, dispatche
|
|||||||
return errors.New("UDP associate with listen port failed")
|
return errors.New("UDP associate with listen port failed")
|
||||||
}
|
}
|
||||||
tempUDPConn.SetTimeout(plcy.Timeouts.ConnectionIdle)
|
tempUDPConn.SetTimeout(plcy.Timeouts.ConnectionIdle)
|
||||||
var udpConn stat.Connection = tempUDPConn
|
|
||||||
if counters, ok := conn.(*stat.CounterConnection); ok {
|
|
||||||
udpConn = &stat.CounterConnection{
|
|
||||||
Connection: tempUDPConn,
|
|
||||||
ReadCounter: counters.ReadCounter,
|
|
||||||
WriteCounter: counters.WriteCounter,
|
|
||||||
}
|
|
||||||
}
|
|
||||||
errCh := make(chan error, 1)
|
errCh := make(chan error, 1)
|
||||||
go func() {
|
go func() {
|
||||||
errCh <- s.handleUDPPayload(ctx, udpConn, dispatcher)
|
errCh <- s.handleUDPPayload(ctx, tempUDPConn, dispatcher)
|
||||||
}()
|
}()
|
||||||
// Associated TCP keeps the UDP alive
|
// Associated TCP keeps the UDP alive
|
||||||
// Close UDP if TCP connection is closed
|
// Close UDP if TCP connection is closed
|
||||||
|
|||||||
@@ -17,7 +17,6 @@ import (
|
|||||||
"github.com/xtls/xray-core/common/errors"
|
"github.com/xtls/xray-core/common/errors"
|
||||||
"github.com/xtls/xray-core/common/net"
|
"github.com/xtls/xray-core/common/net"
|
||||||
"github.com/xtls/xray-core/common/net/cnc"
|
"github.com/xtls/xray-core/common/net/cnc"
|
||||||
"github.com/xtls/xray-core/core"
|
|
||||||
"github.com/xtls/xray-core/transport/internet"
|
"github.com/xtls/xray-core/transport/internet"
|
||||||
"github.com/xtls/xray-core/transport/internet/finalmask"
|
"github.com/xtls/xray-core/transport/internet/finalmask"
|
||||||
"github.com/xtls/xray-core/transport/internet/hysteria/congestion"
|
"github.com/xtls/xray-core/transport/internet/hysteria/congestion"
|
||||||
@@ -29,8 +28,6 @@ import (
|
|||||||
type client struct {
|
type client struct {
|
||||||
sync.Mutex
|
sync.Mutex
|
||||||
|
|
||||||
instance *core.Instance
|
|
||||||
forced bool
|
|
||||||
dest net.Destination
|
dest net.Destination
|
||||||
config *Config
|
config *Config
|
||||||
tlsConfig *gotls.Config
|
tlsConfig *gotls.Config
|
||||||
@@ -67,14 +64,11 @@ func (c *client) close() {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (c *client) dial(ctx context.Context) error {
|
func (c *client) dial(ctx context.Context) error {
|
||||||
if c.forced {
|
status := c.status()
|
||||||
return errors.New("client is closed")
|
if status == StatusActive {
|
||||||
}
|
|
||||||
|
|
||||||
switch c.status() {
|
|
||||||
case StatusActive:
|
|
||||||
return nil
|
return nil
|
||||||
case StatusInactive:
|
}
|
||||||
|
if status == StatusInactive {
|
||||||
c.close()
|
c.close()
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -96,7 +90,7 @@ func (c *client) dial(ctx context.Context) error {
|
|||||||
ChromeParrot: !quicParams.DisableChromeParrot,
|
ChromeParrot: !quicParams.DisableChromeParrot,
|
||||||
EnableDatagrams: true,
|
EnableDatagrams: true,
|
||||||
MaxDatagramFrameSize: MaxDatagramFrameSize,
|
MaxDatagramFrameSize: MaxDatagramFrameSize,
|
||||||
OmitMaxDatagramFrameSize: true,
|
OmitMaxDatagramFrameSize: time.Now().After(time.Date(2026, 9, 1, 0, 0, 0, 0, time.UTC)),
|
||||||
DisablePathManager: true,
|
DisablePathManager: true,
|
||||||
}
|
}
|
||||||
if quicParams.InitStreamReceiveWindow == 0 {
|
if quicParams.InitStreamReceiveWindow == 0 {
|
||||||
@@ -262,12 +256,9 @@ func (c *client) udp(ctx context.Context) (stat.Connection, error) {
|
|||||||
return c.udpSM.udp()
|
return c.udpSM.udp()
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *client) clean(force bool) {
|
func (c *client) clean() {
|
||||||
c.Lock()
|
c.Lock()
|
||||||
if force {
|
if c.status() == StatusInactive {
|
||||||
c.forced = true
|
|
||||||
}
|
|
||||||
if status := c.status(); force && status != StatusNull || status == StatusInactive {
|
|
||||||
c.close()
|
c.close()
|
||||||
}
|
}
|
||||||
c.Unlock()
|
c.Unlock()
|
||||||
@@ -286,23 +277,11 @@ type clientManager struct {
|
|||||||
func (m *clientManager) clean() {
|
func (m *clientManager) clean() {
|
||||||
ticker := time.NewTicker(idleCleanupInterval)
|
ticker := time.NewTicker(idleCleanupInterval)
|
||||||
for range ticker.C {
|
for range ticker.C {
|
||||||
var forced []dialerConf
|
|
||||||
|
|
||||||
m.RLock()
|
m.RLock()
|
||||||
for k, c := range m.m {
|
for _, c := range m.m {
|
||||||
force := c.instance != nil && !c.instance.IsRunning()
|
c.clean()
|
||||||
c.clean(force)
|
|
||||||
if force {
|
|
||||||
forced = append(forced, k)
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
m.RUnlock()
|
m.RUnlock()
|
||||||
|
|
||||||
for i := range forced {
|
|
||||||
m.Lock()
|
|
||||||
delete(m.m, forced[i])
|
|
||||||
m.Unlock()
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -327,18 +306,15 @@ func Dial(ctx context.Context, dest net.Destination, streamSettings *internet.Me
|
|||||||
go manager.clean()
|
go manager.clean()
|
||||||
})
|
})
|
||||||
|
|
||||||
dialerConfKey := dialerConf{dest, streamSettings}
|
|
||||||
|
|
||||||
manager.RLock()
|
manager.RLock()
|
||||||
c := manager.m[dialerConfKey]
|
c := manager.m[dialerConf{dest, streamSettings}]
|
||||||
manager.RUnlock()
|
manager.RUnlock()
|
||||||
|
|
||||||
if c == nil {
|
if c == nil {
|
||||||
manager.Lock()
|
manager.Lock()
|
||||||
c = manager.m[dialerConfKey]
|
c = manager.m[dialerConf{dest, streamSettings}]
|
||||||
if c == nil {
|
if c == nil {
|
||||||
c = &client{
|
c = &client{
|
||||||
instance: core.FromContext(ctx),
|
|
||||||
dest: dest,
|
dest: dest,
|
||||||
config: streamSettings.ProtocolSettings.(*Config),
|
config: streamSettings.ProtocolSettings.(*Config),
|
||||||
tlsConfig: tlsConfig.GetTLSConfig(tls.WithDestination(dest)),
|
tlsConfig: tlsConfig.GetTLSConfig(tls.WithDestination(dest)),
|
||||||
@@ -346,7 +322,7 @@ func Dial(ctx context.Context, dest net.Destination, streamSettings *internet.Me
|
|||||||
finalMask: streamSettings.FinalMask,
|
finalMask: streamSettings.FinalMask,
|
||||||
quicParams: streamSettings.QuicParams,
|
quicParams: streamSettings.QuicParams,
|
||||||
}
|
}
|
||||||
manager.m[dialerConfKey] = c
|
manager.m[dialerConf{dest, streamSettings}] = c
|
||||||
}
|
}
|
||||||
manager.Unlock()
|
manager.Unlock()
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -54,12 +54,8 @@ func useWarp(config *Config, tlsConfig *gotls.Config) error {
|
|||||||
tlsConfig.GetClientCertificate = func(*gotls.CertificateRequestInfo) (*gotls.Certificate, error) {
|
tlsConfig.GetClientCertificate = func(*gotls.CertificateRequestInfo) (*gotls.Certificate, error) {
|
||||||
return cert, nil
|
return cert, nil
|
||||||
}
|
}
|
||||||
if tlsConfig.InsecureSkipVerify == true {
|
if publicKey := config.Warp.PublicKey; len(publicKey) > 0 {
|
||||||
return nil // pcs or vcn or both
|
verify := tlsConfig.VerifyPeerCertificate
|
||||||
}
|
|
||||||
if _, err = x509.ParsePKIXPublicKey(config.Warp.PublicKey); err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
tlsConfig.InsecureSkipVerify = true
|
tlsConfig.InsecureSkipVerify = true
|
||||||
tlsConfig.VerifyPeerCertificate = func(raw [][]byte, chains [][]*x509.Certificate) error {
|
tlsConfig.VerifyPeerCertificate = func(raw [][]byte, chains [][]*x509.Certificate) error {
|
||||||
if len(raw) == 0 {
|
if len(raw) == 0 {
|
||||||
@@ -69,10 +65,14 @@ func useWarp(config *Config, tlsConfig *gotls.Config) error {
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
if !bytes.Equal(leaf.RawSubjectPublicKeyInfo, config.Warp.PublicKey) {
|
if !bytes.Equal(leaf.RawSubjectPublicKeyInfo, publicKey) {
|
||||||
return errors.New("the WARP endpoint's key doesn't match \"publicKey\"")
|
return errors.New("the WARP endpoint's key doesn't match \"publicKey\"")
|
||||||
}
|
}
|
||||||
|
if verify != nil {
|
||||||
|
return verify(raw, chains)
|
||||||
|
}
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
}
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -17,7 +17,6 @@ import (
|
|||||||
|
|
||||||
utls "github.com/refraction-networking/utls"
|
utls "github.com/refraction-networking/utls"
|
||||||
"github.com/xtls/xray-core/common/crypto"
|
"github.com/xtls/xray-core/common/crypto"
|
||||||
"github.com/xtls/xray-core/common/session"
|
|
||||||
"golang.org/x/net/http2"
|
"golang.org/x/net/http2"
|
||||||
|
|
||||||
"github.com/miekg/dns"
|
"github.com/miekg/dns"
|
||||||
@@ -199,27 +198,23 @@ func dnsQuery(server string, domain string, sockopt *internet.SocketConfig) ([]b
|
|||||||
IdleConnTimeout: net.ConnIdleTimeout,
|
IdleConnTimeout: net.ConnIdleTimeout,
|
||||||
ReadIdleTimeout: net.ChromeH2KeepAlivePeriod,
|
ReadIdleTimeout: net.ChromeH2KeepAlivePeriod,
|
||||||
DialTLSContext: func(ctx context.Context, network, addr string, cfg *tls.Config) (net.Conn, error) {
|
DialTLSContext: func(ctx context.Context, network, addr string, cfg *tls.Config) (net.Conn, error) {
|
||||||
host, _, err := net.SplitHostPort(addr)
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
dest, err := net.ParseDestination(network + ":" + addr)
|
dest, err := net.ParseDestination(network + ":" + addr)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
dnsCtx := ctx
|
|
||||||
if h2c {
|
|
||||||
dnsCtx = session.ContextWithMitmAlpn11(dnsCtx, false) // for insurance
|
|
||||||
dnsCtx = session.ContextWithMitmServerName(dnsCtx, host)
|
|
||||||
}
|
|
||||||
var conn net.Conn
|
var conn net.Conn
|
||||||
conn, err = internet.DialSystem(dnsCtx, dest, sockopt)
|
|
||||||
|
conn, err = internet.DialSystem(ctx, dest, sockopt)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
if !h2c {
|
if !h2c {
|
||||||
conn = utls.UClient(conn, &utls.Config{ServerName: host}, utls.HelloChrome_Auto)
|
u, err := url.Parse(server)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
conn = utls.UClient(conn, &utls.Config{ServerName: u.Hostname()}, utls.HelloChrome_Auto)
|
||||||
if err := conn.(*utls.UConn).HandshakeContext(ctx); err != nil {
|
if err := conn.(*utls.UConn).HandshakeContext(ctx); err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user