mirror of
https://github.com/XTLS/Xray-core.git
synced 2026-10-05 21:38:12 +03:00
Compare commits
15
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
14f4bcaf29 | ||
|
|
30a78f1563 | ||
|
|
46d49adc5c | ||
|
|
b71975abed | ||
|
|
6b0f06fb2b | ||
|
|
7da5dae650 | ||
|
|
a3c6cc6aa8 | ||
|
|
1e08d48d88 | ||
|
|
5121c28b85 | ||
|
|
92f7e72490 | ||
|
|
eddad37c69 | ||
|
|
00cde72a9b | ||
|
|
08775afd65 | ||
|
|
e0bae21201 | ||
|
|
9d9a7a1c00 |
@@ -470,9 +470,6 @@ func (d *DefaultDispatcher) routedDispatch(ctx context.Context, link *transport.
|
||||
return // DO NOT CHANGE: the traffic shouldn't be processed by default outbound if the specified outbound tag doesn't exist (yet), e.g., VLESS Reverse Proxy
|
||||
}
|
||||
} else {
|
||||
if err != common.ErrNoClue {
|
||||
errors.LogErrorInner(ctx, err, "failed to pick route for ", destination)
|
||||
}
|
||||
errors.LogInfo(ctx, "default route for ", destination)
|
||||
}
|
||||
}
|
||||
|
||||
+6
-25
@@ -93,7 +93,6 @@ type NameServer struct {
|
||||
UnexpectedIp []*geodata.IPRule `protobuf:"bytes,13,rep,name=unexpected_ip,json=unexpectedIp,proto3" json:"unexpected_ip,omitempty"`
|
||||
ActUnprior bool `protobuf:"varint,14,opt,name=actUnprior,proto3" json:"actUnprior,omitempty"`
|
||||
PolicyID uint32 `protobuf:"varint,17,opt,name=policyID,proto3" json:"policyID,omitempty"`
|
||||
Id string `protobuf:"bytes,18,opt,name=id,proto3" json:"id,omitempty"`
|
||||
unknownFields protoimpl.UnknownFields
|
||||
sizeCache protoimpl.SizeCache
|
||||
}
|
||||
@@ -240,13 +239,6 @@ func (x *NameServer) GetPolicyID() uint32 {
|
||||
return 0
|
||||
}
|
||||
|
||||
func (x *NameServer) GetId() string {
|
||||
if x != nil {
|
||||
return x.Id
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
type Config struct {
|
||||
state protoimpl.MessageState `protogen:"open.v1"`
|
||||
// NameServer list used by this DNS client.
|
||||
@@ -266,10 +258,8 @@ type Config struct {
|
||||
DisableFallback bool `protobuf:"varint,10,opt,name=disableFallback,proto3" json:"disableFallback,omitempty"`
|
||||
DisableFallbackIfMatch bool `protobuf:"varint,11,opt,name=disableFallbackIfMatch,proto3" json:"disableFallbackIfMatch,omitempty"`
|
||||
EnableParallelQuery bool `protobuf:"varint,14,opt,name=enableParallelQuery,proto3" json:"enableParallelQuery,omitempty"`
|
||||
// Absolute path to the Lua DNS query script.
|
||||
Script string `protobuf:"bytes,15,opt,name=script,proto3" json:"script,omitempty"`
|
||||
unknownFields protoimpl.UnknownFields
|
||||
sizeCache protoimpl.SizeCache
|
||||
unknownFields protoimpl.UnknownFields
|
||||
sizeCache protoimpl.SizeCache
|
||||
}
|
||||
|
||||
func (x *Config) Reset() {
|
||||
@@ -379,13 +369,6 @@ func (x *Config) GetEnableParallelQuery() bool {
|
||||
return false
|
||||
}
|
||||
|
||||
func (x *Config) GetScript() string {
|
||||
if x != nil {
|
||||
return x.Script
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
type Config_HostMapping struct {
|
||||
state protoimpl.MessageState `protogen:"open.v1"`
|
||||
Domain *geodata.DomainRule `protobuf:"bytes,2,opt,name=domain,proto3" json:"domain,omitempty"`
|
||||
@@ -452,7 +435,7 @@ var File_app_dns_config_proto protoreflect.FileDescriptor
|
||||
|
||||
const file_app_dns_config_proto_rawDesc = "" +
|
||||
"\n" +
|
||||
"\x14app/dns/config.proto\x12\fxray.app.dns\x1a\x1ccommon/net/destination.proto\x1a\x1bcommon/geodata/geodat.proto\"\xee\x05\n" +
|
||||
"\x14app/dns/config.proto\x12\fxray.app.dns\x1a\x1ccommon/net/destination.proto\x1a\x1bcommon/geodata/geodat.proto\"\xde\x05\n" +
|
||||
"\n" +
|
||||
"NameServer\x123\n" +
|
||||
"\aaddress\x18\x01 \x01(\v2\x19.xray.common.net.EndpointR\aaddress\x12\x1b\n" +
|
||||
@@ -478,11 +461,10 @@ const file_app_dns_config_proto_rawDesc = "" +
|
||||
"\n" +
|
||||
"actUnprior\x18\x0e \x01(\bR\n" +
|
||||
"actUnprior\x12\x1a\n" +
|
||||
"\bpolicyID\x18\x11 \x01(\rR\bpolicyID\x12\x0e\n" +
|
||||
"\x02id\x18\x12 \x01(\tR\x02idB\x0f\n" +
|
||||
"\bpolicyID\x18\x11 \x01(\rR\bpolicyIDB\x0f\n" +
|
||||
"\r_disableCacheB\r\n" +
|
||||
"\v_serveStaleB\x12\n" +
|
||||
"\x10_serveExpiredTTLJ\x04\b\x04\x10\x05\"\x9a\x05\n" +
|
||||
"\x10_serveExpiredTTLJ\x04\b\x04\x10\x05\"\x82\x05\n" +
|
||||
"\x06Config\x129\n" +
|
||||
"\vname_server\x18\x05 \x03(\v2\x18.xray.app.dns.NameServerR\n" +
|
||||
"nameServer\x12\x1b\n" +
|
||||
@@ -498,8 +480,7 @@ const file_app_dns_config_proto_rawDesc = "" +
|
||||
"\x0fdisableFallback\x18\n" +
|
||||
" \x01(\bR\x0fdisableFallback\x126\n" +
|
||||
"\x16disableFallbackIfMatch\x18\v \x01(\bR\x16disableFallbackIfMatch\x120\n" +
|
||||
"\x13enableParallelQuery\x18\x0e \x01(\bR\x13enableParallelQuery\x12\x16\n" +
|
||||
"\x06script\x18\x0f \x01(\tR\x06script\x1a}\n" +
|
||||
"\x13enableParallelQuery\x18\x0e \x01(\bR\x13enableParallelQuery\x1a}\n" +
|
||||
"\vHostMapping\x127\n" +
|
||||
"\x06domain\x18\x02 \x01(\v2\x1f.xray.common.geodata.DomainRuleR\x06domain\x12\x0e\n" +
|
||||
"\x02ip\x18\x03 \x03(\fR\x02ip\x12%\n" +
|
||||
|
||||
@@ -27,7 +27,6 @@ message NameServer {
|
||||
repeated xray.common.geodata.IPRule unexpected_ip = 13;
|
||||
bool actUnprior = 14;
|
||||
uint32 policyID = 17;
|
||||
string id = 18;
|
||||
}
|
||||
|
||||
enum QueryStrategy {
|
||||
@@ -74,7 +73,4 @@ message Config {
|
||||
bool disableFallbackIfMatch = 11;
|
||||
|
||||
bool enableParallelQuery = 14;
|
||||
|
||||
// Absolute path to the Lua DNS query script.
|
||||
string script = 15;
|
||||
}
|
||||
|
||||
@@ -31,8 +31,6 @@ type DNS struct {
|
||||
domainMatcher geodata.DomainMatcher
|
||||
matcherInfos []*DomainMatcherInfo
|
||||
checkSystem bool
|
||||
script *scriptEngine
|
||||
scriptPath string
|
||||
}
|
||||
|
||||
// DomainMatcherInfo contains information attached to index returned by Server.domainMatcher.
|
||||
@@ -182,7 +180,6 @@ func New(ctx context.Context, config *Config) (*DNS, error) {
|
||||
disableFallbackIfMatch: config.DisableFallbackIfMatch,
|
||||
enableParallelQuery: config.EnableParallelQuery,
|
||||
checkSystem: checkSystem,
|
||||
scriptPath: config.Script,
|
||||
}, nil
|
||||
}
|
||||
|
||||
@@ -193,21 +190,11 @@ func (*DNS) Type() interface{} {
|
||||
|
||||
// Start implements common.Runnable.
|
||||
func (s *DNS) Start() error {
|
||||
if s.scriptPath != "" {
|
||||
engine, err := newScriptEngine(s.scriptPath, s)
|
||||
if err != nil {
|
||||
return errors.New("failed to initialize DNS script").Base(err)
|
||||
}
|
||||
s.script = engine
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// Close implements common.Closable.
|
||||
func (s *DNS) Close() error {
|
||||
if s.script != nil {
|
||||
s.script.close()
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -292,9 +279,6 @@ func (s *DNS) LookupIP(domain string, option dns.IPOption) ([]net.IP, uint32, er
|
||||
}
|
||||
|
||||
// Name servers lookup
|
||||
if s.script != nil {
|
||||
return s.script.query(domain, option)
|
||||
}
|
||||
if s.enableParallelQuery {
|
||||
return s.parallelQuery(domain, option)
|
||||
} else {
|
||||
|
||||
-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.
|
||||
type Client struct {
|
||||
id string
|
||||
server Server
|
||||
skipFallback bool
|
||||
expectedIPs geodata.IPMatcher
|
||||
@@ -98,7 +97,7 @@ func NewClient(
|
||||
ipOption dns.IPOption,
|
||||
updateRules func(bool),
|
||||
) (*Client, error) {
|
||||
client := &Client{id: ns.Id}
|
||||
client := &Client{}
|
||||
err := core.RequireFeatures(ctx, func(dispatcher routing.Dispatcher) error {
|
||||
// Create a new server for each client for now
|
||||
server, err := NewServer(ctx, ns.Address.AsDestination(), dispatcher, disableCache, serveStale, serveExpiredTTL, clientIP)
|
||||
|
||||
@@ -49,5 +49,5 @@ func NewLocalNameServer() *LocalNameServer {
|
||||
|
||||
// NewLocalDNSClient creates localdns client object for directly lookup in system DNS.
|
||||
func NewLocalDNSClient(ipOption dns.IPOption) *Client {
|
||||
return &Client{id: "localhost", server: NewLocalNameServer(), ipOption: &ipOption}
|
||||
return &Client{server: NewLocalNameServer(), ipOption: &ipOption}
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
+4
-14
@@ -587,10 +587,8 @@ type Config struct {
|
||||
DomainStrategy Config_DomainStrategy `protobuf:"varint,1,opt,name=domain_strategy,json=domainStrategy,proto3,enum=xray.app.router.Config_DomainStrategy" json:"domain_strategy,omitempty"`
|
||||
Rule []*RoutingRule `protobuf:"bytes,2,rep,name=rule,proto3" json:"rule,omitempty"`
|
||||
BalancingRule []*BalancingRule `protobuf:"bytes,3,rep,name=balancing_rule,json=balancingRule,proto3" json:"balancing_rule,omitempty"`
|
||||
// Absolute path to the Lua routing script.
|
||||
Script string `protobuf:"bytes,4,opt,name=script,proto3" json:"script,omitempty"`
|
||||
unknownFields protoimpl.UnknownFields
|
||||
sizeCache protoimpl.SizeCache
|
||||
unknownFields protoimpl.UnknownFields
|
||||
sizeCache protoimpl.SizeCache
|
||||
}
|
||||
|
||||
func (x *Config) Reset() {
|
||||
@@ -644,13 +642,6 @@ func (x *Config) GetBalancingRule() []*BalancingRule {
|
||||
return nil
|
||||
}
|
||||
|
||||
func (x *Config) GetScript() string {
|
||||
if x != nil {
|
||||
return x.Script
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
var File_app_router_config_proto protoreflect.FileDescriptor
|
||||
|
||||
const file_app_router_config_proto_rawDesc = "" +
|
||||
@@ -708,12 +699,11 @@ const file_app_router_config_proto_rawDesc = "" +
|
||||
"\tbaselines\x18\x03 \x03(\x03R\tbaselines\x12\x1a\n" +
|
||||
"\bexpected\x18\x04 \x01(\x05R\bexpected\x12\x16\n" +
|
||||
"\x06maxRTT\x18\x05 \x01(\x03R\x06maxRTT\x12\x1c\n" +
|
||||
"\ttolerance\x18\x06 \x01(\x02R\ttolerance\"\xae\x02\n" +
|
||||
"\ttolerance\x18\x06 \x01(\x02R\ttolerance\"\x96\x02\n" +
|
||||
"\x06Config\x12O\n" +
|
||||
"\x0fdomain_strategy\x18\x01 \x01(\x0e2&.xray.app.router.Config.DomainStrategyR\x0edomainStrategy\x120\n" +
|
||||
"\x04rule\x18\x02 \x03(\v2\x1c.xray.app.router.RoutingRuleR\x04rule\x12E\n" +
|
||||
"\x0ebalancing_rule\x18\x03 \x03(\v2\x1e.xray.app.router.BalancingRuleR\rbalancingRule\x12\x16\n" +
|
||||
"\x06script\x18\x04 \x01(\tR\x06script\"B\n" +
|
||||
"\x0ebalancing_rule\x18\x03 \x03(\v2\x1e.xray.app.router.BalancingRuleR\rbalancingRule\"B\n" +
|
||||
"\x0eDomainStrategy\x12\b\n" +
|
||||
"\x04AsIs\x10\x00\x12\x10\n" +
|
||||
"\fIpIfNonMatch\x10\x02\x12\x0e\n" +
|
||||
|
||||
@@ -110,6 +110,4 @@ message Config {
|
||||
DomainStrategy domain_strategy = 1;
|
||||
repeated RoutingRule rule = 2;
|
||||
repeated BalancingRule balancing_rule = 3;
|
||||
// Absolute path to the Lua routing script.
|
||||
string script = 4;
|
||||
}
|
||||
|
||||
@@ -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 {
|
||||
domainStrategy Config_DomainStrategy
|
||||
rules atomic.Pointer[[]*Rule]
|
||||
scriptPath string
|
||||
script *scriptEngine
|
||||
balancers atomic.Pointer[map[string]*Balancer]
|
||||
dns dns.Client
|
||||
|
||||
@@ -42,7 +40,6 @@ type Route struct {
|
||||
// Init initializes the Router.
|
||||
func (r *Router) Init(ctx context.Context, config *Config, d dns.Client, ohm outbound.Manager, dispatcher routing.Dispatcher) error {
|
||||
r.domainStrategy = config.DomainStrategy
|
||||
r.scriptPath = config.Script
|
||||
r.dns = d
|
||||
r.ctx = ctx
|
||||
r.ohm = ohm
|
||||
@@ -55,10 +52,6 @@ func (r *Router) Init(ctx context.Context, config *Config, d dns.Client, ohm out
|
||||
|
||||
// PickRoute implements routing.Router.
|
||||
func (r *Router) PickRoute(ctx routing.Context) (routing.Route, error) {
|
||||
if r.script != nil {
|
||||
return r.script.pickRoute(ctx)
|
||||
}
|
||||
|
||||
originalCtx := ctx
|
||||
rule, ctx, err := r.pickRouteInternal(ctx)
|
||||
if err != nil {
|
||||
@@ -228,13 +221,6 @@ func (r *Router) pickRouteInternal(ctx routing.Context) (*Rule, routing.Context,
|
||||
|
||||
// Start implements common.Runnable.
|
||||
func (r *Router) Start() error {
|
||||
if r.scriptPath != "" {
|
||||
engine, err := newScriptEngine(r.scriptPath, r)
|
||||
if err != nil {
|
||||
return errors.New("failed to initialize routing script").Base(err)
|
||||
}
|
||||
r.script = engine
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -249,9 +235,6 @@ func closeWebhooks(rules []*Rule) {
|
||||
|
||||
// Close implements common.Closable.
|
||||
func (r *Router) Close() error {
|
||||
if r.script != nil {
|
||||
r.script.close()
|
||||
}
|
||||
r.mu.Lock()
|
||||
defer r.mu.Unlock()
|
||||
closeWebhooks(*r.rules.Load())
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -314,12 +314,10 @@ func (m *ClientWorker) Dispatch(ctx context.Context, link *transport.Link) bool
|
||||
}
|
||||
|
||||
sm := m.sessionManager
|
||||
s := sm.Allocate(&m.strategy)
|
||||
s := sm.Allocate(&m.strategy, link.Reader, link.Writer)
|
||||
if s == nil {
|
||||
return false
|
||||
}
|
||||
s.input = link.Reader
|
||||
s.output = link.Writer
|
||||
go fetchInput(ctx, s, m.link.Writer)
|
||||
if _, ok := link.Reader.(*pipe.Reader); !ok {
|
||||
select {
|
||||
|
||||
@@ -51,7 +51,7 @@ func (m *SessionManager) Count() int {
|
||||
return int(m.count)
|
||||
}
|
||||
|
||||
func (m *SessionManager) Allocate(Strategy *ClientStrategy) *Session {
|
||||
func (m *SessionManager) Allocate(Strategy *ClientStrategy, input buf.Reader, output buf.Writer) *Session {
|
||||
m.Lock()
|
||||
defer m.Unlock()
|
||||
|
||||
@@ -64,6 +64,8 @@ func (m *SessionManager) Allocate(Strategy *ClientStrategy) *Session {
|
||||
|
||||
m.count++
|
||||
s := &Session{
|
||||
input: input,
|
||||
output: output,
|
||||
ID: m.count,
|
||||
parent: m,
|
||||
done: done.New(),
|
||||
|
||||
@@ -9,7 +9,7 @@ import (
|
||||
func TestSessionManagerAdd(t *testing.T) {
|
||||
m := NewSessionManager()
|
||||
|
||||
s := m.Allocate(&ClientStrategy{})
|
||||
s := m.Allocate(&ClientStrategy{}, nil, nil)
|
||||
if s.ID != 1 {
|
||||
t.Error("id: ", s.ID)
|
||||
}
|
||||
@@ -17,7 +17,7 @@ func TestSessionManagerAdd(t *testing.T) {
|
||||
t.Error("size: ", m.Size())
|
||||
}
|
||||
|
||||
s = m.Allocate(&ClientStrategy{})
|
||||
s = m.Allocate(&ClientStrategy{}, nil, nil)
|
||||
if s.ID != 2 {
|
||||
t.Error("id: ", s.ID)
|
||||
}
|
||||
@@ -39,7 +39,7 @@ func TestSessionManagerAdd(t *testing.T) {
|
||||
|
||||
func TestSessionManagerClose(t *testing.T) {
|
||||
m := NewSessionManager()
|
||||
s := m.Allocate(&ClientStrategy{})
|
||||
s := m.Allocate(&ClientStrategy{}, nil, nil)
|
||||
|
||||
if m.CloseIfNoSessionAndIdle(m.Size(), m.Count()) {
|
||||
t.Error("able to close")
|
||||
|
||||
@@ -1,8 +1,6 @@
|
||||
package platform // import "github.com/xtls/xray-core/common/platform"
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strconv"
|
||||
@@ -92,49 +90,3 @@ func GetConfDirPath() string {
|
||||
configPath := NewEnvFlag(ConfdirLocation).GetValue(func() string { return "" })
|
||||
return configPath
|
||||
}
|
||||
|
||||
// ResolveLuaFile finds a local Lua script and returns its absolute path.
|
||||
// Relative paths: XRAY_LOCATION_CONFDIR > XRAY_LOCATION_CONFIG > working dir > executable dir.
|
||||
func ResolveLuaFile(path string) (string, error) {
|
||||
if path == "" {
|
||||
return "", errors.New("Lua file path is empty")
|
||||
}
|
||||
paths := []string{path}
|
||||
if !filepath.IsAbs(path) {
|
||||
paths = nil
|
||||
for _, dir := range []string{
|
||||
GetConfDirPath(),
|
||||
NewEnvFlag(ConfigLocation).GetValue(func() string { return "" }),
|
||||
".",
|
||||
getExecutableDir(),
|
||||
} {
|
||||
if dir != "" {
|
||||
paths = append(paths, filepath.Join(dir, path))
|
||||
}
|
||||
}
|
||||
}
|
||||
return resolveFile(paths)
|
||||
}
|
||||
|
||||
func resolveFile(paths []string) (string, error) {
|
||||
var tried []string
|
||||
for _, path := range paths {
|
||||
path, err := filepath.Abs(path)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("failed to resolve file path: %w", err)
|
||||
}
|
||||
tried = append(tried, path)
|
||||
info, err := os.Stat(path)
|
||||
if errors.Is(err, os.ErrNotExist) {
|
||||
continue
|
||||
}
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("failed to inspect file %q: %w", path, err)
|
||||
}
|
||||
if !info.Mode().IsRegular() {
|
||||
return "", fmt.Errorf("file is not a regular file: %s", path)
|
||||
}
|
||||
return path, nil
|
||||
}
|
||||
return "", fmt.Errorf("file not found; tried %q: %w", tried, os.ErrNotExist)
|
||||
}
|
||||
|
||||
@@ -1,7 +1,6 @@
|
||||
package platform_test
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"runtime"
|
||||
@@ -65,53 +64,3 @@ func TestGetAssetLocation(t *testing.T) {
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestResolveLuaFile(t *testing.T) {
|
||||
workingDir := t.TempDir()
|
||||
t.Chdir(workingDir)
|
||||
executable, err := os.Executable()
|
||||
common.Must(err)
|
||||
file, err := os.CreateTemp(filepath.Dir(executable), "lua-*.lua")
|
||||
common.Must(err)
|
||||
common.Must(file.Close())
|
||||
defer os.Remove(file.Name())
|
||||
|
||||
name := filepath.Base(file.Name())
|
||||
paths := []string{
|
||||
filepath.Join(t.TempDir(), name),
|
||||
filepath.Join(t.TempDir(), name),
|
||||
filepath.Join(workingDir, name),
|
||||
file.Name(),
|
||||
}
|
||||
t.Setenv(ConfdirLocation, filepath.Dir(paths[0]))
|
||||
t.Setenv(ConfigLocation, filepath.Dir(paths[1]))
|
||||
for _, path := range paths[:3] {
|
||||
common.Must(os.WriteFile(path, nil, 0o600))
|
||||
}
|
||||
if got, err := ResolveLuaFile(paths[2]); err != nil || got != paths[2] {
|
||||
t.Fatalf("absolute path = %q, %v; want %q", got, err, paths[2])
|
||||
}
|
||||
for i, want := range paths {
|
||||
if i == 2 {
|
||||
t.Setenv(ConfdirLocation, "")
|
||||
t.Setenv(ConfigLocation, "")
|
||||
}
|
||||
if got, err := ResolveLuaFile(name); err != nil || got != want {
|
||||
t.Fatalf("resolved path = %q, %v; want %q", got, err, want)
|
||||
}
|
||||
common.Must(os.Remove(want))
|
||||
}
|
||||
if _, err := ResolveLuaFile(name); !errors.Is(err, os.ErrNotExist) {
|
||||
t.Fatalf("missing file error = %v", err)
|
||||
}
|
||||
|
||||
t.Setenv(ConfdirLocation, filepath.Dir(paths[0]))
|
||||
t.Setenv(ConfigLocation, filepath.Dir(paths[1]))
|
||||
common.Must(os.Mkdir(paths[0], 0o700))
|
||||
common.Must(os.WriteFile(paths[1], nil, 0o600))
|
||||
for _, path := range []string{"", name, filepath.Join(t.TempDir(), name)} {
|
||||
if _, err := ResolveLuaFile(path); err == nil {
|
||||
t.Fatalf("accepted invalid path %q", path)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -146,6 +146,10 @@ func SniffQUIC(b []byte) (*SniffHeader, error) {
|
||||
}
|
||||
|
||||
restPayload := b[hdrLen+int(packetLen):]
|
||||
// cachedReader can concatenate zero-padded UDP datagrams.
|
||||
for len(restPayload) > 0 && restPayload[0] == 0 {
|
||||
restPayload = restPayload[1:]
|
||||
}
|
||||
if !isQUICInitial { // Skip this packet if it's not initial packet
|
||||
b = restPayload
|
||||
continue
|
||||
|
||||
File diff suppressed because one or more lines are too long
@@ -21,7 +21,6 @@ require (
|
||||
github.com/stretchr/testify v1.12.1
|
||||
github.com/vishvananda/netlink v1.3.1
|
||||
github.com/xtls/reality v0.0.0-20260908062103-8cdf7bf9c7f0
|
||||
github.com/yuin/gopher-lua v1.1.2
|
||||
go4.org/netipx v0.0.0-20231129151722-fdeea329fbba
|
||||
golang.org/x/crypto v0.57.0
|
||||
golang.org/x/exp v0.0.0-20240506185415-9bf2ced13842
|
||||
@@ -35,7 +34,6 @@ require (
|
||||
google.golang.org/protobuf v1.36.12
|
||||
gvisor.dev/gvisor v0.0.0-20260122175437-89a5d21be8f0
|
||||
h12.io/socks v1.0.3
|
||||
layeh.com/gopher-luar v1.0.11
|
||||
lukechampine.com/blake3 v1.4.1
|
||||
mvdan.cc/gofumpt v0.12.0
|
||||
)
|
||||
|
||||
@@ -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/apernet/quic-go v0.61.1-0.20260806010916-184d081eef3e h1:5mgtR5gwIgBKMiGI1QdXldZZ+SNor06Nbu1wCBulQBg=
|
||||
github.com/apernet/quic-go v0.61.1-0.20260806010916-184d081eef3e/go.mod h1:x7qxEvX6MCVtDuBKHj3E+88+BtrbEMuAL5qGUKItjW8=
|
||||
github.com/chzyer/logex v1.1.10/go.mod h1:+Ywpsq7O8HXn0nuIou7OrIPyXbp3wmkHB+jjWRnGsAI=
|
||||
github.com/chzyer/readline v0.0.0-20180603132655-2972be24d48e/go.mod h1:nSuG5e5PlCu98SY8svDHJxuZscDgtXS6KTTbou5AhLI=
|
||||
github.com/chzyer/test v0.0.0-20180213035817-a1ea475d72b1/go.mod h1:Q3SI9o4m/ZMnBNeIyt5eFwwo7qiLfzFZmjNmxjkiQlU=
|
||||
github.com/cloudflare/circl v1.6.5 h1:O64F26HEqNhznd/hrC5KZXVKYuKM2rx4deZDTc4ihQA=
|
||||
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=
|
||||
@@ -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/go.mod h1:DsJblcWDGt76+FVqBVwbwRhxyyNJsGV48gJLch0OOWI=
|
||||
github.com/yuin/goldmark v1.4.1/go.mod h1:mwnBkeHKe2W/ZEtQ+71ViKU8L12m81fl3OWwC1Zlc8k=
|
||||
github.com/yuin/gopher-lua v0.0.0-20190206043414-8bfc7677f583/go.mod h1:gqRgreBUhTSL0GeU64rtZ3Uq3wtjOa/TB2YfrtkCbVQ=
|
||||
github.com/yuin/gopher-lua v1.1.2 h1:yF/FjE3hD65tBbt0VXLE13HWS9h34fdzJmrWRXwobGA=
|
||||
github.com/yuin/gopher-lua v1.1.2/go.mod h1:7aRmXIWl37SqRf0koeyylBEzJ+aPt8A+mmkQ4f1ntR8=
|
||||
go.uber.org/mock v0.5.2 h1:LbtPTcP8A5k9WPXj54PPPbjcI4Y6lhyOZXn+VS7wNko=
|
||||
go.uber.org/mock v0.5.2/go.mod h1:wLlUxC2vVTPTaE3UD51E0BGOAElKrILxhVSDYQLld5o=
|
||||
go.yaml.in/yaml/v3 v3.0.5 h1:N6y/pJk8buWs9NY5ERU2HSMfm+IuD/OtfdAnq6kESPw=
|
||||
@@ -111,7 +105,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.23.0 h1:KameEIfc1IkluZyXWLn39Wd4tURc6GbCiISGiZm2bQk=
|
||||
golang.org/x/sync v0.23.0/go.mod h1:sUUOizhqBxiL6pEWpqNLUiaJn1ShEbZ6BBqskPbjZm0=
|
||||
golang.org/x/sys v0.0.0-20190204203706-41f3e6584952/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY=
|
||||
golang.org/x/sys v0.0.0-20190215142949-d0b11bdaac8a/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY=
|
||||
golang.org/x/sys v0.0.0-20190412213103-97732733099d/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs=
|
||||
golang.org/x/sys v0.0.0-20201119102817-f84b799fce68/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs=
|
||||
@@ -162,8 +155,6 @@ gvisor.dev/gvisor v0.0.0-20260122175437-89a5d21be8f0 h1:Lk6hARj5UPY47dBep70OD/TI
|
||||
gvisor.dev/gvisor v0.0.0-20260122175437-89a5d21be8f0/go.mod h1:QkHjoMIBaYtpVufgwv3keYAbln78mBoCuShZrPrer1Q=
|
||||
h12.io/socks v1.0.3 h1:Ka3qaQewws4j4/eDQnOdpr4wXsC//dXtWvftlIcCQUo=
|
||||
h12.io/socks v1.0.3/go.mod h1:AIhxy1jOId/XCz9BO+EIgNL2rQiPTBNnOfnVnQ+3Eck=
|
||||
layeh.com/gopher-luar v1.0.11 h1:8zJudpKI6HWkoh9eyyNFaTM79PY6CAPcIr6X/KTiliw=
|
||||
layeh.com/gopher-luar v1.0.11/go.mod h1:TPnIVCZ2RJBndm7ohXyaqfhzjlZ+OA2SZR/YwL8tECk=
|
||||
lukechampine.com/blake3 v1.4.1 h1:I3Smz7gso8w4/TunLKec6K2fn+kyKtDxr/xcQEN84Wg=
|
||||
lukechampine.com/blake3 v1.4.1/go.mod h1:QFosUxmjB8mnrWFSNwKmvxHpfY72bmD2tQ0kBMM3kwo=
|
||||
mvdan.cc/gofumpt v0.12.0 h1:1Lbudkz2kpM9Cjz2pL4M19u7q+GaEhCTNf7N9mfpcho=
|
||||
|
||||
@@ -14,11 +14,9 @@ import (
|
||||
"github.com/xtls/xray-core/common/errors"
|
||||
"github.com/xtls/xray-core/common/geodata"
|
||||
"github.com/xtls/xray-core/common/net"
|
||||
"github.com/xtls/xray-core/common/platform"
|
||||
)
|
||||
|
||||
type NameServerConfig struct {
|
||||
ID string `json:"id"`
|
||||
Address *Address `json:"address"`
|
||||
ClientIP *Address `json:"clientIp"`
|
||||
Port uint16 `json:"port"`
|
||||
@@ -45,7 +43,6 @@ func (c *NameServerConfig) UnmarshalJSON(data []byte) error {
|
||||
}
|
||||
|
||||
var advanced struct {
|
||||
ID string `json:"id"`
|
||||
Address *Address `json:"address"`
|
||||
ClientIP *Address `json:"clientIp"`
|
||||
Port uint16 `json:"port"`
|
||||
@@ -63,7 +60,6 @@ func (c *NameServerConfig) UnmarshalJSON(data []byte) error {
|
||||
UnexpectedIPs StringList `json:"unexpectedIPs"`
|
||||
}
|
||||
if err := json.Unmarshal(data, &advanced); err == nil {
|
||||
c.ID = advanced.ID
|
||||
c.Address = advanced.Address
|
||||
c.ClientIP = advanced.ClientIP
|
||||
c.Port = advanced.Port
|
||||
@@ -138,7 +134,6 @@ func (c *NameServerConfig) Build() (*dns.NameServer, error) {
|
||||
}
|
||||
|
||||
return &dns.NameServer{
|
||||
Id: c.ID,
|
||||
Address: &net.Endpoint{
|
||||
Network: net.Network_UDP,
|
||||
Address: c.Address.Build(),
|
||||
@@ -164,7 +159,6 @@ func (c *NameServerConfig) Build() (*dns.NameServer, error) {
|
||||
// DNSConfig is a JSON serializable object for dns.Config
|
||||
type DNSConfig struct {
|
||||
Servers []*NameServerConfig `json:"servers"`
|
||||
Script string `json:"script"`
|
||||
Hosts *HostsWrapper `json:"hosts"`
|
||||
ClientIP *Address `json:"clientIp"`
|
||||
Tag string `json:"tag"`
|
||||
@@ -284,14 +278,6 @@ func (c *DNSConfig) Build() (*dns.Config, error) {
|
||||
QueryStrategy: resolveQueryStrategy(c.QueryStrategy),
|
||||
}
|
||||
|
||||
if c.Script != "" {
|
||||
path, err := platform.ResolveLuaFile(c.Script)
|
||||
if err != nil {
|
||||
return nil, errors.New("failed to resolve DNS script: ", c.Script).Base(err)
|
||||
}
|
||||
config.Script = path
|
||||
}
|
||||
|
||||
if c.ClientIP != nil {
|
||||
if !c.ClientIP.Family().IsIP() {
|
||||
return nil, errors.New("not an IP address:", c.ClientIP.String())
|
||||
|
||||
@@ -2,8 +2,6 @@ package conf_test
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
|
||||
"github.com/google/go-cmp/cmp"
|
||||
@@ -124,51 +122,3 @@ func TestDNSConfigParsing(t *testing.T) {
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestDNSScriptConfig(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
t.Setenv("xray.location.confdir", dir)
|
||||
path := filepath.Join(dir, "lookup.lua")
|
||||
if err := os.WriteFile(path, []byte("function HandleDNSQuery(domain, ipv4, ipv6, fake) end"), 0o600); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
for _, tc := range []struct {
|
||||
name string
|
||||
script string
|
||||
wantError bool
|
||||
}{
|
||||
{"relative", "lookup.lua", false},
|
||||
{"absolute", path, false},
|
||||
{"missing", "missing.lua", true},
|
||||
{"directory", dir, true},
|
||||
} {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
built, err := (&DNSConfig{Script: tc.script}).Build()
|
||||
if tc.wantError {
|
||||
if err == nil {
|
||||
t.Fatal("Build accepted an invalid script path")
|
||||
}
|
||||
return
|
||||
}
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if built.Script != path {
|
||||
t.Fatalf("script path = %q, want %q", built.Script, path)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
var parsed DNSConfig
|
||||
if err := json.Unmarshal([]byte(`{"servers":[{"id":"primary","address":"1.1.1.1"}]}`), &parsed); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
built, err := parsed.Build()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(built.NameServer) != 1 || built.NameServer[0].Id != "primary" {
|
||||
t.Fatalf("nameserver IDs = %v, want primary", built.NameServer)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -7,7 +7,6 @@ import (
|
||||
"github.com/xtls/xray-core/app/router"
|
||||
"github.com/xtls/xray-core/common/errors"
|
||||
"github.com/xtls/xray-core/common/geodata"
|
||||
"github.com/xtls/xray-core/common/platform"
|
||||
"github.com/xtls/xray-core/common/serial"
|
||||
|
||||
"google.golang.org/protobuf/proto"
|
||||
@@ -73,7 +72,6 @@ type RouterConfig struct {
|
||||
RuleList []json.RawMessage `json:"rules"`
|
||||
DomainStrategy *string `json:"domainStrategy"`
|
||||
Balancers []*BalancingRule `json:"balancers"`
|
||||
Script string `json:"script"`
|
||||
}
|
||||
|
||||
func (c *RouterConfig) getDomainStrategy() router.Config_DomainStrategy {
|
||||
@@ -94,15 +92,6 @@ func (c *RouterConfig) getDomainStrategy() router.Config_DomainStrategy {
|
||||
|
||||
func (c *RouterConfig) Build() (*router.Config, error) {
|
||||
config := new(router.Config)
|
||||
|
||||
if c.Script != "" {
|
||||
path, err := platform.ResolveLuaFile(c.Script)
|
||||
if err != nil {
|
||||
return nil, errors.New("failed to resolve routing script").Base(err)
|
||||
}
|
||||
config.Script = path
|
||||
}
|
||||
|
||||
config.DomainStrategy = c.getDomainStrategy()
|
||||
|
||||
var rawRuleList []json.RawMessage
|
||||
|
||||
@@ -2,8 +2,6 @@ package conf_test
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
"time"
|
||||
_ "unsafe"
|
||||
@@ -238,39 +236,3 @@ func TestRouterConfig(t *testing.T) {
|
||||
},
|
||||
})
|
||||
}
|
||||
|
||||
func TestRouterScriptConfig(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
t.Setenv("xray.location.confdir", dir)
|
||||
path := filepath.Join(dir, "route.lua")
|
||||
if err := os.WriteFile(path, []byte("function HandleRoute() end"), 0o600); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
for _, tc := range []struct {
|
||||
name string
|
||||
script string
|
||||
wantError bool
|
||||
}{
|
||||
{"relative", "route.lua", false},
|
||||
{"absolute", path, false},
|
||||
{"missing", "missing.lua", true},
|
||||
{"directory", dir, true},
|
||||
} {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
built, err := (&RouterConfig{Script: tc.script}).Build()
|
||||
if tc.wantError {
|
||||
if err == nil {
|
||||
t.Fatal("Build accepted invalid script path")
|
||||
}
|
||||
return
|
||||
}
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if built.Script != path {
|
||||
t.Fatalf("script path = %q, want %q", built.Script, path)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -15,7 +15,6 @@ import (
|
||||
googleuuid "github.com/google/uuid"
|
||||
"github.com/xtls/xray-core/common/errors"
|
||||
"github.com/xtls/xray-core/common/net"
|
||||
"github.com/xtls/xray-core/common/serial"
|
||||
"github.com/xtls/xray-core/transport/internet/finalmask/fragment"
|
||||
"github.com/xtls/xray-core/transport/internet/finalmask/header/custom"
|
||||
"github.com/xtls/xray-core/transport/internet/finalmask/mkcp/aes128gcm"
|
||||
@@ -792,37 +791,15 @@ func (c *Sudoku) Build() (proto.Message, error) {
|
||||
}
|
||||
|
||||
type XDNSDomain struct {
|
||||
Name string `json:"name"`
|
||||
LenLimit int32 `json:"lenLimit"`
|
||||
LabelLimit int32 `json:"labelLimit"`
|
||||
Types []int32 `json:"types"`
|
||||
Edns0 int32 `json:"edns0"`
|
||||
Names []string `json:"names"`
|
||||
LenLimit int32 `json:"lenLimit"`
|
||||
LabelLimit int32 `json:"labelLimit"`
|
||||
Types []int32 `json:"types"`
|
||||
Edns0 int32 `json:"edns0"`
|
||||
}
|
||||
|
||||
type XDNSResolverTCP struct {
|
||||
Addr string `json:"addr"`
|
||||
}
|
||||
|
||||
func (c *XDNSResolverTCP) Build() (proto.Message, error) {
|
||||
return &xdns.TCPResolverProto{Addr: c.Addr}, nil
|
||||
}
|
||||
|
||||
type XDNSResolverUDP struct {
|
||||
Addr string `json:"addr"`
|
||||
}
|
||||
|
||||
func (c *XDNSResolverUDP) Build() (proto.Message, error) {
|
||||
return &xdns.UDPResolverProto{Addr: c.Addr}, nil
|
||||
}
|
||||
|
||||
var xdnsLoader = NewJSONConfigLoader(ConfigCreatorCache{
|
||||
"tcp": func() interface{} { return new(XDNSResolverTCP) },
|
||||
"udp": func() interface{} { return new(XDNSResolverUDP) },
|
||||
}, "type", "settings")
|
||||
|
||||
type XDNSResolver struct {
|
||||
Type string `json:"type"`
|
||||
Settings json.RawMessage `json:"settings"`
|
||||
Addrs []string `json:"addrs"`
|
||||
}
|
||||
|
||||
type XDNS struct {
|
||||
@@ -833,7 +810,7 @@ type XDNS struct {
|
||||
|
||||
func (c *XDNS) Build() (proto.Message, error) {
|
||||
var domains []*xdns.DomainProto
|
||||
var resolvers []*serial.TypedMessage
|
||||
var resolvers []*xdns.ResolverProto
|
||||
for i := range c.Domains {
|
||||
if c.Domains[i].LenLimit == 0 {
|
||||
c.Domains[i].LenLimit = 255
|
||||
@@ -841,33 +818,46 @@ func (c *XDNS) Build() (proto.Message, error) {
|
||||
if c.Domains[i].LabelLimit == 0 {
|
||||
c.Domains[i].LabelLimit = 63
|
||||
}
|
||||
types := make([]uint16, 0, len(c.Domains[i].Types))
|
||||
for j := range c.Domains[i].Types {
|
||||
types = append(types, uint16(c.Domains[i].Types[j]))
|
||||
for j := range c.Domains[i].Names {
|
||||
domain, err := xdns.NewDomain(c.Domains[i].Names[j], int(c.Domains[i].LenLimit), int(c.Domains[i].LabelLimit), []uint16{1, 5, 16, 28}, uint16(c.Domains[i].Edns0))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
errors.LogInfo(context.Background(), domain.Show())
|
||||
domains = append(domains, &xdns.DomainProto{
|
||||
Name: c.Domains[i].Names[j],
|
||||
LenLimit: c.Domains[i].LenLimit,
|
||||
LabelLimit: c.Domains[i].LabelLimit,
|
||||
Types: c.Domains[i].Types,
|
||||
Edns0: c.Domains[i].Edns0,
|
||||
})
|
||||
}
|
||||
domain, err := xdns.NewDomain(c.Domains[i].Name, int(c.Domains[i].LenLimit), int(c.Domains[i].LabelLimit), types, uint16(c.Domains[i].Edns0))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
errors.LogInfo(context.Background(), domain.Show())
|
||||
domains = append(domains, &xdns.DomainProto{
|
||||
Name: c.Domains[i].Name,
|
||||
LenLimit: c.Domains[i].LenLimit,
|
||||
LabelLimit: c.Domains[i].LabelLimit,
|
||||
Types: c.Domains[i].Types,
|
||||
Edns0: c.Domains[i].Edns0,
|
||||
})
|
||||
}
|
||||
for i := range c.Resolvers {
|
||||
config, err := xdnsLoader.LoadWithID(c.Resolvers[i].Settings, c.Resolvers[i].Type)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
for j := range c.Resolvers[i].Addrs {
|
||||
var u *url.URL
|
||||
var e error
|
||||
if !strings.Contains(c.Resolvers[i].Addrs[j], "://") {
|
||||
u, e = url.Parse("udp://" + c.Resolvers[i].Addrs[j])
|
||||
} else {
|
||||
u, e = url.Parse(c.Resolvers[i].Addrs[j])
|
||||
}
|
||||
if e != nil {
|
||||
return nil, e
|
||||
}
|
||||
switch u.Scheme {
|
||||
case "tcp", "udp":
|
||||
default:
|
||||
return nil, errors.New("invalid protocol")
|
||||
}
|
||||
var host, port string
|
||||
host = u.Hostname()
|
||||
port = u.Port()
|
||||
if port == "" {
|
||||
port = "53"
|
||||
}
|
||||
resolvers = append(resolvers, &xdns.ResolverProto{Type: u.Scheme, Addr: net.JoinHostPort(host, port)})
|
||||
}
|
||||
pm, err := config.(interface{ Build() (proto.Message, error) }).Build()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
resolvers = append(resolvers, serial.ToTypedMessage(pm))
|
||||
}
|
||||
if c.ExtraPoll < 0 || c.ExtraPoll > 3 {
|
||||
return nil, errors.New("c.ExtraPoll < 0 || c.ExtraPoll > 3")
|
||||
|
||||
@@ -213,6 +213,8 @@ If the filters cannot be added, Xray does not start. They are removed when Xray
|
||||
|
||||
`autoSystemWfpBlockLeak` (Windows only) is empty by default, as the filters break some setups: with `"dns"`, a local DNS resolver other programs use (e.g. on `127.0.0.1:53`), the DNS of another VPN on its own interface, virtual machines whose NAT resolves names on the host, or signing in to a captive portal; with `"misconfigtun"`, IPv4 or IPv6 on the local network while no route of that version leads to the TUN. Without the filters, DNS may leak as described above. To keep an IP version out of the TUN on purpose while still blocking DNS leaks, use only `["dns"]`.
|
||||
|
||||
`autoOutboundsInterface` (the default with `autoSystemRoutingTable`) keeps Xray's own connections out of the TUN by binding them to another interface, which Windows only honors while that interface has weak host send and forwarding off for the IP versions routed to the TUN. Otherwise, Windows sends them into the TUN, from that interface's address, and they stall. While the TUN runs, Xray therefore turns weak host send off on that interface, and on again when it stops or another interface takes over. Forwarding cannot be turned off this way, as Mobile Hotspot and Internet Connection Sharing need it, so a warning is logged while it is on. Having the hotspot share the TUN instead of that interface (Settings, Mobile hotspot, Share my internet connection from) moves forwarding to the TUN, where it does no harm, and sends the hotspot's devices through Xray as well.
|
||||
|
||||
You can give the adapter ip address manually, you can live Windows to give it autogenerated ip address (which take few seconds), it doesn't matter, the traffic going _through_ the interface will be forwarded into the app for proxying. \
|
||||
Minimal configuration that will work for local machine is routing passing the traffic on-link through the interface.
|
||||
You will need the interface id for that, unfortunately it is going to change with every Xray start due to implementation ambiguity between Xray and wintun driver.
|
||||
|
||||
@@ -101,7 +101,7 @@ func (t *stackGVisor) Start() error {
|
||||
// Use custom UDP packet handler, instead of strict gVisor forwarder, for FullCone NAT support
|
||||
udpForwarder := newUdpConnectionHandler(t.handler.HandleConnection, t.writeRawUDPPacket)
|
||||
ipStack.SetTransportProtocolHandler(udp.ProtocolNumber, func(id stack.TransportEndpointID, pkt *stack.PacketBuffer) bool {
|
||||
data := pkt.Clone().Data().AsRange().ToSlice()
|
||||
data := pkt.Data().AsRange().ToSlice()
|
||||
// if len(data) == 0 {
|
||||
// return false
|
||||
// }
|
||||
|
||||
@@ -46,6 +46,7 @@ type WindowsTun struct {
|
||||
luid winipcfg.LUID
|
||||
cbr winipcfg.ChangeCallback
|
||||
cbi winipcfg.ChangeCallback
|
||||
guard outboundGuard
|
||||
wfp windows.Handle
|
||||
resolver *savedResolver
|
||||
skipStop chan struct{}
|
||||
@@ -178,6 +179,11 @@ startOver:
|
||||
}
|
||||
ipif, err := t.luid.IPInterface(family)
|
||||
if err != nil {
|
||||
// With IPv6 disabled system-wide (DisabledComponents), the adapter has no
|
||||
// IPv6 interface at all. Skip the family unless the config asks for it.
|
||||
if err == windows.ERROR_NOT_FOUND && family == windows.AF_INET6 && !address6 && !route6 {
|
||||
continue
|
||||
}
|
||||
return err
|
||||
}
|
||||
ipif.RouterDiscoveryBehavior = winipcfg.RouterDiscoveryDisabled
|
||||
@@ -292,10 +298,21 @@ startOver:
|
||||
}
|
||||
|
||||
if updater != nil {
|
||||
// Xray's own connections have to stay out of the IP versions routed
|
||||
// to the TUN, which needs Windows to honor the binding to updater's
|
||||
// interface.
|
||||
if route4 {
|
||||
t.guard.families = append(t.guard.families, windows.AF_INET)
|
||||
}
|
||||
if route6 {
|
||||
t.guard.families = append(t.guard.families, windows.AF_INET6)
|
||||
}
|
||||
t.guard.check()
|
||||
// Only a registered callback goes into the fields: a nil pointer in
|
||||
// them would not compare equal to nil in Close.
|
||||
cbr, err := winipcfg.RegisterRouteChangeCallback(func(notificationType winipcfg.MibNotificationType, route *winipcfg.MibIPforwardRow2) {
|
||||
updater.Update()
|
||||
t.guard.check()
|
||||
})
|
||||
if err != nil {
|
||||
return err
|
||||
@@ -303,6 +320,7 @@ startOver:
|
||||
t.cbr = cbr
|
||||
cbi, err := winipcfg.RegisterInterfaceChangeCallback(func(notificationType winipcfg.MibNotificationType, iface *winipcfg.MibIPInterfaceRow) {
|
||||
updater.Update()
|
||||
t.guard.check()
|
||||
})
|
||||
if err != nil {
|
||||
return err
|
||||
@@ -326,6 +344,7 @@ func (t *WindowsTun) Close() error {
|
||||
if t.cbi != nil {
|
||||
t.cbi.Unregister()
|
||||
}
|
||||
t.guard.restore()
|
||||
if t.luid != 0 {
|
||||
t.luid.FlushRoutes(windows.AF_INET)
|
||||
t.luid.FlushIPAddresses(windows.AF_INET)
|
||||
|
||||
@@ -0,0 +1,120 @@
|
||||
//go:build windows
|
||||
|
||||
package tun
|
||||
|
||||
import (
|
||||
"context"
|
||||
"slices"
|
||||
"strings"
|
||||
"sync"
|
||||
|
||||
"github.com/xtls/xray-core/common/errors"
|
||||
"golang.org/x/sys/windows"
|
||||
"golang.zx2c4.com/wireguard/windows/tunnel/winipcfg"
|
||||
)
|
||||
|
||||
// outboundGuard keeps Windows to the binding of autoOutboundsInterface, which
|
||||
// keeps Xray's own connections out of the TUN. With weak host send or
|
||||
// forwarding on for an IP version on the bound interface, Windows sends them
|
||||
// where the routes lead, into the TUN, from that interface's address, and
|
||||
// drops what comes back to that address through the TUN, so they stall.
|
||||
//
|
||||
// For the IP versions routed to the TUN, weak host send is turned off on the
|
||||
// bound interface while the TUN runs, and turned on again when the TUN stops
|
||||
// or another interface takes over. Forwarding is what Mobile Hotspot and
|
||||
// Internet Connection Sharing need, so it is only reported.
|
||||
type outboundGuard struct {
|
||||
sync.Mutex
|
||||
families []winipcfg.AddressFamily
|
||||
luid winipcfg.LUID // of the interface last checked
|
||||
name string // of that interface
|
||||
turnedOff []winipcfg.AddressFamily // where weak host send was turned off on it
|
||||
forwarding bool // whether forwarding was on there
|
||||
stopped bool
|
||||
}
|
||||
|
||||
// check turns weak host send off on the bound interface, and warns when
|
||||
// forwarding comes on there, but not again while it stays on.
|
||||
func (g *outboundGuard) check() {
|
||||
g.Lock()
|
||||
defer g.Unlock()
|
||||
if g.stopped {
|
||||
return
|
||||
}
|
||||
var luid winipcfg.LUID
|
||||
var name string
|
||||
if iface := updater.Get(); iface != nil {
|
||||
luid, _ = winipcfg.LUIDFromIndex(uint32(iface.Index))
|
||||
name = iface.Name
|
||||
}
|
||||
if luid != g.luid {
|
||||
g.restoreLocked()
|
||||
g.luid, g.name = luid, name
|
||||
g.forwarding = false // to warn about the new interface as well
|
||||
}
|
||||
if luid == 0 {
|
||||
return
|
||||
}
|
||||
var forwarding []string
|
||||
for _, family := range g.families {
|
||||
row, err := luid.IPInterface(family)
|
||||
if err != nil {
|
||||
continue // the interface lacks that IP version
|
||||
}
|
||||
if row.ForwardingEnabled {
|
||||
forwarding = append(forwarding, familyName(family))
|
||||
}
|
||||
if !row.WeakHostSend {
|
||||
continue
|
||||
}
|
||||
if err := setWeakHostSend(row, false); err != nil {
|
||||
errors.LogWarningInner(context.Background(), err, "[tun] unable to turn weak host send off for ", familyName(family), " on ", name)
|
||||
continue
|
||||
}
|
||||
if !slices.Contains(g.turnedOff, family) {
|
||||
g.turnedOff = append(g.turnedOff, family)
|
||||
errors.LogInfo(context.Background(), "[tun] weak host send turned off for ", familyName(family), " on ", name, " while the TUN runs, as Windows would ignore autoOutboundsInterface")
|
||||
}
|
||||
}
|
||||
wasOn := g.forwarding
|
||||
g.forwarding = len(forwarding) > 0
|
||||
if g.forwarding && !wasOn {
|
||||
errors.LogWarning(context.Background(), "[tun] forwarding is on for ", strings.Join(forwarding, " and "), " on ", name, " (Mobile Hotspot and Internet Connection Sharing turn it on), so Windows ignores autoOutboundsInterface there, and Xray's own connections go into the TUN and stall: turn the hotspot off, or have it share the TUN instead of ", name)
|
||||
}
|
||||
}
|
||||
|
||||
// restore turns weak host send on again where check turned it off, for good.
|
||||
func (g *outboundGuard) restore() {
|
||||
g.Lock()
|
||||
defer g.Unlock()
|
||||
g.restoreLocked()
|
||||
g.stopped = true
|
||||
}
|
||||
|
||||
func (g *outboundGuard) restoreLocked() {
|
||||
for _, family := range g.turnedOff {
|
||||
row, err := g.luid.IPInterface(family)
|
||||
if err == nil {
|
||||
err = setWeakHostSend(row, true)
|
||||
}
|
||||
if err != nil {
|
||||
errors.LogWarningInner(context.Background(), err, "[tun] unable to turn weak host send on again for ", familyName(family), " on ", g.name)
|
||||
}
|
||||
}
|
||||
g.turnedOff = nil
|
||||
}
|
||||
|
||||
func setWeakHostSend(row *winipcfg.MibIPInterfaceRow, on bool) error {
|
||||
row.WeakHostSend = on
|
||||
if row.Family == windows.AF_INET {
|
||||
row.SitePrefixLength = 0 // as SetIpInterfaceEntry requires for IPv4
|
||||
}
|
||||
return row.Set()
|
||||
}
|
||||
|
||||
func familyName(family winipcfg.AddressFamily) string {
|
||||
if family == windows.AF_INET {
|
||||
return "IPv4"
|
||||
}
|
||||
return "IPv6"
|
||||
}
|
||||
@@ -0,0 +1,91 @@
|
||||
package wireguard
|
||||
|
||||
import (
|
||||
"context"
|
||||
|
||||
"github.com/xtls/xray-core/common/errors"
|
||||
tunicmp "github.com/xtls/xray-core/proxy/tun/icmp"
|
||||
"gvisor.dev/gvisor/pkg/buffer"
|
||||
"gvisor.dev/gvisor/pkg/tcpip"
|
||||
"gvisor.dev/gvisor/pkg/tcpip/header"
|
||||
"gvisor.dev/gvisor/pkg/tcpip/stack"
|
||||
"gvisor.dev/gvisor/pkg/tcpip/transport/icmp"
|
||||
)
|
||||
|
||||
// CreateICMPEchoResponder answers ICMP echo requests from peers locally, the way
|
||||
// the TUN inbound does: ICMP is not proxied, but ping and connectivity checks
|
||||
// through the tunnel get a reply instead of timing out.
|
||||
//
|
||||
// In promiscuous mode gVisor skips its own IPv4 echo reply for addresses that are
|
||||
// not assigned to the NIC and leaves it to a custom handler; IPv6 is registered
|
||||
// too so both families behave the same.
|
||||
func CreateICMPEchoResponder(gstack *stack.Stack) {
|
||||
gstack.SetTransportProtocolHandler(icmp.ProtocolNumber4, func(id stack.TransportEndpointID, pkt *stack.PacketBuffer) bool {
|
||||
return handleICMPEcho(gstack, header.IPv4ProtocolNumber, id, pkt)
|
||||
})
|
||||
gstack.SetTransportProtocolHandler(icmp.ProtocolNumber6, func(id stack.TransportEndpointID, pkt *stack.PacketBuffer) bool {
|
||||
return handleICMPEcho(gstack, header.IPv6ProtocolNumber, id, pkt)
|
||||
})
|
||||
}
|
||||
|
||||
func handleICMPEcho(gstack *stack.Stack, netProto tcpip.NetworkProtocolNumber, id stack.TransportEndpointID, pkt *stack.PacketBuffer) bool {
|
||||
srcIP := id.RemoteAddress
|
||||
dstIP := id.LocalAddress
|
||||
if srcIP.Len() == 0 || dstIP.Len() == 0 {
|
||||
return true
|
||||
}
|
||||
|
||||
headerBytes := pkt.TransportHeader().Slice()
|
||||
payloadBytes := pkt.Data().AsRange().ToSlice()
|
||||
message := make([]byte, len(headerBytes)+len(payloadBytes))
|
||||
copy(message, headerBytes)
|
||||
copy(message[len(headerBytes):], payloadBytes)
|
||||
|
||||
if _, _, ok := tunicmp.ParseEchoRequest(netProto, message); !ok {
|
||||
return true
|
||||
}
|
||||
|
||||
reply, err := tunicmp.BuildLocalEchoReply(netProto, message, dstIP, srcIP)
|
||||
if err != nil {
|
||||
errors.LogInfoInner(context.Background(), err, "failed to build local icmp echo reply")
|
||||
return true
|
||||
}
|
||||
if err := writeRawICMPPacket(gstack, netProto, reply, dstIP, srcIP); err != nil {
|
||||
errors.LogInfoInner(context.Background(), err, "failed to write local icmp echo reply")
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
func writeRawICMPPacket(gstack *stack.Stack, netProto tcpip.NetworkProtocolNumber, message []byte, srcIP, dstIP tcpip.Address) error {
|
||||
pkt := stack.NewPacketBuffer(stack.PacketBufferOptions{
|
||||
ReserveHeaderBytes: header.IPv6MinimumSize,
|
||||
Payload: buffer.MakeWithData(message),
|
||||
})
|
||||
defer pkt.DecRef()
|
||||
|
||||
if netProto == header.IPv4ProtocolNumber {
|
||||
ipHdr := header.IPv4(pkt.NetworkHeader().Push(header.IPv4MinimumSize))
|
||||
ipHdr.Encode(&header.IPv4Fields{
|
||||
TotalLength: uint16(header.IPv4MinimumSize + len(message)),
|
||||
TTL: 64,
|
||||
Protocol: uint8(header.ICMPv4ProtocolNumber),
|
||||
SrcAddr: srcIP,
|
||||
DstAddr: dstIP,
|
||||
})
|
||||
ipHdr.SetChecksum(^ipHdr.CalculateChecksum())
|
||||
} else {
|
||||
ipHdr := header.IPv6(pkt.NetworkHeader().Push(header.IPv6MinimumSize))
|
||||
ipHdr.Encode(&header.IPv6Fields{
|
||||
PayloadLength: uint16(len(message)),
|
||||
TransportProtocol: header.ICMPv6ProtocolNumber,
|
||||
HopLimit: 64,
|
||||
SrcAddr: srcIP,
|
||||
DstAddr: dstIP,
|
||||
})
|
||||
}
|
||||
|
||||
if err := gstack.WriteRawPacket(1, netProto, buffer.MakeWithView(pkt.ToView())); err != nil {
|
||||
return errors.New("failed to write raw icmp packet back to stack ", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,176 @@
|
||||
package wireguard
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"net/netip"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/xtls/xray-core/common/net"
|
||||
"gvisor.dev/gvisor/pkg/tcpip"
|
||||
"gvisor.dev/gvisor/pkg/tcpip/checksum"
|
||||
"gvisor.dev/gvisor/pkg/tcpip/header"
|
||||
)
|
||||
|
||||
func newICMPTestStack(t *testing.T) *netTun {
|
||||
t.Helper()
|
||||
dev, _, gstack, err := CreateNetTUN([]netip.Addr{
|
||||
netip.MustParseAddr("10.66.0.1"),
|
||||
netip.MustParseAddr("fd00::1"),
|
||||
}, nil, 1420, false)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
t.Cleanup(func() { dev.Close() })
|
||||
CreateForwarder(gstack, func(conn net.Conn, dest net.Destination) { conn.Close() })
|
||||
CreateICMPEchoResponder(gstack)
|
||||
return dev.(*netTun)
|
||||
}
|
||||
|
||||
// startReader must run before the request is written: the stack may answer
|
||||
// synchronously inside Write, and netTun hands packets over an unbuffered channel.
|
||||
func startReader(dev *netTun) <-chan []byte {
|
||||
got := make(chan []byte, 1)
|
||||
go func() {
|
||||
buf := make([]byte, 2048)
|
||||
sizes := make([]int, 1)
|
||||
if _, err := dev.Read([][]byte{buf}, sizes, 0); err == nil {
|
||||
got <- buf[:sizes[0]]
|
||||
}
|
||||
}()
|
||||
return got
|
||||
}
|
||||
|
||||
func awaitPacket(t *testing.T, got <-chan []byte) []byte {
|
||||
t.Helper()
|
||||
select {
|
||||
case p := <-got:
|
||||
return p
|
||||
case <-time.After(2 * time.Second):
|
||||
t.Fatal("no echo reply from the stack")
|
||||
return nil
|
||||
}
|
||||
}
|
||||
|
||||
func TestICMPv4EchoReply(t *testing.T) {
|
||||
dev := newICMPTestStack(t)
|
||||
src := tcpip.AddrFrom4([4]byte{10, 66, 0, 2})
|
||||
dst := tcpip.AddrFrom4([4]byte{1, 1, 1, 1})
|
||||
payload := []byte("xray wireguard ping")
|
||||
|
||||
icmpMsg := make([]byte, header.ICMPv4MinimumSize+len(payload))
|
||||
req := header.ICMPv4(icmpMsg)
|
||||
req.SetType(header.ICMPv4Echo)
|
||||
req.SetIdent(0x1234)
|
||||
req.SetSequence(7)
|
||||
copy(req.Payload(), payload)
|
||||
req.SetChecksum(header.ICMPv4Checksum(req[:header.ICMPv4MinimumSize], checksum.Checksum(payload, 0)))
|
||||
|
||||
pkt := make([]byte, header.IPv4MinimumSize+len(icmpMsg))
|
||||
ip := header.IPv4(pkt)
|
||||
ip.Encode(&header.IPv4Fields{
|
||||
TotalLength: uint16(len(pkt)),
|
||||
TTL: 64,
|
||||
Protocol: uint8(header.ICMPv4ProtocolNumber),
|
||||
SrcAddr: src,
|
||||
DstAddr: dst,
|
||||
})
|
||||
ip.SetChecksum(^ip.CalculateChecksum())
|
||||
copy(pkt[header.IPv4MinimumSize:], icmpMsg)
|
||||
|
||||
got := startReader(dev)
|
||||
if _, err := dev.Write([][]byte{pkt}, 0); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
reply := header.IPv4(awaitPacket(t, got))
|
||||
if !reply.IsValid(len(reply)) {
|
||||
t.Fatal("invalid ipv4 reply")
|
||||
}
|
||||
if reply.SourceAddress() != dst || reply.DestinationAddress() != src {
|
||||
t.Fatalf("reply addresses %v -> %v, want %v -> %v", reply.SourceAddress(), reply.DestinationAddress(), dst, src)
|
||||
}
|
||||
if reply.TransportProtocol() != header.ICMPv4ProtocolNumber {
|
||||
t.Fatalf("reply protocol %v, want icmpv4", reply.TransportProtocol())
|
||||
}
|
||||
echo := header.ICMPv4(reply.Payload())
|
||||
if echo.Type() != header.ICMPv4EchoReply {
|
||||
t.Fatalf("reply type %v, want echo reply", echo.Type())
|
||||
}
|
||||
if echo.Ident() != 0x1234 || echo.Sequence() != 7 {
|
||||
t.Fatalf("reply ident/seq %#x/%d, want 0x1234/7", echo.Ident(), echo.Sequence())
|
||||
}
|
||||
if !bytes.Equal(echo.Payload(), payload) {
|
||||
t.Fatalf("reply payload %q, want %q", echo.Payload(), payload)
|
||||
}
|
||||
if checksum.Checksum(echo, 0) != 0xffff {
|
||||
t.Fatal("bad icmpv4 checksum")
|
||||
}
|
||||
}
|
||||
|
||||
func TestICMPv6EchoReply(t *testing.T) {
|
||||
dev := newICMPTestStack(t)
|
||||
src := tcpip.AddrFrom16([16]byte{0xfd, 15: 2})
|
||||
dst := tcpip.AddrFrom16([16]byte{0x26, 0x06, 0x47, 0x00, 0x47, 0x00, 15: 0x11})
|
||||
payload := []byte("xray wireguard ping6")
|
||||
|
||||
icmpMsg := make([]byte, header.ICMPv6MinimumSize+len(payload))
|
||||
req := header.ICMPv6(icmpMsg)
|
||||
req.SetType(header.ICMPv6EchoRequest)
|
||||
req.SetIdent(0x4321)
|
||||
req.SetSequence(9)
|
||||
copy(req.Payload(), payload)
|
||||
req.SetChecksum(header.ICMPv6Checksum(header.ICMPv6ChecksumParams{
|
||||
Header: req[:header.ICMPv6MinimumSize],
|
||||
Src: src,
|
||||
Dst: dst,
|
||||
PayloadCsum: checksum.Checksum(payload, 0),
|
||||
PayloadLen: len(payload),
|
||||
}))
|
||||
|
||||
pkt := make([]byte, header.IPv6MinimumSize+len(icmpMsg))
|
||||
ip := header.IPv6(pkt)
|
||||
ip.Encode(&header.IPv6Fields{
|
||||
PayloadLength: uint16(len(icmpMsg)),
|
||||
TransportProtocol: header.ICMPv6ProtocolNumber,
|
||||
HopLimit: 64,
|
||||
SrcAddr: src,
|
||||
DstAddr: dst,
|
||||
})
|
||||
copy(pkt[header.IPv6MinimumSize:], icmpMsg)
|
||||
|
||||
got := startReader(dev)
|
||||
if _, err := dev.Write([][]byte{pkt}, 0); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
reply := header.IPv6(awaitPacket(t, got))
|
||||
if !reply.IsValid(len(reply)) {
|
||||
t.Fatal("invalid ipv6 reply")
|
||||
}
|
||||
if reply.SourceAddress() != dst || reply.DestinationAddress() != src {
|
||||
t.Fatalf("reply addresses %v -> %v, want %v -> %v", reply.SourceAddress(), reply.DestinationAddress(), dst, src)
|
||||
}
|
||||
echo := header.ICMPv6(reply.Payload())
|
||||
if echo.Type() != header.ICMPv6EchoReply {
|
||||
t.Fatalf("reply type %v, want echo reply", echo.Type())
|
||||
}
|
||||
if echo.Ident() != 0x4321 || echo.Sequence() != 9 {
|
||||
t.Fatalf("reply ident/seq %#x/%d, want 0x4321/9", echo.Ident(), echo.Sequence())
|
||||
}
|
||||
if !bytes.Equal(echo.Payload(), payload) {
|
||||
t.Fatalf("reply payload %q, want %q", echo.Payload(), payload)
|
||||
}
|
||||
zeroed := header.ICMPv6(append([]byte(nil), echo[:header.ICMPv6MinimumSize]...))
|
||||
zeroed.SetChecksum(0)
|
||||
want := header.ICMPv6Checksum(header.ICMPv6ChecksumParams{
|
||||
Header: zeroed,
|
||||
Src: dst,
|
||||
Dst: src,
|
||||
PayloadCsum: checksum.Checksum(echo.Payload(), 0),
|
||||
PayloadLen: len(echo.Payload()),
|
||||
})
|
||||
if echo.Checksum() != want {
|
||||
t.Fatalf("icmpv6 checksum %#x, want %#x", echo.Checksum(), want)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,80 @@
|
||||
package wireguard
|
||||
|
||||
import (
|
||||
"runtime"
|
||||
"testing"
|
||||
)
|
||||
|
||||
const benchBatch = 64
|
||||
|
||||
// Raw cost of queueing and draining a small burst, as one flow's reader does.
|
||||
func BenchmarkQueueBurstChan(b *testing.B) {
|
||||
ch := make(chan *packet, udpQueueLimit)
|
||||
p := &packet{}
|
||||
b.ReportAllocs()
|
||||
for i := 0; i < b.N; i++ {
|
||||
for j := 0; j < benchBatch; j++ {
|
||||
ch <- p
|
||||
}
|
||||
for j := 0; j < benchBatch; j++ {
|
||||
<-ch
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func BenchmarkQueueBurstPacketQueue(b *testing.B) {
|
||||
q := newPacketQueue(udpQueueLimit)
|
||||
p := &packet{}
|
||||
b.ReportAllocs()
|
||||
for i := 0; i < b.N; i++ {
|
||||
for j := 0; j < benchBatch; j++ {
|
||||
q.push(p)
|
||||
}
|
||||
for j := 0; j < benchBatch; j++ {
|
||||
q.pop()
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Producer and consumer on different goroutines; the producer yields when the
|
||||
// queue is full instead of spinning, like a blocking channel send would.
|
||||
func BenchmarkQueueStreamChan(b *testing.B) {
|
||||
ch := make(chan *packet, udpQueueLimit)
|
||||
p := &packet{}
|
||||
done := make(chan struct{})
|
||||
go func() {
|
||||
for range ch {
|
||||
}
|
||||
close(done)
|
||||
}()
|
||||
b.ReportAllocs()
|
||||
b.ResetTimer()
|
||||
for i := 0; i < b.N; i++ {
|
||||
ch <- p
|
||||
}
|
||||
close(ch)
|
||||
<-done
|
||||
}
|
||||
|
||||
func BenchmarkQueueStreamPacketQueue(b *testing.B) {
|
||||
q := newPacketQueue(udpQueueLimit)
|
||||
p := &packet{}
|
||||
done := make(chan struct{})
|
||||
go func() {
|
||||
for {
|
||||
if _, ok := q.pop(); !ok {
|
||||
break
|
||||
}
|
||||
}
|
||||
close(done)
|
||||
}()
|
||||
b.ReportAllocs()
|
||||
b.ResetTimer()
|
||||
for i := 0; i < b.N; i++ {
|
||||
for !q.push(p) {
|
||||
runtime.Gosched()
|
||||
}
|
||||
}
|
||||
q.close()
|
||||
<-done
|
||||
}
|
||||
@@ -134,6 +134,7 @@ func NewServer(ctx context.Context, conf *DeviceConfig) (*Server, error) {
|
||||
}
|
||||
// Install the stack's protocol handlers before the device can deliver packets to it (Start -> dev.Up).
|
||||
CreateForwarder(stack, s.HandleConnection)
|
||||
CreateICMPEchoResponder(stack)
|
||||
return s, nil
|
||||
}
|
||||
|
||||
|
||||
+73
-18
@@ -85,7 +85,7 @@ func CreateForwarder(gstack *stack.Stack, handler func(conn net.Conn, dest net.D
|
||||
}
|
||||
|
||||
gstack.SetTransportProtocolHandler(udp.ProtocolNumber, func(id stack.TransportEndpointID, pkt *stack.PacketBuffer) bool {
|
||||
data := pkt.Clone().Data().AsRange().ToSlice()
|
||||
data := pkt.Data().AsRange().ToSlice()
|
||||
// if len(data) == 0 {
|
||||
// return false
|
||||
// }
|
||||
@@ -112,12 +112,7 @@ func (m *udpManager) feed(src net.Destination, dst net.Destination, data []byte)
|
||||
m.mutex.RLock()
|
||||
uc, ok := m.m[src.NetAddr()]
|
||||
if ok {
|
||||
select {
|
||||
case uc.queue <- &packet{
|
||||
p: data,
|
||||
dest: &dst,
|
||||
}:
|
||||
default:
|
||||
if !uc.queue.push(&packet{p: data, dest: &dst}) {
|
||||
errors.LogDebug(context.Background(), "drop udp with size ", len(data), " to ", dst.NetAddr(), " original ", uc.dst.NetAddr(), " > queue full")
|
||||
}
|
||||
m.mutex.RUnlock()
|
||||
@@ -131,7 +126,7 @@ func (m *udpManager) feed(src net.Destination, dst net.Destination, data []byte)
|
||||
uc, ok = m.m[src.NetAddr()]
|
||||
if !ok {
|
||||
uc = &udpConn{
|
||||
queue: make(chan *packet, 1024),
|
||||
queue: newPacketQueue(udpQueueLimit),
|
||||
src: src,
|
||||
dst: dst,
|
||||
}
|
||||
@@ -145,12 +140,7 @@ func (m *udpManager) feed(src net.Destination, dst net.Destination, data []byte)
|
||||
go m.handler(uc, dst)
|
||||
}
|
||||
|
||||
select {
|
||||
case uc.queue <- &packet{
|
||||
p: data,
|
||||
dest: &dst,
|
||||
}:
|
||||
default:
|
||||
if !uc.queue.push(&packet{p: data, dest: &dst}) {
|
||||
errors.LogDebug(context.Background(), "drop udp with size ", len(data), " to ", dst.NetAddr(), " original ", uc.dst.NetAddr(), " > queue full 2")
|
||||
}
|
||||
}
|
||||
@@ -158,7 +148,7 @@ func (m *udpManager) feed(src net.Destination, dst net.Destination, data []byte)
|
||||
func (m *udpManager) close(uc *udpConn) {
|
||||
if !uc.closed {
|
||||
uc.closed = true
|
||||
close(uc.queue)
|
||||
uc.queue.close()
|
||||
delete(m.m, uc.src.NetAddr())
|
||||
}
|
||||
}
|
||||
@@ -232,7 +222,7 @@ type packet struct {
|
||||
}
|
||||
|
||||
type udpConn struct {
|
||||
queue chan *packet
|
||||
queue *packetQueue
|
||||
src net.Destination
|
||||
dst net.Destination
|
||||
writeFunc func(payload []byte, src net.Destination, dst net.Destination) error
|
||||
@@ -242,7 +232,7 @@ type udpConn struct {
|
||||
|
||||
func (c *udpConn) ReadMultiBuffer() (buf.MultiBuffer, error) {
|
||||
for {
|
||||
q, ok := <-c.queue
|
||||
q, ok := c.queue.pop()
|
||||
if !ok {
|
||||
return nil, io.EOF
|
||||
}
|
||||
@@ -261,7 +251,7 @@ func (c *udpConn) ReadMultiBuffer() (buf.MultiBuffer, error) {
|
||||
}
|
||||
|
||||
func (c *udpConn) Read(p []byte) (int, error) {
|
||||
q, ok := <-c.queue
|
||||
q, ok := c.queue.pop()
|
||||
if !ok {
|
||||
return 0, io.EOF
|
||||
}
|
||||
@@ -324,3 +314,68 @@ func (c *udpConn) SetReadDeadline(t time.Time) error {
|
||||
func (c *udpConn) SetWriteDeadline(t time.Time) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
// udpQueueLimit bounds the packets waiting for one UDP flow; more are dropped.
|
||||
const udpQueueLimit = 1024
|
||||
|
||||
// packetQueue holds the packets waiting for one UDP flow. Unlike a buffered
|
||||
// channel of the same bound it only allocates for packets actually queued, so
|
||||
// the many idle flows kept until the idle timeout cost next to nothing.
|
||||
type packetQueue struct {
|
||||
mu sync.Mutex
|
||||
items []*packet
|
||||
limit int
|
||||
notify chan struct{}
|
||||
closed bool
|
||||
}
|
||||
|
||||
func newPacketQueue(limit int) *packetQueue {
|
||||
return &packetQueue{limit: limit, notify: make(chan struct{}, 1)}
|
||||
}
|
||||
|
||||
// push queues p and reports whether it was accepted.
|
||||
func (q *packetQueue) push(p *packet) bool {
|
||||
q.mu.Lock()
|
||||
defer q.mu.Unlock()
|
||||
if q.closed || len(q.items) >= q.limit {
|
||||
return false
|
||||
}
|
||||
q.items = append(q.items, p)
|
||||
select {
|
||||
case q.notify <- struct{}{}:
|
||||
default:
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
// pop blocks until a packet is queued or the queue is closed and drained.
|
||||
func (q *packetQueue) pop() (*packet, bool) {
|
||||
for {
|
||||
q.mu.Lock()
|
||||
if len(q.items) > 0 {
|
||||
p := q.items[0]
|
||||
q.items[0] = nil
|
||||
q.items = q.items[1:]
|
||||
if len(q.items) == 0 {
|
||||
q.items = nil
|
||||
}
|
||||
q.mu.Unlock()
|
||||
return p, true
|
||||
}
|
||||
if q.closed {
|
||||
q.mu.Unlock()
|
||||
return nil, false
|
||||
}
|
||||
q.mu.Unlock()
|
||||
<-q.notify
|
||||
}
|
||||
}
|
||||
|
||||
func (q *packetQueue) close() {
|
||||
q.mu.Lock()
|
||||
defer q.mu.Unlock()
|
||||
if !q.closed {
|
||||
q.closed = true
|
||||
close(q.notify)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -9,7 +9,7 @@ import (
|
||||
"net"
|
||||
"net/netip"
|
||||
"os"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"syscall"
|
||||
|
||||
"golang.org/x/sys/unix"
|
||||
@@ -20,21 +20,10 @@ import (
|
||||
"golang.zx2c4.com/wireguard/tun"
|
||||
)
|
||||
|
||||
var (
|
||||
tableIndex int = 10230
|
||||
mu sync.Mutex
|
||||
)
|
||||
var tableIndex atomic.Uint32
|
||||
|
||||
func allocateIPv6TableIndex() int {
|
||||
mu.Lock()
|
||||
defer mu.Unlock()
|
||||
|
||||
if tableIndex > 10230 {
|
||||
errors.LogInfo(context.Background(), "allocate new ipv6 table index: ", tableIndex)
|
||||
}
|
||||
currentIndex := tableIndex
|
||||
tableIndex++
|
||||
return currentIndex
|
||||
func init() {
|
||||
tableIndex.Store(10230)
|
||||
}
|
||||
|
||||
type kernelTun struct {
|
||||
@@ -111,17 +100,23 @@ func createKernelTun(localAddresses, dnsServers []netip.Addr, mtu int) (tdev tun
|
||||
}
|
||||
}
|
||||
|
||||
ipv6TableIndex := allocateIPv6TableIndex()
|
||||
var ipv6TableIndex int
|
||||
if v6 != nil {
|
||||
r := &netlink.Route{Table: ipv6TableIndex}
|
||||
r := &netlink.Route{}
|
||||
for {
|
||||
ipv6TableIndex = int(tableIndex.Add(1)) - 1
|
||||
r.Table = ipv6TableIndex
|
||||
routeList, fErr := netlink.RouteListFiltered(netlink.FAMILY_V6, r, netlink.RT_FILTER_TABLE)
|
||||
if len(routeList) == 0 || fErr != nil {
|
||||
if fErr != nil {
|
||||
return nil, nil, errors.New("failed to pre check routes for table: ", ipv6TableIndex).Base(fErr)
|
||||
}
|
||||
if len(routeList) == 0 {
|
||||
errors.LogInfo(context.Background(), "allocate new ipv6 table index: ", ipv6TableIndex)
|
||||
break
|
||||
}
|
||||
ipv6TableIndex--
|
||||
if ipv6TableIndex < 0 {
|
||||
return nil, nil, fmt.Errorf("failed to find available ipv6 table index")
|
||||
// to prevent infinite loop
|
||||
if ipv6TableIndex > 65535 {
|
||||
return nil, nil, errors.New("failed to find available ipv6 table index")
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,91 @@
|
||||
package wireguard
|
||||
|
||||
import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/xtls/xray-core/common/net"
|
||||
)
|
||||
|
||||
// BenchmarkUDPManagerNewSession measures what one new UDP flow costs the
|
||||
// inbound while it stays open: QUIC and DNS open many short flows, and each
|
||||
// one lives until the connection idle timeout.
|
||||
func BenchmarkUDPManagerNewSession(b *testing.B) {
|
||||
m := &udpManager{
|
||||
handler: func(conn net.Conn, dest net.Destination) {},
|
||||
m: make(map[string]*udpConn),
|
||||
}
|
||||
dst := net.UDPDestination(net.ParseAddress("1.1.1.1"), 443)
|
||||
payload := make([]byte, 1200)
|
||||
b.ReportAllocs()
|
||||
b.ResetTimer()
|
||||
for i := 0; i < b.N; i++ {
|
||||
src := net.UDPDestination(net.IPAddress([]byte{10, byte(i >> 16), byte(i >> 8), byte(i)}), net.Port(1024+i%60000))
|
||||
m.feed(src, dst, payload)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPacketQueueOrderAndClose(t *testing.T) {
|
||||
q := newPacketQueue(udpQueueLimit)
|
||||
for i := 0; i < 3; i++ {
|
||||
if !q.push(&packet{p: []byte{byte(i)}}) {
|
||||
t.Fatalf("push %d rejected", i)
|
||||
}
|
||||
}
|
||||
for i := 0; i < 3; i++ {
|
||||
p, ok := q.pop()
|
||||
if !ok || p.p[0] != byte(i) {
|
||||
t.Fatalf("pop %d: got %v, %v", i, p, ok)
|
||||
}
|
||||
}
|
||||
q.close()
|
||||
if _, ok := q.pop(); ok {
|
||||
t.Fatal("pop after close returned a packet")
|
||||
}
|
||||
if q.push(&packet{}) {
|
||||
t.Fatal("push after close accepted")
|
||||
}
|
||||
}
|
||||
|
||||
func TestPacketQueueLimit(t *testing.T) {
|
||||
q := newPacketQueue(udpQueueLimit)
|
||||
for i := 0; i < udpQueueLimit; i++ {
|
||||
if !q.push(&packet{}) {
|
||||
t.Fatalf("push %d rejected below the limit", i)
|
||||
}
|
||||
}
|
||||
if q.push(&packet{}) {
|
||||
t.Fatal("push above the limit accepted")
|
||||
}
|
||||
}
|
||||
|
||||
func TestPacketQueueCloseUnblocksReader(t *testing.T) {
|
||||
q := newPacketQueue(udpQueueLimit)
|
||||
done := make(chan bool)
|
||||
go func() {
|
||||
_, ok := q.pop()
|
||||
done <- ok
|
||||
}()
|
||||
q.close()
|
||||
select {
|
||||
case ok := <-done:
|
||||
if ok {
|
||||
t.Fatal("blocked pop returned a packet after close")
|
||||
}
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("close did not wake the reader")
|
||||
}
|
||||
}
|
||||
|
||||
func TestPacketQueueDropsDrainedStorage(t *testing.T) {
|
||||
q := newPacketQueue(udpQueueLimit)
|
||||
for i := 0; i < 100; i++ {
|
||||
q.push(&packet{})
|
||||
}
|
||||
for i := 0; i < 100; i++ {
|
||||
q.pop()
|
||||
}
|
||||
if q.items != nil {
|
||||
t.Fatalf("drained queue still holds %d slots", cap(q.items))
|
||||
}
|
||||
}
|
||||
@@ -115,6 +115,9 @@ func TestDokodemoTCP(t *testing.T) {
|
||||
defer CloseServer(server)
|
||||
break
|
||||
}
|
||||
if server != nil {
|
||||
CloseServer(server)
|
||||
}
|
||||
retry++
|
||||
if retry > 5 {
|
||||
t.Fatal("All attempts failed to start client")
|
||||
@@ -209,6 +212,9 @@ func TestDokodemoUDP(t *testing.T) {
|
||||
defer CloseServer(server)
|
||||
break
|
||||
}
|
||||
if server != nil {
|
||||
CloseServer(server)
|
||||
}
|
||||
retry++
|
||||
if retry > 5 {
|
||||
t.Fatal("All attempts failed to start client")
|
||||
|
||||
@@ -227,6 +227,9 @@ func TestSocksBridageUDP(t *testing.T) {
|
||||
defer CloseServer(server)
|
||||
break
|
||||
}
|
||||
if server != nil {
|
||||
CloseServer(server)
|
||||
}
|
||||
retry++
|
||||
if retry > 5 {
|
||||
t.Fatal("All attempts failed to start server")
|
||||
@@ -342,6 +345,9 @@ func TestSocksBridageUDPWithRouting(t *testing.T) {
|
||||
defer CloseServer(server)
|
||||
break
|
||||
}
|
||||
if server != nil {
|
||||
CloseServer(server)
|
||||
}
|
||||
retry++
|
||||
if retry > 5 {
|
||||
t.Fatal("All attempts failed to start server")
|
||||
|
||||
@@ -70,6 +70,9 @@ func NewClient(c *Config, dialer *finalmask.Dialer) (net.PacketConn, error) {
|
||||
for j := range c.Domains[i].Types {
|
||||
types = append(types, uint16(c.Domains[i].Types[j]))
|
||||
}
|
||||
if len(types) == 0 {
|
||||
types = []uint16{16}
|
||||
}
|
||||
domain, err := NewDomain(c.Domains[i].Name, int(c.Domains[i].LenLimit), int(c.Domains[i].LabelLimit), types, uint16(c.Domains[i].Edns0))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
@@ -80,6 +83,9 @@ func NewClient(c *Config, dialer *finalmask.Dialer) (net.PacketConn, error) {
|
||||
for i := range c.Resolvers {
|
||||
resolver, err := NewResolver(c.Resolvers[i], dialer)
|
||||
if err != nil {
|
||||
for _, resolver := range resolvers {
|
||||
resolver.Close()
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
resolvers = append(resolvers, resolver)
|
||||
|
||||
@@ -7,7 +7,7 @@
|
||||
package xdns
|
||||
|
||||
import (
|
||||
serial "github.com/xtls/xray-core/common/serial"
|
||||
_ "github.com/xtls/xray-core/common/serial"
|
||||
protoreflect "google.golang.org/protobuf/reflect/protoreflect"
|
||||
protoimpl "google.golang.org/protobuf/runtime/protoimpl"
|
||||
reflect "reflect"
|
||||
@@ -98,10 +98,62 @@ func (x *DomainProto) GetEdns0() int32 {
|
||||
return 0
|
||||
}
|
||||
|
||||
type ResolverProto struct {
|
||||
state protoimpl.MessageState `protogen:"open.v1"`
|
||||
Type string `protobuf:"bytes,1,opt,name=type,proto3" json:"type,omitempty"`
|
||||
Addr string `protobuf:"bytes,2,opt,name=addr,proto3" json:"addr,omitempty"`
|
||||
unknownFields protoimpl.UnknownFields
|
||||
sizeCache protoimpl.SizeCache
|
||||
}
|
||||
|
||||
func (x *ResolverProto) Reset() {
|
||||
*x = ResolverProto{}
|
||||
mi := &file_transport_internet_finalmask_xdns_config_proto_msgTypes[1]
|
||||
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
|
||||
ms.StoreMessageInfo(mi)
|
||||
}
|
||||
|
||||
func (x *ResolverProto) String() string {
|
||||
return protoimpl.X.MessageStringOf(x)
|
||||
}
|
||||
|
||||
func (*ResolverProto) ProtoMessage() {}
|
||||
|
||||
func (x *ResolverProto) ProtoReflect() protoreflect.Message {
|
||||
mi := &file_transport_internet_finalmask_xdns_config_proto_msgTypes[1]
|
||||
if x != nil {
|
||||
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
|
||||
if ms.LoadMessageInfo() == nil {
|
||||
ms.StoreMessageInfo(mi)
|
||||
}
|
||||
return ms
|
||||
}
|
||||
return mi.MessageOf(x)
|
||||
}
|
||||
|
||||
// Deprecated: Use ResolverProto.ProtoReflect.Descriptor instead.
|
||||
func (*ResolverProto) Descriptor() ([]byte, []int) {
|
||||
return file_transport_internet_finalmask_xdns_config_proto_rawDescGZIP(), []int{1}
|
||||
}
|
||||
|
||||
func (x *ResolverProto) GetType() string {
|
||||
if x != nil {
|
||||
return x.Type
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func (x *ResolverProto) GetAddr() string {
|
||||
if x != nil {
|
||||
return x.Addr
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
type Config struct {
|
||||
state protoimpl.MessageState `protogen:"open.v1"`
|
||||
Domains []*DomainProto `protobuf:"bytes,1,rep,name=domains,proto3" json:"domains,omitempty"`
|
||||
Resolvers []*serial.TypedMessage `protobuf:"bytes,2,rep,name=resolvers,proto3" json:"resolvers,omitempty"`
|
||||
Resolvers []*ResolverProto `protobuf:"bytes,2,rep,name=resolvers,proto3" json:"resolvers,omitempty"`
|
||||
ExtraPoll int32 `protobuf:"varint,3,opt,name=extra_poll,json=extraPoll,proto3" json:"extra_poll,omitempty"`
|
||||
unknownFields protoimpl.UnknownFields
|
||||
sizeCache protoimpl.SizeCache
|
||||
@@ -109,7 +161,7 @@ type Config struct {
|
||||
|
||||
func (x *Config) Reset() {
|
||||
*x = Config{}
|
||||
mi := &file_transport_internet_finalmask_xdns_config_proto_msgTypes[1]
|
||||
mi := &file_transport_internet_finalmask_xdns_config_proto_msgTypes[2]
|
||||
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
|
||||
ms.StoreMessageInfo(mi)
|
||||
}
|
||||
@@ -121,7 +173,7 @@ func (x *Config) String() string {
|
||||
func (*Config) ProtoMessage() {}
|
||||
|
||||
func (x *Config) ProtoReflect() protoreflect.Message {
|
||||
mi := &file_transport_internet_finalmask_xdns_config_proto_msgTypes[1]
|
||||
mi := &file_transport_internet_finalmask_xdns_config_proto_msgTypes[2]
|
||||
if x != nil {
|
||||
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
|
||||
if ms.LoadMessageInfo() == nil {
|
||||
@@ -134,7 +186,7 @@ func (x *Config) ProtoReflect() protoreflect.Message {
|
||||
|
||||
// Deprecated: Use Config.ProtoReflect.Descriptor instead.
|
||||
func (*Config) Descriptor() ([]byte, []int) {
|
||||
return file_transport_internet_finalmask_xdns_config_proto_rawDescGZIP(), []int{1}
|
||||
return file_transport_internet_finalmask_xdns_config_proto_rawDescGZIP(), []int{2}
|
||||
}
|
||||
|
||||
func (x *Config) GetDomains() []*DomainProto {
|
||||
@@ -144,7 +196,7 @@ func (x *Config) GetDomains() []*DomainProto {
|
||||
return nil
|
||||
}
|
||||
|
||||
func (x *Config) GetResolvers() []*serial.TypedMessage {
|
||||
func (x *Config) GetResolvers() []*ResolverProto {
|
||||
if x != nil {
|
||||
return x.Resolvers
|
||||
}
|
||||
@@ -158,94 +210,6 @@ func (x *Config) GetExtraPoll() int32 {
|
||||
return 0
|
||||
}
|
||||
|
||||
type TCPResolverProto struct {
|
||||
state protoimpl.MessageState `protogen:"open.v1"`
|
||||
Addr string `protobuf:"bytes,1,opt,name=addr,proto3" json:"addr,omitempty"`
|
||||
unknownFields protoimpl.UnknownFields
|
||||
sizeCache protoimpl.SizeCache
|
||||
}
|
||||
|
||||
func (x *TCPResolverProto) Reset() {
|
||||
*x = TCPResolverProto{}
|
||||
mi := &file_transport_internet_finalmask_xdns_config_proto_msgTypes[2]
|
||||
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
|
||||
ms.StoreMessageInfo(mi)
|
||||
}
|
||||
|
||||
func (x *TCPResolverProto) String() string {
|
||||
return protoimpl.X.MessageStringOf(x)
|
||||
}
|
||||
|
||||
func (*TCPResolverProto) ProtoMessage() {}
|
||||
|
||||
func (x *TCPResolverProto) ProtoReflect() protoreflect.Message {
|
||||
mi := &file_transport_internet_finalmask_xdns_config_proto_msgTypes[2]
|
||||
if x != nil {
|
||||
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
|
||||
if ms.LoadMessageInfo() == nil {
|
||||
ms.StoreMessageInfo(mi)
|
||||
}
|
||||
return ms
|
||||
}
|
||||
return mi.MessageOf(x)
|
||||
}
|
||||
|
||||
// Deprecated: Use TCPResolverProto.ProtoReflect.Descriptor instead.
|
||||
func (*TCPResolverProto) Descriptor() ([]byte, []int) {
|
||||
return file_transport_internet_finalmask_xdns_config_proto_rawDescGZIP(), []int{2}
|
||||
}
|
||||
|
||||
func (x *TCPResolverProto) GetAddr() string {
|
||||
if x != nil {
|
||||
return x.Addr
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
type UDPResolverProto struct {
|
||||
state protoimpl.MessageState `protogen:"open.v1"`
|
||||
Addr string `protobuf:"bytes,1,opt,name=addr,proto3" json:"addr,omitempty"`
|
||||
unknownFields protoimpl.UnknownFields
|
||||
sizeCache protoimpl.SizeCache
|
||||
}
|
||||
|
||||
func (x *UDPResolverProto) Reset() {
|
||||
*x = UDPResolverProto{}
|
||||
mi := &file_transport_internet_finalmask_xdns_config_proto_msgTypes[3]
|
||||
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
|
||||
ms.StoreMessageInfo(mi)
|
||||
}
|
||||
|
||||
func (x *UDPResolverProto) String() string {
|
||||
return protoimpl.X.MessageStringOf(x)
|
||||
}
|
||||
|
||||
func (*UDPResolverProto) ProtoMessage() {}
|
||||
|
||||
func (x *UDPResolverProto) ProtoReflect() protoreflect.Message {
|
||||
mi := &file_transport_internet_finalmask_xdns_config_proto_msgTypes[3]
|
||||
if x != nil {
|
||||
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
|
||||
if ms.LoadMessageInfo() == nil {
|
||||
ms.StoreMessageInfo(mi)
|
||||
}
|
||||
return ms
|
||||
}
|
||||
return mi.MessageOf(x)
|
||||
}
|
||||
|
||||
// Deprecated: Use UDPResolverProto.ProtoReflect.Descriptor instead.
|
||||
func (*UDPResolverProto) Descriptor() ([]byte, []int) {
|
||||
return file_transport_internet_finalmask_xdns_config_proto_rawDescGZIP(), []int{3}
|
||||
}
|
||||
|
||||
func (x *UDPResolverProto) GetAddr() string {
|
||||
if x != nil {
|
||||
return x.Addr
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
var File_transport_internet_finalmask_xdns_config_proto protoreflect.FileDescriptor
|
||||
|
||||
const file_transport_internet_finalmask_xdns_config_proto_rawDesc = "" +
|
||||
@@ -257,16 +221,15 @@ const file_transport_internet_finalmask_xdns_config_proto_rawDesc = "" +
|
||||
"\vlabel_limit\x18\x03 \x01(\x05R\n" +
|
||||
"labelLimit\x12\x14\n" +
|
||||
"\x05types\x18\x04 \x03(\x05R\x05types\x12\x14\n" +
|
||||
"\x05edns0\x18\x05 \x01(\x05R\x05edns0\"\xb6\x01\n" +
|
||||
"\x05edns0\x18\x05 \x01(\x05R\x05edns0\"7\n" +
|
||||
"\rResolverProto\x12\x12\n" +
|
||||
"\x04type\x18\x01 \x01(\tR\x04type\x12\x12\n" +
|
||||
"\x04addr\x18\x02 \x01(\tR\x04addr\"\xcb\x01\n" +
|
||||
"\x06Config\x12M\n" +
|
||||
"\adomains\x18\x01 \x03(\v23.xray.transport.internet.finalmask.xdns.DomainProtoR\adomains\x12>\n" +
|
||||
"\tresolvers\x18\x02 \x03(\v2 .xray.common.serial.TypedMessageR\tresolvers\x12\x1d\n" +
|
||||
"\adomains\x18\x01 \x03(\v23.xray.transport.internet.finalmask.xdns.DomainProtoR\adomains\x12S\n" +
|
||||
"\tresolvers\x18\x02 \x03(\v25.xray.transport.internet.finalmask.xdns.ResolverProtoR\tresolvers\x12\x1d\n" +
|
||||
"\n" +
|
||||
"extra_poll\x18\x03 \x01(\x05R\textraPoll\"&\n" +
|
||||
"\x10TCPResolverProto\x12\x12\n" +
|
||||
"\x04addr\x18\x01 \x01(\tR\x04addr\"&\n" +
|
||||
"\x10UDPResolverProto\x12\x12\n" +
|
||||
"\x04addr\x18\x01 \x01(\tR\x04addrB\x94\x01\n" +
|
||||
"extra_poll\x18\x03 \x01(\x05R\textraPollB\x94\x01\n" +
|
||||
"*com.xray.transport.internet.finalmask.xdnsP\x01Z;github.com/xtls/xray-core/transport/internet/finalmask/xdns\xaa\x02&Xray.Transport.Internet.Finalmask.Xdnsb\x06proto3"
|
||||
|
||||
var (
|
||||
@@ -281,17 +244,15 @@ func file_transport_internet_finalmask_xdns_config_proto_rawDescGZIP() []byte {
|
||||
return file_transport_internet_finalmask_xdns_config_proto_rawDescData
|
||||
}
|
||||
|
||||
var file_transport_internet_finalmask_xdns_config_proto_msgTypes = make([]protoimpl.MessageInfo, 4)
|
||||
var file_transport_internet_finalmask_xdns_config_proto_msgTypes = make([]protoimpl.MessageInfo, 3)
|
||||
var file_transport_internet_finalmask_xdns_config_proto_goTypes = []any{
|
||||
(*DomainProto)(nil), // 0: xray.transport.internet.finalmask.xdns.DomainProto
|
||||
(*Config)(nil), // 1: xray.transport.internet.finalmask.xdns.Config
|
||||
(*TCPResolverProto)(nil), // 2: xray.transport.internet.finalmask.xdns.TCPResolverProto
|
||||
(*UDPResolverProto)(nil), // 3: xray.transport.internet.finalmask.xdns.UDPResolverProto
|
||||
(*serial.TypedMessage)(nil), // 4: xray.common.serial.TypedMessage
|
||||
(*DomainProto)(nil), // 0: xray.transport.internet.finalmask.xdns.DomainProto
|
||||
(*ResolverProto)(nil), // 1: xray.transport.internet.finalmask.xdns.ResolverProto
|
||||
(*Config)(nil), // 2: xray.transport.internet.finalmask.xdns.Config
|
||||
}
|
||||
var file_transport_internet_finalmask_xdns_config_proto_depIdxs = []int32{
|
||||
0, // 0: xray.transport.internet.finalmask.xdns.Config.domains:type_name -> xray.transport.internet.finalmask.xdns.DomainProto
|
||||
4, // 1: xray.transport.internet.finalmask.xdns.Config.resolvers:type_name -> xray.common.serial.TypedMessage
|
||||
1, // 1: xray.transport.internet.finalmask.xdns.Config.resolvers:type_name -> xray.transport.internet.finalmask.xdns.ResolverProto
|
||||
2, // [2:2] is the sub-list for method output_type
|
||||
2, // [2:2] is the sub-list for method input_type
|
||||
2, // [2:2] is the sub-list for extension type_name
|
||||
@@ -310,7 +271,7 @@ func file_transport_internet_finalmask_xdns_config_proto_init() {
|
||||
GoPackagePath: reflect.TypeOf(x{}).PkgPath(),
|
||||
RawDescriptor: unsafe.Slice(unsafe.StringData(file_transport_internet_finalmask_xdns_config_proto_rawDesc), len(file_transport_internet_finalmask_xdns_config_proto_rawDesc)),
|
||||
NumEnums: 0,
|
||||
NumMessages: 4,
|
||||
NumMessages: 3,
|
||||
NumExtensions: 0,
|
||||
NumServices: 0,
|
||||
},
|
||||
|
||||
@@ -16,16 +16,13 @@ message DomainProto {
|
||||
int32 edns0 = 5;
|
||||
}
|
||||
|
||||
message ResolverProto {
|
||||
string type = 1;
|
||||
string addr = 2;
|
||||
}
|
||||
|
||||
message Config {
|
||||
repeated DomainProto domains = 1;
|
||||
repeated xray.common.serial.TypedMessage resolvers = 2;
|
||||
repeated ResolverProto resolvers = 2;
|
||||
int32 extra_poll = 3;
|
||||
}
|
||||
|
||||
message TCPResolverProto {
|
||||
string addr = 1;
|
||||
}
|
||||
|
||||
message UDPResolverProto {
|
||||
string addr = 1;
|
||||
}
|
||||
@@ -4,7 +4,6 @@ import (
|
||||
"errors"
|
||||
"net"
|
||||
|
||||
"github.com/xtls/xray-core/common/serial"
|
||||
"github.com/xtls/xray-core/transport/internet/finalmask"
|
||||
)
|
||||
|
||||
@@ -15,17 +14,13 @@ type Resolver interface {
|
||||
Close()
|
||||
}
|
||||
|
||||
func NewResolver(proto *serial.TypedMessage, dialer *finalmask.Dialer) (Resolver, error) {
|
||||
config, err := proto.GetInstance()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
switch v := config.(type) {
|
||||
case *TCPResolverProto:
|
||||
return NewTCPResolver(v, dialer)
|
||||
case *UDPResolverProto:
|
||||
return NewUDPResolver(v, dialer)
|
||||
func NewResolver(config *ResolverProto, dialer *finalmask.Dialer) (Resolver, error) {
|
||||
switch config.Type {
|
||||
case "tcp":
|
||||
return NewTCPResolver(config, dialer)
|
||||
case "udp":
|
||||
return NewUDPResolver(config, dialer)
|
||||
default:
|
||||
return nil, errors.New("unknown proto")
|
||||
return nil, errors.New("unknown type")
|
||||
}
|
||||
}
|
||||
|
||||
@@ -24,7 +24,7 @@ type TCPResolver struct {
|
||||
mu sync.Mutex
|
||||
}
|
||||
|
||||
func NewTCPResolver(config *TCPResolverProto, dialer *finalmask.Dialer) (Resolver, error) {
|
||||
func NewTCPResolver(config *ResolverProto, dialer *finalmask.Dialer) (Resolver, error) {
|
||||
dest, err := net.ParseDestination("tcp:" + config.Addr)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
@@ -130,13 +130,15 @@ func (r *TCPResolver) Send(p []byte) {
|
||||
|
||||
func (r *TCPResolver) Close() {
|
||||
r.mu.Lock()
|
||||
defer r.mu.Unlock()
|
||||
if r.closed() {
|
||||
r.mu.Unlock()
|
||||
return
|
||||
}
|
||||
close(r.closeCh)
|
||||
if r.conn != nil {
|
||||
_ = r.conn.Close()
|
||||
conn := r.conn
|
||||
r.mu.Unlock()
|
||||
if conn != nil {
|
||||
_ = conn.Close()
|
||||
}
|
||||
r.wg.Wait()
|
||||
close(r.readCh)
|
||||
|
||||
@@ -22,7 +22,7 @@ type UDPResolver struct {
|
||||
mu sync.Mutex
|
||||
}
|
||||
|
||||
func NewUDPResolver(config *UDPResolverProto, dialer *finalmask.Dialer) (Resolver, error) {
|
||||
func NewUDPResolver(config *ResolverProto, dialer *finalmask.Dialer) (Resolver, error) {
|
||||
dest, err := net.ParseDestination("udp:" + config.Addr)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
@@ -117,13 +117,15 @@ func (r *UDPResolver) Send(p []byte) {
|
||||
|
||||
func (r *UDPResolver) Close() {
|
||||
r.mu.Lock()
|
||||
defer r.mu.Unlock()
|
||||
if r.closed() {
|
||||
r.mu.Unlock()
|
||||
return
|
||||
}
|
||||
close(r.closeCh)
|
||||
if r.conn != nil {
|
||||
_ = r.conn.Close()
|
||||
conn := r.conn
|
||||
r.mu.Unlock()
|
||||
if conn != nil {
|
||||
_ = conn.Close()
|
||||
}
|
||||
r.wg.Wait()
|
||||
close(r.readCh)
|
||||
|
||||
@@ -52,6 +52,9 @@ func NewServer(c *Config, raw net.PacketConn) (net.PacketConn, error) {
|
||||
for j := range c.Domains[i].Types {
|
||||
types = append(types, uint16(c.Domains[i].Types[j]))
|
||||
}
|
||||
if len(types) == 0 {
|
||||
types = []uint16{1, 5, 16, 28}
|
||||
}
|
||||
domain, err := NewDomain(c.Domains[i].Name, int(c.Domains[i].LenLimit), int(c.Domains[i].LabelLimit), types, uint16(c.Domains[i].Edns0))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
|
||||
@@ -127,7 +127,7 @@ func getGrpcClient(ctx context.Context, dest net.Destination, streamSettings *in
|
||||
if streamSettings.FinalMask != nil {
|
||||
c, err = streamSettings.FinalMask.DialTCP(gctx, net.TCPDestination(address, port))
|
||||
} else {
|
||||
c, err = internet.DialSystem(ctx, dest, streamSettings.SocketSettings)
|
||||
c, err = internet.DialSystem(gctx, net.TCPDestination(address, port), streamSettings.SocketSettings)
|
||||
}
|
||||
if err == nil {
|
||||
if tlsConfig != nil {
|
||||
|
||||
Reference in New Issue
Block a user