mirror of
https://github.com/XTLS/Xray-core.git
synced 2026-10-05 05:18:15 +03:00
Compare commits
42
Commits
windows-tun-fix
...
lua
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
8d80e9924f | ||
|
|
103f5cbfdb | ||
|
|
52be8f8170 | ||
|
|
248cb666f3 | ||
|
|
54610a39a5 | ||
|
|
37d9894dde | ||
|
|
e73fb0d500 | ||
|
|
d74b8b1ba5 | ||
|
|
edb8e7477e | ||
|
|
6cd6c61578 | ||
|
|
db2dc8840a | ||
|
|
7ab6930f27 | ||
|
|
73fb3e8f4a | ||
|
|
745526f14c | ||
|
|
399563b6d9 | ||
|
|
e38794ed88 | ||
|
|
2610e57ecf | ||
|
|
5afe260f10 | ||
|
|
5d1d8200d9 | ||
|
|
2440f53cdd | ||
|
|
b26a91de4f | ||
|
|
1f304916bd | ||
|
|
0086362663 | ||
|
|
e51b3c3621 | ||
|
|
6243d2a26e | ||
|
|
35e616d3b9 | ||
|
|
08cb6e6bca | ||
|
|
48ad0300ea | ||
|
|
0fc379203f | ||
|
|
fc8f8a451d | ||
|
|
1c52c65872 | ||
|
|
5724db08f4 | ||
|
|
3d3306503d | ||
|
|
5e1bb92b98 | ||
|
|
e5e85ca9da | ||
|
|
459301d42e | ||
|
|
3982028a9c | ||
|
|
70b8e9a61d | ||
|
|
219f758060 | ||
|
|
9628003594 | ||
|
|
72d9ab50b9 | ||
|
|
235843c5d2 |
@@ -470,6 +470,9 @@ func (d *DefaultDispatcher) routedDispatch(ctx context.Context, link *transport.
|
|||||||
return // DO NOT CHANGE: the traffic shouldn't be processed by default outbound if the specified outbound tag doesn't exist (yet), e.g., VLESS Reverse Proxy
|
return // DO NOT CHANGE: the traffic shouldn't be processed by default outbound if the specified outbound tag doesn't exist (yet), e.g., VLESS Reverse Proxy
|
||||||
}
|
}
|
||||||
} else {
|
} else {
|
||||||
|
if err != common.ErrNoClue {
|
||||||
|
errors.LogErrorInner(ctx, err, "failed to pick route for ", destination)
|
||||||
|
}
|
||||||
errors.LogInfo(ctx, "default route for ", destination)
|
errors.LogInfo(ctx, "default route for ", destination)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
+25
-6
@@ -93,6 +93,7 @@ type NameServer struct {
|
|||||||
UnexpectedIp []*geodata.IPRule `protobuf:"bytes,13,rep,name=unexpected_ip,json=unexpectedIp,proto3" json:"unexpected_ip,omitempty"`
|
UnexpectedIp []*geodata.IPRule `protobuf:"bytes,13,rep,name=unexpected_ip,json=unexpectedIp,proto3" json:"unexpected_ip,omitempty"`
|
||||||
ActUnprior bool `protobuf:"varint,14,opt,name=actUnprior,proto3" json:"actUnprior,omitempty"`
|
ActUnprior bool `protobuf:"varint,14,opt,name=actUnprior,proto3" json:"actUnprior,omitempty"`
|
||||||
PolicyID uint32 `protobuf:"varint,17,opt,name=policyID,proto3" json:"policyID,omitempty"`
|
PolicyID uint32 `protobuf:"varint,17,opt,name=policyID,proto3" json:"policyID,omitempty"`
|
||||||
|
Id string `protobuf:"bytes,18,opt,name=id,proto3" json:"id,omitempty"`
|
||||||
unknownFields protoimpl.UnknownFields
|
unknownFields protoimpl.UnknownFields
|
||||||
sizeCache protoimpl.SizeCache
|
sizeCache protoimpl.SizeCache
|
||||||
}
|
}
|
||||||
@@ -239,6 +240,13 @@ func (x *NameServer) GetPolicyID() uint32 {
|
|||||||
return 0
|
return 0
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (x *NameServer) GetId() string {
|
||||||
|
if x != nil {
|
||||||
|
return x.Id
|
||||||
|
}
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
|
||||||
type Config struct {
|
type Config struct {
|
||||||
state protoimpl.MessageState `protogen:"open.v1"`
|
state protoimpl.MessageState `protogen:"open.v1"`
|
||||||
// NameServer list used by this DNS client.
|
// NameServer list used by this DNS client.
|
||||||
@@ -258,8 +266,10 @@ type Config struct {
|
|||||||
DisableFallback bool `protobuf:"varint,10,opt,name=disableFallback,proto3" json:"disableFallback,omitempty"`
|
DisableFallback bool `protobuf:"varint,10,opt,name=disableFallback,proto3" json:"disableFallback,omitempty"`
|
||||||
DisableFallbackIfMatch bool `protobuf:"varint,11,opt,name=disableFallbackIfMatch,proto3" json:"disableFallbackIfMatch,omitempty"`
|
DisableFallbackIfMatch bool `protobuf:"varint,11,opt,name=disableFallbackIfMatch,proto3" json:"disableFallbackIfMatch,omitempty"`
|
||||||
EnableParallelQuery bool `protobuf:"varint,14,opt,name=enableParallelQuery,proto3" json:"enableParallelQuery,omitempty"`
|
EnableParallelQuery bool `protobuf:"varint,14,opt,name=enableParallelQuery,proto3" json:"enableParallelQuery,omitempty"`
|
||||||
unknownFields protoimpl.UnknownFields
|
// Absolute path to the Lua DNS query script.
|
||||||
sizeCache protoimpl.SizeCache
|
Script string `protobuf:"bytes,15,opt,name=script,proto3" json:"script,omitempty"`
|
||||||
|
unknownFields protoimpl.UnknownFields
|
||||||
|
sizeCache protoimpl.SizeCache
|
||||||
}
|
}
|
||||||
|
|
||||||
func (x *Config) Reset() {
|
func (x *Config) Reset() {
|
||||||
@@ -369,6 +379,13 @@ func (x *Config) GetEnableParallelQuery() bool {
|
|||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (x *Config) GetScript() string {
|
||||||
|
if x != nil {
|
||||||
|
return x.Script
|
||||||
|
}
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
|
||||||
type Config_HostMapping struct {
|
type Config_HostMapping struct {
|
||||||
state protoimpl.MessageState `protogen:"open.v1"`
|
state protoimpl.MessageState `protogen:"open.v1"`
|
||||||
Domain *geodata.DomainRule `protobuf:"bytes,2,opt,name=domain,proto3" json:"domain,omitempty"`
|
Domain *geodata.DomainRule `protobuf:"bytes,2,opt,name=domain,proto3" json:"domain,omitempty"`
|
||||||
@@ -435,7 +452,7 @@ var File_app_dns_config_proto protoreflect.FileDescriptor
|
|||||||
|
|
||||||
const file_app_dns_config_proto_rawDesc = "" +
|
const file_app_dns_config_proto_rawDesc = "" +
|
||||||
"\n" +
|
"\n" +
|
||||||
"\x14app/dns/config.proto\x12\fxray.app.dns\x1a\x1ccommon/net/destination.proto\x1a\x1bcommon/geodata/geodat.proto\"\xde\x05\n" +
|
"\x14app/dns/config.proto\x12\fxray.app.dns\x1a\x1ccommon/net/destination.proto\x1a\x1bcommon/geodata/geodat.proto\"\xee\x05\n" +
|
||||||
"\n" +
|
"\n" +
|
||||||
"NameServer\x123\n" +
|
"NameServer\x123\n" +
|
||||||
"\aaddress\x18\x01 \x01(\v2\x19.xray.common.net.EndpointR\aaddress\x12\x1b\n" +
|
"\aaddress\x18\x01 \x01(\v2\x19.xray.common.net.EndpointR\aaddress\x12\x1b\n" +
|
||||||
@@ -461,10 +478,11 @@ const file_app_dns_config_proto_rawDesc = "" +
|
|||||||
"\n" +
|
"\n" +
|
||||||
"actUnprior\x18\x0e \x01(\bR\n" +
|
"actUnprior\x18\x0e \x01(\bR\n" +
|
||||||
"actUnprior\x12\x1a\n" +
|
"actUnprior\x12\x1a\n" +
|
||||||
"\bpolicyID\x18\x11 \x01(\rR\bpolicyIDB\x0f\n" +
|
"\bpolicyID\x18\x11 \x01(\rR\bpolicyID\x12\x0e\n" +
|
||||||
|
"\x02id\x18\x12 \x01(\tR\x02idB\x0f\n" +
|
||||||
"\r_disableCacheB\r\n" +
|
"\r_disableCacheB\r\n" +
|
||||||
"\v_serveStaleB\x12\n" +
|
"\v_serveStaleB\x12\n" +
|
||||||
"\x10_serveExpiredTTLJ\x04\b\x04\x10\x05\"\x82\x05\n" +
|
"\x10_serveExpiredTTLJ\x04\b\x04\x10\x05\"\x9a\x05\n" +
|
||||||
"\x06Config\x129\n" +
|
"\x06Config\x129\n" +
|
||||||
"\vname_server\x18\x05 \x03(\v2\x18.xray.app.dns.NameServerR\n" +
|
"\vname_server\x18\x05 \x03(\v2\x18.xray.app.dns.NameServerR\n" +
|
||||||
"nameServer\x12\x1b\n" +
|
"nameServer\x12\x1b\n" +
|
||||||
@@ -480,7 +498,8 @@ const file_app_dns_config_proto_rawDesc = "" +
|
|||||||
"\x0fdisableFallback\x18\n" +
|
"\x0fdisableFallback\x18\n" +
|
||||||
" \x01(\bR\x0fdisableFallback\x126\n" +
|
" \x01(\bR\x0fdisableFallback\x126\n" +
|
||||||
"\x16disableFallbackIfMatch\x18\v \x01(\bR\x16disableFallbackIfMatch\x120\n" +
|
"\x16disableFallbackIfMatch\x18\v \x01(\bR\x16disableFallbackIfMatch\x120\n" +
|
||||||
"\x13enableParallelQuery\x18\x0e \x01(\bR\x13enableParallelQuery\x1a}\n" +
|
"\x13enableParallelQuery\x18\x0e \x01(\bR\x13enableParallelQuery\x12\x16\n" +
|
||||||
|
"\x06script\x18\x0f \x01(\tR\x06script\x1a}\n" +
|
||||||
"\vHostMapping\x127\n" +
|
"\vHostMapping\x127\n" +
|
||||||
"\x06domain\x18\x02 \x01(\v2\x1f.xray.common.geodata.DomainRuleR\x06domain\x12\x0e\n" +
|
"\x06domain\x18\x02 \x01(\v2\x1f.xray.common.geodata.DomainRuleR\x06domain\x12\x0e\n" +
|
||||||
"\x02ip\x18\x03 \x03(\fR\x02ip\x12%\n" +
|
"\x02ip\x18\x03 \x03(\fR\x02ip\x12%\n" +
|
||||||
|
|||||||
@@ -27,6 +27,7 @@ message NameServer {
|
|||||||
repeated xray.common.geodata.IPRule unexpected_ip = 13;
|
repeated xray.common.geodata.IPRule unexpected_ip = 13;
|
||||||
bool actUnprior = 14;
|
bool actUnprior = 14;
|
||||||
uint32 policyID = 17;
|
uint32 policyID = 17;
|
||||||
|
string id = 18;
|
||||||
}
|
}
|
||||||
|
|
||||||
enum QueryStrategy {
|
enum QueryStrategy {
|
||||||
@@ -73,4 +74,7 @@ message Config {
|
|||||||
bool disableFallbackIfMatch = 11;
|
bool disableFallbackIfMatch = 11;
|
||||||
|
|
||||||
bool enableParallelQuery = 14;
|
bool enableParallelQuery = 14;
|
||||||
|
|
||||||
|
// Absolute path to the Lua DNS query script.
|
||||||
|
string script = 15;
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -31,6 +31,8 @@ type DNS struct {
|
|||||||
domainMatcher geodata.DomainMatcher
|
domainMatcher geodata.DomainMatcher
|
||||||
matcherInfos []*DomainMatcherInfo
|
matcherInfos []*DomainMatcherInfo
|
||||||
checkSystem bool
|
checkSystem bool
|
||||||
|
script *scriptEngine
|
||||||
|
scriptPath string
|
||||||
}
|
}
|
||||||
|
|
||||||
// DomainMatcherInfo contains information attached to index returned by Server.domainMatcher.
|
// DomainMatcherInfo contains information attached to index returned by Server.domainMatcher.
|
||||||
@@ -180,6 +182,7 @@ func New(ctx context.Context, config *Config) (*DNS, error) {
|
|||||||
disableFallbackIfMatch: config.DisableFallbackIfMatch,
|
disableFallbackIfMatch: config.DisableFallbackIfMatch,
|
||||||
enableParallelQuery: config.EnableParallelQuery,
|
enableParallelQuery: config.EnableParallelQuery,
|
||||||
checkSystem: checkSystem,
|
checkSystem: checkSystem,
|
||||||
|
scriptPath: config.Script,
|
||||||
}, nil
|
}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -190,11 +193,21 @@ func (*DNS) Type() interface{} {
|
|||||||
|
|
||||||
// Start implements common.Runnable.
|
// Start implements common.Runnable.
|
||||||
func (s *DNS) Start() error {
|
func (s *DNS) Start() error {
|
||||||
|
if s.scriptPath != "" {
|
||||||
|
engine, err := newScriptEngine(s.scriptPath, s)
|
||||||
|
if err != nil {
|
||||||
|
return errors.New("failed to initialize DNS script").Base(err)
|
||||||
|
}
|
||||||
|
s.script = engine
|
||||||
|
}
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// Close implements common.Closable.
|
// Close implements common.Closable.
|
||||||
func (s *DNS) Close() error {
|
func (s *DNS) Close() error {
|
||||||
|
if s.script != nil {
|
||||||
|
s.script.close()
|
||||||
|
}
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -279,6 +292,9 @@ func (s *DNS) LookupIP(domain string, option dns.IPOption) ([]net.IP, uint32, er
|
|||||||
}
|
}
|
||||||
|
|
||||||
// Name servers lookup
|
// Name servers lookup
|
||||||
|
if s.script != nil {
|
||||||
|
return s.script.query(domain, option)
|
||||||
|
}
|
||||||
if s.enableParallelQuery {
|
if s.enableParallelQuery {
|
||||||
return s.parallelQuery(domain, option)
|
return s.parallelQuery(domain, option)
|
||||||
} else {
|
} else {
|
||||||
|
|||||||
+169
@@ -0,0 +1,169 @@
|
|||||||
|
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
|
||||||
|
}
|
||||||
@@ -0,0 +1,328 @@
|
|||||||
|
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
|
||||||
|
}
|
||||||
|
|
||||||
|
// BenchmarkLuaDNSQuery measures a preloaded DNS script using server:Query.
|
||||||
|
// The direct case measures the same DNS client without Lua.
|
||||||
|
func BenchmarkLuaDNSQuery(b *testing.B) {
|
||||||
|
option := featureDNS.IPOption{IPv4Enable: true}
|
||||||
|
ip := net.ParseIP("127.0.0.1")
|
||||||
|
upstream := &benchmarkLuaNameServer{ips: []net.IP{ip}}
|
||||||
|
client := &Client{server: upstream, ipOption: &option, timeoutMs: time.Second}
|
||||||
|
server := &DNS{clients: []*Client{client}}
|
||||||
|
L := lua.NewState()
|
||||||
|
defer L.Close()
|
||||||
|
server.registerLua(L)
|
||||||
|
if err := L.DoString(`
|
||||||
|
local server = require("xray.dns").Servers[1]
|
||||||
|
function HandleDNSQuery(domain, ipv4, ipv6, fake)
|
||||||
|
return server:Query(domain, ipv4, ipv6, fake)
|
||||||
|
end
|
||||||
|
`); err != nil {
|
||||||
|
b.Fatal(err)
|
||||||
|
}
|
||||||
|
|
||||||
|
ctx := context.Background()
|
||||||
|
L.SetContext(ctx)
|
||||||
|
for _, bench := range []struct {
|
||||||
|
name string
|
||||||
|
query func() ([]net.IP, uint32, error)
|
||||||
|
}{
|
||||||
|
{"direct", func() ([]net.IP, uint32, error) { return client.QueryIP(ctx, "example.com", option) }},
|
||||||
|
{"lua_script", 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
|
||||||
|
}},
|
||||||
|
} {
|
||||||
|
b.Run(bench.name, func(b *testing.B) {
|
||||||
|
b.ReportAllocs()
|
||||||
|
b.ResetTimer()
|
||||||
|
var ips []net.IP
|
||||||
|
var ttl uint32
|
||||||
|
var err error
|
||||||
|
for i := 0; i < b.N; i++ {
|
||||||
|
ips, ttl, err = bench.query()
|
||||||
|
if err != nil {
|
||||||
|
b.Fatal(err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
b.StopTimer()
|
||||||
|
if ttl != 60 || len(ips) != 1 || !ips[0].Equal(ip) {
|
||||||
|
b.Fatalf("query() = %v, TTL %d; want %v, TTL 60", ips, ttl, ip)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -29,6 +29,7 @@ type Server interface {
|
|||||||
|
|
||||||
// Client is the interface for DNS client.
|
// Client is the interface for DNS client.
|
||||||
type Client struct {
|
type Client struct {
|
||||||
|
id string
|
||||||
server Server
|
server Server
|
||||||
skipFallback bool
|
skipFallback bool
|
||||||
expectedIPs geodata.IPMatcher
|
expectedIPs geodata.IPMatcher
|
||||||
@@ -97,7 +98,7 @@ func NewClient(
|
|||||||
ipOption dns.IPOption,
|
ipOption dns.IPOption,
|
||||||
updateRules func(bool),
|
updateRules func(bool),
|
||||||
) (*Client, error) {
|
) (*Client, error) {
|
||||||
client := &Client{}
|
client := &Client{id: ns.Id}
|
||||||
err := core.RequireFeatures(ctx, func(dispatcher routing.Dispatcher) error {
|
err := core.RequireFeatures(ctx, func(dispatcher routing.Dispatcher) error {
|
||||||
// Create a new server for each client for now
|
// Create a new server for each client for now
|
||||||
server, err := NewServer(ctx, ns.Address.AsDestination(), dispatcher, disableCache, serveStale, serveExpiredTTL, clientIP)
|
server, err := NewServer(ctx, ns.Address.AsDestination(), dispatcher, disableCache, serveStale, serveExpiredTTL, clientIP)
|
||||||
|
|||||||
@@ -49,5 +49,5 @@ func NewLocalNameServer() *LocalNameServer {
|
|||||||
|
|
||||||
// NewLocalDNSClient creates localdns client object for directly lookup in system DNS.
|
// NewLocalDNSClient creates localdns client object for directly lookup in system DNS.
|
||||||
func NewLocalDNSClient(ipOption dns.IPOption) *Client {
|
func NewLocalDNSClient(ipOption dns.IPOption) *Client {
|
||||||
return &Client{server: NewLocalNameServer(), ipOption: &ipOption}
|
return &Client{id: "localhost", server: NewLocalNameServer(), ipOption: &ipOption}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -0,0 +1,63 @@
|
|||||||
|
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
|
||||||
|
}
|
||||||
@@ -0,0 +1,280 @@
|
|||||||
|
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)
|
||||||
|
}
|
||||||
|
}
|
||||||
+14
-4
@@ -587,8 +587,10 @@ type Config struct {
|
|||||||
DomainStrategy Config_DomainStrategy `protobuf:"varint,1,opt,name=domain_strategy,json=domainStrategy,proto3,enum=xray.app.router.Config_DomainStrategy" json:"domain_strategy,omitempty"`
|
DomainStrategy Config_DomainStrategy `protobuf:"varint,1,opt,name=domain_strategy,json=domainStrategy,proto3,enum=xray.app.router.Config_DomainStrategy" json:"domain_strategy,omitempty"`
|
||||||
Rule []*RoutingRule `protobuf:"bytes,2,rep,name=rule,proto3" json:"rule,omitempty"`
|
Rule []*RoutingRule `protobuf:"bytes,2,rep,name=rule,proto3" json:"rule,omitempty"`
|
||||||
BalancingRule []*BalancingRule `protobuf:"bytes,3,rep,name=balancing_rule,json=balancingRule,proto3" json:"balancing_rule,omitempty"`
|
BalancingRule []*BalancingRule `protobuf:"bytes,3,rep,name=balancing_rule,json=balancingRule,proto3" json:"balancing_rule,omitempty"`
|
||||||
unknownFields protoimpl.UnknownFields
|
// Absolute path to the Lua routing script.
|
||||||
sizeCache protoimpl.SizeCache
|
Script string `protobuf:"bytes,4,opt,name=script,proto3" json:"script,omitempty"`
|
||||||
|
unknownFields protoimpl.UnknownFields
|
||||||
|
sizeCache protoimpl.SizeCache
|
||||||
}
|
}
|
||||||
|
|
||||||
func (x *Config) Reset() {
|
func (x *Config) Reset() {
|
||||||
@@ -642,6 +644,13 @@ func (x *Config) GetBalancingRule() []*BalancingRule {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (x *Config) GetScript() string {
|
||||||
|
if x != nil {
|
||||||
|
return x.Script
|
||||||
|
}
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
|
||||||
var File_app_router_config_proto protoreflect.FileDescriptor
|
var File_app_router_config_proto protoreflect.FileDescriptor
|
||||||
|
|
||||||
const file_app_router_config_proto_rawDesc = "" +
|
const file_app_router_config_proto_rawDesc = "" +
|
||||||
@@ -699,11 +708,12 @@ const file_app_router_config_proto_rawDesc = "" +
|
|||||||
"\tbaselines\x18\x03 \x03(\x03R\tbaselines\x12\x1a\n" +
|
"\tbaselines\x18\x03 \x03(\x03R\tbaselines\x12\x1a\n" +
|
||||||
"\bexpected\x18\x04 \x01(\x05R\bexpected\x12\x16\n" +
|
"\bexpected\x18\x04 \x01(\x05R\bexpected\x12\x16\n" +
|
||||||
"\x06maxRTT\x18\x05 \x01(\x03R\x06maxRTT\x12\x1c\n" +
|
"\x06maxRTT\x18\x05 \x01(\x03R\x06maxRTT\x12\x1c\n" +
|
||||||
"\ttolerance\x18\x06 \x01(\x02R\ttolerance\"\x96\x02\n" +
|
"\ttolerance\x18\x06 \x01(\x02R\ttolerance\"\xae\x02\n" +
|
||||||
"\x06Config\x12O\n" +
|
"\x06Config\x12O\n" +
|
||||||
"\x0fdomain_strategy\x18\x01 \x01(\x0e2&.xray.app.router.Config.DomainStrategyR\x0edomainStrategy\x120\n" +
|
"\x0fdomain_strategy\x18\x01 \x01(\x0e2&.xray.app.router.Config.DomainStrategyR\x0edomainStrategy\x120\n" +
|
||||||
"\x04rule\x18\x02 \x03(\v2\x1c.xray.app.router.RoutingRuleR\x04rule\x12E\n" +
|
"\x04rule\x18\x02 \x03(\v2\x1c.xray.app.router.RoutingRuleR\x04rule\x12E\n" +
|
||||||
"\x0ebalancing_rule\x18\x03 \x03(\v2\x1e.xray.app.router.BalancingRuleR\rbalancingRule\"B\n" +
|
"\x0ebalancing_rule\x18\x03 \x03(\v2\x1e.xray.app.router.BalancingRuleR\rbalancingRule\x12\x16\n" +
|
||||||
|
"\x06script\x18\x04 \x01(\tR\x06script\"B\n" +
|
||||||
"\x0eDomainStrategy\x12\b\n" +
|
"\x0eDomainStrategy\x12\b\n" +
|
||||||
"\x04AsIs\x10\x00\x12\x10\n" +
|
"\x04AsIs\x10\x00\x12\x10\n" +
|
||||||
"\fIpIfNonMatch\x10\x02\x12\x0e\n" +
|
"\fIpIfNonMatch\x10\x02\x12\x0e\n" +
|
||||||
|
|||||||
@@ -110,4 +110,6 @@ message Config {
|
|||||||
DomainStrategy domain_strategy = 1;
|
DomainStrategy domain_strategy = 1;
|
||||||
repeated RoutingRule rule = 2;
|
repeated RoutingRule rule = 2;
|
||||||
repeated BalancingRule balancing_rule = 3;
|
repeated BalancingRule balancing_rule = 3;
|
||||||
|
// Absolute path to the Lua routing script.
|
||||||
|
string script = 4;
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -0,0 +1,175 @@
|
|||||||
|
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)
|
||||||
|
}
|
||||||
@@ -0,0 +1,349 @@
|
|||||||
|
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)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// BenchmarkLuaRoute measures a preloaded routing script using its context bridge.
|
||||||
|
// The direct case runs an equivalent native routing rule.
|
||||||
|
func BenchmarkLuaRoute(b *testing.B) {
|
||||||
|
r := new(Router)
|
||||||
|
if err := r.Init(context.Background(), &Config{Rule: []*RoutingRule{{
|
||||||
|
TargetTag: &RoutingRule_Tag{Tag: "out"},
|
||||||
|
RuleTag: "rule",
|
||||||
|
InboundTag: []string{"in"},
|
||||||
|
Networks: []net.Network{net.Network_TCP},
|
||||||
|
Ip: []*geodata.IPRule{{
|
||||||
|
Value: &geodata.IPRule_Custom{Custom: &geodata.CIDRRule{
|
||||||
|
Cidr: &geodata.CIDR{Ip: []byte{127, 0, 0, 0}, Prefix: 8},
|
||||||
|
}},
|
||||||
|
}},
|
||||||
|
}}}, nil, nil, nil); err != nil {
|
||||||
|
b.Fatal(err)
|
||||||
|
}
|
||||||
|
L := lua.NewState()
|
||||||
|
defer L.Close()
|
||||||
|
r.RegisterLua(L)
|
||||||
|
geodata.RegisterLua(L)
|
||||||
|
if err := L.DoString(`
|
||||||
|
local router = require("xray.router")
|
||||||
|
local matcher = require("xray.geodata").BuildIPMatcher("127.0.0.0/8")
|
||||||
|
function HandleRoute(ctx, inboundTag, sourcePort, targetPort, localPort,
|
||||||
|
targetDomain, network, protocol, user, vlessRoute, skipDNSResolve)
|
||||||
|
if inboundTag == "in" and network == router.NetworkTCP and matcher:AnyMatch(ctx:GetTargetIPs()) then
|
||||||
|
return "out", "rule"
|
||||||
|
end
|
||||||
|
end
|
||||||
|
`); err != nil {
|
||||||
|
b.Fatal(err)
|
||||||
|
}
|
||||||
|
|
||||||
|
L.SetContext(context.Background())
|
||||||
|
ctx := newLuaRouteTestContext()
|
||||||
|
for _, benchmark := range []struct {
|
||||||
|
name string
|
||||||
|
route func() (string, string, error)
|
||||||
|
}{
|
||||||
|
{"direct", func() (string, string, error) {
|
||||||
|
route, err := r.PickRoute(ctx)
|
||||||
|
if err != nil {
|
||||||
|
return "", "", err
|
||||||
|
}
|
||||||
|
return route.GetOutboundTag(), route.GetRuleTag(), nil
|
||||||
|
}},
|
||||||
|
{"lua_script", func() (string, string, error) {
|
||||||
|
if err := callLuaRoute(L, ctx); err != nil {
|
||||||
|
return "", "", err
|
||||||
|
}
|
||||||
|
outboundTag, ruleTag, err := readLuaRouteResult(L)
|
||||||
|
L.Pop(3)
|
||||||
|
return outboundTag, ruleTag, err
|
||||||
|
}},
|
||||||
|
} {
|
||||||
|
b.Run(benchmark.name, func(b *testing.B) {
|
||||||
|
b.ReportAllocs()
|
||||||
|
b.ResetTimer()
|
||||||
|
var outboundTag, ruleTag string
|
||||||
|
var err error
|
||||||
|
for i := 0; i < b.N; i++ {
|
||||||
|
outboundTag, ruleTag, err = benchmark.route()
|
||||||
|
if err != nil {
|
||||||
|
b.Fatal(err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
b.StopTimer()
|
||||||
|
if outboundTag != "out" || ruleTag != "rule" {
|
||||||
|
b.Fatalf("route() = %q, %q; want out, rule", outboundTag, ruleTag)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
var _ routing.Context = (*luaRouteTestContext)(nil)
|
||||||
@@ -20,6 +20,8 @@ import (
|
|||||||
type Router struct {
|
type Router struct {
|
||||||
domainStrategy Config_DomainStrategy
|
domainStrategy Config_DomainStrategy
|
||||||
rules atomic.Pointer[[]*Rule]
|
rules atomic.Pointer[[]*Rule]
|
||||||
|
scriptPath string
|
||||||
|
script *scriptEngine
|
||||||
balancers atomic.Pointer[map[string]*Balancer]
|
balancers atomic.Pointer[map[string]*Balancer]
|
||||||
dns dns.Client
|
dns dns.Client
|
||||||
|
|
||||||
@@ -40,6 +42,7 @@ type Route struct {
|
|||||||
// Init initializes the Router.
|
// Init initializes the Router.
|
||||||
func (r *Router) Init(ctx context.Context, config *Config, d dns.Client, ohm outbound.Manager, dispatcher routing.Dispatcher) error {
|
func (r *Router) Init(ctx context.Context, config *Config, d dns.Client, ohm outbound.Manager, dispatcher routing.Dispatcher) error {
|
||||||
r.domainStrategy = config.DomainStrategy
|
r.domainStrategy = config.DomainStrategy
|
||||||
|
r.scriptPath = config.Script
|
||||||
r.dns = d
|
r.dns = d
|
||||||
r.ctx = ctx
|
r.ctx = ctx
|
||||||
r.ohm = ohm
|
r.ohm = ohm
|
||||||
@@ -52,6 +55,10 @@ func (r *Router) Init(ctx context.Context, config *Config, d dns.Client, ohm out
|
|||||||
|
|
||||||
// PickRoute implements routing.Router.
|
// PickRoute implements routing.Router.
|
||||||
func (r *Router) PickRoute(ctx routing.Context) (routing.Route, error) {
|
func (r *Router) PickRoute(ctx routing.Context) (routing.Route, error) {
|
||||||
|
if r.script != nil {
|
||||||
|
return r.script.pickRoute(ctx)
|
||||||
|
}
|
||||||
|
|
||||||
originalCtx := ctx
|
originalCtx := ctx
|
||||||
rule, ctx, err := r.pickRouteInternal(ctx)
|
rule, ctx, err := r.pickRouteInternal(ctx)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -221,6 +228,13 @@ func (r *Router) pickRouteInternal(ctx routing.Context) (*Rule, routing.Context,
|
|||||||
|
|
||||||
// Start implements common.Runnable.
|
// Start implements common.Runnable.
|
||||||
func (r *Router) Start() error {
|
func (r *Router) Start() error {
|
||||||
|
if r.scriptPath != "" {
|
||||||
|
engine, err := newScriptEngine(r.scriptPath, r)
|
||||||
|
if err != nil {
|
||||||
|
return errors.New("failed to initialize routing script").Base(err)
|
||||||
|
}
|
||||||
|
r.script = engine
|
||||||
|
}
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -235,6 +249,9 @@ func closeWebhooks(rules []*Rule) {
|
|||||||
|
|
||||||
// Close implements common.Closable.
|
// Close implements common.Closable.
|
||||||
func (r *Router) Close() error {
|
func (r *Router) Close() error {
|
||||||
|
if r.script != nil {
|
||||||
|
r.script.close()
|
||||||
|
}
|
||||||
r.mu.Lock()
|
r.mu.Lock()
|
||||||
defer r.mu.Unlock()
|
defer r.mu.Unlock()
|
||||||
closeWebhooks(*r.rules.Load())
|
closeWebhooks(*r.rules.Load())
|
||||||
|
|||||||
@@ -0,0 +1,76 @@
|
|||||||
|
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
|
||||||
|
}
|
||||||
@@ -0,0 +1,382 @@
|
|||||||
|
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")
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -82,19 +82,10 @@ func (f *MphDomainMatcherFactory) BuildMatcher(rules []*DomainRule) (DomainMatch
|
|||||||
}
|
}
|
||||||
g.Add(m, uint32(i))
|
g.Add(m, uint32(i))
|
||||||
case *DomainRule_Geosite:
|
case *DomainRule_Geosite:
|
||||||
domains, err := loadSiteWithAttrs(v.Geosite.File, v.Geosite.Code, v.Geosite.Attrs)
|
err := loadSiteMatchers(v.Geosite, func(m strmatcher.Matcher) { g.Add(m, uint32(i)) })
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
for j, d := range domains {
|
|
||||||
domains[j] = nil // peak mem
|
|
||||||
m, err := parseDomain(d)
|
|
||||||
if err != nil {
|
|
||||||
errors.LogError(context.Background(), "ignore invalid geosite entry in ", v.Geosite.File, ":", v.Geosite.Code, " at index ", j, ", ", err)
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
g.Add(m, uint32(i))
|
|
||||||
}
|
|
||||||
default:
|
default:
|
||||||
panic("unknown domain rule type")
|
panic("unknown domain rule type")
|
||||||
}
|
}
|
||||||
@@ -108,12 +99,12 @@ func (f *MphDomainMatcherFactory) BuildMatcher(rules []*DomainRule) (DomainMatch
|
|||||||
return g, nil
|
return g, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
type CompactDomainMatcherFactory struct {
|
type CompactMphDomainMatcherFactory struct {
|
||||||
sync.Mutex
|
sync.Mutex
|
||||||
shared *utils.WeakCacheMap[string, strmatcher.LinearAnyMatcher]
|
shared *utils.WeakCacheMap[string, strmatcher.MphValueMatcher]
|
||||||
}
|
}
|
||||||
|
|
||||||
func (f *CompactDomainMatcherFactory) getOrCreateFrom(rule *GeoSiteRule) (strmatcher.MatcherSet, error) {
|
func (f *CompactMphDomainMatcherFactory) getOrCreateFrom(rule *GeoSiteRule) (*strmatcher.MphValueMatcher, error) {
|
||||||
key := rule.File + ":" + rule.Code + "@" + rule.Attrs
|
key := rule.File + ":" + rule.Code + "@" + rule.Attrs
|
||||||
|
|
||||||
f.Lock()
|
f.Lock()
|
||||||
@@ -125,33 +116,23 @@ func (f *CompactDomainMatcherFactory) getOrCreateFrom(rule *GeoSiteRule) (strmat
|
|||||||
}
|
}
|
||||||
errors.LogDebug(context.Background(), "geodata geosite matcher cache MISS ", key)
|
errors.LogDebug(context.Background(), "geodata geosite matcher cache MISS ", key)
|
||||||
|
|
||||||
s := strmatcher.NewLinearAnyMatcher()
|
s := strmatcher.NewMphValueMatcher()
|
||||||
domains, err := loadSiteWithAttrs(rule.File, rule.Code, rule.Attrs)
|
if err := loadSiteMatchers(rule, func(m strmatcher.Matcher) { s.Add(m, 0) }); err != nil {
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
for i, d := range domains {
|
if err := s.Build(); err != nil {
|
||||||
domains[i] = nil // peak mem
|
return nil, err
|
||||||
m, err := parseDomain(d)
|
|
||||||
if err != nil {
|
|
||||||
errors.LogError(context.Background(), "ignore invalid geosite entry in ", rule.File, ":", rule.Code, " at index ", i, ", ", err)
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
s.Add(m)
|
|
||||||
}
|
}
|
||||||
f.shared.Store(key, s)
|
f.shared.Store(key, s)
|
||||||
return s, err
|
return s, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// BuildMatcher implements DomainMatcherFactory.
|
// BuildMatcher implements DomainMatcherFactory.
|
||||||
func (f *CompactDomainMatcherFactory) BuildMatcher(rules []*DomainRule) (DomainMatcher, error) {
|
func (f *CompactMphDomainMatcherFactory) BuildMatcher(rules []*DomainRule) (DomainMatcher, error) {
|
||||||
if len(rules) == 0 {
|
if len(rules) == 0 {
|
||||||
return nil, errors.New("empty domain rule list")
|
return nil, errors.New("empty domain rule list")
|
||||||
}
|
}
|
||||||
compact := &CompactDomainMatcher{
|
compact := new(CompactMphDomainMatcher)
|
||||||
matchers: make([]strmatcher.MatcherSet, 0, len(rules)),
|
|
||||||
values: make([]uint32, 0, len(rules)),
|
|
||||||
}
|
|
||||||
for i, r := range rules {
|
for i, r := range rules {
|
||||||
switch v := r.Value.(type) {
|
switch v := r.Value.(type) {
|
||||||
case *DomainRule_Custom:
|
case *DomainRule_Custom:
|
||||||
@@ -168,8 +149,7 @@ func (f *CompactDomainMatcherFactory) BuildMatcher(rules []*DomainRule) (DomainM
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
compact.matchers = append(compact.matchers, m)
|
compact.combiner.Add(m, uint32(i))
|
||||||
compact.values = append(compact.values, uint32(i))
|
|
||||||
default:
|
default:
|
||||||
panic("unknown domain rule type")
|
panic("unknown domain rule type")
|
||||||
}
|
}
|
||||||
@@ -177,37 +157,40 @@ func (f *CompactDomainMatcherFactory) BuildMatcher(rules []*DomainRule) (DomainM
|
|||||||
return compact, nil
|
return compact, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
type CompactDomainMatcher struct {
|
type CompactMphDomainMatcher struct {
|
||||||
custom strmatcher.ValueMatcher
|
custom strmatcher.ValueMatcher
|
||||||
matchers []strmatcher.MatcherSet
|
combiner strmatcher.MphValueMatcherCombiner
|
||||||
values []uint32
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// Match implements DomainMatcher.
|
// Match implements DomainMatcher.
|
||||||
func (c *CompactDomainMatcher) Match(input string) []uint32 {
|
func (c *CompactMphDomainMatcher) Match(input string) []uint32 {
|
||||||
var result []uint32
|
result := c.combiner.Match(input)
|
||||||
if c.custom != nil {
|
if c.custom != nil {
|
||||||
result = append(result, c.custom.Match(input)...)
|
result = append(c.custom.Match(input), result...)
|
||||||
}
|
|
||||||
for i, m := range c.matchers {
|
|
||||||
if m.MatchAny(input) {
|
|
||||||
result = append(result, c.values[i])
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
return result
|
return result
|
||||||
}
|
}
|
||||||
|
|
||||||
// MatchAny implements DomainMatcher.
|
// MatchAny implements DomainMatcher.
|
||||||
func (c *CompactDomainMatcher) MatchAny(input string) bool {
|
func (c *CompactMphDomainMatcher) MatchAny(input string) bool {
|
||||||
if c.custom != nil && c.custom.MatchAny(input) {
|
if c.custom != nil && c.custom.MatchAny(input) {
|
||||||
return true
|
return true
|
||||||
}
|
}
|
||||||
for _, m := range c.matchers {
|
return c.combiner.MatchAny(input)
|
||||||
if m.MatchAny(input) {
|
}
|
||||||
return true
|
|
||||||
|
// loadSiteMatchers calls add with a matcher for every domain of the geosite rule and logs the invalid ones.
|
||||||
|
func loadSiteMatchers(rule *GeoSiteRule, add func(strmatcher.Matcher)) error {
|
||||||
|
i := 0
|
||||||
|
return loadSite(rule.File, rule.Code, rule.Attrs, func(t Domain_Type, value []byte) {
|
||||||
|
m, err := parseDomain(&Domain{Type: t, Value: string(value)})
|
||||||
|
if err != nil {
|
||||||
|
errors.LogError(context.Background(), "ignore invalid geosite entry in ", rule.File, ":", rule.Code, " at index ", i, ", ", err)
|
||||||
|
} else {
|
||||||
|
add(m)
|
||||||
}
|
}
|
||||||
}
|
i++
|
||||||
return false
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
func parseDomain(d *Domain) (strmatcher.Matcher, error) {
|
func parseDomain(d *Domain) (strmatcher.Matcher, error) {
|
||||||
@@ -231,7 +214,7 @@ func parseDomain(d *Domain) (strmatcher.Matcher, error) {
|
|||||||
func newDomainMatcherFactory() DomainMatcherFactory {
|
func newDomainMatcherFactory() DomainMatcherFactory {
|
||||||
switch runtime.GOOS {
|
switch runtime.GOOS {
|
||||||
case "ios", "android":
|
case "ios", "android":
|
||||||
return &CompactDomainMatcherFactory{shared: utils.NewWeakCacheMap[string, strmatcher.LinearAnyMatcher]()}
|
return &CompactMphDomainMatcherFactory{shared: utils.NewWeakCacheMap[string, strmatcher.MphValueMatcher]()}
|
||||||
default:
|
default:
|
||||||
return &MphDomainMatcherFactory{shared: utils.NewWeakCacheMap[string, strmatcher.MphValueMatcher]()}
|
return &MphDomainMatcherFactory{shared: utils.NewWeakCacheMap[string, strmatcher.MphValueMatcher]()}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -4,6 +4,7 @@ import (
|
|||||||
"path/filepath"
|
"path/filepath"
|
||||||
"reflect"
|
"reflect"
|
||||||
"slices"
|
"slices"
|
||||||
|
"sync"
|
||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
"github.com/xtls/xray-core/common/geodata/strmatcher"
|
"github.com/xtls/xray-core/common/geodata/strmatcher"
|
||||||
@@ -11,7 +12,7 @@ import (
|
|||||||
)
|
)
|
||||||
|
|
||||||
func TestCompactDomainMatcher_PreservesCustomRuleIndices(t *testing.T) {
|
func TestCompactDomainMatcher_PreservesCustomRuleIndices(t *testing.T) {
|
||||||
factory := &CompactDomainMatcherFactory{shared: utils.NewWeakCacheMap[string, strmatcher.LinearAnyMatcher]()}
|
factory := &CompactMphDomainMatcherFactory{shared: utils.NewWeakCacheMap[string, strmatcher.MphValueMatcher]()}
|
||||||
matcher, err := factory.BuildMatcher([]*DomainRule{
|
matcher, err := factory.BuildMatcher([]*DomainRule{
|
||||||
{Value: &DomainRule_Custom{Custom: &Domain{Type: Domain_Full, Value: "example.com"}}},
|
{Value: &DomainRule_Custom{Custom: &Domain{Type: Domain_Full, Value: "example.com"}}},
|
||||||
{Value: &DomainRule_Custom{Custom: &Domain{Type: Domain_Domain, Value: "example.com"}}},
|
{Value: &DomainRule_Custom{Custom: &Domain{Type: Domain_Domain, Value: "example.com"}}},
|
||||||
@@ -32,7 +33,7 @@ func TestCompactDomainMatcher_PreservesCustomRuleIndices(t *testing.T) {
|
|||||||
func TestCompactDomainMatcher_PreservesMixedRuleIndices(t *testing.T) {
|
func TestCompactDomainMatcher_PreservesMixedRuleIndices(t *testing.T) {
|
||||||
t.Setenv("xray.location.asset", filepath.Join("..", "..", "resources"))
|
t.Setenv("xray.location.asset", filepath.Join("..", "..", "resources"))
|
||||||
|
|
||||||
factory := &CompactDomainMatcherFactory{shared: utils.NewWeakCacheMap[string, strmatcher.LinearAnyMatcher]()}
|
factory := &CompactMphDomainMatcherFactory{shared: utils.NewWeakCacheMap[string, strmatcher.MphValueMatcher]()}
|
||||||
matcher, err := factory.BuildMatcher([]*DomainRule{
|
matcher, err := factory.BuildMatcher([]*DomainRule{
|
||||||
{Value: &DomainRule_Geosite{Geosite: &GeoSiteRule{File: DefaultGeoSiteDat, Code: "CN"}}},
|
{Value: &DomainRule_Geosite{Geosite: &GeoSiteRule{File: DefaultGeoSiteDat, Code: "CN"}}},
|
||||||
{Value: &DomainRule_Custom{Custom: &Domain{Type: Domain_Full, Value: "163.com"}}},
|
{Value: &DomainRule_Custom{Custom: &Domain{Type: Domain_Full, Value: "163.com"}}},
|
||||||
@@ -72,3 +73,76 @@ func TestMphDomainMatcher_MatchReturnsDetachedSlice(t *testing.T) {
|
|||||||
t.Fatalf("Match() after caller mutation = %v, want %v", gotAgain, []uint32{0, 1})
|
t.Fatalf("Match() after caller mutation = %v, want %v", gotAgain, []uint32{0, 1})
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// DNS sorts every Match result in place, so a matcher must never hand out a
|
||||||
|
// slice it keeps, also when only its keyword or regex part matches.
|
||||||
|
func TestDomainMatcher_MatchResultsCanBeSortedConcurrently(t *testing.T) {
|
||||||
|
t.Setenv("xray.location.asset", filepath.Join("..", "..", "resources"))
|
||||||
|
|
||||||
|
rules := []*DomainRule{
|
||||||
|
{Value: &DomainRule_Custom{Custom: &Domain{Type: Domain_Full, Value: "example.com"}}},
|
||||||
|
{Value: &DomainRule_Custom{Custom: &Domain{Type: Domain_Domain, Value: "example.com"}}},
|
||||||
|
{Value: &DomainRule_Custom{Custom: &Domain{Type: Domain_Substr, Value: "exam"}}},
|
||||||
|
{Value: &DomainRule_Custom{Custom: &Domain{Type: Domain_Regex, Value: `^ex.*\.org$`}}},
|
||||||
|
{Value: &DomainRule_Custom{Custom: &Domain{Type: Domain_Substr, Value: "exam"}}},
|
||||||
|
{Value: &DomainRule_Geosite{Geosite: &GeoSiteRule{File: DefaultGeoSiteDat, Code: "CN"}}},
|
||||||
|
{Value: &DomainRule_Custom{Custom: &Domain{Type: Domain_Full, Value: "only.full.test"}}},
|
||||||
|
}
|
||||||
|
cases := []struct {
|
||||||
|
input string
|
||||||
|
want []uint32
|
||||||
|
}{
|
||||||
|
{"example.com", []uint32{0, 1, 2, 4}},
|
||||||
|
{"www.example.com", []uint32{1, 2, 4}},
|
||||||
|
{"exam.net", []uint32{2, 4}}, // keyword part only
|
||||||
|
{"example.org", []uint32{2, 3, 4}},
|
||||||
|
{"163.com", []uint32{5}},
|
||||||
|
{"www.163.com", []uint32{5}},
|
||||||
|
{"only.full.test", []uint32{6}}, // full part only
|
||||||
|
{"nomatch.test", nil},
|
||||||
|
}
|
||||||
|
factories := map[string]DomainMatcherFactory{
|
||||||
|
"mph": &MphDomainMatcherFactory{shared: utils.NewWeakCacheMap[string, strmatcher.MphValueMatcher]()},
|
||||||
|
"compact": &CompactMphDomainMatcherFactory{shared: utils.NewWeakCacheMap[string, strmatcher.MphValueMatcher]()},
|
||||||
|
}
|
||||||
|
for name, factory := range factories {
|
||||||
|
t.Run(name, func(t *testing.T) {
|
||||||
|
matcher, err := factory.BuildMatcher(rules)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("BuildMatcher() failed: %v", err)
|
||||||
|
}
|
||||||
|
for _, c := range cases {
|
||||||
|
got := matcher.Match(c.input)
|
||||||
|
if sorted := slices.Sorted(slices.Values(got)); !slices.Equal(sorted, c.want) {
|
||||||
|
t.Fatalf("Match(%q) = %v, want %v", c.input, sorted, c.want)
|
||||||
|
}
|
||||||
|
got = got[:cap(got)]
|
||||||
|
for j := range got {
|
||||||
|
got[j] = ^uint32(0)
|
||||||
|
}
|
||||||
|
if again := slices.Sorted(slices.Values(matcher.Match(c.input))); !slices.Equal(again, c.want) {
|
||||||
|
t.Fatalf("Match(%q) after caller mutation = %v, want %v", c.input, again, c.want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
var wg sync.WaitGroup
|
||||||
|
for range 8 {
|
||||||
|
wg.Add(1)
|
||||||
|
go func() {
|
||||||
|
defer wg.Done()
|
||||||
|
for range 500 {
|
||||||
|
for _, c := range cases {
|
||||||
|
got := matcher.Match(c.input)
|
||||||
|
slices.Sort(got)
|
||||||
|
if !slices.Equal(got, c.want) {
|
||||||
|
t.Errorf("Match(%q) = %v, want %v", c.input, got, c.want)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
}
|
||||||
|
wg.Wait()
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
+212
-61
@@ -5,11 +5,14 @@ import (
|
|||||||
"bytes"
|
"bytes"
|
||||||
"io"
|
"io"
|
||||||
"runtime"
|
"runtime"
|
||||||
|
"slices"
|
||||||
"strings"
|
"strings"
|
||||||
|
"unicode/utf8"
|
||||||
|
|
||||||
"github.com/xtls/xray-core/common/errors"
|
"github.com/xtls/xray-core/common/errors"
|
||||||
"github.com/xtls/xray-core/common/platform/filesystem"
|
"github.com/xtls/xray-core/common/platform/filesystem"
|
||||||
|
|
||||||
|
"google.golang.org/protobuf/encoding/protowire"
|
||||||
"google.golang.org/protobuf/proto"
|
"google.golang.org/protobuf/proto"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -52,17 +55,56 @@ func loadIP(file, code string) ([]*CIDR, error) {
|
|||||||
return geoip.Cidr, nil
|
return geoip.Cidr, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func loadSite(file, code string) ([]*Domain, error) {
|
// loadSite calls fn, in file order, with the type and value of every domain of the geosite code
|
||||||
bs, err := loadFile(file, code)
|
// that has all the "@"-separated attrs. It decodes the entry while reading the file instead of
|
||||||
|
// unmarshalling it into a []*Domain, so value is only valid during fn.
|
||||||
|
func loadSite(file, code, attrs string, fn func(Domain_Type, []byte)) error {
|
||||||
|
runtime.GC() // peak mem
|
||||||
|
r, err := filesystem.OpenAsset(file)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return errors.New("failed to open ", file).Base(err)
|
||||||
}
|
}
|
||||||
defer runtime.GC() // peak mem
|
defer r.Close()
|
||||||
var geosite GeoSite
|
br := bufio.NewReaderSize(r, 64*1024)
|
||||||
if err := proto.Unmarshal(bs, &geosite); err != nil {
|
n, err := seek(br, []byte(code))
|
||||||
return nil, errors.New("error unmarshal Site in ", file, ":", code).Base(err)
|
if err != nil {
|
||||||
|
return errors.New("failed to load code ", code, " from ", file).Base(err)
|
||||||
}
|
}
|
||||||
return geosite.Domain, nil
|
loadErr := func(err error) error {
|
||||||
|
if err == io.EOF {
|
||||||
|
err = io.ErrUnexpectedEOF
|
||||||
|
}
|
||||||
|
return errors.New("failed to load code ", code, " from ", file).Base(err)
|
||||||
|
}
|
||||||
|
unmarshalErr := func(err error) error {
|
||||||
|
return errors.New("error unmarshal Site in ", file, ":", code).Base(err)
|
||||||
|
}
|
||||||
|
d := newSiteDecoder(attrs, fn)
|
||||||
|
for n > 0 {
|
||||||
|
w, err := br.Peek(min(n, br.Size()))
|
||||||
|
if err != nil {
|
||||||
|
return loadErr(err)
|
||||||
|
}
|
||||||
|
used, err := d.decode(w, len(w) < n)
|
||||||
|
if err != nil {
|
||||||
|
return unmarshalErr(err)
|
||||||
|
}
|
||||||
|
if used == 0 {
|
||||||
|
break // a field longer than the buffer
|
||||||
|
}
|
||||||
|
br.Discard(used)
|
||||||
|
n -= used
|
||||||
|
}
|
||||||
|
if n > 0 {
|
||||||
|
w := make([]byte, n)
|
||||||
|
if _, err := io.ReadFull(br, w); err != nil {
|
||||||
|
return loadErr(err)
|
||||||
|
}
|
||||||
|
if _, err := d.decode(w, false); err != nil {
|
||||||
|
return unmarshalErr(err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func decodeVarint(br *bufio.Reader) (uint64, error) {
|
func decodeVarint(br *bufio.Reader) (uint64, error) {
|
||||||
@@ -82,68 +124,63 @@ func decodeVarint(br *bufio.Reader) (uint64, error) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func find(r io.Reader, code []byte, readBody bool) ([]byte, error) {
|
func find(r io.Reader, code []byte, readBody bool) ([]byte, error) {
|
||||||
|
br := bufio.NewReaderSize(r, 64*1024)
|
||||||
|
bodyL, err := seek(br, code)
|
||||||
|
if err != nil || !readBody {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
out := make([]byte, bodyL)
|
||||||
|
if _, err := io.ReadFull(br, out); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
return out, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// seek advances br to the body of the entry for code and returns the body length.
|
||||||
|
func seek(br *bufio.Reader, code []byte) (int, error) {
|
||||||
codeL := len(code)
|
codeL := len(code)
|
||||||
if codeL == 0 {
|
if codeL == 0 {
|
||||||
return nil, errors.New("empty code")
|
return 0, errors.New("empty code")
|
||||||
}
|
}
|
||||||
|
|
||||||
br := bufio.NewReaderSize(r, 64*1024)
|
|
||||||
need := 2 + codeL // TODO: if code too long
|
need := 2 + codeL // TODO: if code too long
|
||||||
prefixBuf := make([]byte, need)
|
|
||||||
|
|
||||||
for {
|
for {
|
||||||
if _, err := br.ReadByte(); err != nil {
|
if _, err := br.ReadByte(); err != nil {
|
||||||
return nil, err
|
return 0, err
|
||||||
}
|
}
|
||||||
|
|
||||||
x, err := decodeVarint(br)
|
x, err := decodeVarint(br)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return 0, err
|
||||||
}
|
}
|
||||||
bodyL := int(x)
|
bodyL := int(x)
|
||||||
if bodyL <= 0 {
|
if bodyL <= 0 {
|
||||||
return nil, errors.New("invalid body length: ", bodyL)
|
return 0, errors.New("invalid body length: ", bodyL)
|
||||||
}
|
}
|
||||||
|
|
||||||
prefixL := bodyL
|
// Peek no more than the buffer holds: a code longer than the buffer cannot match a single
|
||||||
if prefixL > need {
|
// length byte anyway, so a short peek only skips it, as base find (io.ReadFull) does.
|
||||||
prefixL = need
|
prefix, err := br.Peek(min(bodyL, need, br.Size()))
|
||||||
}
|
if err != nil {
|
||||||
prefix := prefixBuf[:prefixL]
|
if err == io.EOF && len(prefix) > 0 {
|
||||||
if _, err := io.ReadFull(br, prefix); err != nil {
|
err = io.ErrUnexpectedEOF // as io.ReadFull
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
match := false
|
|
||||||
if bodyL >= need {
|
|
||||||
if int(prefix[1]) == codeL && bytes.Equal(prefix[2:need], code) {
|
|
||||||
if !readBody {
|
|
||||||
return nil, nil
|
|
||||||
}
|
|
||||||
match = true
|
|
||||||
}
|
}
|
||||||
|
return 0, err
|
||||||
}
|
}
|
||||||
|
if bodyL >= need && len(prefix) >= need && int(prefix[1]) == codeL && bytes.Equal(prefix[2:], code) {
|
||||||
remain := bodyL - prefixL
|
return bodyL, nil
|
||||||
if match {
|
|
||||||
out := make([]byte, bodyL)
|
|
||||||
copy(out, prefix)
|
|
||||||
if remain > 0 {
|
|
||||||
if _, err := io.ReadFull(br, out[prefixL:]); err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return out, nil
|
|
||||||
}
|
}
|
||||||
|
if _, err := br.Discard(bodyL); err != nil {
|
||||||
if remain > 0 {
|
return 0, err
|
||||||
if _, err := br.Discard(remain); err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// AttributeMatcher, HasAttrMatcher, AllAttrsMatcher and NewAllAttrsMatcher are the exported
|
||||||
|
// attribute helpers that have been part of this package's API since #5814. The streaming loader
|
||||||
|
// above filters attributes itself without building a *Domain, so it does not use them, but they
|
||||||
|
// are kept for external callers. Their behaviour is unchanged.
|
||||||
|
|
||||||
type AttributeMatcher interface {
|
type AttributeMatcher interface {
|
||||||
Match(*Domain) bool
|
Match(*Domain) bool
|
||||||
}
|
}
|
||||||
@@ -185,23 +222,137 @@ func NewAllAttrsMatcher(attrs string) AttributeMatcher {
|
|||||||
return m
|
return m
|
||||||
}
|
}
|
||||||
|
|
||||||
func loadSiteWithAttrs(file, code, attrs string) ([]*Domain, error) {
|
var errInvalidUTF8 = errors.New("string field contains invalid UTF-8")
|
||||||
domains, err := loadSite(file, code)
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
matcher := NewAllAttrsMatcher(attrs)
|
type siteDecoder struct {
|
||||||
if matcher == nil {
|
want []string
|
||||||
return domains, nil
|
has []bool
|
||||||
}
|
fn func(Domain_Type, []byte)
|
||||||
|
}
|
||||||
|
|
||||||
filtered := make([]*Domain, 0, len(domains))
|
func newSiteDecoder(attrs string, fn func(Domain_Type, []byte)) *siteDecoder {
|
||||||
for _, d := range domains {
|
d := &siteDecoder{fn: fn}
|
||||||
if matcher.Match(d) {
|
if attrs != "" {
|
||||||
filtered = append(filtered, d)
|
d.want = strings.Split(attrs, "@")
|
||||||
|
d.has = make([]bool, len(d.want))
|
||||||
|
}
|
||||||
|
return d
|
||||||
|
}
|
||||||
|
|
||||||
|
// decode walks the whole fields at the start of b, a part of an encoded GeoSite (see geodat.proto),
|
||||||
|
// calls fn for every domain that has all attrs and returns how many bytes it used. A field cut off
|
||||||
|
// by the end of b is an error unless more is set. It accepts and rejects what proto.Unmarshal does.
|
||||||
|
func (d *siteDecoder) decode(b []byte, more bool) (int, error) {
|
||||||
|
used := 0
|
||||||
|
for used < len(b) {
|
||||||
|
f, n, err := consumeField(b[used:])
|
||||||
|
if err == io.ErrUnexpectedEOF && more {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
if err != nil {
|
||||||
|
return used, err
|
||||||
|
}
|
||||||
|
used += n
|
||||||
|
if f.typ != protowire.BytesType {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
switch f.num {
|
||||||
|
case 1: // code
|
||||||
|
if !utf8.Valid(f.v) {
|
||||||
|
return used, errInvalidUTF8
|
||||||
|
}
|
||||||
|
case 2: // domain
|
||||||
|
t, value, err := decodeDomain(f.v, d.want, d.has)
|
||||||
|
if err != nil {
|
||||||
|
return used, err
|
||||||
|
}
|
||||||
|
if !slices.Contains(d.has, false) {
|
||||||
|
d.fn(t, value)
|
||||||
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
return used, nil
|
||||||
|
}
|
||||||
|
|
||||||
return filtered, nil
|
// decodeDomain decodes an encoded Domain and sets has[i] if one of its attributes has the key want[i].
|
||||||
|
func decodeDomain(b []byte, want []string, has []bool) (t Domain_Type, value []byte, err error) {
|
||||||
|
clear(has)
|
||||||
|
for len(b) > 0 {
|
||||||
|
f, n, err := consumeField(b)
|
||||||
|
if err != nil {
|
||||||
|
return 0, nil, err
|
||||||
|
}
|
||||||
|
b = b[n:]
|
||||||
|
switch {
|
||||||
|
case f.num == 1 && f.typ == protowire.VarintType: // type
|
||||||
|
t = Domain_Type(f.x)
|
||||||
|
case f.num == 2 && f.typ == protowire.BytesType: // value
|
||||||
|
if !utf8.Valid(f.v) {
|
||||||
|
return 0, nil, errInvalidUTF8
|
||||||
|
}
|
||||||
|
value = f.v
|
||||||
|
case f.num == 3 && f.typ == protowire.BytesType: // attribute
|
||||||
|
key, err := decodeAttributeKey(f.v)
|
||||||
|
if err != nil {
|
||||||
|
return 0, nil, err
|
||||||
|
}
|
||||||
|
for i, w := range want {
|
||||||
|
if string(key) == w {
|
||||||
|
has[i] = true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return t, value, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// decodeAttributeKey returns the key of an encoded Domain.Attribute.
|
||||||
|
func decodeAttributeKey(b []byte) ([]byte, error) {
|
||||||
|
var key []byte
|
||||||
|
for len(b) > 0 {
|
||||||
|
f, n, err := consumeField(b)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
b = b[n:]
|
||||||
|
if f.num == 1 && f.typ == protowire.BytesType {
|
||||||
|
if !utf8.Valid(f.v) {
|
||||||
|
return nil, errInvalidUTF8
|
||||||
|
}
|
||||||
|
key = f.v
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return key, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
type protoField struct {
|
||||||
|
num protowire.Number
|
||||||
|
typ protowire.Type
|
||||||
|
v []byte // payload of a length-delimited field
|
||||||
|
x uint64 // value of a varint field
|
||||||
|
}
|
||||||
|
|
||||||
|
// consumeField parses the first field of an encoded message and returns it with its length.
|
||||||
|
func consumeField(b []byte) (protoField, int, error) {
|
||||||
|
num, typ, n := protowire.ConsumeTag(b)
|
||||||
|
if n < 0 {
|
||||||
|
return protoField{}, 0, protowire.ParseError(n)
|
||||||
|
}
|
||||||
|
if num > protowire.MaxValidNumber {
|
||||||
|
return protoField{}, 0, errors.New("invalid field number ", num)
|
||||||
|
}
|
||||||
|
f := protoField{num: num, typ: typ}
|
||||||
|
var m int
|
||||||
|
switch typ {
|
||||||
|
case protowire.BytesType:
|
||||||
|
f.v, m = protowire.ConsumeBytes(b[n:])
|
||||||
|
case protowire.VarintType:
|
||||||
|
f.x, m = protowire.ConsumeVarint(b[n:])
|
||||||
|
default:
|
||||||
|
m = protowire.ConsumeFieldValue(num, typ, b[n:])
|
||||||
|
}
|
||||||
|
if m < 0 {
|
||||||
|
return protoField{}, 0, protowire.ParseError(m)
|
||||||
|
}
|
||||||
|
return f, n + m, nil
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -0,0 +1,283 @@
|
|||||||
|
package geodata
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"slices"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"google.golang.org/protobuf/encoding/protowire"
|
||||||
|
"google.golang.org/protobuf/proto"
|
||||||
|
)
|
||||||
|
|
||||||
|
type siteEntry struct {
|
||||||
|
Type Domain_Type
|
||||||
|
Value string
|
||||||
|
}
|
||||||
|
|
||||||
|
// unmarshalSite is what loadSite used to do: proto.Unmarshal, then keep the domains that have all attrs.
|
||||||
|
func unmarshalSite(b []byte, attrs string) ([]siteEntry, error) {
|
||||||
|
var site GeoSite
|
||||||
|
if err := proto.Unmarshal(b, &site); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
var entries []siteEntry
|
||||||
|
for _, d := range site.Domain {
|
||||||
|
ok := true
|
||||||
|
for _, key := range strings.Split(attrs, "@") {
|
||||||
|
ok = ok && (attrs == "" || slices.ContainsFunc(d.Attribute, func(a *Domain_Attribute) bool { return a.Key == key }))
|
||||||
|
}
|
||||||
|
if ok {
|
||||||
|
entries = append(entries, siteEntry{d.Type, d.Value})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return entries, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func checkDecodeSite(t *testing.T, name string, b []byte, attrs string) {
|
||||||
|
t.Helper()
|
||||||
|
want, wantErr := unmarshalSite(b, attrs)
|
||||||
|
var got []siteEntry
|
||||||
|
_, err := newSiteDecoder(attrs, func(typ Domain_Type, value []byte) {
|
||||||
|
got = append(got, siteEntry{typ, string(value)})
|
||||||
|
}).decode(b, false)
|
||||||
|
if (err == nil) != (wantErr == nil) {
|
||||||
|
t.Fatalf("%s@%s: error %v, proto.Unmarshal: %v", name, attrs, err, wantErr)
|
||||||
|
}
|
||||||
|
if err == nil && !slices.Equal(got, want) {
|
||||||
|
t.Fatalf("%s@%s: got %v, want %v", name, attrs, got, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestDecodeSiteMatchesUnmarshal(t *testing.T) {
|
||||||
|
bs, err := os.ReadFile(filepath.Join("..", "..", "resources", DefaultGeoSiteDat))
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
for len(bs) > 0 {
|
||||||
|
num, typ, n := protowire.ConsumeTag(bs)
|
||||||
|
if n < 0 || num != 1 || typ != protowire.BytesType {
|
||||||
|
t.Fatal("unexpected GeoSiteList field")
|
||||||
|
}
|
||||||
|
entry, m := protowire.ConsumeBytes(bs[n:])
|
||||||
|
if m < 0 {
|
||||||
|
t.Fatal(protowire.ParseError(m))
|
||||||
|
}
|
||||||
|
bs = bs[n+m:]
|
||||||
|
|
||||||
|
var site GeoSite
|
||||||
|
if err := proto.Unmarshal(entry, &site); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
queries := []string{"", "none"}
|
||||||
|
for _, d := range site.Domain {
|
||||||
|
for _, a := range d.Attribute {
|
||||||
|
if !slices.Contains(queries, a.Key) {
|
||||||
|
queries = append(queries, a.Key, a.Key+"@none")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
for _, attrs := range queries {
|
||||||
|
checkDecodeSite(t, site.Code, entry, attrs)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestDecodeSiteUnusualEncodings(t *testing.T) {
|
||||||
|
field := func(num protowire.Number, v []byte) []byte {
|
||||||
|
return protowire.AppendBytes(protowire.AppendTag(nil, num, protowire.BytesType), v)
|
||||||
|
}
|
||||||
|
typ := func(v Domain_Type) []byte {
|
||||||
|
return protowire.AppendVarint(protowire.AppendTag(nil, 1, protowire.VarintType), uint64(v))
|
||||||
|
}
|
||||||
|
value := func(s string) []byte { return field(2, []byte(s)) }
|
||||||
|
attr := func(keys ...string) []byte {
|
||||||
|
var b []byte
|
||||||
|
for _, k := range keys {
|
||||||
|
b = append(b, field(1, []byte(k))...)
|
||||||
|
}
|
||||||
|
return field(3, b)
|
||||||
|
}
|
||||||
|
domain := func(fields ...[]byte) []byte { return field(2, slices.Concat(fields...)) }
|
||||||
|
unknown := protowire.AppendFixed32(protowire.AppendTag(nil, 9, protowire.Fixed32Type), 1)
|
||||||
|
|
||||||
|
for name, b := range map[string][]byte{
|
||||||
|
"unknown field": domain(typ(Domain_Full), unknown, value("example.com")),
|
||||||
|
"repeated value": domain(value("a.com"), typ(Domain_Full), value("b.com")),
|
||||||
|
"repeated type": domain(typ(Domain_Full), value("a.com"), typ(Domain_Regex)),
|
||||||
|
"repeated key": domain(value("a.com"), attr("cn", "ads")),
|
||||||
|
"type as bytes": domain(field(1, []byte("x")), value("a.com")),
|
||||||
|
"no value": domain(typ(Domain_Domain), attr("cn")),
|
||||||
|
"truncated": domain(typ(Domain_Full), value("example.com"))[:10],
|
||||||
|
"invalid utf8": domain(value("example.\xff")),
|
||||||
|
"invalid key": domain(value("a.com"), attr("\xff")),
|
||||||
|
"bad field": protowire.AppendVarint(protowire.AppendTag(nil, protowire.MaxValidNumber+1, protowire.VarintType), 1),
|
||||||
|
"stray end group": protowire.AppendTag(nil, 5, protowire.EndGroupType),
|
||||||
|
} {
|
||||||
|
for _, attrs := range []string{"", "cn", "ads", "cn@ads"} {
|
||||||
|
checkDecodeSite(t, name, b, attrs)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestLoadSiteReadsInPieces covers what real lists never do: an entry far longer than the read
|
||||||
|
// buffer, with a field longer than the buffer in the middle, and a file cut short.
|
||||||
|
func TestLoadSiteReadsInPieces(t *testing.T) {
|
||||||
|
site := &GeoSite{Code: "BIG"}
|
||||||
|
for i := range 5000 {
|
||||||
|
d := &Domain{Type: Domain_Domain, Value: strings.Repeat("x", i%40) + ".example.com"}
|
||||||
|
if i%3 == 0 {
|
||||||
|
d.Attribute = []*Domain_Attribute{{Key: "cn"}}
|
||||||
|
}
|
||||||
|
if i == 2500 {
|
||||||
|
d = &Domain{Type: Domain_Regex, Value: strings.Repeat("a", 100_000)}
|
||||||
|
}
|
||||||
|
site.Domain = append(site.Domain, d)
|
||||||
|
}
|
||||||
|
list := &GeoSiteList{Entry: []*GeoSite{{Code: "SMALL", Domain: []*Domain{{Type: Domain_Full, Value: "a.com"}}}, site}}
|
||||||
|
bs, err := proto.Marshal(list)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
entry, err := proto.Marshal(site)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
dir := t.TempDir()
|
||||||
|
t.Setenv("xray.location.asset", dir)
|
||||||
|
write := func(b []byte) {
|
||||||
|
if err := os.WriteFile(filepath.Join(dir, "big.dat"), b, 0o644); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
for _, attrs := range []string{"", "cn"} {
|
||||||
|
want, _ := unmarshalSite(entry, attrs)
|
||||||
|
var got []siteEntry
|
||||||
|
write(bs)
|
||||||
|
err := loadSite("big.dat", "BIG", attrs, func(typ Domain_Type, value []byte) {
|
||||||
|
got = append(got, siteEntry{typ, string(value)})
|
||||||
|
})
|
||||||
|
if err != nil || !slices.Equal(got, want) {
|
||||||
|
t.Fatalf("attrs %q: %d entries, want %d, error %v", attrs, len(got), len(want), err)
|
||||||
|
}
|
||||||
|
for _, cut := range []int{30_000, len(bs) - 150_000, len(bs) - 1} {
|
||||||
|
write(bs[:cut])
|
||||||
|
if err := loadSite("big.dat", "BIG", attrs, func(Domain_Type, []byte) {}); err == nil {
|
||||||
|
t.Fatalf("file cut at %d of %d: no error", cut, len(bs))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// oneEntryGeoSiteFile wraps an encoded GeoSite as a one-entry GeoSiteList, the file loadSite reads.
|
||||||
|
func oneEntryGeoSiteFile(entry []byte) []byte {
|
||||||
|
return protowire.AppendBytes(protowire.AppendTag(nil, 1, protowire.BytesType), entry)
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestLoadSiteWindowedMatchesSingleShot checks that the windowed reader in loadSite (its Peek/Discard
|
||||||
|
// loop, the more-break when a field is cut by a window edge, the used==0 fallback for a field longer
|
||||||
|
// than the buffer, and the tail path) reaches exactly the same result as decoding the whole entry at
|
||||||
|
// once, for a category several 64 KiB windows long, valid and then mutated near a window edge and
|
||||||
|
// early in the file: same error-or-not, and the same emitted (type, value) sequence when both accept.
|
||||||
|
func TestLoadSiteWindowedMatchesSingleShot(t *testing.T) {
|
||||||
|
const window = 64 * 1024
|
||||||
|
site := &GeoSite{Code: "BIG"}
|
||||||
|
for i := range 12000 { // ~250 KiB, four windows
|
||||||
|
d := &Domain{Type: Domain_Domain, Value: fmt.Sprintf("host%d.%s.example.com", i, strings.Repeat("y", i%30))}
|
||||||
|
if i%3 == 0 {
|
||||||
|
d.Attribute = []*Domain_Attribute{{Key: "cn"}}
|
||||||
|
}
|
||||||
|
site.Domain = append(site.Domain, d)
|
||||||
|
}
|
||||||
|
// a field longer than the buffer, straddling the third window, to force the used==0 fallback
|
||||||
|
site.Domain = slices.Insert(site.Domain, 8000, &Domain{Type: Domain_Regex, Value: strings.Repeat("a", 90_000)})
|
||||||
|
entry, err := proto.Marshal(site)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
dir := t.TempDir()
|
||||||
|
t.Setenv("xray.location.asset", dir)
|
||||||
|
|
||||||
|
// mutations of the encoded entry: unchanged, a byte flipped at several offsets (early windows and
|
||||||
|
// either side of a window edge), and truncations at the same places.
|
||||||
|
type mut struct {
|
||||||
|
name string
|
||||||
|
make func([]byte) []byte
|
||||||
|
}
|
||||||
|
muts := []mut{{"valid", func(b []byte) []byte { return b }}}
|
||||||
|
for _, off := range []int{3, 40, 4000, window - 1, window, window + 1, 2*window - 2, 2 * window} {
|
||||||
|
if off < len(entry) {
|
||||||
|
off := off
|
||||||
|
muts = append(muts, mut{fmt.Sprintf("flip@%d", off), func(b []byte) []byte {
|
||||||
|
c := slices.Clone(b)
|
||||||
|
c[off] ^= 0xff
|
||||||
|
return c
|
||||||
|
}})
|
||||||
|
muts = append(muts, mut{fmt.Sprintf("cut@%d", off), func(b []byte) []byte { return slices.Clone(b[:off]) }})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, attrs := range []string{"", "cn"} {
|
||||||
|
for _, m := range muts {
|
||||||
|
e := m.make(entry)
|
||||||
|
// single-shot reference: decode the whole entry in one call
|
||||||
|
var want []siteEntry
|
||||||
|
_, wantErr := newSiteDecoder(attrs, func(typ Domain_Type, value []byte) {
|
||||||
|
want = append(want, siteEntry{typ, string(value)})
|
||||||
|
}).decode(e, false)
|
||||||
|
// windowed: loadSite reads the file 64 KiB at a time
|
||||||
|
if err := os.WriteFile(filepath.Join(dir, "w.dat"), oneEntryGeoSiteFile(e), 0o644); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
var got []siteEntry
|
||||||
|
gotErr := loadSite("w.dat", "BIG", attrs, func(typ Domain_Type, value []byte) {
|
||||||
|
got = append(got, siteEntry{typ, string(value)})
|
||||||
|
})
|
||||||
|
if (gotErr == nil) != (wantErr == nil) {
|
||||||
|
t.Fatalf("%s attrs=%q: windowed err %v, single-shot err %v", m.name, attrs, gotErr, wantErr)
|
||||||
|
}
|
||||||
|
if gotErr == nil && !slices.Equal(got, want) {
|
||||||
|
t.Fatalf("%s attrs=%q: windowed got %d entries, single-shot %d", m.name, attrs, len(got), len(want))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestLoadSiteLongCode covers a geosite entry whose code is longer than the 64 KiB read buffer. seek
|
||||||
|
// must skip it (find compares a single length byte, so it never matches such a code) and still find a
|
||||||
|
// later entry, and looking the long code up must fail cleanly, like a missing code, not panic.
|
||||||
|
func TestLoadSiteLongCode(t *testing.T) {
|
||||||
|
longCode := strings.Repeat("Z", 70000)
|
||||||
|
list := &GeoSiteList{Entry: []*GeoSite{
|
||||||
|
{Code: "FIRST", Domain: []*Domain{{Type: Domain_Full, Value: "first.com"}}},
|
||||||
|
{Code: longCode, Domain: []*Domain{{Type: Domain_Full, Value: "huge.com"}}},
|
||||||
|
{Code: "AFTER", Domain: []*Domain{{Type: Domain_Domain, Value: "after.com"}}},
|
||||||
|
}}
|
||||||
|
bs, err := proto.Marshal(list)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
dir := t.TempDir()
|
||||||
|
t.Setenv("xray.location.asset", dir)
|
||||||
|
if err := os.WriteFile(filepath.Join(dir, "lc.dat"), bs, 0o644); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
collect := func(code string) ([]siteEntry, error) {
|
||||||
|
var got []siteEntry
|
||||||
|
err := loadSite("lc.dat", code, "", func(typ Domain_Type, value []byte) {
|
||||||
|
got = append(got, siteEntry{typ, string(value)})
|
||||||
|
})
|
||||||
|
return got, err
|
||||||
|
}
|
||||||
|
if got, err := collect("FIRST"); err != nil || !slices.Equal(got, []siteEntry{{Domain_Full, "first.com"}}) {
|
||||||
|
t.Fatalf("FIRST: %v %v", got, err)
|
||||||
|
}
|
||||||
|
if got, err := collect("AFTER"); err != nil || !slices.Equal(got, []siteEntry{{Domain_Domain, "after.com"}}) {
|
||||||
|
t.Fatalf("AFTER (past the oversized entry): %v %v", got, err)
|
||||||
|
}
|
||||||
|
if _, err := collect(longCode); err == nil {
|
||||||
|
t.Fatal("oversized code: expected a not-found error, got nil")
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,163 @@
|
|||||||
|
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
|
||||||
|
}
|
||||||
@@ -0,0 +1,172 @@
|
|||||||
|
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)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -52,7 +52,9 @@ func (g *MphIndexMatcher) Add(matcher Matcher) uint32 {
|
|||||||
func (g *MphIndexMatcher) Build() error {
|
func (g *MphIndexMatcher) Build() error {
|
||||||
if g.mph != nil {
|
if g.mph != nil {
|
||||||
runtime.GC() // peak mem
|
runtime.GC() // peak mem
|
||||||
g.mph.Build()
|
if err := g.mph.Build(); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
}
|
}
|
||||||
runtime.GC() // peak mem
|
runtime.GC() // peak mem
|
||||||
if g.ac != nil {
|
if g.ac != nil {
|
||||||
@@ -64,23 +66,17 @@ func (g *MphIndexMatcher) Build() error {
|
|||||||
|
|
||||||
// Match implements IndexMatcher.Match.
|
// Match implements IndexMatcher.Match.
|
||||||
func (g *MphIndexMatcher) Match(input string) []uint32 {
|
func (g *MphIndexMatcher) Match(input string) []uint32 {
|
||||||
result := make([][]uint32, 0, 5)
|
var result []uint32
|
||||||
if g.mph != nil {
|
if g.mph != nil {
|
||||||
if matches := g.mph.Match(input); len(matches) > 0 {
|
result = g.mph.Match(input) // a new slice, returned without another copy
|
||||||
result = append(result, matches)
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
if g.ac != nil {
|
if g.ac != nil {
|
||||||
if matches := g.ac.Match(input); len(matches) > 0 {
|
result = append(result, g.ac.Match(input)...)
|
||||||
result = append(result, matches)
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
if g.regex != nil {
|
if g.regex != nil {
|
||||||
if matches := g.regex.Match(input); len(matches) > 0 {
|
result = append(result, g.regex.Match(input)...)
|
||||||
result = append(result, matches)
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
return CompositeMatches(result)
|
return result
|
||||||
}
|
}
|
||||||
|
|
||||||
// MatchAny implements IndexMatcher.MatchAny.
|
// MatchAny implements IndexMatcher.MatchAny.
|
||||||
|
|||||||
@@ -78,6 +78,10 @@ func TestMphIndexMatcher(t *testing.T) {
|
|||||||
Input: "example.com",
|
Input: "example.com",
|
||||||
Output: []uint32{10, 4},
|
Output: []uint32{10, 4},
|
||||||
},
|
},
|
||||||
|
{
|
||||||
|
Input: "apis.org",
|
||||||
|
Output: []uint32{2, 6},
|
||||||
|
},
|
||||||
}
|
}
|
||||||
matcherGroup := NewMphIndexMatcher()
|
matcherGroup := NewMphIndexMatcher()
|
||||||
for _, rule := range rules {
|
for _, rule := range rules {
|
||||||
@@ -87,8 +91,13 @@ func TestMphIndexMatcher(t *testing.T) {
|
|||||||
}
|
}
|
||||||
matcherGroup.Build()
|
matcherGroup.Build()
|
||||||
for _, test := range cases {
|
for _, test := range cases {
|
||||||
if m := matcherGroup.Match(test.Input); !reflect.DeepEqual(m, test.Output) {
|
m := matcherGroup.Match(test.Input)
|
||||||
|
if !reflect.DeepEqual(m, test.Output) {
|
||||||
t.Error("unexpected output: ", m, " for test case ", test)
|
t.Error("unexpected output: ", m, " for test case ", test)
|
||||||
}
|
}
|
||||||
|
clear(m) // the caller owns the result, so this must not change the next one
|
||||||
|
if m := matcherGroup.Match(test.Input); !reflect.DeepEqual(m, test.Output) {
|
||||||
|
t.Error("unexpected output after clearing the previous one: ", m, " for test case ", test)
|
||||||
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1,231 +1,440 @@
|
|||||||
package strmatcher
|
package strmatcher
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"bytes"
|
||||||
|
"cmp"
|
||||||
|
"encoding/binary"
|
||||||
"errors"
|
"errors"
|
||||||
"math"
|
"math"
|
||||||
"math/bits"
|
"slices"
|
||||||
"runtime"
|
|
||||||
"sort"
|
|
||||||
"strings"
|
"strings"
|
||||||
"unsafe"
|
"unsafe"
|
||||||
)
|
)
|
||||||
|
|
||||||
// PrimeRK is the prime base used in Rabin-Karp algorithm.
|
// Flags of a level1 slot, stored above the record offset.
|
||||||
const PrimeRK = 16777619
|
|
||||||
|
|
||||||
// RollingHash calculates the rolling murmurHash of given string based on a provided suffix hash.
|
|
||||||
func RollingHash(hash uint32, input string) uint32 {
|
|
||||||
for i := len(input) - 1; i >= 0; i-- {
|
|
||||||
hash = hash*PrimeRK + uint32(input[i])
|
|
||||||
}
|
|
||||||
return hash
|
|
||||||
}
|
|
||||||
|
|
||||||
// MemHash is the hash function used by go map, it utilizes available hardware instructions(behaves
|
|
||||||
// as aeshash if aes instruction is available).
|
|
||||||
// With different seed, each MemHash<seed> performs as distinct hash functions.
|
|
||||||
func MemHash(seed uint32, input string) uint32 {
|
|
||||||
return uint32(strhash(unsafe.Pointer(&input), uintptr(seed))) // nosemgrep
|
|
||||||
}
|
|
||||||
|
|
||||||
const (
|
const (
|
||||||
mphMatchTypeCount = 2 // Full and Domain
|
mphDomain = 1 << 31 // matches the pattern and its subdomains
|
||||||
|
mphFull = 1 << 30 // matches the pattern only
|
||||||
|
mphParent = 1 << 29 // matches subdomains only, from a pattern with a leading dot
|
||||||
|
mphOffMask = mphParent - 1
|
||||||
)
|
)
|
||||||
|
|
||||||
type mphRuleInfo struct {
|
// Kinds of an added pattern, indexes of mphKinds.
|
||||||
rollingHash uint32
|
const (
|
||||||
matchers [mphMatchTypeCount][]uint32
|
mphKindFull = iota
|
||||||
|
mphKindParent
|
||||||
|
mphKindDomain
|
||||||
|
)
|
||||||
|
|
||||||
|
// mphKinds are the slot flags in the order Match reports their values.
|
||||||
|
var mphKinds = [...]uint32{mphFull, mphParent, mphDomain}
|
||||||
|
|
||||||
|
// mphMultipliers are odd multipliers for the suffix hash. Build moves to the next one if two patterns collide.
|
||||||
|
var mphMultipliers = [...]uint64{0x9e3779b97f4a7c15, 0xc2b2ae3d27d4eb4f, 0x165667b19e3779f9, 0x27d4eb2f165667c5}
|
||||||
|
|
||||||
|
var (
|
||||||
|
errMphCollision = errors.New("strmatcher: suffix hash collision in MphMatcherGroup")
|
||||||
|
errMphBuilt = errors.New("strmatcher: MphMatcherGroup is already built")
|
||||||
|
)
|
||||||
|
|
||||||
|
type mphEntry struct {
|
||||||
|
off uint32 // pattern start in buf
|
||||||
|
value uint32
|
||||||
|
n uint32 // pattern length
|
||||||
|
kind uint8
|
||||||
}
|
}
|
||||||
|
|
||||||
// MphMatcherGroup is an implementation of MatcherGroup.
|
// MphMatcherGroup is an implementation of MatcherGroup for Full and Domain matchers.
|
||||||
// It implements Rabin-Karp algorithm and minimal perfect hash table for Full and Domain matcher.
|
// Each distinct pattern is stored once as a record in arena: its length (255 means a uvarint length follows),
|
||||||
|
// its bytes and, if the group holds more than one distinct value, its values. A minimal perfect hash table
|
||||||
|
// built with hash, displace and compress (http://cmph.sourceforge.net/papers/esa09.pdf) maps a pattern to its
|
||||||
|
// record. Patterns are hashed from the right, so one pass over the input hashes all its parent domains.
|
||||||
type MphMatcherGroup struct {
|
type MphMatcherGroup struct {
|
||||||
patterns string // All rule patterns concatenated
|
arena string
|
||||||
patternOffs []uint32 // RuleIdx -> patterns[patternOffs[i]:patternOffs[i+1]], index 0 reserved for failed lookup
|
level0 []uint16 // bucket -> seed
|
||||||
values []uint32 // All registered matcher values concatenated
|
level1 []uint32 // slot -> flags | record offset
|
||||||
valueOffs []uint32 // RuleIdx -> values[valueOffs[i]:valueOffs[i+1]] (Full Matcher takes precedence)
|
fp []uint8 // slot -> low byte of its pattern's hash, rejects most misses without reading arena
|
||||||
level0 []uint32 // RollingHash & Mask -> seed for Memhash
|
n0, n1 uint32
|
||||||
level0Mask uint32 // Mask restricting RollingHash to 0 ~ len(level0)
|
mul uint64 // multiplier of the suffix hash
|
||||||
level1 []uint32 // Memhash<seed> & Mask -> stored index for rules
|
single uint32 // the only value if !multi
|
||||||
level1Mask uint32 // Mask for restricting Memhash<seed> to 0 ~ len(level1)
|
multi bool
|
||||||
rules []string // RuleIdx -> pattern string, only used for building
|
|
||||||
ruleInfos *map[string]mphRuleInfo
|
buf []byte // build only, patterns in Add order
|
||||||
|
entries []mphEntry
|
||||||
}
|
}
|
||||||
|
|
||||||
func NewMphMatcherGroup() *MphMatcherGroup {
|
func NewMphMatcherGroup() *MphMatcherGroup {
|
||||||
return &MphMatcherGroup{
|
return new(MphMatcherGroup)
|
||||||
rules: []string{""},
|
|
||||||
level0: nil,
|
|
||||||
level0Mask: 0,
|
|
||||||
level1: nil,
|
|
||||||
level1Mask: 0,
|
|
||||||
ruleInfos: &map[string]mphRuleInfo{}, // Only used for building, destroyed after build complete
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// AddFullMatcher implements MatcherGroupForFull.
|
// AddFullMatcher implements MatcherGroupForFull.
|
||||||
func (g *MphMatcherGroup) AddFullMatcher(matcher FullMatcher, value uint32) {
|
func (g *MphMatcherGroup) AddFullMatcher(matcher FullMatcher, value uint32) {
|
||||||
pattern := strings.ToLower(matcher.Pattern())
|
g.add(matcher.Pattern(), mphKindFull, value)
|
||||||
g.addPattern(0, "", pattern, matcher.Type(), value)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// AddDomainMatcher implements MatcherGroupForDomain.
|
// AddDomainMatcher implements MatcherGroupForDomain.
|
||||||
func (g *MphMatcherGroup) AddDomainMatcher(matcher DomainMatcher, value uint32) {
|
func (g *MphMatcherGroup) AddDomainMatcher(matcher DomainMatcher, value uint32) {
|
||||||
pattern := strings.ToLower(matcher.Pattern())
|
g.add(matcher.Pattern(), mphKindDomain, value)
|
||||||
hash := g.addPattern(0, "", pattern, matcher.Type(), value) // For full domain match
|
|
||||||
g.addPattern(hash, pattern, ".", matcher.Type(), value) // For partial domain match
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (g *MphMatcherGroup) addPattern(suffixHash uint32, suffixPattern string, pattern string, matcherType Type, value uint32) uint32 {
|
func (g *MphMatcherGroup) add(pattern string, kind uint8, value uint32) {
|
||||||
fullPattern := pattern + suffixPattern
|
if g.arena != "" {
|
||||||
info, found := (*g.ruleInfos)[fullPattern]
|
panic(errMphBuilt)
|
||||||
if !found {
|
}
|
||||||
info = mphRuleInfo{rollingHash: RollingHash(suffixHash, pattern)}
|
pattern = strings.ToLower(pattern)
|
||||||
g.rules = append(g.rules, fullPattern)
|
off := uint32(len(g.buf))
|
||||||
|
g.buf = append(g.buf, pattern...)
|
||||||
|
g.entries = append(g.entries, mphEntry{off: off, value: value, n: uint32(len(pattern)), kind: kind})
|
||||||
|
if len(pattern) > 0 && pattern[0] == '.' {
|
||||||
|
// ".x" has always matched "*.x" as well, so it also gets a parent-only record for "x"
|
||||||
|
g.entries = append(g.entries, mphEntry{off: off + 1, value: value, n: uint32(len(pattern) - 1), kind: mphKindParent})
|
||||||
}
|
}
|
||||||
info.matchers[matcherType] = append(info.matchers[matcherType], value)
|
|
||||||
(*g.ruleInfos)[fullPattern] = info
|
|
||||||
return info.rollingHash
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// Build builds a minimal perfect hash table for insert rules.
|
func (g *MphMatcherGroup) key(i uint32) []byte {
|
||||||
// Algorithm used: Hash, displace, and compress. See http://cmph.sourceforge.net/papers/esa09.pdf
|
e := &g.entries[i]
|
||||||
|
return g.buf[e.off : e.off+e.n]
|
||||||
|
}
|
||||||
|
|
||||||
|
// Build builds the hash table. It must be called once, after the last Add.
|
||||||
func (g *MphMatcherGroup) Build() error {
|
func (g *MphMatcherGroup) Build() error {
|
||||||
ruleCount := len(*g.ruleInfos)
|
if g.arena != "" {
|
||||||
g.level0 = make([]uint32, nextPow2(ruleCount/4))
|
return errMphBuilt
|
||||||
g.level0Mask = uint32(len(g.level0) - 1)
|
|
||||||
g.level1 = make([]uint32, nextPow2(ruleCount))
|
|
||||||
g.level1Mask = uint32(len(g.level1) - 1)
|
|
||||||
|
|
||||||
// Flatten patterns and values so the built group has no per-rule objects
|
|
||||||
valueCount := 0
|
|
||||||
for _, ruleInfo := range *g.ruleInfos {
|
|
||||||
valueCount += len(ruleInfo.matchers[Full]) + len(ruleInfo.matchers[Domain])
|
|
||||||
}
|
}
|
||||||
g.patterns = strings.Join(g.rules, "")
|
if uint64(len(g.buf)) > math.MaxUint32 {
|
||||||
if uint64(len(g.patterns)) > math.MaxUint32 || uint64(valueCount) > math.MaxUint32 {
|
|
||||||
return errors.New("too many rules for MphMatcherGroup")
|
return errors.New("too many rules for MphMatcherGroup")
|
||||||
}
|
}
|
||||||
g.patternOffs = make([]uint32, len(g.rules)+1)
|
recs := g.writeRecords()
|
||||||
g.values = make([]uint32, 0, valueCount)
|
if len(g.arena) > mphOffMask {
|
||||||
g.valueOffs = make([]uint32, len(g.rules)+1)
|
return errors.New("too many rules for MphMatcherGroup")
|
||||||
|
|
||||||
// Create buckets based on all rule's rolling hash
|
|
||||||
buckets := make([][]uint32, len(g.level0))
|
|
||||||
for ruleIdx := 1; ruleIdx < len(g.rules); ruleIdx++ { // Traverse rules starting from index 1 (0 reserved for failed lookup)
|
|
||||||
ruleInfo := (*g.ruleInfos)[g.rules[ruleIdx]]
|
|
||||||
bucketIdx := ruleInfo.rollingHash & g.level0Mask
|
|
||||||
buckets[bucketIdx] = append(buckets[bucketIdx], uint32(ruleIdx))
|
|
||||||
g.patternOffs[ruleIdx+1] = g.patternOffs[ruleIdx] + uint32(len(g.rules[ruleIdx]))
|
|
||||||
g.values = append(append(g.values, ruleInfo.matchers[Full]...), ruleInfo.matchers[Domain]...)
|
|
||||||
g.valueOffs[ruleIdx+1] = uint32(len(g.values))
|
|
||||||
}
|
}
|
||||||
g.rules = nil
|
hashes := make([]uint64, len(recs))
|
||||||
g.ruleInfos = nil // Set ruleInfos nil to release memory
|
for _, mul := range mphMultipliers {
|
||||||
runtime.GC() // peak mem
|
for i, rec := range recs {
|
||||||
|
hashes[i] = mphMix(mphHash(mul, g.recKey(rec)))
|
||||||
// Sort buckets in descending order with respect to each bucket's size
|
}
|
||||||
bucketIdxs := make([]int, len(buckets))
|
g.mul = mul
|
||||||
for bucketIdx := range buckets {
|
if err := g.place(recs, hashes); err != errMphCollision {
|
||||||
bucketIdxs[bucketIdx] = bucketIdx
|
return err
|
||||||
|
}
|
||||||
}
|
}
|
||||||
sort.Slice(bucketIdxs, func(i, j int) bool { return len(buckets[bucketIdxs[i]]) > len(buckets[bucketIdxs[j]]) })
|
return errMphCollision
|
||||||
|
}
|
||||||
|
|
||||||
// Exercise Hash, Displace, and Compress algorithm to construct minimal perfect hash table
|
// writeRecords writes one record per distinct pattern to arena and returns flags | offset of each.
|
||||||
occupied := make([]bool, len(g.level1)) // Whether a second-level hash has been already used
|
func (g *MphMatcherGroup) writeRecords() []uint32 {
|
||||||
hashedBucket := make([]uint32, 0, 4) // Second-level hashes for each rule in a specific bucket
|
g.multi = false
|
||||||
for _, bucketIdx := range bucketIdxs {
|
if len(g.entries) > 0 {
|
||||||
bucket := buckets[bucketIdx]
|
g.single = g.entries[0].value
|
||||||
hashedBucket = hashedBucket[:0]
|
for _, e := range g.entries {
|
||||||
seed := uint32(0)
|
if e.value != g.single {
|
||||||
for len(hashedBucket) != len(bucket) {
|
g.multi = true
|
||||||
for _, ruleIdx := range bucket {
|
break
|
||||||
memHash := MemHash(seed, g.pattern(ruleIdx)) & g.level1Mask
|
|
||||||
if occupied[memHash] { // Collision occurred with this seed
|
|
||||||
for _, hash := range hashedBucket { // Revert all values in this hashed bucket
|
|
||||||
occupied[hash] = false
|
|
||||||
g.level1[hash] = 0
|
|
||||||
}
|
|
||||||
hashedBucket = hashedBucket[:0]
|
|
||||||
seed++ // Try next seed
|
|
||||||
break
|
|
||||||
}
|
|
||||||
occupied[memHash] = true
|
|
||||||
g.level1[memHash] = ruleIdx // The final value in the hash table
|
|
||||||
hashedBucket = append(hashedBucket, memHash)
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
g.level0[bucketIdx] = seed // Displacement value for this bucket
|
}
|
||||||
|
// Equal patterns become neighbours in Add order, so their values keep their priority
|
||||||
|
order := make([]uint32, len(g.entries))
|
||||||
|
for i := range order {
|
||||||
|
order[i] = uint32(i)
|
||||||
|
}
|
||||||
|
slices.SortFunc(order, func(a, b uint32) int {
|
||||||
|
return cmp.Or(bytes.Compare(g.key(a), g.key(b)), cmp.Compare(a, b))
|
||||||
|
})
|
||||||
|
|
||||||
|
size := len(g.buf) + len(g.entries) + 2
|
||||||
|
if g.multi {
|
||||||
|
size += 3 * len(g.entries)
|
||||||
|
}
|
||||||
|
arena := make([]byte, 0, size)
|
||||||
|
recs := make([]uint32, 0, len(order))
|
||||||
|
var vals [len(mphKinds)][]uint32
|
||||||
|
for i := 0; i < len(order); {
|
||||||
|
k := g.key(order[i])
|
||||||
|
for t := range vals {
|
||||||
|
vals[t] = vals[t][:0]
|
||||||
|
}
|
||||||
|
for ; i < len(order) && bytes.Equal(g.key(order[i]), k); i++ {
|
||||||
|
e := &g.entries[order[i]]
|
||||||
|
if !slices.Contains(vals[e.kind], e.value) {
|
||||||
|
vals[e.kind] = append(vals[e.kind], e.value)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
rec := uint32(len(arena))
|
||||||
|
if len(k) < 255 {
|
||||||
|
arena = append(arena, byte(len(k)))
|
||||||
|
} else {
|
||||||
|
arena = binary.AppendUvarint(append(arena, 255), uint64(len(k)))
|
||||||
|
}
|
||||||
|
arena = append(arena, k...)
|
||||||
|
for t, v := range vals {
|
||||||
|
if len(v) == 0 {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
rec |= mphKinds[t]
|
||||||
|
if g.multi {
|
||||||
|
arena = binary.AppendUvarint(arena, uint64(len(v)))
|
||||||
|
for _, x := range v {
|
||||||
|
arena = binary.AppendUvarint(arena, uint64(x))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
recs = append(recs, rec)
|
||||||
|
}
|
||||||
|
// Lookups may point one byte past a pattern, and an empty group needs a record at offset 0 for empty slots
|
||||||
|
arena = append(arena, 0)
|
||||||
|
if len(recs) == 0 {
|
||||||
|
arena = append(arena, 0)
|
||||||
|
}
|
||||||
|
g.buf, g.entries = nil, nil
|
||||||
|
if cap(arena)-len(arena) > len(arena)/32 {
|
||||||
|
arena = slices.Clone(arena)
|
||||||
|
}
|
||||||
|
g.arena = unsafe.String(unsafe.SliceData(arena), len(arena)) // arena is not written after this
|
||||||
|
return recs
|
||||||
|
}
|
||||||
|
|
||||||
|
// place fills level0, level1 and fp: records are bucketed by hash, and each bucket, largest first, gets
|
||||||
|
// the first seed that puts all its records in free slots.
|
||||||
|
func (g *MphMatcherGroup) place(recs []uint32, hashes []uint64) error {
|
||||||
|
r := len(recs)
|
||||||
|
n0, n1 := max(1, r/3), max(1, r+r/99)
|
||||||
|
g.n0, g.n1 = uint32(n0), uint32(n1)
|
||||||
|
g.level0 = make([]uint16, n0)
|
||||||
|
g.level1 = make([]uint32, n1)
|
||||||
|
g.fp = make([]uint8, n1)
|
||||||
|
|
||||||
|
start := make([]uint32, n0+1)
|
||||||
|
for _, h := range hashes {
|
||||||
|
start[g.bucket(h)+1]++
|
||||||
|
}
|
||||||
|
for b := range n0 {
|
||||||
|
start[b+1] += start[b]
|
||||||
|
}
|
||||||
|
members := make([]uint32, r)
|
||||||
|
fill := slices.Clone(start[:n0])
|
||||||
|
for i, h := range hashes {
|
||||||
|
b := g.bucket(h)
|
||||||
|
members[fill[b]] = uint32(i)
|
||||||
|
fill[b]++
|
||||||
|
}
|
||||||
|
fill = nil
|
||||||
|
buckets := make([]uint32, n0)
|
||||||
|
for b := range buckets {
|
||||||
|
buckets[b] = uint32(b)
|
||||||
|
}
|
||||||
|
slices.SortStableFunc(buckets, func(a, b uint32) int {
|
||||||
|
return cmp.Compare(start[b+1]-start[b], start[a+1]-start[a])
|
||||||
|
})
|
||||||
|
|
||||||
|
occupied := make([]uint64, (n1+63)/64)
|
||||||
|
var slots []uint32
|
||||||
|
next:
|
||||||
|
for _, b := range buckets {
|
||||||
|
m := members[start[b]:start[b+1]]
|
||||||
|
if len(m) == 0 {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
for i := range m {
|
||||||
|
for j := range i {
|
||||||
|
if hashes[m[i]] == hashes[m[j]] {
|
||||||
|
return errMphCollision // no seed can separate them
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
search:
|
||||||
|
for seed := range math.MaxUint16 + 1 {
|
||||||
|
slots = slots[:0]
|
||||||
|
for _, ri := range m {
|
||||||
|
s := g.slot(hashes[ri], uint16(seed))
|
||||||
|
if occupied[s/64]&(1<<(s%64)) != 0 || slices.Contains(slots, s) {
|
||||||
|
continue search
|
||||||
|
}
|
||||||
|
slots = append(slots, s)
|
||||||
|
}
|
||||||
|
for k, ri := range m {
|
||||||
|
s := slots[k]
|
||||||
|
occupied[s/64] |= 1 << (s % 64)
|
||||||
|
g.level1[s] = recs[ri]
|
||||||
|
g.fp[s] = uint8(hashes[ri])
|
||||||
|
}
|
||||||
|
g.level0[b] = uint16(seed)
|
||||||
|
continue next
|
||||||
|
}
|
||||||
|
return errors.New("strmatcher: no seed found for a bucket in MphMatcherGroup")
|
||||||
}
|
}
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (g *MphMatcherGroup) pattern(ruleIdx uint32) string {
|
// mphHash is the suffix hash of s, taken from the right: the hash of s[i:] is the state after reading s[i].
|
||||||
return g.patterns[g.patternOffs[ruleIdx]:g.patternOffs[ruleIdx+1]]
|
func mphHash(mul uint64, s string) uint64 {
|
||||||
|
h := uint64(0)
|
||||||
|
for i := len(s) - 1; i >= 0; i-- {
|
||||||
|
h = h*mul + uint64(s[i])
|
||||||
|
}
|
||||||
|
return h
|
||||||
}
|
}
|
||||||
|
|
||||||
// valuesOf caps the capacity, so appending to a Match result can't overwrite the next rule's values.
|
// mphMix spreads the weak low bits of a suffix hash.
|
||||||
func (g *MphMatcherGroup) valuesOf(ruleIdx uint32) []uint32 {
|
func mphMix(h uint64) uint64 {
|
||||||
start, end := g.valueOffs[ruleIdx], g.valueOffs[ruleIdx+1]
|
h ^= h >> 32
|
||||||
return g.values[start:end:end]
|
h *= 0xd6e8feb86659fd93
|
||||||
|
return h ^ h>>32
|
||||||
}
|
}
|
||||||
|
|
||||||
// Lookup searches for input in minimal perfect hash table and returns its index. 0 indicates not found.
|
func (g *MphMatcherGroup) bucket(f uint64) uint32 {
|
||||||
func (g *MphMatcherGroup) Lookup(rollingHash uint32, input string) uint32 {
|
return uint32(((f >> 32) * uint64(g.n0)) >> 32)
|
||||||
i0 := rollingHash & g.level0Mask
|
}
|
||||||
seed := g.level0[i0]
|
|
||||||
i1 := MemHash(seed, input) & g.level1Mask
|
func (g *MphMatcherGroup) slot(f uint64, seed uint16) uint32 {
|
||||||
n := g.level1[i1]
|
x := ((f ^ uint64(seed)*0x9e3779b97f4a7c15) * 0xc4ceb9fe1a85ec53) >> 32
|
||||||
// Build only puts valid rule indices in level1, so n+1 < len(patternOffs) and the span is inside patterns.
|
return uint32((x * uint64(g.n1)) >> 32)
|
||||||
// Skip the bounds checks, they made this hot path measurably slower than indexing a []string
|
}
|
||||||
offs := (*[2]uint32)(unsafe.Add(unsafe.Pointer(unsafe.SliceData(g.patternOffs)), uintptr(n)*4))
|
|
||||||
if start := offs[0]; int(offs[1]-start) == len(input) && unsafe.String((*byte)(unsafe.Add(unsafe.Pointer(unsafe.StringData(g.patterns)), start)), len(input)) == input {
|
func (g *MphMatcherGroup) uvarint(p uint32) (x, next uint32) {
|
||||||
return n
|
for shift := 0; ; shift += 7 {
|
||||||
|
c := g.arena[p]
|
||||||
|
p++
|
||||||
|
x |= uint32(c&0x7f) << shift
|
||||||
|
if c < 0x80 {
|
||||||
|
return x, p
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// recSpan returns where the pattern of the record at off starts and how long it is.
|
||||||
|
func (g *MphMatcherGroup) recSpan(off uint32) (p, n uint32) {
|
||||||
|
n, p = uint32(g.arena[off]), off+1
|
||||||
|
if n == 255 {
|
||||||
|
n, p = g.uvarint(p)
|
||||||
|
}
|
||||||
|
return p, n
|
||||||
|
}
|
||||||
|
|
||||||
|
func (g *MphMatcherGroup) recKey(rec uint32) string {
|
||||||
|
p, n := g.recSpan(rec & mphOffMask)
|
||||||
|
return g.arena[p : p+n]
|
||||||
|
}
|
||||||
|
|
||||||
|
// lookup returns the level1 entry of s, or 0 if s is not a pattern. h is the suffix hash of s.
|
||||||
|
func (g *MphMatcherGroup) lookup(h uint64, s string) uint32 {
|
||||||
|
f := mphMix(h)
|
||||||
|
// bucket < n0 == len(level0) and slot < n1 == len(level1) == len(fp), skip the bounds checks
|
||||||
|
seed := *(*uint16)(unsafe.Add(unsafe.Pointer(unsafe.SliceData(g.level0)), uintptr(g.bucket(f))*2))
|
||||||
|
slot := uintptr(g.slot(f, seed))
|
||||||
|
if *(*uint8)(unsafe.Add(unsafe.Pointer(unsafe.SliceData(g.fp)), slot)) != uint8(f) {
|
||||||
|
return 0
|
||||||
|
}
|
||||||
|
e := *(*uint32)(unsafe.Add(unsafe.Pointer(unsafe.SliceData(g.level1)), slot*4))
|
||||||
|
if len(s) < 255 {
|
||||||
|
// A record whose length byte is len(s) has len(s) pattern bytes after it
|
||||||
|
p := unsafe.Add(unsafe.Pointer(unsafe.StringData(g.arena)), e&mphOffMask)
|
||||||
|
if int(*(*byte)(p)) == len(s) && unsafe.String((*byte)(unsafe.Add(p, 1)), len(s)) == s {
|
||||||
|
return e
|
||||||
|
}
|
||||||
|
return 0
|
||||||
|
}
|
||||||
|
if g.recKey(e) == s {
|
||||||
|
return e
|
||||||
}
|
}
|
||||||
return 0
|
return 0
|
||||||
}
|
}
|
||||||
|
|
||||||
// Match implements MatcherGroup.Match.
|
// appendValues appends the values of record e for the flags in want, in mphKinds order.
|
||||||
func (g *MphMatcherGroup) Match(input string) []uint32 {
|
func (g *MphMatcherGroup) appendValues(dst []uint32, e, want uint32) []uint32 {
|
||||||
matches := make([][]uint32, 0, 5)
|
if !g.multi {
|
||||||
hash := uint32(0)
|
for _, flag := range mphKinds {
|
||||||
for i := len(input) - 1; i >= 0; i-- {
|
if e&want&flag != 0 {
|
||||||
hash = hash*PrimeRK + uint32(input[i])
|
dst = append(dst, g.single)
|
||||||
if input[i] == '.' {
|
}
|
||||||
if mphIdx := g.Lookup(hash, input[i:]); mphIdx != 0 {
|
}
|
||||||
matches = append(matches, g.valuesOf(mphIdx))
|
return dst
|
||||||
|
}
|
||||||
|
if e&want == 0 {
|
||||||
|
return dst
|
||||||
|
}
|
||||||
|
p, n := g.recSpan(e & mphOffMask)
|
||||||
|
p += n
|
||||||
|
for _, flag := range mphKinds {
|
||||||
|
if e&flag == 0 {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
var count, v uint32
|
||||||
|
for count, p = g.uvarint(p); count > 0; count-- {
|
||||||
|
v, p = g.uvarint(p)
|
||||||
|
if want&flag != 0 {
|
||||||
|
dst = append(dst, v)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
if mphIdx := g.Lookup(hash, input); mphIdx != 0 {
|
return dst
|
||||||
matches = append(matches, g.valuesOf(mphIdx))
|
}
|
||||||
|
|
||||||
|
// Match implements MatcherGroup.Match. Values of an exact match come first (Full, then Domain), then those of
|
||||||
|
// the parent domains, nearest first.
|
||||||
|
func (g *MphMatcherGroup) Match(input string) []uint32 {
|
||||||
|
var stack [8]uint32
|
||||||
|
parents := stack[:0] // TLD side first
|
||||||
|
h, mul := uint64(0), g.mul
|
||||||
|
for i := len(input) - 1; i >= 0; i-- {
|
||||||
|
if input[i] == '.' {
|
||||||
|
if e := g.lookup(h, input[i+1:]); e&(mphDomain|mphParent) != 0 {
|
||||||
|
parents = append(parents, e)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
h = h*mul + uint64(input[i])
|
||||||
}
|
}
|
||||||
return CompositeMatchesReverse(matches)
|
exact := g.lookup(h, input)
|
||||||
|
if exact&(mphFull|mphDomain) == 0 && len(parents) == 0 {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
result := g.appendValues(make([]uint32, 0, len(parents)+1), exact, mphFull|mphDomain)
|
||||||
|
for k := len(parents) - 1; k >= 0; k-- {
|
||||||
|
result = g.appendValues(result, parents[k], mphParent|mphDomain)
|
||||||
|
}
|
||||||
|
return result
|
||||||
}
|
}
|
||||||
|
|
||||||
// MatchAny implements MatcherGroup.MatchAny.
|
// MatchAny implements MatcherGroup.MatchAny.
|
||||||
func (g *MphMatcherGroup) MatchAny(input string) bool {
|
func (g *MphMatcherGroup) MatchAny(input string) bool {
|
||||||
hash := uint32(0)
|
h, mul := uint64(0), g.mul
|
||||||
|
for i := len(input) - 1; i >= 0; i-- {
|
||||||
|
if input[i] == '.' && g.lookup(h, input[i+1:])&(mphDomain|mphParent) != 0 {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
h = h*mul + uint64(input[i])
|
||||||
|
}
|
||||||
|
return g.lookup(h, input)&(mphFull|mphDomain) != 0
|
||||||
|
}
|
||||||
|
|
||||||
|
// mphSuffix is the suffix hash of input[off:], a parent domain of the input.
|
||||||
|
type mphSuffix struct {
|
||||||
|
h uint64
|
||||||
|
off int
|
||||||
|
}
|
||||||
|
|
||||||
|
// mphSuffixes appends the suffix hashes of the parent domains of input to dst, TLD side first, and returns them
|
||||||
|
// with the hash of input itself: what MatchAny computes, computed once for several groups.
|
||||||
|
func mphSuffixes(dst []mphSuffix, mul uint64, input string) ([]mphSuffix, uint64) {
|
||||||
|
h := uint64(0)
|
||||||
for i := len(input) - 1; i >= 0; i-- {
|
for i := len(input) - 1; i >= 0; i-- {
|
||||||
hash = hash*PrimeRK + uint32(input[i])
|
|
||||||
if input[i] == '.' {
|
if input[i] == '.' {
|
||||||
if g.Lookup(hash, input[i:]) != 0 {
|
dst = append(dst, mphSuffix{h, i + 1})
|
||||||
return true
|
}
|
||||||
}
|
h = h*mul + uint64(input[i])
|
||||||
|
}
|
||||||
|
return dst, h
|
||||||
|
}
|
||||||
|
|
||||||
|
// matchAnyHashed is MatchAny with parents and h from mphSuffixes(_, mul, input).
|
||||||
|
func (g *MphMatcherGroup) matchAnyHashed(input string, parents []mphSuffix, h, mul uint64) bool {
|
||||||
|
if g.mul != mul {
|
||||||
|
return g.MatchAny(input) // built with a later multiplier after a collision
|
||||||
|
}
|
||||||
|
for _, p := range parents {
|
||||||
|
if g.lookup(p.h, input[p.off:])&(mphDomain|mphParent) != 0 {
|
||||||
|
return true
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
return g.Lookup(hash, input) != 0
|
return g.lookup(h, input)&(mphFull|mphDomain) != 0
|
||||||
}
|
}
|
||||||
|
|
||||||
func nextPow2(v int) int {
|
|
||||||
if v <= 1 {
|
|
||||||
return 1
|
|
||||||
}
|
|
||||||
const MaxUInt = ^uint(0)
|
|
||||||
n := (MaxUInt >> bits.LeadingZeros(uint(v))) + 1
|
|
||||||
return int(n)
|
|
||||||
}
|
|
||||||
|
|
||||||
//go:noescape
|
|
||||||
//go:linkname strhash runtime.strhash
|
|
||||||
func strhash(p unsafe.Pointer, h uintptr) uintptr
|
|
||||||
|
|||||||
@@ -0,0 +1,108 @@
|
|||||||
|
package strmatcher
|
||||||
|
|
||||||
|
import (
|
||||||
|
"slices"
|
||||||
|
"testing"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestMphMatcherGroupHashCollision(t *testing.T) {
|
||||||
|
saved := mphMultipliers
|
||||||
|
defer func() { mphMultipliers = saved }()
|
||||||
|
|
||||||
|
mphMultipliers[0] = 1 // anagrams collide
|
||||||
|
g := NewMphMatcherGroup()
|
||||||
|
g.AddFullMatcher(FullMatcher("ab.com"), 1)
|
||||||
|
g.AddDomainMatcher(DomainMatcher("ba.com"), 2)
|
||||||
|
g.AddDomainMatcher(DomainMatcher("com"), 3)
|
||||||
|
if err := g.Build(); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if g.mul != saved[1] {
|
||||||
|
t.Errorf("multiplier %#x, want the second one %#x", g.mul, saved[1])
|
||||||
|
}
|
||||||
|
for input, want := range map[string][]uint32{"ab.com": {1, 3}, "x.ba.com": {2, 3}, "x.ab.com": {3}, "ba.com": {2, 3}} {
|
||||||
|
if m := g.Match(input); !slices.Equal(m, want) {
|
||||||
|
t.Errorf("Match(%q) = %v, want %v", input, m, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Thue-Morse strings of 2048 bytes and their complements collide for every odd multiplier
|
||||||
|
mphMultipliers = saved
|
||||||
|
a, b := make([]byte, 2048), make([]byte, 2048)
|
||||||
|
for i := range a {
|
||||||
|
a[i], b[i] = "ab"[bitsOnes(i)%2], "ba"[bitsOnes(i)%2]
|
||||||
|
}
|
||||||
|
g = NewMphMatcherGroup()
|
||||||
|
g.AddFullMatcher(FullMatcher(a), 1)
|
||||||
|
g.AddFullMatcher(FullMatcher(b), 1)
|
||||||
|
if err := g.Build(); err != errMphCollision {
|
||||||
|
t.Errorf("Build() = %v, want %v", err, errMphCollision)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func bitsOnes(i int) int {
|
||||||
|
n := 0
|
||||||
|
for ; i > 0; i &= i - 1 {
|
||||||
|
n++
|
||||||
|
}
|
||||||
|
return n
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMphValueMatcherCombiner(t *testing.T) {
|
||||||
|
build := func(matchers ...Matcher) *MphValueMatcher {
|
||||||
|
m := NewMphValueMatcher()
|
||||||
|
for _, x := range matchers {
|
||||||
|
m.Add(x, 0)
|
||||||
|
}
|
||||||
|
if err := m.Build(); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
return m
|
||||||
|
}
|
||||||
|
regex, err := Regex.New(`^a\d+\.net$`)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
saved := mphMultipliers
|
||||||
|
t.Cleanup(func() { mphMultipliers = saved })
|
||||||
|
mphMultipliers[0] = 1 // anagrams collide, so this one falls back to its own hash pass
|
||||||
|
collided := build(FullMatcher("ab.com"), DomainMatcher("ba.com"))
|
||||||
|
mphMultipliers = saved
|
||||||
|
if collided.mph.mul == mphMultipliers[0] {
|
||||||
|
t.Fatal("collided matcher uses the first multiplier")
|
||||||
|
}
|
||||||
|
matchers := []*MphValueMatcher{
|
||||||
|
build(DomainMatcher("example.com"), FullMatcher("full.org"), DomainMatcher(".dot.io")),
|
||||||
|
collided,
|
||||||
|
build(regex, SubstrMatcher("keyword")),
|
||||||
|
build(),
|
||||||
|
build(DomainMatcher("com"), DomainMatcher("a.b.c.d.e.f.g.h.i.j.k.l.m.n.o.p.q.r.s")),
|
||||||
|
}
|
||||||
|
var s MphValueMatcherCombiner
|
||||||
|
for i, m := range matchers {
|
||||||
|
s.Add(m, uint32(10+i))
|
||||||
|
}
|
||||||
|
inputs := []string{
|
||||||
|
"", ".", "..", "com", "example.com", "www.example.com", "xexample.com", "example.com.", "full.org", "x.full.org",
|
||||||
|
"dot.io", "x.dot.io", ".dot.io", "ab.com", "x.ab.com", "ba.com", "x.ba.com", "a12.net", "a12.net.x", "my-keyword.org",
|
||||||
|
"a.b.c.d.e.f.g.h.i.j.k.l.m.n.o.p.q.r.s", "0.a.b.c.d.e.f.g.h.i.j.k.l.m.n.o.p.q.r.s", "b.c.d.e.f.g.h.i.j.k.l.m.n.o.p.q.r.s",
|
||||||
|
"x.y.z.1.2.3.4.5.6.7.8.9.10.11.12.13.14.15.16.17.ab.com", "x.y.z.1.2.3.4.5.6.7.8.9.10.11.12.13.14.15.16.17.org",
|
||||||
|
}
|
||||||
|
for _, input := range inputs {
|
||||||
|
var want []uint32
|
||||||
|
for i, m := range matchers {
|
||||||
|
if m.MatchAny(input) {
|
||||||
|
want = append(want, uint32(10+i))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if got := s.Match(input); !slices.Equal(got, want) {
|
||||||
|
t.Errorf("Match(%q) = %v, want %v", input, got, want)
|
||||||
|
}
|
||||||
|
if got := s.MatchAny(input); got != (len(want) > 0) {
|
||||||
|
t.Errorf("MatchAny(%q) = %v", input, got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if n := testing.AllocsPerRun(100, func() { s.MatchAny("www.a.b.c.example.org") }); n != 0 {
|
||||||
|
t.Errorf("MatchAny allocates %v times", n)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -4,6 +4,7 @@ import (
|
|||||||
"math/rand"
|
"math/rand"
|
||||||
"reflect"
|
"reflect"
|
||||||
"slices"
|
"slices"
|
||||||
|
"strings"
|
||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
"github.com/xtls/xray-core/common"
|
"github.com/xtls/xray-core/common"
|
||||||
@@ -304,7 +305,7 @@ func TestMphMatcherGroupRandom(t *testing.T) {
|
|||||||
domain["."+p] = append(domain["."+p], value)
|
domain["."+p] = append(domain["."+p], value)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
g.Build()
|
common.Must(g.Build())
|
||||||
for _, input := range inputs {
|
for _, input := range inputs {
|
||||||
keys := []string{input} // Whole input first, then "." suffixes from longest to shortest
|
keys := []string{input} // Whole input first, then "." suffixes from longest to shortest
|
||||||
for i := range len(input) {
|
for i := range len(input) {
|
||||||
@@ -316,7 +317,10 @@ func TestMphMatcherGroupRandom(t *testing.T) {
|
|||||||
for _, k := range keys {
|
for _, k := range keys {
|
||||||
want = append(append(want, full[k]...), domain[k]...)
|
want = append(append(want, full[k]...), domain[k]...)
|
||||||
}
|
}
|
||||||
if m := g.Match(input); !slices.Equal(m, want) {
|
// Compared as sets: Match reports a value once per matching pattern, and orders them differently
|
||||||
|
// from want for patterns and inputs with a leading dot
|
||||||
|
m := g.Match(input)
|
||||||
|
if !slices.Equal(sortedSet(m), sortedSet(want)) {
|
||||||
t.Fatalf("seed %d: Match(%q) = %v, want %v", seed, input, m, want)
|
t.Fatalf("seed %d: Match(%q) = %v, want %v", seed, input, m, want)
|
||||||
}
|
}
|
||||||
if m := g.MatchAny(input); m != (len(want) > 0) {
|
if m := g.MatchAny(input); m != (len(want) > 0) {
|
||||||
@@ -338,3 +342,79 @@ func TestMphMatcherGroupAppend(t *testing.T) {
|
|||||||
t.Error("expect [2], but ", m)
|
t.Error("expect [2], but ", m)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func sortedSet(v []uint32) []uint32 {
|
||||||
|
v = slices.Clone(v)
|
||||||
|
slices.Sort(v)
|
||||||
|
return slices.Compact(v)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMphMatcherGroupLongPattern(t *testing.T) {
|
||||||
|
long := strings.Repeat("a", 300) + ".com"
|
||||||
|
for _, values := range [][4]uint32{{1, 2, 3, 4}, {7, 7, 7, 7}} {
|
||||||
|
g := NewMphMatcherGroup()
|
||||||
|
g.AddDomainMatcher(DomainMatcher(long), values[0])
|
||||||
|
g.AddFullMatcher(FullMatcher("x."+long), values[1])
|
||||||
|
g.AddFullMatcher(FullMatcher(long[:255]), values[2]) // the shortest pattern stored with a long length
|
||||||
|
g.AddFullMatcher(FullMatcher(long[:254]), values[3])
|
||||||
|
common.Must(g.Build())
|
||||||
|
cases := []struct {
|
||||||
|
input string
|
||||||
|
want []uint32
|
||||||
|
}{
|
||||||
|
{long, []uint32{values[0]}},
|
||||||
|
{"www." + long, []uint32{values[0]}},
|
||||||
|
{"x." + long, []uint32{values[1], values[0]}},
|
||||||
|
{long[1:], nil},
|
||||||
|
{"a" + long, nil},
|
||||||
|
{long[:255], []uint32{values[2]}},
|
||||||
|
{long[:254], []uint32{values[3]}},
|
||||||
|
{long[:256], nil},
|
||||||
|
{long[:253], nil},
|
||||||
|
}
|
||||||
|
for _, c := range cases {
|
||||||
|
if m := g.Match(c.input); !slices.Equal(m, c.want) {
|
||||||
|
t.Errorf("Match(%d bytes) = %v, want %v", len(c.input), m, c.want)
|
||||||
|
}
|
||||||
|
if m := g.MatchAny(c.input); m != (c.want != nil) {
|
||||||
|
t.Errorf("MatchAny(%d bytes) = %v", len(c.input), m)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// A pattern longer than 65535 bytes builds and matches: a record's length is a uvarint,
|
||||||
|
// so the only cap was the build-time length field, now widened to uint32.
|
||||||
|
huge := strings.Repeat("a", 70000)
|
||||||
|
g := NewMphMatcherGroup()
|
||||||
|
g.AddFullMatcher(FullMatcher(strings.Repeat("a", 65535)), 1)
|
||||||
|
g.AddDomainMatcher(DomainMatcher(huge+".com"), 2)
|
||||||
|
g.AddFullMatcher(FullMatcher("a.com"), 3)
|
||||||
|
common.Must(g.Build())
|
||||||
|
if !g.MatchAny(strings.Repeat("a", 65535)) || g.MatchAny(strings.Repeat("a", 65534)) {
|
||||||
|
t.Error("wrong answer for a 65535-byte pattern")
|
||||||
|
}
|
||||||
|
if m := g.Match(huge + ".com"); !slices.Equal(m, []uint32{2}) {
|
||||||
|
t.Errorf("Match(%d-byte input) = %v, want [2]", len(huge)+4, m)
|
||||||
|
}
|
||||||
|
if m := g.Match("x." + huge + ".com"); !slices.Equal(m, []uint32{2}) {
|
||||||
|
t.Errorf("Match(subdomain of a %d-byte pattern) = %v, want [2]", len(huge)+4, m)
|
||||||
|
}
|
||||||
|
if g.MatchAny(huge) { // the 70000-byte label on its own is not a rule
|
||||||
|
t.Error("unexpected match for the bare 70000-byte label")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMphMatcherGroupBuildOnce(t *testing.T) {
|
||||||
|
g := NewMphMatcherGroup()
|
||||||
|
g.AddFullMatcher(FullMatcher("a.com"), 1)
|
||||||
|
common.Must(g.Build())
|
||||||
|
if err := g.Build(); err == nil || !g.MatchAny("a.com") {
|
||||||
|
t.Errorf("second Build() = %v, MatchAny(a.com) = %v", err, g.MatchAny("a.com"))
|
||||||
|
}
|
||||||
|
defer func() {
|
||||||
|
if recover() == nil {
|
||||||
|
t.Error("Add after Build did not panic")
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
g.AddDomainMatcher(DomainMatcher("b.com"), 2)
|
||||||
|
}
|
||||||
|
|||||||
@@ -2,10 +2,12 @@ package strmatcher
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"errors"
|
"errors"
|
||||||
|
"math/bits"
|
||||||
"regexp"
|
"regexp"
|
||||||
"regexp/syntax"
|
"regexp/syntax"
|
||||||
"slices"
|
"slices"
|
||||||
"strings"
|
"strings"
|
||||||
|
"unicode"
|
||||||
"unicode/utf8"
|
"unicode/utf8"
|
||||||
|
|
||||||
"golang.org/x/net/idna"
|
"golang.org/x/net/idna"
|
||||||
@@ -75,7 +77,9 @@ func (m SubstrMatcher) Match(s string) bool {
|
|||||||
// RegexMatcher is an implementation of Matcher.
|
// RegexMatcher is an implementation of Matcher.
|
||||||
type RegexMatcher struct {
|
type RegexMatcher struct {
|
||||||
pattern *regexp.Regexp
|
pattern *regexp.Regexp
|
||||||
literals []string // every match contains all of them, longest first
|
literals []string // every match contains all of them, longest first
|
||||||
|
tail []byteSet // tail[i] holds the bytes a matching input can have i bytes before its end
|
||||||
|
rest *byteSet // the bytes it can have further before, nil if any
|
||||||
}
|
}
|
||||||
|
|
||||||
func newRegexMatcher(pattern string) (Matcher, error) {
|
func newRegexMatcher(pattern string) (Matcher, error) {
|
||||||
@@ -87,10 +91,239 @@ func newRegexMatcher(pattern string) (Matcher, error) {
|
|||||||
if re, err := syntax.Parse(pattern, syntax.Perl); err == nil { // same flags as regexp.Compile
|
if re, err := syntax.Parse(pattern, syntax.Perl); err == nil { // same flags as regexp.Compile
|
||||||
m.literals = requiredLiterals(re, nil)
|
m.literals = requiredLiterals(re, nil)
|
||||||
slices.SortStableFunc(m.literals, func(a, b string) int { return len(b) - len(a) })
|
slices.SortStableFunc(m.literals, func(a, b string) int { return len(b) - len(a) })
|
||||||
|
m.tail, m.rest = tailGuard(re)
|
||||||
}
|
}
|
||||||
return m, nil
|
return m, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// byteSet is a set of bytes. The bytes >= 0x80 share one bit with 0x7f.
|
||||||
|
type byteSet [4]uint32
|
||||||
|
|
||||||
|
func (s *byteSet) add(c byte) { c = min(c, 0x7f); s[c>>5] |= 1 << (c & 31) }
|
||||||
|
func (s *byteSet) has(c byte) bool { c = min(c, 0x7f); return s[c>>5]&(1<<(c&31)) != 0 }
|
||||||
|
func (s *byteSet) or(t *byteSet) {
|
||||||
|
for i := range s {
|
||||||
|
s[i] |= t[i]
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
var allBytes = byteSet{^uint32(0), ^uint32(0), ^uint32(0), ^uint32(0)}
|
||||||
|
|
||||||
|
// tailLen is how many positions before the end of the input tailGuard tells apart.
|
||||||
|
const tailLen = 8
|
||||||
|
|
||||||
|
// tailBudget caps how many repetition steps tailGuard walks. Only nested repeats can make the
|
||||||
|
// walk explode, so only they are charged: a flat pattern, however long, is walked once and keeps
|
||||||
|
// its guard.
|
||||||
|
const tailBudget = 100000
|
||||||
|
|
||||||
|
// tailWalk is a set of positions in the input, counted in bytes before its end.
|
||||||
|
type tailWalk struct {
|
||||||
|
at uint32 // bit i: exactly i bytes before the end, for i < tailLen
|
||||||
|
far bool // tailLen or more bytes before the end
|
||||||
|
free bool // not tied to the end of the input yet
|
||||||
|
}
|
||||||
|
|
||||||
|
func (w tailWalk) union(v tailWalk) tailWalk {
|
||||||
|
return tailWalk{w.at | v.at, w.far || v.far, w.free || v.free}
|
||||||
|
}
|
||||||
|
|
||||||
|
type tailBuilder struct {
|
||||||
|
tail [tailLen]byteSet
|
||||||
|
rest byteSet
|
||||||
|
void bool
|
||||||
|
work int
|
||||||
|
}
|
||||||
|
|
||||||
|
// tailGuard walks re backwards from the end of the input and collects the bytes an input
|
||||||
|
// matching re can have at each position before its end. It returns nil, nil when a branch
|
||||||
|
// of re does not end with $ or when nested repeats push the walk past tailBudget.
|
||||||
|
func tailGuard(re *syntax.Regexp) ([]byteSet, *byteSet) {
|
||||||
|
var b tailBuilder
|
||||||
|
w := b.walk(re, tailWalk{free: true})
|
||||||
|
b.stop(w)
|
||||||
|
if b.void {
|
||||||
|
return nil, nil
|
||||||
|
}
|
||||||
|
if w.at != 0 { // a match can start here, so any bytes can come before
|
||||||
|
for i := bits.TrailingZeros32(w.at); i < tailLen; i++ {
|
||||||
|
b.tail[i] = allBytes
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if w.at != 0 || w.far {
|
||||||
|
b.rest = allBytes
|
||||||
|
}
|
||||||
|
n := tailLen
|
||||||
|
for n > 0 && b.tail[n-1] == b.rest {
|
||||||
|
n--
|
||||||
|
}
|
||||||
|
var tail []byteSet
|
||||||
|
if n > 0 {
|
||||||
|
tail = slices.Clone(b.tail[:n])
|
||||||
|
}
|
||||||
|
if b.rest != allBytes {
|
||||||
|
rest := b.rest
|
||||||
|
return tail, &rest
|
||||||
|
}
|
||||||
|
return tail, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// stop ends the paths of w. One that never met $ lets its match be followed by anything.
|
||||||
|
func (b *tailBuilder) stop(w tailWalk) {
|
||||||
|
if w.free {
|
||||||
|
b.void = true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (b *tailBuilder) walk(re *syntax.Regexp, w tailWalk) tailWalk {
|
||||||
|
if w == (tailWalk{}) || b.void {
|
||||||
|
return w
|
||||||
|
}
|
||||||
|
switch re.Op {
|
||||||
|
case syntax.OpNoMatch:
|
||||||
|
return tailWalk{}
|
||||||
|
case syntax.OpLiteral:
|
||||||
|
for i := len(re.Rune) - 1; i >= 0; i-- {
|
||||||
|
var set byteSet
|
||||||
|
set.add(byte(min(re.Rune[i], utf8.RuneSelf)))
|
||||||
|
if re.Flags&syntax.FoldCase != 0 {
|
||||||
|
for f := unicode.SimpleFold(re.Rune[i]); f != re.Rune[i]; f = unicode.SimpleFold(f) {
|
||||||
|
set.add(byte(min(f, utf8.RuneSelf)))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
w = b.step(w, &set)
|
||||||
|
}
|
||||||
|
return w
|
||||||
|
case syntax.OpCharClass:
|
||||||
|
var set byteSet
|
||||||
|
for i := 0; i+1 < len(re.Rune); i += 2 {
|
||||||
|
for r := min(re.Rune[i], utf8.RuneSelf); r <= min(re.Rune[i+1], utf8.RuneSelf); r++ {
|
||||||
|
set.add(byte(r))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return b.step(w, &set)
|
||||||
|
case syntax.OpAnyChar, syntax.OpAnyCharNotNL: // a domain has no \n to reject
|
||||||
|
return b.step(w, &allBytes)
|
||||||
|
case syntax.OpBeginText: // nothing comes before
|
||||||
|
b.stop(w)
|
||||||
|
return tailWalk{}
|
||||||
|
case syntax.OpEndText:
|
||||||
|
out := tailWalk{at: w.at & 1}
|
||||||
|
if w.free {
|
||||||
|
out.at = 1
|
||||||
|
}
|
||||||
|
return out
|
||||||
|
case syntax.OpCapture:
|
||||||
|
return b.walk(re.Sub[0], w)
|
||||||
|
case syntax.OpConcat:
|
||||||
|
for i := len(re.Sub) - 1; i >= 0; i-- {
|
||||||
|
w = b.walk(re.Sub[i], w)
|
||||||
|
}
|
||||||
|
return w
|
||||||
|
case syntax.OpAlternate:
|
||||||
|
var out tailWalk
|
||||||
|
for _, sub := range re.Sub {
|
||||||
|
out = out.union(b.walk(sub, w))
|
||||||
|
}
|
||||||
|
return out
|
||||||
|
case syntax.OpQuest:
|
||||||
|
return b.repeat(re.Sub[0], w, 1)
|
||||||
|
case syntax.OpStar:
|
||||||
|
return b.repeat(re.Sub[0], w, -1)
|
||||||
|
case syntax.OpPlus:
|
||||||
|
return b.repeat(re.Sub[0], b.walk(re.Sub[0], w), -1)
|
||||||
|
case syntax.OpRepeat:
|
||||||
|
for i := 0; i < re.Min; i++ {
|
||||||
|
if b.charge() {
|
||||||
|
return w
|
||||||
|
}
|
||||||
|
w = b.walk(re.Sub[0], w)
|
||||||
|
}
|
||||||
|
if re.Max < 0 {
|
||||||
|
return b.repeat(re.Sub[0], w, -1)
|
||||||
|
}
|
||||||
|
return b.repeat(re.Sub[0], w, re.Max-re.Min)
|
||||||
|
}
|
||||||
|
return w // empty match, line and word boundaries: no constraint
|
||||||
|
}
|
||||||
|
|
||||||
|
// charge counts one repetition step and reports whether the walk has run out of budget. Only
|
||||||
|
// repeats re-walk their body, so charging them alone bounds the blow-up of nested repeats while
|
||||||
|
// leaving a single linear pass, of any length, free.
|
||||||
|
func (b *tailBuilder) charge() bool {
|
||||||
|
b.work++
|
||||||
|
if b.work > tailBudget {
|
||||||
|
b.void = true
|
||||||
|
}
|
||||||
|
return b.void
|
||||||
|
}
|
||||||
|
|
||||||
|
// repeat walks back over up to n more repetitions of re, any number if n < 0.
|
||||||
|
func (b *tailBuilder) repeat(re *syntax.Regexp, w tailWalk, n int) tailWalk {
|
||||||
|
for ; n != 0; n-- {
|
||||||
|
if b.charge() {
|
||||||
|
return w
|
||||||
|
}
|
||||||
|
next := w.union(b.walk(re, w))
|
||||||
|
if next == w {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
w = next
|
||||||
|
}
|
||||||
|
return w
|
||||||
|
}
|
||||||
|
|
||||||
|
// step walks back over one character whose last byte is in set. A character that can be
|
||||||
|
// non-ASCII can take up to 4 bytes, all >= 0x80; regexp matches an invalid byte as U+FFFD.
|
||||||
|
func (b *tailBuilder) step(w tailWalk, set *byteSet) tailWalk {
|
||||||
|
out := tailWalk{far: w.far, free: w.free}
|
||||||
|
if w.far {
|
||||||
|
b.rest.or(set)
|
||||||
|
}
|
||||||
|
width := 1
|
||||||
|
if set.has(0x80) {
|
||||||
|
width = utf8.UTFMax
|
||||||
|
}
|
||||||
|
for i := 0; i < tailLen; i++ {
|
||||||
|
if w.at&(1<<i) == 0 {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
b.tail[i].or(set)
|
||||||
|
for n := 1; n <= width; n++ {
|
||||||
|
if j := i + n; j < tailLen {
|
||||||
|
out.at |= 1 << j
|
||||||
|
if n < width {
|
||||||
|
b.tail[j].add(0x80)
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
out.far = true
|
||||||
|
if n < width {
|
||||||
|
b.rest.add(0x80)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return out
|
||||||
|
}
|
||||||
|
|
||||||
|
// mayMatch reports whether s passes the tail guard.
|
||||||
|
func (m *RegexMatcher) mayMatch(s string) bool {
|
||||||
|
n := len(s)
|
||||||
|
if m.rest == nil {
|
||||||
|
n = min(n, len(m.tail))
|
||||||
|
}
|
||||||
|
for i := 0; i < n; i++ {
|
||||||
|
set := m.rest
|
||||||
|
if i < len(m.tail) {
|
||||||
|
set = &m.tail[i]
|
||||||
|
}
|
||||||
|
if !set.has(s[len(s)-1-i]) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
// requiredLiterals appends to dst the case-sensitive strings that every match of re contains.
|
// requiredLiterals appends to dst the case-sensitive strings that every match of re contains.
|
||||||
func requiredLiterals(re *syntax.Regexp, dst []string) []string {
|
func requiredLiterals(re *syntax.Regexp, dst []string) []string {
|
||||||
switch re.Op {
|
switch re.Op {
|
||||||
@@ -126,6 +359,9 @@ func (m *RegexMatcher) String() string {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (m *RegexMatcher) Match(s string) bool {
|
func (m *RegexMatcher) Match(s string) bool {
|
||||||
|
if !m.mayMatch(s) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
for _, l := range m.literals {
|
for _, l := range m.literals {
|
||||||
if !strings.Contains(s, l) {
|
if !strings.Contains(s, l) {
|
||||||
return false
|
return false
|
||||||
|
|||||||
@@ -1,9 +1,16 @@
|
|||||||
package strmatcher
|
package strmatcher
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"hash/fnv"
|
||||||
|
"math/rand/v2"
|
||||||
"regexp"
|
"regexp"
|
||||||
|
"regexp/syntax"
|
||||||
"slices"
|
"slices"
|
||||||
|
"strconv"
|
||||||
|
"strings"
|
||||||
"testing"
|
"testing"
|
||||||
|
"unicode"
|
||||||
|
"unicode/utf8"
|
||||||
)
|
)
|
||||||
|
|
||||||
var regexLiteralCases = []struct {
|
var regexLiteralCases = []struct {
|
||||||
@@ -37,6 +44,147 @@ func TestRegexRequiredLiterals(t *testing.T) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
var regexTailCases = []struct {
|
||||||
|
pattern string
|
||||||
|
guard bool
|
||||||
|
match []string // inputs the pattern matches
|
||||||
|
reject []string // inputs the tail guard alone rejects
|
||||||
|
}{
|
||||||
|
{`^[a-z]([a-z0-9-]{0,61}[a-z0-9])?$`, true, []string{"a", "localhost", "x-1"}, []string{"www.example.com", "localhost.", "LOCALHOST", "a b"}},
|
||||||
|
{`(^|\.)[a-z][1-9][0-9][a-z]\.com$`, true, []string{"a12b.com", "x.q10z.com"}, []string{"google.com", "a12b.co", "a12b.com.", "ab12.com"}},
|
||||||
|
{`^hses[1-7]?\.akamaized\.net$`, true, []string{"hses.akamaized.net", "hses3.akamaized.net"}, []string{"xhses.akamaized.net", "www.hses.akamaized.net"}},
|
||||||
|
{`(?i)k\.net$`, true, []string{"k.net", "K.NET", "\u212a.net"}, []string{"x.net", "k.nex"}},
|
||||||
|
{`[^.]+\.cn$`, true, []string{"a.cn", "\xff.cn", "\u4e2d.cn"}, []string{"a.cnn", "a.c"}},
|
||||||
|
{`\x{FFFD}$`, true, []string{"\xff", "a\xc3", "\uFFFD"}, []string{"a", "\xff."}},
|
||||||
|
{`^.\.cn$`, true, []string{"a.cn", "\u4E2D.cn", "\xff.cn"}, []string{"ab.cn"}},
|
||||||
|
{`^$`, true, []string{""}, []string{"a"}},
|
||||||
|
{`(^|\.)youyuapi\..+$`, false, []string{"youyuapi.com"}, nil},
|
||||||
|
{`abc`, false, []string{"abc", "xabcx"}, nil},
|
||||||
|
{`^ab`, false, []string{"ab", "abc"}, nil},
|
||||||
|
{`a$|b`, false, []string{"a", "bx"}, nil},
|
||||||
|
{`(?m)a$`, false, []string{"a", "a\nb"}, nil},
|
||||||
|
{strings.Repeat(`(?:abcdefgh(?:a`, 20) + strings.Repeat(`)*)*`, 20) + `\.com$`, false, []string{".com", "abcdefgha.com"}, nil}, // over tailBudget
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRegexTailGuard(t *testing.T) {
|
||||||
|
for _, test := range regexTailCases {
|
||||||
|
m, err := newRegexMatcher(test.pattern)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
rm := m.(*RegexMatcher)
|
||||||
|
if guard := rm.tail != nil || rm.rest != nil; guard != test.guard {
|
||||||
|
t.Errorf("%s: guard %v, want %v", test.pattern, guard, test.guard)
|
||||||
|
}
|
||||||
|
for _, s := range test.match {
|
||||||
|
if !rm.pattern.MatchString(s) || !rm.Match(s) {
|
||||||
|
t.Errorf("%s: %q does not match", test.pattern, s)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
for _, s := range test.reject {
|
||||||
|
if rm.pattern.MatchString(s) || rm.mayMatch(s) {
|
||||||
|
t.Errorf("%s: %q passes the guard", test.pattern, s)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestRegexTailGuardFlatAlternation checks that a long but non-recursive pattern keeps its
|
||||||
|
// guard. Only nested repeats are charged against tailBudget, so a flat alternation of many
|
||||||
|
// names, however large, is walked once and guarded; its guard is checked against regexp.
|
||||||
|
func TestRegexTailGuardFlatAlternation(t *testing.T) {
|
||||||
|
var sb strings.Builder
|
||||||
|
sb.WriteString("(?:")
|
||||||
|
for i := 0; i < 20000; i++ {
|
||||||
|
if i > 0 {
|
||||||
|
sb.WriteByte('|')
|
||||||
|
}
|
||||||
|
sb.WriteString("name")
|
||||||
|
sb.WriteString(strconv.Itoa(i))
|
||||||
|
}
|
||||||
|
sb.WriteString(`)\.example\.com$`)
|
||||||
|
m, err := newRegexMatcher(sb.String())
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
rm := m.(*RegexMatcher)
|
||||||
|
if rm.tail == nil && rm.rest == nil {
|
||||||
|
t.Fatal("flat alternation of 20000 names lost its guard")
|
||||||
|
}
|
||||||
|
for _, s := range []string{"name0.example.com", "name19999.example.com", "x.name12345.example.com"} {
|
||||||
|
if !rm.pattern.MatchString(s) || !rm.Match(s) {
|
||||||
|
t.Errorf("%q should match", s)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
for _, s := range []string{"name0.example.org", "name0.example.com.", "name0.example.con", "google.com"} {
|
||||||
|
if rm.pattern.MatchString(s) {
|
||||||
|
t.Fatalf("test bug: %q matches the pattern", s)
|
||||||
|
}
|
||||||
|
if rm.mayMatch(s) {
|
||||||
|
t.Errorf("%q should be rejected by the guard", s)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// sampleMatch appends a string that re matches, assertions aside, unless it runs out of
|
||||||
|
// budget, which it spends one per call so that nested repeats stay cheap.
|
||||||
|
func sampleMatch(sb *strings.Builder, re *syntax.Regexp, rnd *rand.Rand, budget *int) {
|
||||||
|
if *budget <= 0 {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
*budget--
|
||||||
|
switch re.Op {
|
||||||
|
case syntax.OpLiteral:
|
||||||
|
for _, r := range re.Rune {
|
||||||
|
if re.Flags&syntax.FoldCase != 0 {
|
||||||
|
for n := rnd.IntN(4); n > 0; n-- {
|
||||||
|
r = unicode.SimpleFold(r)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
sampleRune(sb, r, rnd)
|
||||||
|
}
|
||||||
|
case syntax.OpCharClass:
|
||||||
|
if len(re.Rune) > 0 {
|
||||||
|
i := rnd.IntN(len(re.Rune)/2) * 2
|
||||||
|
sampleRune(sb, re.Rune[i]+rnd.Int32N(min(re.Rune[i+1]-re.Rune[i]+1, 300)), rnd)
|
||||||
|
}
|
||||||
|
case syntax.OpAnyChar, syntax.OpAnyCharNotNL:
|
||||||
|
sampleRune(sb, []rune{'a', '.', '\n', 0xe9, 0x212a, utf8.RuneError}[rnd.IntN(6)], rnd)
|
||||||
|
case syntax.OpCapture:
|
||||||
|
sampleMatch(sb, re.Sub[0], rnd, budget)
|
||||||
|
case syntax.OpConcat:
|
||||||
|
for _, sub := range re.Sub {
|
||||||
|
sampleMatch(sb, sub, rnd, budget)
|
||||||
|
}
|
||||||
|
case syntax.OpAlternate:
|
||||||
|
sampleMatch(sb, re.Sub[rnd.IntN(len(re.Sub))], rnd, budget)
|
||||||
|
case syntax.OpQuest, syntax.OpStar, syntax.OpPlus, syntax.OpRepeat:
|
||||||
|
lo, hi := 0, 3
|
||||||
|
switch re.Op {
|
||||||
|
case syntax.OpQuest:
|
||||||
|
hi = 1
|
||||||
|
case syntax.OpPlus:
|
||||||
|
lo = 1
|
||||||
|
case syntax.OpRepeat:
|
||||||
|
lo, hi = re.Min, re.Min+3
|
||||||
|
if re.Max >= 0 {
|
||||||
|
hi = min(hi, re.Max)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
for n := lo + rnd.IntN(hi-lo+1); n > 0; n-- {
|
||||||
|
sampleMatch(sb, re.Sub[0], rnd, budget)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func sampleRune(sb *strings.Builder, r rune, rnd *rand.Rand) {
|
||||||
|
if r == utf8.RuneError && rnd.IntN(2) == 0 {
|
||||||
|
sb.WriteByte(0x80 | byte(rnd.IntN(0x80))) // regexp matches an invalid byte as U+FFFD
|
||||||
|
return
|
||||||
|
}
|
||||||
|
sb.WriteRune(r)
|
||||||
|
}
|
||||||
|
|
||||||
func FuzzRegexMatcher(f *testing.F) {
|
func FuzzRegexMatcher(f *testing.F) {
|
||||||
inputs := []string{
|
inputs := []string{
|
||||||
"", "x", "yy", "abd", "ccc", "ABC", "abCDef", "abcdef", "abababcc", "a.b", "a\xffb", "a\uFFFDb",
|
"", "x", "yy", "abd", "ccc", "ABC", "abCDef", "abcdef", "abababcc", "a.b", "a\xffb", "a\uFFFDb",
|
||||||
@@ -47,14 +195,39 @@ func FuzzRegexMatcher(f *testing.F) {
|
|||||||
f.Add(test.pattern, s)
|
f.Add(test.pattern, s)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
for _, test := range regexTailCases {
|
||||||
|
for _, s := range append(test.match, test.reject...) {
|
||||||
|
f.Add(test.pattern, s)
|
||||||
|
}
|
||||||
|
}
|
||||||
f.Fuzz(func(t *testing.T, pattern, s string) {
|
f.Fuzz(func(t *testing.T, pattern, s string) {
|
||||||
re, err := regexp.Compile(pattern)
|
re, err := regexp.Compile(pattern)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
m, _ := newRegexMatcher(pattern)
|
m, _ := newRegexMatcher(pattern)
|
||||||
if got, want := m.Match(s), re.MatchString(s); got != want {
|
check := func(s string) {
|
||||||
t.Errorf("pattern %q, input %q: got %v, want %v", pattern, s, got, want)
|
if got, want := m.Match(s), re.MatchString(s); got != want {
|
||||||
|
t.Errorf("pattern %q, input %q: got %v, want %v", pattern, s, got, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
check(s)
|
||||||
|
// random inputs seldom match, so also try strings built from the pattern
|
||||||
|
parsed, _ := syntax.Parse(pattern, syntax.Perl)
|
||||||
|
h := fnv.New64a()
|
||||||
|
h.Write([]byte(s))
|
||||||
|
rnd := rand.New(rand.NewPCG(h.Sum64(), 1))
|
||||||
|
for range 8 {
|
||||||
|
var sb strings.Builder
|
||||||
|
budget := 256
|
||||||
|
sampleMatch(&sb, parsed, rnd, &budget)
|
||||||
|
sample := sb.String()
|
||||||
|
check(sample)
|
||||||
|
check(s + sample)
|
||||||
|
if len(sample) > 0 && len(s) > 0 {
|
||||||
|
i := rnd.IntN(len(sample))
|
||||||
|
check(sample[:i] + s[:1] + sample[i+1:])
|
||||||
|
}
|
||||||
}
|
}
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -46,7 +46,9 @@ func (g *MphValueMatcher) Add(matcher Matcher, value uint32) {
|
|||||||
func (g *MphValueMatcher) Build() error {
|
func (g *MphValueMatcher) Build() error {
|
||||||
if g.mph != nil {
|
if g.mph != nil {
|
||||||
runtime.GC() // peak mem
|
runtime.GC() // peak mem
|
||||||
g.mph.Build()
|
if err := g.mph.Build(); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
}
|
}
|
||||||
runtime.GC() // peak mem
|
runtime.GC() // peak mem
|
||||||
if g.ac != nil {
|
if g.ac != nil {
|
||||||
@@ -58,23 +60,17 @@ func (g *MphValueMatcher) Build() error {
|
|||||||
|
|
||||||
// Match implements ValueMatcher.Match.
|
// Match implements ValueMatcher.Match.
|
||||||
func (g *MphValueMatcher) Match(input string) []uint32 {
|
func (g *MphValueMatcher) Match(input string) []uint32 {
|
||||||
result := make([][]uint32, 0, 5)
|
var result []uint32
|
||||||
if g.mph != nil {
|
if g.mph != nil {
|
||||||
if matches := g.mph.Match(input); len(matches) > 0 {
|
result = g.mph.Match(input) // a new slice, returned without another copy
|
||||||
result = append(result, matches)
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
if g.ac != nil {
|
if g.ac != nil {
|
||||||
if matches := g.ac.Match(input); len(matches) > 0 {
|
result = append(result, g.ac.Match(input)...)
|
||||||
result = append(result, matches)
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
if g.regex != nil {
|
if g.regex != nil {
|
||||||
if matches := g.regex.Match(input); len(matches) > 0 {
|
result = append(result, g.regex.Match(input)...)
|
||||||
result = append(result, matches)
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
return CompositeMatches(result)
|
return result
|
||||||
}
|
}
|
||||||
|
|
||||||
// MatchAny implements ValueMatcher.MatchAny.
|
// MatchAny implements ValueMatcher.MatchAny.
|
||||||
@@ -87,3 +83,62 @@ func (g *MphValueMatcher) MatchAny(input string) bool {
|
|||||||
}
|
}
|
||||||
return g.regex != nil && g.regex.MatchAny(input)
|
return g.regex != nil && g.regex.MatchAny(input)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (g *MphValueMatcher) matchAnyHashed(input string, parents []mphSuffix, h, mul uint64) bool {
|
||||||
|
if g.mph != nil && g.mph.matchAnyHashed(input, parents, h, mul) {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
if g.ac != nil && g.ac.MatchAny(input) {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
return g.regex != nil && g.regex.MatchAny(input)
|
||||||
|
}
|
||||||
|
|
||||||
|
// MphValueMatcherCombiner combines several built MphValueMatchers, each bound to one value, and matches an input
|
||||||
|
// against them as their MatchAny would, hashing the input once for all of them.
|
||||||
|
type MphValueMatcherCombiner struct {
|
||||||
|
matchers []*MphValueMatcher
|
||||||
|
values []uint32
|
||||||
|
}
|
||||||
|
|
||||||
|
// Add adds a built matcher that stands for value.
|
||||||
|
func (s *MphValueMatcherCombiner) Add(m *MphValueMatcher, value uint32) {
|
||||||
|
s.matchers = append(s.matchers, m)
|
||||||
|
s.values = append(s.values, value)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Match returns the values of the matchers that match input, in Add order.
|
||||||
|
func (s *MphValueMatcherCombiner) Match(input string) []uint32 {
|
||||||
|
if len(s.matchers) == 0 {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
var stack [16]mphSuffix
|
||||||
|
mul := mphMultipliers[0]
|
||||||
|
parents, h := mphSuffixes(stack[:0], mul, input)
|
||||||
|
var result []uint32
|
||||||
|
for i, m := range s.matchers {
|
||||||
|
if m.matchAnyHashed(input, parents, h, mul) {
|
||||||
|
result = append(result, s.values[i])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return result
|
||||||
|
}
|
||||||
|
|
||||||
|
// MatchAny returns true as soon as one matcher matches input.
|
||||||
|
func (s *MphValueMatcherCombiner) MatchAny(input string) bool {
|
||||||
|
switch len(s.matchers) {
|
||||||
|
case 0:
|
||||||
|
return false
|
||||||
|
case 1:
|
||||||
|
return s.matchers[0].MatchAny(input) // nothing to share, and it stops at the first matching suffix
|
||||||
|
}
|
||||||
|
var stack [16]mphSuffix
|
||||||
|
mul := mphMultipliers[0]
|
||||||
|
parents, h := mphSuffixes(stack[:0], mul, input)
|
||||||
|
for _, m := range s.matchers {
|
||||||
|
if m.matchAnyHashed(input, parents, h, mul) {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|||||||
@@ -0,0 +1,61 @@
|
|||||||
|
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()
|
||||||
|
}
|
||||||
@@ -0,0 +1,213 @@
|
|||||||
|
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)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,3 @@
|
|||||||
|
// Package lua provides shared GopherLua programs, state management, and value
|
||||||
|
// conversion and validation helpers for Xray scripts.
|
||||||
|
package lua
|
||||||
@@ -0,0 +1,65 @@
|
|||||||
|
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()
|
||||||
|
}
|
||||||
@@ -0,0 +1,84 @@
|
|||||||
|
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)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,151 @@
|
|||||||
|
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()
|
||||||
|
}
|
||||||
@@ -0,0 +1,466 @@
|
|||||||
|
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)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,77 @@
|
|||||||
|
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)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,76 @@
|
|||||||
|
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())
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,95 @@
|
|||||||
|
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)
|
||||||
|
}
|
||||||
@@ -0,0 +1,121 @@
|
|||||||
|
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)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,20 @@
|
|||||||
|
package net
|
||||||
|
|
||||||
|
// PacketConnWrapper wraps a PacketConn into a Conn with a fixed destination address.
|
||||||
|
type PacketConnWrapper struct {
|
||||||
|
PacketConn
|
||||||
|
Dest Addr
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *PacketConnWrapper) Read(p []byte) (int, error) {
|
||||||
|
n, _, err := c.PacketConn.ReadFrom(p)
|
||||||
|
return n, err
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *PacketConnWrapper) Write(p []byte) (int, error) {
|
||||||
|
return c.PacketConn.WriteTo(p, c.Dest)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *PacketConnWrapper) RemoteAddr() Addr {
|
||||||
|
return c.Dest
|
||||||
|
}
|
||||||
@@ -1,6 +1,8 @@
|
|||||||
package platform // import "github.com/xtls/xray-core/common/platform"
|
package platform // import "github.com/xtls/xray-core/common/platform"
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"errors"
|
||||||
|
"fmt"
|
||||||
"os"
|
"os"
|
||||||
"path/filepath"
|
"path/filepath"
|
||||||
"strconv"
|
"strconv"
|
||||||
@@ -90,3 +92,49 @@ func GetConfDirPath() string {
|
|||||||
configPath := NewEnvFlag(ConfdirLocation).GetValue(func() string { return "" })
|
configPath := NewEnvFlag(ConfdirLocation).GetValue(func() string { return "" })
|
||||||
return configPath
|
return configPath
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// ResolveLuaFile finds a local Lua script and returns its absolute path.
|
||||||
|
// Relative paths: XRAY_LOCATION_CONFDIR > XRAY_LOCATION_CONFIG > working dir > executable dir.
|
||||||
|
func ResolveLuaFile(path string) (string, error) {
|
||||||
|
if path == "" {
|
||||||
|
return "", errors.New("Lua file path is empty")
|
||||||
|
}
|
||||||
|
paths := []string{path}
|
||||||
|
if !filepath.IsAbs(path) {
|
||||||
|
paths = nil
|
||||||
|
for _, dir := range []string{
|
||||||
|
GetConfDirPath(),
|
||||||
|
NewEnvFlag(ConfigLocation).GetValue(func() string { return "" }),
|
||||||
|
".",
|
||||||
|
getExecutableDir(),
|
||||||
|
} {
|
||||||
|
if dir != "" {
|
||||||
|
paths = append(paths, filepath.Join(dir, path))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return resolveFile(paths)
|
||||||
|
}
|
||||||
|
|
||||||
|
func resolveFile(paths []string) (string, error) {
|
||||||
|
var tried []string
|
||||||
|
for _, path := range paths {
|
||||||
|
path, err := filepath.Abs(path)
|
||||||
|
if err != nil {
|
||||||
|
return "", fmt.Errorf("failed to resolve file path: %w", err)
|
||||||
|
}
|
||||||
|
tried = append(tried, path)
|
||||||
|
info, err := os.Stat(path)
|
||||||
|
if errors.Is(err, os.ErrNotExist) {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if err != nil {
|
||||||
|
return "", fmt.Errorf("failed to inspect file %q: %w", path, err)
|
||||||
|
}
|
||||||
|
if !info.Mode().IsRegular() {
|
||||||
|
return "", fmt.Errorf("file is not a regular file: %s", path)
|
||||||
|
}
|
||||||
|
return path, nil
|
||||||
|
}
|
||||||
|
return "", fmt.Errorf("file not found; tried %q: %w", tried, os.ErrNotExist)
|
||||||
|
}
|
||||||
|
|||||||
@@ -1,6 +1,7 @@
|
|||||||
package platform_test
|
package platform_test
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"errors"
|
||||||
"os"
|
"os"
|
||||||
"path/filepath"
|
"path/filepath"
|
||||||
"runtime"
|
"runtime"
|
||||||
@@ -64,3 +65,53 @@ func TestGetAssetLocation(t *testing.T) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestResolveLuaFile(t *testing.T) {
|
||||||
|
workingDir := t.TempDir()
|
||||||
|
t.Chdir(workingDir)
|
||||||
|
executable, err := os.Executable()
|
||||||
|
common.Must(err)
|
||||||
|
file, err := os.CreateTemp(filepath.Dir(executable), "lua-*.lua")
|
||||||
|
common.Must(err)
|
||||||
|
common.Must(file.Close())
|
||||||
|
defer os.Remove(file.Name())
|
||||||
|
|
||||||
|
name := filepath.Base(file.Name())
|
||||||
|
paths := []string{
|
||||||
|
filepath.Join(t.TempDir(), name),
|
||||||
|
filepath.Join(t.TempDir(), name),
|
||||||
|
filepath.Join(workingDir, name),
|
||||||
|
file.Name(),
|
||||||
|
}
|
||||||
|
t.Setenv(ConfdirLocation, filepath.Dir(paths[0]))
|
||||||
|
t.Setenv(ConfigLocation, filepath.Dir(paths[1]))
|
||||||
|
for _, path := range paths[:3] {
|
||||||
|
common.Must(os.WriteFile(path, nil, 0o600))
|
||||||
|
}
|
||||||
|
if got, err := ResolveLuaFile(paths[2]); err != nil || got != paths[2] {
|
||||||
|
t.Fatalf("absolute path = %q, %v; want %q", got, err, paths[2])
|
||||||
|
}
|
||||||
|
for i, want := range paths {
|
||||||
|
if i == 2 {
|
||||||
|
t.Setenv(ConfdirLocation, "")
|
||||||
|
t.Setenv(ConfigLocation, "")
|
||||||
|
}
|
||||||
|
if got, err := ResolveLuaFile(name); err != nil || got != want {
|
||||||
|
t.Fatalf("resolved path = %q, %v; want %q", got, err, want)
|
||||||
|
}
|
||||||
|
common.Must(os.Remove(want))
|
||||||
|
}
|
||||||
|
if _, err := ResolveLuaFile(name); !errors.Is(err, os.ErrNotExist) {
|
||||||
|
t.Fatalf("missing file error = %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
t.Setenv(ConfdirLocation, filepath.Dir(paths[0]))
|
||||||
|
t.Setenv(ConfigLocation, filepath.Dir(paths[1]))
|
||||||
|
common.Must(os.Mkdir(paths[0], 0o700))
|
||||||
|
common.Must(os.WriteFile(paths[1], nil, 0o600))
|
||||||
|
for _, path := range []string{"", name, filepath.Join(t.TempDir(), name)} {
|
||||||
|
if _, err := ResolveLuaFile(path); err == nil {
|
||||||
|
t.Fatalf("accepted invalid path %q", path)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
+1
-1
@@ -20,7 +20,7 @@ import (
|
|||||||
var (
|
var (
|
||||||
Version_x byte = 26
|
Version_x byte = 26
|
||||||
Version_y byte = 9
|
Version_y byte = 9
|
||||||
Version_z byte = 9
|
Version_z byte = 30
|
||||||
)
|
)
|
||||||
|
|
||||||
var (
|
var (
|
||||||
|
|||||||
@@ -97,6 +97,9 @@ func New() *Client {
|
|||||||
r := &net.Resolver{
|
r := &net.Resolver{
|
||||||
PreferGo: true,
|
PreferGo: true,
|
||||||
Dial: func(ctx context.Context, network, address string) (net.Conn, error) {
|
Dial: func(ctx context.Context, network, address string) (net.Conn, error) {
|
||||||
|
if internet.IsSkippedDNSServer(address) {
|
||||||
|
return nil, errors.New("skipped DNS server ", address)
|
||||||
|
}
|
||||||
return d.DialContext(ctx, network, address)
|
return d.DialContext(ctx, network, address)
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -0,0 +1,23 @@
|
|||||||
|
package localdns
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"net/netip"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/xtls/xray-core/transport/internet"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestSkippedDNSServers(t *testing.T) {
|
||||||
|
internet.SkipDNSServers([]netip.Addr{netip.MustParseAddr("203.0.113.53")})
|
||||||
|
t.Cleanup(func() { internet.SkipDNSServers(nil) })
|
||||||
|
c := New()
|
||||||
|
if _, err := c.r.Dial(context.Background(), "udp", "203.0.113.53:53"); err == nil {
|
||||||
|
t.Error("a skipped DNS server was dialed")
|
||||||
|
}
|
||||||
|
conn, err := c.r.Dial(context.Background(), "udp", "127.0.0.1:53")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
conn.Close()
|
||||||
|
}
|
||||||
@@ -21,6 +21,7 @@ require (
|
|||||||
github.com/stretchr/testify v1.12.1
|
github.com/stretchr/testify v1.12.1
|
||||||
github.com/vishvananda/netlink v1.3.1
|
github.com/vishvananda/netlink v1.3.1
|
||||||
github.com/xtls/reality v0.0.0-20260908062103-8cdf7bf9c7f0
|
github.com/xtls/reality v0.0.0-20260908062103-8cdf7bf9c7f0
|
||||||
|
github.com/yuin/gopher-lua v1.1.2
|
||||||
go4.org/netipx v0.0.0-20231129151722-fdeea329fbba
|
go4.org/netipx v0.0.0-20231129151722-fdeea329fbba
|
||||||
golang.org/x/crypto v0.57.0
|
golang.org/x/crypto v0.57.0
|
||||||
golang.org/x/exp v0.0.0-20240506185415-9bf2ced13842
|
golang.org/x/exp v0.0.0-20240506185415-9bf2ced13842
|
||||||
@@ -34,6 +35,7 @@ require (
|
|||||||
google.golang.org/protobuf v1.36.12
|
google.golang.org/protobuf v1.36.12
|
||||||
gvisor.dev/gvisor v0.0.0-20260122175437-89a5d21be8f0
|
gvisor.dev/gvisor v0.0.0-20260122175437-89a5d21be8f0
|
||||||
h12.io/socks v1.0.3
|
h12.io/socks v1.0.3
|
||||||
|
layeh.com/gopher-luar v1.0.11
|
||||||
lukechampine.com/blake3 v1.4.1
|
lukechampine.com/blake3 v1.4.1
|
||||||
mvdan.cc/gofumpt v0.12.0
|
mvdan.cc/gofumpt v0.12.0
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -2,6 +2,9 @@ github.com/andybalholm/brotli v1.0.6 h1:Yf9fFpf49Zrxb9NlQaluyE92/+X7UVHlhMNJN2sx
|
|||||||
github.com/andybalholm/brotli v1.0.6/go.mod h1:fO7iG3H7G2nSZ7m0zPUDn85XEX2GTukHGRSepvi9Eig=
|
github.com/andybalholm/brotli v1.0.6/go.mod h1:fO7iG3H7G2nSZ7m0zPUDn85XEX2GTukHGRSepvi9Eig=
|
||||||
github.com/apernet/quic-go v0.61.1-0.20260806010916-184d081eef3e h1:5mgtR5gwIgBKMiGI1QdXldZZ+SNor06Nbu1wCBulQBg=
|
github.com/apernet/quic-go v0.61.1-0.20260806010916-184d081eef3e h1:5mgtR5gwIgBKMiGI1QdXldZZ+SNor06Nbu1wCBulQBg=
|
||||||
github.com/apernet/quic-go v0.61.1-0.20260806010916-184d081eef3e/go.mod h1:x7qxEvX6MCVtDuBKHj3E+88+BtrbEMuAL5qGUKItjW8=
|
github.com/apernet/quic-go v0.61.1-0.20260806010916-184d081eef3e/go.mod h1:x7qxEvX6MCVtDuBKHj3E+88+BtrbEMuAL5qGUKItjW8=
|
||||||
|
github.com/chzyer/logex v1.1.10/go.mod h1:+Ywpsq7O8HXn0nuIou7OrIPyXbp3wmkHB+jjWRnGsAI=
|
||||||
|
github.com/chzyer/readline v0.0.0-20180603132655-2972be24d48e/go.mod h1:nSuG5e5PlCu98SY8svDHJxuZscDgtXS6KTTbou5AhLI=
|
||||||
|
github.com/chzyer/test v0.0.0-20180213035817-a1ea475d72b1/go.mod h1:Q3SI9o4m/ZMnBNeIyt5eFwwo7qiLfzFZmjNmxjkiQlU=
|
||||||
github.com/cloudflare/circl v1.6.5 h1:O64F26HEqNhznd/hrC5KZXVKYuKM2rx4deZDTc4ihQA=
|
github.com/cloudflare/circl v1.6.5 h1:O64F26HEqNhznd/hrC5KZXVKYuKM2rx4deZDTc4ihQA=
|
||||||
github.com/cloudflare/circl v1.6.5/go.mod h1:h5LNyxAc5nTue9DS5jT+48en2PSDYt3zdGnz5OstK6c=
|
github.com/cloudflare/circl v1.6.5/go.mod h1:h5LNyxAc5nTue9DS5jT+48en2PSDYt3zdGnz5OstK6c=
|
||||||
github.com/ghodss/yaml v1.0.1-0.20220118164431-d8423dcdf344 h1:Arcl6UOIS/kgO2nW3A65HN+7CMjSDP/gofXL4CZt1V4=
|
github.com/ghodss/yaml v1.0.1-0.20220118164431-d8423dcdf344 h1:Arcl6UOIS/kgO2nW3A65HN+7CMjSDP/gofXL4CZt1V4=
|
||||||
@@ -81,6 +84,9 @@ github.com/wlynxg/anet v0.0.5/go.mod h1:eay5PRQr7fIVAMbTbchTnO9gG65Hg/uYGdc7mguH
|
|||||||
github.com/xtls/reality v0.0.0-20260908062103-8cdf7bf9c7f0 h1:rb+fKQFhz+5I2PPuQsNYxI5mUU840XWYtRF0ZBjvkws=
|
github.com/xtls/reality v0.0.0-20260908062103-8cdf7bf9c7f0 h1:rb+fKQFhz+5I2PPuQsNYxI5mUU840XWYtRF0ZBjvkws=
|
||||||
github.com/xtls/reality v0.0.0-20260908062103-8cdf7bf9c7f0/go.mod h1:DsJblcWDGt76+FVqBVwbwRhxyyNJsGV48gJLch0OOWI=
|
github.com/xtls/reality v0.0.0-20260908062103-8cdf7bf9c7f0/go.mod h1:DsJblcWDGt76+FVqBVwbwRhxyyNJsGV48gJLch0OOWI=
|
||||||
github.com/yuin/goldmark v1.4.1/go.mod h1:mwnBkeHKe2W/ZEtQ+71ViKU8L12m81fl3OWwC1Zlc8k=
|
github.com/yuin/goldmark v1.4.1/go.mod h1:mwnBkeHKe2W/ZEtQ+71ViKU8L12m81fl3OWwC1Zlc8k=
|
||||||
|
github.com/yuin/gopher-lua v0.0.0-20190206043414-8bfc7677f583/go.mod h1:gqRgreBUhTSL0GeU64rtZ3Uq3wtjOa/TB2YfrtkCbVQ=
|
||||||
|
github.com/yuin/gopher-lua v1.1.2 h1:yF/FjE3hD65tBbt0VXLE13HWS9h34fdzJmrWRXwobGA=
|
||||||
|
github.com/yuin/gopher-lua v1.1.2/go.mod h1:7aRmXIWl37SqRf0koeyylBEzJ+aPt8A+mmkQ4f1ntR8=
|
||||||
go.uber.org/mock v0.5.2 h1:LbtPTcP8A5k9WPXj54PPPbjcI4Y6lhyOZXn+VS7wNko=
|
go.uber.org/mock v0.5.2 h1:LbtPTcP8A5k9WPXj54PPPbjcI4Y6lhyOZXn+VS7wNko=
|
||||||
go.uber.org/mock v0.5.2/go.mod h1:wLlUxC2vVTPTaE3UD51E0BGOAElKrILxhVSDYQLld5o=
|
go.uber.org/mock v0.5.2/go.mod h1:wLlUxC2vVTPTaE3UD51E0BGOAElKrILxhVSDYQLld5o=
|
||||||
go.yaml.in/yaml/v3 v3.0.5 h1:N6y/pJk8buWs9NY5ERU2HSMfm+IuD/OtfdAnq6kESPw=
|
go.yaml.in/yaml/v3 v3.0.5 h1:N6y/pJk8buWs9NY5ERU2HSMfm+IuD/OtfdAnq6kESPw=
|
||||||
@@ -105,6 +111,7 @@ golang.org/x/sync v0.0.0-20190423024810-112230192c58/go.mod h1:RxMgew5VJxzue5/jJ
|
|||||||
golang.org/x/sync v0.0.0-20210220032951-036812b2e83c/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
golang.org/x/sync v0.0.0-20210220032951-036812b2e83c/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
||||||
golang.org/x/sync v0.23.0 h1:KameEIfc1IkluZyXWLn39Wd4tURc6GbCiISGiZm2bQk=
|
golang.org/x/sync v0.23.0 h1:KameEIfc1IkluZyXWLn39Wd4tURc6GbCiISGiZm2bQk=
|
||||||
golang.org/x/sync v0.23.0/go.mod h1:sUUOizhqBxiL6pEWpqNLUiaJn1ShEbZ6BBqskPbjZm0=
|
golang.org/x/sync v0.23.0/go.mod h1:sUUOizhqBxiL6pEWpqNLUiaJn1ShEbZ6BBqskPbjZm0=
|
||||||
|
golang.org/x/sys v0.0.0-20190204203706-41f3e6584952/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY=
|
||||||
golang.org/x/sys v0.0.0-20190215142949-d0b11bdaac8a/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY=
|
golang.org/x/sys v0.0.0-20190215142949-d0b11bdaac8a/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY=
|
||||||
golang.org/x/sys v0.0.0-20190412213103-97732733099d/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs=
|
golang.org/x/sys v0.0.0-20190412213103-97732733099d/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs=
|
||||||
golang.org/x/sys v0.0.0-20201119102817-f84b799fce68/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs=
|
golang.org/x/sys v0.0.0-20201119102817-f84b799fce68/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs=
|
||||||
@@ -155,6 +162,8 @@ gvisor.dev/gvisor v0.0.0-20260122175437-89a5d21be8f0 h1:Lk6hARj5UPY47dBep70OD/TI
|
|||||||
gvisor.dev/gvisor v0.0.0-20260122175437-89a5d21be8f0/go.mod h1:QkHjoMIBaYtpVufgwv3keYAbln78mBoCuShZrPrer1Q=
|
gvisor.dev/gvisor v0.0.0-20260122175437-89a5d21be8f0/go.mod h1:QkHjoMIBaYtpVufgwv3keYAbln78mBoCuShZrPrer1Q=
|
||||||
h12.io/socks v1.0.3 h1:Ka3qaQewws4j4/eDQnOdpr4wXsC//dXtWvftlIcCQUo=
|
h12.io/socks v1.0.3 h1:Ka3qaQewws4j4/eDQnOdpr4wXsC//dXtWvftlIcCQUo=
|
||||||
h12.io/socks v1.0.3/go.mod h1:AIhxy1jOId/XCz9BO+EIgNL2rQiPTBNnOfnVnQ+3Eck=
|
h12.io/socks v1.0.3/go.mod h1:AIhxy1jOId/XCz9BO+EIgNL2rQiPTBNnOfnVnQ+3Eck=
|
||||||
|
layeh.com/gopher-luar v1.0.11 h1:8zJudpKI6HWkoh9eyyNFaTM79PY6CAPcIr6X/KTiliw=
|
||||||
|
layeh.com/gopher-luar v1.0.11/go.mod h1:TPnIVCZ2RJBndm7ohXyaqfhzjlZ+OA2SZR/YwL8tECk=
|
||||||
lukechampine.com/blake3 v1.4.1 h1:I3Smz7gso8w4/TunLKec6K2fn+kyKtDxr/xcQEN84Wg=
|
lukechampine.com/blake3 v1.4.1 h1:I3Smz7gso8w4/TunLKec6K2fn+kyKtDxr/xcQEN84Wg=
|
||||||
lukechampine.com/blake3 v1.4.1/go.mod h1:QFosUxmjB8mnrWFSNwKmvxHpfY72bmD2tQ0kBMM3kwo=
|
lukechampine.com/blake3 v1.4.1/go.mod h1:QFosUxmjB8mnrWFSNwKmvxHpfY72bmD2tQ0kBMM3kwo=
|
||||||
mvdan.cc/gofumpt v0.12.0 h1:1Lbudkz2kpM9Cjz2pL4M19u7q+GaEhCTNf7N9mfpcho=
|
mvdan.cc/gofumpt v0.12.0 h1:1Lbudkz2kpM9Cjz2pL4M19u7q+GaEhCTNf7N9mfpcho=
|
||||||
|
|||||||
@@ -14,9 +14,11 @@ import (
|
|||||||
"github.com/xtls/xray-core/common/errors"
|
"github.com/xtls/xray-core/common/errors"
|
||||||
"github.com/xtls/xray-core/common/geodata"
|
"github.com/xtls/xray-core/common/geodata"
|
||||||
"github.com/xtls/xray-core/common/net"
|
"github.com/xtls/xray-core/common/net"
|
||||||
|
"github.com/xtls/xray-core/common/platform"
|
||||||
)
|
)
|
||||||
|
|
||||||
type NameServerConfig struct {
|
type NameServerConfig struct {
|
||||||
|
ID string `json:"id"`
|
||||||
Address *Address `json:"address"`
|
Address *Address `json:"address"`
|
||||||
ClientIP *Address `json:"clientIp"`
|
ClientIP *Address `json:"clientIp"`
|
||||||
Port uint16 `json:"port"`
|
Port uint16 `json:"port"`
|
||||||
@@ -43,6 +45,7 @@ func (c *NameServerConfig) UnmarshalJSON(data []byte) error {
|
|||||||
}
|
}
|
||||||
|
|
||||||
var advanced struct {
|
var advanced struct {
|
||||||
|
ID string `json:"id"`
|
||||||
Address *Address `json:"address"`
|
Address *Address `json:"address"`
|
||||||
ClientIP *Address `json:"clientIp"`
|
ClientIP *Address `json:"clientIp"`
|
||||||
Port uint16 `json:"port"`
|
Port uint16 `json:"port"`
|
||||||
@@ -60,6 +63,7 @@ func (c *NameServerConfig) UnmarshalJSON(data []byte) error {
|
|||||||
UnexpectedIPs StringList `json:"unexpectedIPs"`
|
UnexpectedIPs StringList `json:"unexpectedIPs"`
|
||||||
}
|
}
|
||||||
if err := json.Unmarshal(data, &advanced); err == nil {
|
if err := json.Unmarshal(data, &advanced); err == nil {
|
||||||
|
c.ID = advanced.ID
|
||||||
c.Address = advanced.Address
|
c.Address = advanced.Address
|
||||||
c.ClientIP = advanced.ClientIP
|
c.ClientIP = advanced.ClientIP
|
||||||
c.Port = advanced.Port
|
c.Port = advanced.Port
|
||||||
@@ -134,6 +138,7 @@ func (c *NameServerConfig) Build() (*dns.NameServer, error) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
return &dns.NameServer{
|
return &dns.NameServer{
|
||||||
|
Id: c.ID,
|
||||||
Address: &net.Endpoint{
|
Address: &net.Endpoint{
|
||||||
Network: net.Network_UDP,
|
Network: net.Network_UDP,
|
||||||
Address: c.Address.Build(),
|
Address: c.Address.Build(),
|
||||||
@@ -159,6 +164,7 @@ func (c *NameServerConfig) Build() (*dns.NameServer, error) {
|
|||||||
// DNSConfig is a JSON serializable object for dns.Config
|
// DNSConfig is a JSON serializable object for dns.Config
|
||||||
type DNSConfig struct {
|
type DNSConfig struct {
|
||||||
Servers []*NameServerConfig `json:"servers"`
|
Servers []*NameServerConfig `json:"servers"`
|
||||||
|
Script string `json:"script"`
|
||||||
Hosts *HostsWrapper `json:"hosts"`
|
Hosts *HostsWrapper `json:"hosts"`
|
||||||
ClientIP *Address `json:"clientIp"`
|
ClientIP *Address `json:"clientIp"`
|
||||||
Tag string `json:"tag"`
|
Tag string `json:"tag"`
|
||||||
@@ -278,6 +284,14 @@ func (c *DNSConfig) Build() (*dns.Config, error) {
|
|||||||
QueryStrategy: resolveQueryStrategy(c.QueryStrategy),
|
QueryStrategy: resolveQueryStrategy(c.QueryStrategy),
|
||||||
}
|
}
|
||||||
|
|
||||||
|
if c.Script != "" {
|
||||||
|
path, err := platform.ResolveLuaFile(c.Script)
|
||||||
|
if err != nil {
|
||||||
|
return nil, errors.New("failed to resolve DNS script: ", c.Script).Base(err)
|
||||||
|
}
|
||||||
|
config.Script = path
|
||||||
|
}
|
||||||
|
|
||||||
if c.ClientIP != nil {
|
if c.ClientIP != nil {
|
||||||
if !c.ClientIP.Family().IsIP() {
|
if !c.ClientIP.Family().IsIP() {
|
||||||
return nil, errors.New("not an IP address:", c.ClientIP.String())
|
return nil, errors.New("not an IP address:", c.ClientIP.String())
|
||||||
|
|||||||
@@ -2,6 +2,8 @@ package conf_test
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"encoding/json"
|
"encoding/json"
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
"github.com/google/go-cmp/cmp"
|
"github.com/google/go-cmp/cmp"
|
||||||
@@ -122,3 +124,51 @@ 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,6 +7,7 @@ import (
|
|||||||
"github.com/xtls/xray-core/app/router"
|
"github.com/xtls/xray-core/app/router"
|
||||||
"github.com/xtls/xray-core/common/errors"
|
"github.com/xtls/xray-core/common/errors"
|
||||||
"github.com/xtls/xray-core/common/geodata"
|
"github.com/xtls/xray-core/common/geodata"
|
||||||
|
"github.com/xtls/xray-core/common/platform"
|
||||||
"github.com/xtls/xray-core/common/serial"
|
"github.com/xtls/xray-core/common/serial"
|
||||||
|
|
||||||
"google.golang.org/protobuf/proto"
|
"google.golang.org/protobuf/proto"
|
||||||
@@ -72,6 +73,7 @@ type RouterConfig struct {
|
|||||||
RuleList []json.RawMessage `json:"rules"`
|
RuleList []json.RawMessage `json:"rules"`
|
||||||
DomainStrategy *string `json:"domainStrategy"`
|
DomainStrategy *string `json:"domainStrategy"`
|
||||||
Balancers []*BalancingRule `json:"balancers"`
|
Balancers []*BalancingRule `json:"balancers"`
|
||||||
|
Script string `json:"script"`
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *RouterConfig) getDomainStrategy() router.Config_DomainStrategy {
|
func (c *RouterConfig) getDomainStrategy() router.Config_DomainStrategy {
|
||||||
@@ -92,6 +94,15 @@ func (c *RouterConfig) getDomainStrategy() router.Config_DomainStrategy {
|
|||||||
|
|
||||||
func (c *RouterConfig) Build() (*router.Config, error) {
|
func (c *RouterConfig) Build() (*router.Config, error) {
|
||||||
config := new(router.Config)
|
config := new(router.Config)
|
||||||
|
|
||||||
|
if c.Script != "" {
|
||||||
|
path, err := platform.ResolveLuaFile(c.Script)
|
||||||
|
if err != nil {
|
||||||
|
return nil, errors.New("failed to resolve routing script").Base(err)
|
||||||
|
}
|
||||||
|
config.Script = path
|
||||||
|
}
|
||||||
|
|
||||||
config.DomainStrategy = c.getDomainStrategy()
|
config.DomainStrategy = c.getDomainStrategy()
|
||||||
|
|
||||||
var rawRuleList []json.RawMessage
|
var rawRuleList []json.RawMessage
|
||||||
|
|||||||
@@ -2,6 +2,8 @@ package conf_test
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"encoding/json"
|
"encoding/json"
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
"testing"
|
"testing"
|
||||||
"time"
|
"time"
|
||||||
_ "unsafe"
|
_ "unsafe"
|
||||||
@@ -236,3 +238,39 @@ func TestRouterConfig(t *testing.T) {
|
|||||||
},
|
},
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestRouterScriptConfig(t *testing.T) {
|
||||||
|
dir := t.TempDir()
|
||||||
|
t.Setenv("xray.location.confdir", dir)
|
||||||
|
path := filepath.Join(dir, "route.lua")
|
||||||
|
if err := os.WriteFile(path, []byte("function HandleRoute() end"), 0o600); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tc := range []struct {
|
||||||
|
name string
|
||||||
|
script string
|
||||||
|
wantError bool
|
||||||
|
}{
|
||||||
|
{"relative", "route.lua", false},
|
||||||
|
{"absolute", path, false},
|
||||||
|
{"missing", "missing.lua", true},
|
||||||
|
{"directory", dir, true},
|
||||||
|
} {
|
||||||
|
t.Run(tc.name, func(t *testing.T) {
|
||||||
|
built, err := (&RouterConfig{Script: tc.script}).Build()
|
||||||
|
if tc.wantError {
|
||||||
|
if err == nil {
|
||||||
|
t.Fatal("Build accepted invalid script path")
|
||||||
|
}
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if built.Script != path {
|
||||||
|
t.Fatalf("script path = %q, want %q", built.Script, path)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
@@ -1,6 +1,7 @@
|
|||||||
package conf
|
package conf
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"context"
|
||||||
"crypto/x509"
|
"crypto/x509"
|
||||||
"encoding/base64"
|
"encoding/base64"
|
||||||
"encoding/hex"
|
"encoding/hex"
|
||||||
@@ -14,6 +15,7 @@ import (
|
|||||||
googleuuid "github.com/google/uuid"
|
googleuuid "github.com/google/uuid"
|
||||||
"github.com/xtls/xray-core/common/errors"
|
"github.com/xtls/xray-core/common/errors"
|
||||||
"github.com/xtls/xray-core/common/net"
|
"github.com/xtls/xray-core/common/net"
|
||||||
|
"github.com/xtls/xray-core/common/serial"
|
||||||
"github.com/xtls/xray-core/transport/internet/finalmask/fragment"
|
"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/header/custom"
|
||||||
"github.com/xtls/xray-core/transport/internet/finalmask/mkcp/aes128gcm"
|
"github.com/xtls/xray-core/transport/internet/finalmask/mkcp/aes128gcm"
|
||||||
@@ -81,7 +83,7 @@ var (
|
|||||||
"noise": func() interface{} { return new(NoiseMask) },
|
"noise": func() interface{} { return new(NoiseMask) },
|
||||||
"salamander": func() interface{} { return new(Salamander) },
|
"salamander": func() interface{} { return new(Salamander) },
|
||||||
"sudoku": func() interface{} { return new(Sudoku) },
|
"sudoku": func() interface{} { return new(Sudoku) },
|
||||||
"xdns": func() interface{} { return new(Xdns) },
|
"xdns": func() interface{} { return new(XDNS) },
|
||||||
"xicmp": func() interface{} { return new(Xicmp) },
|
"xicmp": func() interface{} { return new(Xicmp) },
|
||||||
"realm": func() interface{} { return new(Realm) },
|
"realm": func() interface{} { return new(Realm) },
|
||||||
"udphop": func() interface{} { return new(UDPHop) },
|
"udphop": func() interface{} { return new(UDPHop) },
|
||||||
@@ -308,14 +310,27 @@ type NoiseMask struct {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (c *NoiseMask) Build() (proto.Message, error) {
|
func (c *NoiseMask) Build() (proto.Message, error) {
|
||||||
|
noiseSlice := make([]*noise.Item, 0, len(c.Noise))
|
||||||
for _, item := range c.Noise {
|
for _, item := range c.Noise {
|
||||||
if len(item.Packet) > 0 && item.Rand.To > 0 {
|
if len(item.Packet) > 0 && item.Rand.To > 0 {
|
||||||
return nil, errors.New("len(item.Packet) > 0 && item.Rand.To > 0")
|
return nil, errors.New("len(item.Packet) > 0 && item.Rand.To > 0")
|
||||||
}
|
}
|
||||||
}
|
if strings.ToLower(item.Type) == "exp" {
|
||||||
|
var exp string
|
||||||
noiseSlice := make([]*noise.Item, 0, len(c.Noise))
|
if err := json.Unmarshal(item.Packet, &exp); err != nil {
|
||||||
for _, item := range c.Noise {
|
return nil, errors.New(`"packet" of noise "type": "exp" must be a string`).Base(err)
|
||||||
|
}
|
||||||
|
segments, err := parseNoiseExp(exp)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
noiseSlice = append(noiseSlice, &noise.Item{
|
||||||
|
Segments: segments,
|
||||||
|
DelayMin: int64(item.Delay.From),
|
||||||
|
DelayMax: int64(item.Delay.To),
|
||||||
|
})
|
||||||
|
continue
|
||||||
|
}
|
||||||
if item.RandRange == nil {
|
if item.RandRange == nil {
|
||||||
item.RandRange = &Int32Range{From: 0, To: 255}
|
item.RandRange = &Int32Range{From: 0, To: 255}
|
||||||
}
|
}
|
||||||
@@ -344,6 +359,88 @@ func (c *NoiseMask) Build() (proto.Message, error) {
|
|||||||
}, nil
|
}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
var noiseExpPattern = regexp.MustCompile(`<\s*([a-z]+)(?:\s+([^>]*?))?\s*>`)
|
||||||
|
|
||||||
|
func parseNoiseExp(exp string) ([]*noise.Segment, error) {
|
||||||
|
var segments []*noise.Segment
|
||||||
|
matches := noiseExpPattern.FindAllStringSubmatchIndex(exp, -1)
|
||||||
|
last := 0
|
||||||
|
for _, m := range matches {
|
||||||
|
if strings.TrimSpace(exp[last:m[0]]) != "" {
|
||||||
|
return nil, errors.New("invalid noise exp near ", exp[last:m[0]])
|
||||||
|
}
|
||||||
|
last = m[1]
|
||||||
|
key := exp[m[2]:m[3]]
|
||||||
|
arg := ""
|
||||||
|
if m[4] >= 0 {
|
||||||
|
arg = exp[m[4]:m[5]]
|
||||||
|
}
|
||||||
|
segment, err := buildNoiseSegment(key, arg)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
segments = append(segments, segment)
|
||||||
|
}
|
||||||
|
if strings.TrimSpace(exp[last:]) != "" {
|
||||||
|
return nil, errors.New("invalid noise exp near ", exp[last:])
|
||||||
|
}
|
||||||
|
if len(segments) == 0 {
|
||||||
|
return nil, errors.New("empty noise exp: ", exp)
|
||||||
|
}
|
||||||
|
return segments, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func buildNoiseSegment(key, arg string) (*noise.Segment, error) {
|
||||||
|
sizeSegment := func(kind noise.Segment_Kind) (*noise.Segment, error) {
|
||||||
|
if arg == "" {
|
||||||
|
return nil, errors.New("<", key, "> in noise exp needs a size")
|
||||||
|
}
|
||||||
|
lo, hi, err := ParseRangeString(arg)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
if lo < 0 || hi < lo || hi > 65535 {
|
||||||
|
return nil, errors.New("invalid size in noise exp: ", arg)
|
||||||
|
}
|
||||||
|
return &noise.Segment{Kind: kind, MinSize: int64(lo), MaxSize: int64(hi)}, nil
|
||||||
|
}
|
||||||
|
switch key {
|
||||||
|
case "b":
|
||||||
|
hexStr := strings.TrimPrefix(strings.TrimPrefix(strings.Join(strings.Fields(arg), ""), "0x"), "0X")
|
||||||
|
if len(hexStr) == 0 {
|
||||||
|
return nil, errors.New("empty bytes in noise exp")
|
||||||
|
}
|
||||||
|
raw, err := hex.DecodeString(hexStr)
|
||||||
|
if err != nil {
|
||||||
|
return nil, errors.New("invalid hex in noise exp: ", arg).Base(err)
|
||||||
|
}
|
||||||
|
return &noise.Segment{Kind: noise.Segment_BYTES, Bytes: raw}, nil
|
||||||
|
case "r":
|
||||||
|
return sizeSegment(noise.Segment_RANDOM)
|
||||||
|
case "rc":
|
||||||
|
return sizeSegment(noise.Segment_RANDOM_ASCII)
|
||||||
|
case "rd":
|
||||||
|
return sizeSegment(noise.Segment_RANDOM_DIGIT)
|
||||||
|
case "t":
|
||||||
|
if arg != "" {
|
||||||
|
return nil, errors.New("<t> in noise exp takes no argument")
|
||||||
|
}
|
||||||
|
return &noise.Segment{Kind: noise.Segment_TIMESTAMP}, nil
|
||||||
|
case "c":
|
||||||
|
if arg != "" {
|
||||||
|
return nil, errors.New("<c> in noise exp takes no argument")
|
||||||
|
}
|
||||||
|
return &noise.Segment{Kind: noise.Segment_COUNTER}, nil
|
||||||
|
case "n":
|
||||||
|
if arg != "" {
|
||||||
|
return nil, errors.New("<n> in noise exp takes no argument")
|
||||||
|
}
|
||||||
|
return &noise.Segment{Kind: noise.Segment_NONCE}, nil
|
||||||
|
default:
|
||||||
|
return nil, errors.New("unknown <", key, "> in noise exp")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
type UDPItem struct {
|
type UDPItem struct {
|
||||||
Rand int32 `json:"rand"`
|
Rand int32 `json:"rand"`
|
||||||
RandRange *Int32Range `json:"randRange"`
|
RandRange *Int32Range `json:"randRange"`
|
||||||
@@ -694,32 +791,88 @@ func (c *Sudoku) Build() (proto.Message, error) {
|
|||||||
}, nil
|
}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
type Xdns struct {
|
type XDNSDomain struct {
|
||||||
Domain json.RawMessage `json:"domain"`
|
Name string `json:"name"`
|
||||||
|
LenLimit int32 `json:"lenLimit"`
|
||||||
Domains []string `json:"domains"`
|
LabelLimit int32 `json:"labelLimit"`
|
||||||
Resolvers []string `json:"resolvers"`
|
Types []int32 `json:"types"`
|
||||||
|
Edns0 int32 `json:"edns0"`
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *Xdns) Build() (proto.Message, error) {
|
type XDNSResolverTCP struct {
|
||||||
if c.Domain != nil {
|
Addr string `json:"addr"`
|
||||||
return nil, errors.PrintRemovedFeatureError("domain", "domains(server) & resolvers(client)")
|
}
|
||||||
}
|
|
||||||
|
|
||||||
if len(c.Domains) == 0 && len(c.Resolvers) == 0 {
|
func (c *XDNSResolverTCP) Build() (proto.Message, error) {
|
||||||
return nil, errors.New("empty domains & empty resolvers")
|
return &xdns.TCPResolverProto{Addr: c.Addr}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
for _, r := range c.Resolvers {
|
type XDNSResolverUDP struct {
|
||||||
if !strings.Contains(r, "+udp://") {
|
Addr string `json:"addr"`
|
||||||
return nil, errors.New("invalid resolver ", r)
|
}
|
||||||
|
|
||||||
|
func (c *XDNSResolverUDP) Build() (proto.Message, error) {
|
||||||
|
return &xdns.UDPResolverProto{Addr: c.Addr}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
var xdnsLoader = NewJSONConfigLoader(ConfigCreatorCache{
|
||||||
|
"tcp": func() interface{} { return new(XDNSResolverTCP) },
|
||||||
|
"udp": func() interface{} { return new(XDNSResolverUDP) },
|
||||||
|
}, "type", "settings")
|
||||||
|
|
||||||
|
type XDNSResolver struct {
|
||||||
|
Type string `json:"type"`
|
||||||
|
Settings json.RawMessage `json:"settings"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type XDNS struct {
|
||||||
|
Domains []XDNSDomain `json:"domains"`
|
||||||
|
Resolvers []XDNSResolver `json:"resolvers"`
|
||||||
|
ExtraPoll int32 `json:"extraPoll"`
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *XDNS) Build() (proto.Message, error) {
|
||||||
|
var domains []*xdns.DomainProto
|
||||||
|
var resolvers []*serial.TypedMessage
|
||||||
|
for i := range c.Domains {
|
||||||
|
if c.Domains[i].LenLimit == 0 {
|
||||||
|
c.Domains[i].LenLimit = 255
|
||||||
}
|
}
|
||||||
|
if c.Domains[i].LabelLimit == 0 {
|
||||||
|
c.Domains[i].LabelLimit = 63
|
||||||
|
}
|
||||||
|
types := make([]uint16, 0, len(c.Domains[i].Types))
|
||||||
|
for j := range c.Domains[i].Types {
|
||||||
|
types = append(types, uint16(c.Domains[i].Types[j]))
|
||||||
|
}
|
||||||
|
domain, err := xdns.NewDomain(c.Domains[i].Name, int(c.Domains[i].LenLimit), int(c.Domains[i].LabelLimit), types, uint16(c.Domains[i].Edns0))
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
errors.LogInfo(context.Background(), domain.Show())
|
||||||
|
domains = append(domains, &xdns.DomainProto{
|
||||||
|
Name: c.Domains[i].Name,
|
||||||
|
LenLimit: c.Domains[i].LenLimit,
|
||||||
|
LabelLimit: c.Domains[i].LabelLimit,
|
||||||
|
Types: c.Domains[i].Types,
|
||||||
|
Edns0: c.Domains[i].Edns0,
|
||||||
|
})
|
||||||
}
|
}
|
||||||
|
for i := range c.Resolvers {
|
||||||
return &xdns.Config{
|
config, err := xdnsLoader.LoadWithID(c.Resolvers[i].Settings, c.Resolvers[i].Type)
|
||||||
Domains: c.Domains,
|
if err != nil {
|
||||||
Resolvers: c.Resolvers,
|
return nil, err
|
||||||
}, nil
|
}
|
||||||
|
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")
|
||||||
|
}
|
||||||
|
return &xdns.Config{Domains: domains, Resolvers: resolvers, ExtraPoll: c.ExtraPoll}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
type XMC struct {
|
type XMC struct {
|
||||||
|
|||||||
@@ -0,0 +1,136 @@
|
|||||||
|
package conf
|
||||||
|
|
||||||
|
import (
|
||||||
|
"encoding/json"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/xtls/xray-core/transport/internet/finalmask/noise"
|
||||||
|
)
|
||||||
|
|
||||||
|
func expPacket(exp string) json.RawMessage {
|
||||||
|
b, _ := json.Marshal(exp)
|
||||||
|
return b
|
||||||
|
}
|
||||||
|
|
||||||
|
func buildNoiseExp(exp string) (*noise.Config, error) {
|
||||||
|
msg, err := (&NoiseMask{Noise: []NoiseItem{{Type: "exp", Packet: expPacket(exp)}}}).Build()
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
return msg.(*noise.Config), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestNoiseExp(t *testing.T) {
|
||||||
|
cfg, err := buildNoiseExp("<b 0d0a0d0a><t><r 24><rc 20-40><rd 8><c><n>")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
segments := cfg.Items[0].Segments
|
||||||
|
if len(segments) != 7 {
|
||||||
|
t.Fatalf("got %d segments, want 7", len(segments))
|
||||||
|
}
|
||||||
|
want := []struct {
|
||||||
|
kind noise.Segment_Kind
|
||||||
|
bytes []byte
|
||||||
|
min, max int64
|
||||||
|
}{
|
||||||
|
{noise.Segment_BYTES, []byte{0x0d, 0x0a, 0x0d, 0x0a}, 0, 0},
|
||||||
|
{noise.Segment_TIMESTAMP, nil, 0, 0},
|
||||||
|
{noise.Segment_RANDOM, nil, 24, 24},
|
||||||
|
{noise.Segment_RANDOM_ASCII, nil, 20, 40},
|
||||||
|
{noise.Segment_RANDOM_DIGIT, nil, 8, 8},
|
||||||
|
{noise.Segment_COUNTER, nil, 0, 0},
|
||||||
|
{noise.Segment_NONCE, nil, 0, 0},
|
||||||
|
}
|
||||||
|
for i, w := range want {
|
||||||
|
s := segments[i]
|
||||||
|
if s.Kind != w.kind || s.MinSize != w.min || s.MaxSize != w.max || string(s.Bytes) != string(w.bytes) {
|
||||||
|
t.Errorf("segment %d = %+v, want %+v", i, s, w)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestNoiseExpStripsHexPrefix(t *testing.T) {
|
||||||
|
cfg, err := buildNoiseExp("<b 0x16030100>")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if got := cfg.Items[0].Segments[0].Bytes; string(got) != string([]byte{0x16, 0x03, 0x01, 0x00}) {
|
||||||
|
t.Errorf("got %x", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestNoiseExpWhitespace(t *testing.T) {
|
||||||
|
if _, err := buildNoiseExp(" <b 00> <t> "); err != nil {
|
||||||
|
t.Errorf("surrounding whitespace should be allowed: %v", err)
|
||||||
|
}
|
||||||
|
cfg, err := buildNoiseExp("<b 0d 0a 0d 0a>")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if got := cfg.Items[0].Segments[0].Bytes; string(got) != "\r\n\r\n" {
|
||||||
|
t.Errorf("got %x", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestNoiseExpRejects(t *testing.T) {
|
||||||
|
for _, exp := range []string{
|
||||||
|
"<x 1>",
|
||||||
|
"<b>",
|
||||||
|
"<b zz>",
|
||||||
|
"<b 0d0>",
|
||||||
|
"<r>",
|
||||||
|
"<r -1>",
|
||||||
|
"<r 40-20>",
|
||||||
|
"<r 70000>",
|
||||||
|
"<t 5>",
|
||||||
|
"<n 5>",
|
||||||
|
"garbage<t>",
|
||||||
|
"<t> tail",
|
||||||
|
"<t><b>",
|
||||||
|
} {
|
||||||
|
if _, err := buildNoiseExp(exp); err == nil {
|
||||||
|
t.Errorf("expected an error for %q", exp)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestNoiseExpConflicts(t *testing.T) {
|
||||||
|
if _, err := (&NoiseMask{Noise: []NoiseItem{{Type: "exp", Packet: expPacket("<t>"), Rand: Int32Range{From: 10, To: 20}}}}).Build(); err == nil {
|
||||||
|
t.Error("exp with rand should be rejected")
|
||||||
|
}
|
||||||
|
for _, packet := range []string{``, `[1, 2]`, `5`} {
|
||||||
|
if _, err := (&NoiseMask{Noise: []NoiseItem{{Type: "exp", Packet: json.RawMessage(packet)}}}).Build(); err == nil {
|
||||||
|
t.Errorf("expected an error for packet %q", packet)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestNoiseExpFromJSON(t *testing.T) {
|
||||||
|
var mask NoiseMask
|
||||||
|
if err := json.Unmarshal([]byte(`{"noise": [
|
||||||
|
{"type": "exp", "packet": "<b 504f5354><rd 10-20>", "delay": "1-3"},
|
||||||
|
{"type": "EXP", "packet": "<t>"},
|
||||||
|
{"type": "str", "packet": "<t>"},
|
||||||
|
{"rand": "10-20"}
|
||||||
|
]}`), &mask); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
msg, err := mask.Build()
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
items := msg.(*noise.Config).Items
|
||||||
|
if len(items[0].Segments) != 2 || items[0].DelayMin != 1 || items[0].DelayMax != 3 {
|
||||||
|
t.Errorf("item 0 = %+v", items[0])
|
||||||
|
}
|
||||||
|
if len(items[1].Segments) != 1 || items[1].Segments[0].Kind != noise.Segment_TIMESTAMP {
|
||||||
|
t.Errorf("item 1 = %+v", items[1])
|
||||||
|
}
|
||||||
|
if len(items[2].Segments) != 0 || string(items[2].Packet) != "<t>" {
|
||||||
|
t.Errorf("item 2 = %+v", items[2])
|
||||||
|
}
|
||||||
|
if len(items[3].Segments) != 0 || items[3].RandMin != 10 || items[3].RandMax != 20 {
|
||||||
|
t.Errorf("item 3 = %+v", items[3])
|
||||||
|
}
|
||||||
|
}
|
||||||
+32
-2
@@ -5,8 +5,12 @@ import (
|
|||||||
"fmt"
|
"fmt"
|
||||||
"math/big"
|
"math/big"
|
||||||
"net"
|
"net"
|
||||||
|
"runtime"
|
||||||
|
"slices"
|
||||||
"strconv"
|
"strconv"
|
||||||
|
"strings"
|
||||||
|
|
||||||
|
"github.com/xtls/xray-core/common/errors"
|
||||||
"github.com/xtls/xray-core/proxy/tun"
|
"github.com/xtls/xray-core/proxy/tun"
|
||||||
"google.golang.org/protobuf/proto"
|
"google.golang.org/protobuf/proto"
|
||||||
)
|
)
|
||||||
@@ -20,7 +24,8 @@ type TunConfig struct {
|
|||||||
UserLevel uint32 `json:"userLevel"`
|
UserLevel uint32 `json:"userLevel"`
|
||||||
AutoSystemRoutingTable []string `json:"autoSystemRoutingTable"`
|
AutoSystemRoutingTable []string `json:"autoSystemRoutingTable"`
|
||||||
AutoOutboundsInterface *string `json:"autoOutboundsInterface"`
|
AutoOutboundsInterface *string `json:"autoOutboundsInterface"`
|
||||||
AutoSystemDNS bool `json:"autoSystemDNS"`
|
AutoSystemDnsToGateway bool `json:"autoSystemDnsToGateway"`
|
||||||
|
AutoSystemWfpBlockLeak []string `json:"autoSystemWfpBlockLeak"`
|
||||||
}
|
}
|
||||||
|
|
||||||
func (v *TunConfig) Build() (proto.Message, error) {
|
func (v *TunConfig) Build() (proto.Message, error) {
|
||||||
@@ -32,7 +37,32 @@ func (v *TunConfig) Build() (proto.Message, error) {
|
|||||||
DNS: v.DNS,
|
DNS: v.DNS,
|
||||||
UserLevel: v.UserLevel,
|
UserLevel: v.UserLevel,
|
||||||
AutoSystemRoutingTable: v.AutoSystemRoutingTable,
|
AutoSystemRoutingTable: v.AutoSystemRoutingTable,
|
||||||
AutoSystemDns: v.AutoSystemDNS,
|
AutoSystemDnsToGateway: v.AutoSystemDnsToGateway,
|
||||||
|
}
|
||||||
|
for _, leak := range v.AutoSystemWfpBlockLeak {
|
||||||
|
switch leak := strings.ToLower(leak); leak {
|
||||||
|
case "dns", "misconfigtun":
|
||||||
|
config.AutoSystemWfpBlockLeak = append(config.AutoSystemWfpBlockLeak, leak)
|
||||||
|
default:
|
||||||
|
return nil, errors.New("unknown autoSystemWfpBlockLeak value: ", leak)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
// Each option needs other settings on the system it takes effect on: the
|
||||||
|
// filters go along with the routes of autoSystemRoutingTable, "dns" lets
|
||||||
|
// DNS through the TUN only, and autoSystemDnsToGateway points the system
|
||||||
|
// DNS at the gateway.
|
||||||
|
switch runtime.GOOS {
|
||||||
|
case "windows":
|
||||||
|
if len(config.AutoSystemWfpBlockLeak) > 0 && len(v.AutoSystemRoutingTable) == 0 {
|
||||||
|
return nil, errors.New("autoSystemWfpBlockLeak needs autoSystemRoutingTable to be set")
|
||||||
|
}
|
||||||
|
if slices.Contains(config.AutoSystemWfpBlockLeak, "dns") && len(v.DNS) == 0 {
|
||||||
|
return nil, errors.New(`autoSystemWfpBlockLeak "dns" needs dns to be set`)
|
||||||
|
}
|
||||||
|
case "linux":
|
||||||
|
if v.AutoSystemDnsToGateway && len(v.Gateway) == 0 {
|
||||||
|
return nil, errors.New("autoSystemDnsToGateway needs gateway to be set")
|
||||||
|
}
|
||||||
}
|
}
|
||||||
if v.AutoOutboundsInterface != nil {
|
if v.AutoOutboundsInterface != nil {
|
||||||
config.AutoOutboundsInterface = *v.AutoOutboundsInterface
|
config.AutoOutboundsInterface = *v.AutoOutboundsInterface
|
||||||
|
|||||||
@@ -0,0 +1,71 @@
|
|||||||
|
package conf_test
|
||||||
|
|
||||||
|
import (
|
||||||
|
"encoding/json"
|
||||||
|
"runtime"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
. "github.com/xtls/xray-core/infra/conf"
|
||||||
|
"github.com/xtls/xray-core/proxy/tun"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestTunConfigAutoSystem(t *testing.T) {
|
||||||
|
creator := func() Buildable {
|
||||||
|
return new(TunConfig)
|
||||||
|
}
|
||||||
|
|
||||||
|
runMultiTestCase(t, []TestCase{
|
||||||
|
{
|
||||||
|
Input: `{"name": "xray0"}`,
|
||||||
|
Parser: loadJSON(creator),
|
||||||
|
Output: &tun.Config{Name: "xray0", Desc: "Wintun", MTU: 1500},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
Input: `{"name": "xray0", "gateway": ["10.0.0.1/24"], "autoSystemDnsToGateway": true}`,
|
||||||
|
Parser: loadJSON(creator),
|
||||||
|
Output: &tun.Config{Name: "xray0", Desc: "Wintun", MTU: 1500, Gateway: []string{"10.0.0.1/24"}, AutoSystemDnsToGateway: true},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
Input: `{"name": "xray0", "dns": ["1.1.1.1"], "autoSystemRoutingTable": ["0.0.0.0/0"], "autoSystemWfpBlockLeak": ["dns", "misconfigtun"]}`,
|
||||||
|
Parser: loadJSON(creator),
|
||||||
|
Output: &tun.Config{Name: "xray0", Desc: "Wintun", MTU: 1500, DNS: []string{"1.1.1.1"}, AutoSystemRoutingTable: []string{"0.0.0.0/0"}, AutoOutboundsInterface: "auto", AutoSystemWfpBlockLeak: []string{"dns", "misconfigtun"}},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
Input: `{"name": "xray0", "dns": ["1.1.1.1"], "autoSystemRoutingTable": ["0.0.0.0/0"], "autoSystemWfpBlockLeak": ["DNS"]}`,
|
||||||
|
Parser: loadJSON(creator),
|
||||||
|
Output: &tun.Config{Name: "xray0", Desc: "Wintun", MTU: 1500, DNS: []string{"1.1.1.1"}, AutoSystemRoutingTable: []string{"0.0.0.0/0"}, AutoOutboundsInterface: "auto", AutoSystemWfpBlockLeak: []string{"dns"}},
|
||||||
|
},
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestTunConfigAutoSystemNeeds checks that an option is rejected without the
|
||||||
|
// setting it needs, only on the system it takes effect on.
|
||||||
|
func TestTunConfigAutoSystemNeeds(t *testing.T) {
|
||||||
|
for _, c := range []struct {
|
||||||
|
input string
|
||||||
|
goos string // where it is rejected
|
||||||
|
}{
|
||||||
|
{`{"name": "xray0", "autoSystemWfpBlockLeak": ["misconfigtun"]}`, "windows"},
|
||||||
|
{`{"name": "xray0", "autoSystemRoutingTable": ["0.0.0.0/0"], "autoSystemWfpBlockLeak": ["misconfigtun"]}`, ""},
|
||||||
|
{`{"name": "xray0", "autoSystemRoutingTable": ["0.0.0.0/0"], "autoSystemWfpBlockLeak": ["dns"]}`, "windows"},
|
||||||
|
{`{"name": "xray0", "autoSystemDnsToGateway": true}`, "linux"},
|
||||||
|
} {
|
||||||
|
config := new(TunConfig)
|
||||||
|
if err := json.Unmarshal([]byte(c.input), config); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if _, err := config.Build(); (err != nil) != (runtime.GOOS == c.goos) {
|
||||||
|
t.Errorf("%s on %s: error = %v", c.input, runtime.GOOS, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestTunConfigAutoSystemWfpBlockLeakUnknown(t *testing.T) {
|
||||||
|
config := new(TunConfig)
|
||||||
|
if err := json.Unmarshal([]byte(`{"name": "xray0", "autoSystemWfpBlockLeak": ["dns", "ip"]}`), config); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if _, err := config.Build(); err == nil {
|
||||||
|
t.Error("an unknown autoSystemWfpBlockLeak value was accepted")
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -12,6 +12,7 @@ import (
|
|||||||
"github.com/xtls/xray-core/core"
|
"github.com/xtls/xray-core/core"
|
||||||
"github.com/xtls/xray-core/infra/conf"
|
"github.com/xtls/xray-core/infra/conf"
|
||||||
"github.com/xtls/xray-core/infra/conf/serial"
|
"github.com/xtls/xray-core/infra/conf/serial"
|
||||||
|
"github.com/xtls/xray-core/proxy/hysteria"
|
||||||
"github.com/xtls/xray-core/proxy/masque"
|
"github.com/xtls/xray-core/proxy/masque"
|
||||||
"github.com/xtls/xray-core/proxy/shadowsocks"
|
"github.com/xtls/xray-core/proxy/shadowsocks"
|
||||||
"github.com/xtls/xray-core/proxy/shadowsocks_2022"
|
"github.com/xtls/xray-core/proxy/shadowsocks_2022"
|
||||||
@@ -91,6 +92,8 @@ func extractInboundUsers(inb *core.InboundHandlerConfig) []*protocol.User {
|
|||||||
return ty.Users
|
return ty.Users
|
||||||
case *masque.ServerConfig:
|
case *masque.ServerConfig:
|
||||||
return ty.Users
|
return ty.Users
|
||||||
|
case *hysteria.ServerConfig:
|
||||||
|
return ty.Users
|
||||||
default:
|
default:
|
||||||
fmt.Println("unsupported inbound type")
|
fmt.Println("unsupported inbound type")
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -467,7 +467,7 @@ func NewPacketReader(conn net.Conn, h *Handler, defaultRule *FinalRule, UDPOverr
|
|||||||
if statConn != nil {
|
if statConn != nil {
|
||||||
counter = statConn.ReadCounter
|
counter = statConn.ReadCounter
|
||||||
}
|
}
|
||||||
if c, ok := iConn.(*internet.PacketConnWrapper); ok {
|
if c, ok := iConn.(*net.PacketConnWrapper); ok {
|
||||||
isOverridden := false
|
isOverridden := false
|
||||||
if UDPOverride.Address != nil || UDPOverride.Port != 0 {
|
if UDPOverride.Address != nil || UDPOverride.Port != 0 {
|
||||||
isOverridden = true
|
isOverridden = true
|
||||||
@@ -487,7 +487,7 @@ func NewPacketReader(conn net.Conn, h *Handler, defaultRule *FinalRule, UDPOverr
|
|||||||
}
|
}
|
||||||
|
|
||||||
type PacketReader struct {
|
type PacketReader struct {
|
||||||
*internet.PacketConnWrapper
|
*net.PacketConnWrapper
|
||||||
stats.Counter
|
stats.Counter
|
||||||
Handler *Handler
|
Handler *Handler
|
||||||
DefaultRule *FinalRule
|
DefaultRule *FinalRule
|
||||||
@@ -542,7 +542,7 @@ func NewPacketWriter(conn net.Conn, h *Handler, defaultRule *FinalRule, UDPOverr
|
|||||||
if statConn != nil {
|
if statConn != nil {
|
||||||
counter = statConn.WriteCounter
|
counter = statConn.WriteCounter
|
||||||
}
|
}
|
||||||
if c, ok := iConn.(*internet.PacketConnWrapper); ok {
|
if c, ok := iConn.(*net.PacketConnWrapper); ok {
|
||||||
// If DialDest is a domain, it will be resolved in dialer
|
// If DialDest is a domain, it will be resolved in dialer
|
||||||
// check this behavior and add it to map
|
// check this behavior and add it to map
|
||||||
resolvedUDPAddr := utils.NewTypedSyncMap[string, net.Address]()
|
resolvedUDPAddr := utils.NewTypedSyncMap[string, net.Address]()
|
||||||
@@ -563,7 +563,7 @@ func NewPacketWriter(conn net.Conn, h *Handler, defaultRule *FinalRule, UDPOverr
|
|||||||
}
|
}
|
||||||
|
|
||||||
type PacketWriter struct {
|
type PacketWriter struct {
|
||||||
*internet.PacketConnWrapper
|
*net.PacketConnWrapper
|
||||||
stats.Counter
|
stats.Counter
|
||||||
*Handler
|
*Handler
|
||||||
DefaultRule *FinalRule
|
DefaultRule *FinalRule
|
||||||
|
|||||||
@@ -151,7 +151,7 @@ func (c *Client) Process(ctx context.Context, link *transport.Link, dialer inter
|
|||||||
}
|
}
|
||||||
defer conn.Close()
|
defer conn.Close()
|
||||||
uc := &wireguard.UDPConnClient{
|
uc := &wireguard.UDPConnClient{
|
||||||
PacketConn: conn.(*internet.PacketConnWrapper).PacketConn,
|
PacketConn: conn.(*net.PacketConnWrapper).PacketConn,
|
||||||
Dest: conn.RemoteAddr().(*net.UDPAddr),
|
Dest: conn.RemoteAddr().(*net.UDPAddr),
|
||||||
}
|
}
|
||||||
reader = uc
|
reader = uc
|
||||||
|
|||||||
@@ -277,6 +277,7 @@ func (w *VisionReader) ReadMultiBuffer() (buf.MultiBuffer, error) {
|
|||||||
w.ob.CanSpliceCopy = 1
|
w.ob.CanSpliceCopy = 1
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
SuppressOuterCloseNotify(w.conn)
|
||||||
readerConn, readCounter, _ := UnwrapRawConn(w.conn)
|
readerConn, readCounter, _ := UnwrapRawConn(w.conn)
|
||||||
w.directReadCounter = readCounter
|
w.directReadCounter = readCounter
|
||||||
w.Reader = buf.NewReader(readerConn)
|
w.Reader = buf.NewReader(readerConn)
|
||||||
@@ -340,6 +341,7 @@ func (w *VisionWriter) WriteMultiBuffer(mb buf.MultiBuffer) error {
|
|||||||
// w.ob.CanSpliceCopy = 1
|
// w.ob.CanSpliceCopy = 1
|
||||||
// }
|
// }
|
||||||
}
|
}
|
||||||
|
SuppressOuterCloseNotify(w.conn)
|
||||||
rawConn, _, writerCounter := UnwrapRawConn(w.conn)
|
rawConn, _, writerCounter := UnwrapRawConn(w.conn)
|
||||||
w.Writer = buf.NewWriter(rawConn)
|
w.Writer = buf.NewWriter(rawConn)
|
||||||
w.directWriteCounter = writerCounter
|
w.directWriteCounter = writerCounter
|
||||||
@@ -669,6 +671,19 @@ func XtlsFilterTls(buffer buf.MultiBuffer, trafficState *TrafficState, ctx conte
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
type CloseNotifySuppressor interface {
|
||||||
|
SuppressCloseNotify()
|
||||||
|
}
|
||||||
|
|
||||||
|
// Close our local TLS conn instance might send a incorrect close_notify alert
|
||||||
|
// if the XTLS direct copy mode is enabled and cause TLS BAD_RECORD_MAC on users' browser
|
||||||
|
// Close the underlying connection directly to avoid this issue.
|
||||||
|
func SuppressOuterCloseNotify(conn net.Conn) {
|
||||||
|
if suppressor, ok := stat.TryUnwrapStatsConn(conn).(CloseNotifySuppressor); ok {
|
||||||
|
suppressor.SuppressCloseNotify()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
// UnwrapRawConn support unwrap encryption, stats, mask wrappers, tls, utls, reality, proxyproto, uds-wrapper conn and get raw tcp/uds conn from it
|
// UnwrapRawConn support unwrap encryption, stats, mask wrappers, tls, utls, reality, proxyproto, uds-wrapper conn and get raw tcp/uds conn from it
|
||||||
func UnwrapRawConn(conn net.Conn) (net.Conn, stats.Counter, stats.Counter) {
|
func UnwrapRawConn(conn net.Conn) (net.Conn, stats.Counter, stats.Counter) {
|
||||||
var readCounter, writerCounter stats.Counter
|
var readCounter, writerCounter stats.Counter
|
||||||
|
|||||||
@@ -2,7 +2,6 @@ package shadowsocks_2022
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"io"
|
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/xtls/xray-core/common"
|
"github.com/xtls/xray-core/common"
|
||||||
@@ -13,9 +12,6 @@ import (
|
|||||||
"github.com/xtls/xray-core/common/net"
|
"github.com/xtls/xray-core/common/net"
|
||||||
"github.com/xtls/xray-core/common/protocol"
|
"github.com/xtls/xray-core/common/protocol"
|
||||||
"github.com/xtls/xray-core/common/session"
|
"github.com/xtls/xray-core/common/session"
|
||||||
"github.com/xtls/xray-core/common/signal"
|
|
||||||
"github.com/xtls/xray-core/common/task"
|
|
||||||
"github.com/xtls/xray-core/common/utils"
|
|
||||||
"github.com/xtls/xray-core/core"
|
"github.com/xtls/xray-core/core"
|
||||||
"github.com/xtls/xray-core/features/policy"
|
"github.com/xtls/xray-core/features/policy"
|
||||||
"github.com/xtls/xray-core/features/routing"
|
"github.com/xtls/xray-core/features/routing"
|
||||||
@@ -101,35 +97,29 @@ func (i *Inbound) processTCP(ctx context.Context, conn net.Conn, dispatcher rout
|
|||||||
return errors.New("unable to set read deadline").Base(err)
|
return errors.New("unable to set read deadline").Base(err)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// 1. Single read call for Salt + Fixed-length header chunk per SIP022 §3.1.4
|
||||||
|
headerLen := i.method.KeySaltLength + RequestHeaderFixedChunkLength + AEADTagSize
|
||||||
|
headerBuf := make([]byte, headerLen)
|
||||||
|
n, err := conn.Read(headerBuf)
|
||||||
|
if err != nil || n < headerLen {
|
||||||
|
ResetTCPConn(conn)
|
||||||
|
return errors.New("failed to read complete handshake header")
|
||||||
|
}
|
||||||
|
|
||||||
var salt [32]byte
|
var salt [32]byte
|
||||||
|
copy(salt[:i.method.KeySaltLength], headerBuf[:i.method.KeySaltLength])
|
||||||
saltSlice := salt[:i.method.KeySaltLength]
|
saltSlice := salt[:i.method.KeySaltLength]
|
||||||
if _, err := io.ReadFull(conn, saltSlice); err != nil {
|
fixedChunk := headerBuf[i.method.KeySaltLength:]
|
||||||
return err
|
|
||||||
}
|
|
||||||
|
|
||||||
if !i.saltFilter.Check(salt) {
|
reader, reqHeader, err := InitServerStream(conn, i.method, i.psk, saltSlice, salt, fixedChunk, i.saltFilter)
|
||||||
return ErrSaltNotUnique
|
|
||||||
}
|
|
||||||
|
|
||||||
sessionKey := DeriveSessionSubKey(i.psk, saltSlice, i.method.KeySaltLength)
|
|
||||||
aead, err := i.method.NewAEAD(sessionKey)
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
ResetTCPConn(conn)
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
reader := NewStreamReader(conn, aead)
|
|
||||||
|
|
||||||
reqHeader, err := ReadClientRequestHeader(conn, reader)
|
|
||||||
if err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
conn.SetReadDeadline(time.Time{})
|
|
||||||
dest := reqHeader.Destination
|
dest := reqHeader.Destination
|
||||||
|
|
||||||
writer, err := WriteTCPResponse(conn, i.method, i.psk, saltSlice, nil)
|
writer := NewServerStreamWriter(conn, i.method, i.psk, saltSlice)
|
||||||
if err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
|
|
||||||
ctx = log.ContextWithAccessMessage(ctx, &log.AccessMessage{
|
ctx = log.ContextWithAccessMessage(ctx, &log.AccessMessage{
|
||||||
From: conn.RemoteAddr(),
|
From: conn.RemoteAddr(),
|
||||||
@@ -146,42 +136,17 @@ func (i *Inbound) processTCP(ctx context.Context, conn net.Conn, dispatcher rout
|
|||||||
}
|
}
|
||||||
|
|
||||||
if len(reqHeader.EarlyData) > 0 {
|
if len(reqHeader.EarlyData) > 0 {
|
||||||
earlyBuf := buf.New()
|
mb := buf.MergeBytes(nil, reqHeader.EarlyData)
|
||||||
earlyBuf.Write(reqHeader.EarlyData)
|
if err := link.Writer.WriteMultiBuffer(mb); err != nil {
|
||||||
if err := link.Writer.WriteMultiBuffer(buf.MultiBuffer{earlyBuf}); err != nil {
|
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
sessionPolicy = i.policyManager.ForLevel(uint32(i.user.Level))
|
return TransportTCP(ctx, i.policyManager.ForLevel(uint32(i.user.Level)), reader, writer, link)
|
||||||
ctx, cancel := context.WithCancel(ctx)
|
|
||||||
timer := signal.CancelAfterInactivity(ctx, cancel, sessionPolicy.Timeouts.ConnectionIdle)
|
|
||||||
ctx = policy.ContextWithBufferPolicy(ctx, sessionPolicy.Buffer)
|
|
||||||
|
|
||||||
requestDone := func() error {
|
|
||||||
defer timer.SetTimeout(sessionPolicy.Timeouts.DownlinkOnly)
|
|
||||||
return buf.Copy(reader, link.Writer, buf.UpdateActivity(timer))
|
|
||||||
}
|
|
||||||
|
|
||||||
responseDone := func() error {
|
|
||||||
defer timer.SetTimeout(sessionPolicy.Timeouts.UplinkOnly)
|
|
||||||
return buf.Copy(link.Reader, writer, buf.UpdateActivity(timer))
|
|
||||||
}
|
|
||||||
|
|
||||||
responseDoneAndCloseWriter := task.OnSuccess(responseDone, task.Close(link.Writer))
|
|
||||||
return task.Run(ctx, requestDone, responseDoneAndCloseWriter)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (i *Inbound) processUDP(ctx context.Context, conn stat.Connection, dispatcher routing.Dispatcher) error {
|
func (i *Inbound) processUDP(ctx context.Context, conn stat.Connection, dispatcher routing.Dispatcher) error {
|
||||||
udpConns := utils.NewTypedSyncMap[uint64, *udpConnEntry]()
|
reader := buf.NewPacketReader(conn)
|
||||||
defer func() {
|
|
||||||
udpConns.Range(func(key uint64, entry *udpConnEntry) bool {
|
|
||||||
entry.timer.SetTimeout(0)
|
|
||||||
return true
|
|
||||||
})
|
|
||||||
}()
|
|
||||||
|
|
||||||
reader := buf.NewReader(conn)
|
|
||||||
for {
|
for {
|
||||||
mb, err := reader.ReadMultiBuffer()
|
mb, err := reader.ReadMultiBuffer()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -191,75 +156,30 @@ func (i *Inbound) processUDP(ctx context.Context, conn stat.Connection, dispatch
|
|||||||
|
|
||||||
for _, b := range mb {
|
for _, b := range mb {
|
||||||
decoded, err := i.udpCodec.DecodePacket(b.Bytes())
|
decoded, err := i.udpCodec.DecodePacket(b.Bytes())
|
||||||
if err != nil {
|
b.Release()
|
||||||
b.Release()
|
if err != nil || decoded.HeaderType != HeaderTypeClient {
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
|
||||||
entry, ok := udpConns.Load(decoded.SessionID)
|
sessionItem := i.udpCodec.GetSession(decoded.SessionID)
|
||||||
if !ok {
|
if sessionItem.User == nil {
|
||||||
sessCtx, cancel := context.WithCancel(ctx)
|
sessionItem.Lock()
|
||||||
sessCtx = log.ContextWithAccessMessage(sessCtx, &log.AccessMessage{
|
if sessionItem.User == nil {
|
||||||
From: conn.RemoteAddr(),
|
sessionItem.User = i.user
|
||||||
To: decoded.Destination,
|
|
||||||
Status: log.AccessAccepted,
|
|
||||||
Email: i.user.Email,
|
|
||||||
})
|
|
||||||
|
|
||||||
link, err := dispatcher.Dispatch(sessCtx, decoded.Destination)
|
|
||||||
if err != nil {
|
|
||||||
cancel()
|
|
||||||
b.Release()
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
|
|
||||||
newEntry := &udpConnEntry{
|
|
||||||
link: link,
|
|
||||||
cancel: cancel,
|
|
||||||
}
|
|
||||||
sessionPolicy := i.policyManager.ForLevel(uint32(i.user.Level))
|
|
||||||
newEntry.timer = signal.CancelAfterInactivity(sessCtx, func() {
|
|
||||||
udpConns.Delete(decoded.SessionID)
|
|
||||||
common.Interrupt(link.Reader)
|
|
||||||
common.Interrupt(link.Writer)
|
|
||||||
cancel()
|
|
||||||
}, sessionPolicy.Timeouts.ConnectionIdle)
|
|
||||||
|
|
||||||
actual, loaded := udpConns.LoadOrStore(decoded.SessionID, newEntry)
|
|
||||||
if loaded {
|
|
||||||
// Another goroutine/packet beat us to storing, terminate our redundant link
|
|
||||||
newEntry.timer.SetTimeout(0)
|
|
||||||
entry = actual
|
|
||||||
} else {
|
|
||||||
entry = newEntry
|
|
||||||
go func(sessID uint64, dest net.Destination, cEntry *udpConnEntry) {
|
|
||||||
defer func() {
|
|
||||||
cEntry.timer.SetTimeout(0)
|
|
||||||
}()
|
|
||||||
for {
|
|
||||||
resMb, err := cEntry.link.Reader.ReadMultiBuffer()
|
|
||||||
if err != nil {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
cEntry.timer.Update()
|
|
||||||
for _, rb := range resMb {
|
|
||||||
encPacket, err := i.udpCodec.EncodeServerPacket(sessID, dest, rb.Bytes())
|
|
||||||
rb.Release()
|
|
||||||
if err != nil {
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
_, _ = conn.Write(encPacket)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}(decoded.SessionID, decoded.Destination, entry)
|
|
||||||
}
|
}
|
||||||
|
sessionItem.Unlock()
|
||||||
|
}
|
||||||
|
link, err := sessionItem.EnsureLink(ctx, conn, decoded.Destination, dispatcher, i.policyManager, func(dest net.Destination, payload []byte) ([]byte, error) {
|
||||||
|
return i.udpCodec.EncodeServerPacket(decoded.SessionID, dest, payload)
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
continue
|
||||||
}
|
}
|
||||||
|
|
||||||
entry.timer.Update()
|
|
||||||
payloadBuf := buf.New()
|
payloadBuf := buf.New()
|
||||||
payloadBuf.Write(decoded.Payload)
|
payloadBuf.Write(decoded.Payload)
|
||||||
b.Release()
|
payloadBuf.UDP = &decoded.Destination
|
||||||
_ = entry.link.Writer.WriteMultiBuffer(buf.MultiBuffer{payloadBuf})
|
_ = link.Writer.WriteMultiBuffer(buf.MultiBuffer{payloadBuf})
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -4,7 +4,6 @@ import (
|
|||||||
"context"
|
"context"
|
||||||
"crypto/cipher"
|
"crypto/cipher"
|
||||||
"encoding/binary"
|
"encoding/binary"
|
||||||
"io"
|
|
||||||
"strconv"
|
"strconv"
|
||||||
"strings"
|
"strings"
|
||||||
"sync"
|
"sync"
|
||||||
@@ -19,8 +18,6 @@ import (
|
|||||||
"github.com/xtls/xray-core/common/net"
|
"github.com/xtls/xray-core/common/net"
|
||||||
"github.com/xtls/xray-core/common/protocol"
|
"github.com/xtls/xray-core/common/protocol"
|
||||||
"github.com/xtls/xray-core/common/session"
|
"github.com/xtls/xray-core/common/session"
|
||||||
"github.com/xtls/xray-core/common/signal"
|
|
||||||
"github.com/xtls/xray-core/common/task"
|
|
||||||
"github.com/xtls/xray-core/common/utils"
|
"github.com/xtls/xray-core/common/utils"
|
||||||
"github.com/xtls/xray-core/common/uuid"
|
"github.com/xtls/xray-core/common/uuid"
|
||||||
"github.com/xtls/xray-core/core"
|
"github.com/xtls/xray-core/core"
|
||||||
@@ -207,64 +204,46 @@ func (i *MultiUserInbound) processTCP(ctx context.Context, conn net.Conn, dispat
|
|||||||
return errors.New("unable to set read deadline").Base(err)
|
return errors.New("unable to set read deadline").Base(err)
|
||||||
}
|
}
|
||||||
|
|
||||||
// 1. Read Request Salt (16 or 32 bytes)
|
// 1. Single read call for Salt + EIH + Fixed-length header chunk per SIP022 §3.1.4
|
||||||
|
headerLen := i.method.KeySaltLength + AESBlockSize + RequestHeaderFixedChunkLength + AEADTagSize
|
||||||
|
headerBuf := make([]byte, headerLen)
|
||||||
|
n, err := conn.Read(headerBuf)
|
||||||
|
if err != nil || n < headerLen {
|
||||||
|
ResetTCPConn(conn)
|
||||||
|
return errors.New("failed to read complete handshake header")
|
||||||
|
}
|
||||||
|
|
||||||
var salt [32]byte
|
var salt [32]byte
|
||||||
|
copy(salt[:i.method.KeySaltLength], headerBuf[:i.method.KeySaltLength])
|
||||||
saltSlice := salt[:i.method.KeySaltLength]
|
saltSlice := salt[:i.method.KeySaltLength]
|
||||||
if _, err := io.ReadFull(conn, saltSlice); err != nil {
|
eih := headerBuf[i.method.KeySaltLength : i.method.KeySaltLength+AESBlockSize]
|
||||||
return err
|
fixedChunk := headerBuf[i.method.KeySaltLength+AESBlockSize:]
|
||||||
}
|
|
||||||
|
|
||||||
if !i.saltFilter.Check(salt) {
|
decryptedHash, err := DecryptEIH(i.method, i.masterPSK, saltSlice, eih)
|
||||||
return ErrSaltNotUnique
|
|
||||||
}
|
|
||||||
|
|
||||||
// 2. Read Extended Identity Header (16 bytes)
|
|
||||||
var eih [AESBlockSize]byte
|
|
||||||
if _, err := io.ReadFull(conn, eih[:]); err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
|
|
||||||
// Decrypt EIH with IdentitySubKey derived from masterPSK and salt
|
|
||||||
identitySubkey := DeriveIdentitySubKey(i.masterPSK, saltSlice, i.method.KeySaltLength)
|
|
||||||
block, err := i.method.NewBlock(identitySubkey)
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
ResetTCPConn(conn)
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
var decryptedHash [AESBlockSize]byte
|
|
||||||
block.Decrypt(decryptedHash[:], eih[:])
|
|
||||||
|
|
||||||
// Lookup user
|
// Lookup user
|
||||||
user, ok := i.usersByHash.Load(decryptedHash)
|
user, ok := i.usersByHash.Load(decryptedHash)
|
||||||
if !ok || user == nil {
|
if !ok {
|
||||||
|
ResetTCPConn(conn)
|
||||||
return ErrInvalidRequest
|
return ErrInvalidRequest
|
||||||
}
|
}
|
||||||
userPSK := user.Account.(*MemoryAccount).Key
|
userPSK := user.Account.(*MemoryAccount).Key
|
||||||
|
|
||||||
// 3. Derive Session Subkey using matched user's PSK
|
reader, reqHeader, err := InitServerStream(conn, i.method, userPSK, saltSlice, salt, fixedChunk, i.saltFilter)
|
||||||
sessionKey := DeriveSessionSubKey(userPSK, saltSlice, i.method.KeySaltLength)
|
|
||||||
aead, err := i.method.NewAEAD(sessionKey)
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
ResetTCPConn(conn)
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
reader := NewStreamReader(conn, aead)
|
|
||||||
|
|
||||||
// 4 & 5. Read Client Request Header
|
|
||||||
reqHeader, err := ReadClientRequestHeader(conn, reader)
|
|
||||||
if err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
conn.SetReadDeadline(time.Time{})
|
|
||||||
dest := reqHeader.Destination
|
dest := reqHeader.Destination
|
||||||
|
|
||||||
// 6. Send Server Response Handshake
|
writer := NewServerStreamWriter(conn, i.method, userPSK, saltSlice)
|
||||||
writer, err := WriteTCPResponse(conn, i.method, userPSK, saltSlice, nil)
|
|
||||||
if err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
|
|
||||||
// 7. Dispatch Connection to Xray routing with matched User
|
// Dispatch Connection to Xray routing with matched User
|
||||||
inbound := session.InboundFromContext(ctx)
|
inbound := session.InboundFromContext(ctx)
|
||||||
inbound.User = user
|
inbound.User = user
|
||||||
|
|
||||||
@@ -283,42 +262,17 @@ func (i *MultiUserInbound) processTCP(ctx context.Context, conn net.Conn, dispat
|
|||||||
}
|
}
|
||||||
|
|
||||||
if len(reqHeader.EarlyData) > 0 {
|
if len(reqHeader.EarlyData) > 0 {
|
||||||
earlyBuf := buf.New()
|
mb := buf.MergeBytes(nil, reqHeader.EarlyData)
|
||||||
earlyBuf.Write(reqHeader.EarlyData)
|
if err := link.Writer.WriteMultiBuffer(mb); err != nil {
|
||||||
if err := link.Writer.WriteMultiBuffer(buf.MultiBuffer{earlyBuf}); err != nil {
|
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
sessionPolicy = i.policyManager.ForLevel(user.Level)
|
return TransportTCP(ctx, i.policyManager.ForLevel(user.Level), reader, writer, link)
|
||||||
ctx, cancel := context.WithCancel(ctx)
|
|
||||||
timer := signal.CancelAfterInactivity(ctx, cancel, sessionPolicy.Timeouts.ConnectionIdle)
|
|
||||||
ctx = policy.ContextWithBufferPolicy(ctx, sessionPolicy.Buffer)
|
|
||||||
|
|
||||||
requestDone := func() error {
|
|
||||||
defer timer.SetTimeout(sessionPolicy.Timeouts.DownlinkOnly)
|
|
||||||
return buf.Copy(reader, link.Writer, buf.UpdateActivity(timer))
|
|
||||||
}
|
|
||||||
|
|
||||||
responseDone := func() error {
|
|
||||||
defer timer.SetTimeout(sessionPolicy.Timeouts.UplinkOnly)
|
|
||||||
return buf.Copy(link.Reader, writer, buf.UpdateActivity(timer))
|
|
||||||
}
|
|
||||||
|
|
||||||
responseDoneAndCloseWriter := task.OnSuccess(responseDone, task.Close(link.Writer))
|
|
||||||
return task.Run(ctx, requestDone, responseDoneAndCloseWriter)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (i *MultiUserInbound) processUDP(ctx context.Context, conn stat.Connection, dispatcher routing.Dispatcher) error {
|
func (i *MultiUserInbound) processUDP(ctx context.Context, conn stat.Connection, dispatcher routing.Dispatcher) error {
|
||||||
udpConns := utils.NewTypedSyncMap[uint64, *udpConnEntry]()
|
reader := buf.NewPacketReader(conn)
|
||||||
defer func() {
|
|
||||||
udpConns.Range(func(key uint64, entry *udpConnEntry) bool {
|
|
||||||
entry.timer.SetTimeout(0)
|
|
||||||
return true
|
|
||||||
})
|
|
||||||
}()
|
|
||||||
|
|
||||||
reader := buf.NewReader(conn)
|
|
||||||
for {
|
for {
|
||||||
mb, err := reader.ReadMultiBuffer()
|
mb, err := reader.ReadMultiBuffer()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -342,168 +296,61 @@ func (i *MultiUserInbound) processUDP(ctx context.Context, conn stat.Connection,
|
|||||||
sessionID := binary.BigEndian.Uint64(rawHeader[:8])
|
sessionID := binary.BigEndian.Uint64(rawHeader[:8])
|
||||||
packetID := binary.BigEndian.Uint64(rawHeader[8:16])
|
packetID := binary.BigEndian.Uint64(rawHeader[8:16])
|
||||||
|
|
||||||
// Replay protection & session lookup
|
|
||||||
sessionItem := i.udpSessions.GetOrCreate(sessionID)
|
sessionItem := i.udpSessions.GetOrCreate(sessionID)
|
||||||
|
|
||||||
sessionItem.Lock()
|
if !sessionItem.CheckPacketID(packetID) {
|
||||||
if !sessionItem.Window.Check(packetID) {
|
|
||||||
sessionItem.Unlock()
|
|
||||||
b.Release()
|
b.Release()
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
|
||||||
var userPSK []byte
|
var userPSK []byte
|
||||||
var currentUser *protocol.MemoryUser
|
var currentUser *protocol.MemoryUser
|
||||||
if sessionItem.User != nil {
|
sessionItem.Lock()
|
||||||
currentUser = sessionItem.User
|
currentUser = sessionItem.User
|
||||||
userPSK = sessionItem.UserPSK
|
userPSK = sessionItem.UserPSK
|
||||||
sessionItem.Unlock()
|
sessionItem.Unlock()
|
||||||
} else {
|
|
||||||
sessionItem.Unlock()
|
|
||||||
// Decrypt EIH
|
|
||||||
identitySubkey := DeriveIdentitySubKey(i.masterPSK, rawHeader[:8], i.method.KeySaltLength)
|
|
||||||
idBlock, err := i.method.NewBlock(identitySubkey)
|
|
||||||
if err != nil {
|
|
||||||
b.Release()
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
|
|
||||||
var decryptedHash [16]byte
|
if currentUser == nil {
|
||||||
idBlock.Decrypt(decryptedHash[:], packetBytes[16:32])
|
// Decrypt EIH
|
||||||
|
decryptedHash := DecryptUDPEIH(i.udpMasterCipher, rawHeader[:], packetBytes[16:32])
|
||||||
|
|
||||||
user, ok := i.usersByHash.Load(decryptedHash)
|
user, ok := i.usersByHash.Load(decryptedHash)
|
||||||
if !ok || user == nil {
|
if !ok {
|
||||||
b.Release()
|
b.Release()
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
currentUser = user
|
currentUser = user
|
||||||
userPSK = user.Account.(*MemoryAccount).Key
|
userPSK = user.Account.(*MemoryAccount).Key
|
||||||
|
|
||||||
sessionItem.Lock()
|
|
||||||
sessionItem.User = user
|
|
||||||
sessionItem.UserPSK = userPSK
|
|
||||||
sessionItem.Unlock()
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// Decrypt Body (with AEAD caching per session)
|
decoded, err := sessionItem.DecryptAESPayload(i.method, userPSK, sessionID, packetID, rawHeader[:], packetBytes[32:])
|
||||||
bodyAead := sessionItem.GetRemoteCipher()
|
|
||||||
if bodyAead == nil {
|
|
||||||
bodyKey := DeriveSessionSubKey(userPSK, rawHeader[:8], i.method.KeySaltLength)
|
|
||||||
var err error
|
|
||||||
bodyAead, err = i.method.NewAEAD(bodyKey)
|
|
||||||
if err != nil {
|
|
||||||
b.Release()
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
sessionItem.SetRemoteCipher(bodyAead)
|
|
||||||
}
|
|
||||||
|
|
||||||
bodyNonce := rawHeader[4:16]
|
|
||||||
bodyCipher := packetBytes[32:]
|
|
||||||
bodyPlain, err := bodyAead.Open(nil, bodyNonce, bodyCipher, nil)
|
|
||||||
b.Release()
|
b.Release()
|
||||||
if err != nil || len(bodyPlain) < 1+8+2 {
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
|
|
||||||
sessionItem.Lock()
|
|
||||||
sessionItem.Window.Add(packetID)
|
|
||||||
sessionItem.Unlock()
|
|
||||||
|
|
||||||
if bodyPlain[0] != HeaderTypeClient {
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
epoch := binary.BigEndian.Uint64(bodyPlain[1:9])
|
|
||||||
diff := time.Now().Unix() - int64(epoch)
|
|
||||||
if diff < -30 || diff > 30 {
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
|
|
||||||
paddingLen := int(binary.BigEndian.Uint16(bodyPlain[9:11]))
|
|
||||||
offset := 11 + paddingLen
|
|
||||||
if len(bodyPlain) < offset {
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
|
|
||||||
dest, addrLen, err := parseAddressPort(bodyPlain[offset:])
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
|
||||||
payload := bodyPlain[offset+addrLen:]
|
sessionItem.Lock()
|
||||||
|
if sessionItem.User == nil {
|
||||||
|
sessionItem.User = currentUser
|
||||||
|
sessionItem.UserPSK = userPSK
|
||||||
|
}
|
||||||
|
sessionItem.Unlock()
|
||||||
|
|
||||||
entry, ok := udpConns.Load(sessionID)
|
link, err := sessionItem.EnsureLink(ctx, conn, decoded.Destination, dispatcher, i.policyManager, func(replyDest net.Destination, payload []byte) ([]byte, error) {
|
||||||
if !ok {
|
return i.encodeServerUDPPacket(sessionID, userPSK, replyDest, payload)
|
||||||
sessCtx, cancel := context.WithCancel(ctx)
|
})
|
||||||
inbound := session.InboundFromContext(sessCtx)
|
if err != nil {
|
||||||
inbound.User = currentUser
|
continue
|
||||||
|
|
||||||
sessCtx = log.ContextWithAccessMessage(sessCtx, &log.AccessMessage{
|
|
||||||
From: conn.RemoteAddr(),
|
|
||||||
To: dest,
|
|
||||||
Status: log.AccessAccepted,
|
|
||||||
Email: currentUser.Email,
|
|
||||||
})
|
|
||||||
|
|
||||||
link, err := dispatcher.Dispatch(sessCtx, dest)
|
|
||||||
if err != nil {
|
|
||||||
cancel()
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
|
|
||||||
newEntry := &udpConnEntry{
|
|
||||||
link: link,
|
|
||||||
cancel: cancel,
|
|
||||||
}
|
|
||||||
sessionPolicy := i.policyManager.ForLevel(currentUser.Level)
|
|
||||||
newEntry.timer = signal.CancelAfterInactivity(sessCtx, func() {
|
|
||||||
udpConns.Delete(sessionID)
|
|
||||||
common.Interrupt(link.Reader)
|
|
||||||
common.Interrupt(link.Writer)
|
|
||||||
cancel()
|
|
||||||
}, sessionPolicy.Timeouts.ConnectionIdle)
|
|
||||||
|
|
||||||
actual, loaded := udpConns.LoadOrStore(sessionID, newEntry)
|
|
||||||
if loaded {
|
|
||||||
newEntry.timer.SetTimeout(0)
|
|
||||||
entry = actual
|
|
||||||
} else {
|
|
||||||
entry = newEntry
|
|
||||||
go func(sessID uint64, uPSK []byte, d net.Destination, cEntry *udpConnEntry) {
|
|
||||||
defer func() {
|
|
||||||
cEntry.timer.SetTimeout(0)
|
|
||||||
}()
|
|
||||||
for {
|
|
||||||
resMb, err := cEntry.link.Reader.ReadMultiBuffer()
|
|
||||||
if err != nil {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
cEntry.timer.Update()
|
|
||||||
for _, rb := range resMb {
|
|
||||||
encPacket, err := i.encodeServerUDPPacket(sessID, uPSK, d, rb.Bytes())
|
|
||||||
rb.Release()
|
|
||||||
if err != nil {
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
_, _ = conn.Write(encPacket)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}(sessionID, userPSK, dest, entry)
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
entry.timer.Update()
|
|
||||||
pBuf := buf.New()
|
pBuf := buf.New()
|
||||||
pBuf.Write(payload)
|
pBuf.Write(decoded.Payload)
|
||||||
_ = entry.link.Writer.WriteMultiBuffer(buf.MultiBuffer{pBuf})
|
pBuf.UDP = &decoded.Destination
|
||||||
|
_ = link.Writer.WriteMultiBuffer(buf.MultiBuffer{pBuf})
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (i *MultiUserInbound) encodeServerUDPPacket(clientSessionID uint64, userPSK []byte, dest net.Destination, payload []byte) ([]byte, error) {
|
func (i *MultiUserInbound) encodeServerUDPPacket(clientSessionID uint64, userPSK []byte, dest net.Destination, payload []byte) ([]byte, error) {
|
||||||
sessionItem := i.udpSessions.GetOrCreate(clientSessionID)
|
return i.udpSessions.EncodeServerPacket(i.method, userPSK, clientSessionID, dest, payload)
|
||||||
if err := sessionItem.EnsureServerState(i.method, i.udpMasterCipher, nil, userPSK); err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
return sessionItem.EncodeServerPacket(i.method, clientSessionID, dest, payload)
|
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -4,7 +4,6 @@ import (
|
|||||||
"context"
|
"context"
|
||||||
"crypto/cipher"
|
"crypto/cipher"
|
||||||
"encoding/binary"
|
"encoding/binary"
|
||||||
"io"
|
|
||||||
"strconv"
|
"strconv"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
@@ -15,9 +14,6 @@ import (
|
|||||||
"github.com/xtls/xray-core/common/net"
|
"github.com/xtls/xray-core/common/net"
|
||||||
"github.com/xtls/xray-core/common/protocol"
|
"github.com/xtls/xray-core/common/protocol"
|
||||||
"github.com/xtls/xray-core/common/session"
|
"github.com/xtls/xray-core/common/session"
|
||||||
"github.com/xtls/xray-core/common/signal"
|
|
||||||
"github.com/xtls/xray-core/common/task"
|
|
||||||
"github.com/xtls/xray-core/common/utils"
|
|
||||||
"github.com/xtls/xray-core/common/uuid"
|
"github.com/xtls/xray-core/common/uuid"
|
||||||
"github.com/xtls/xray-core/core"
|
"github.com/xtls/xray-core/core"
|
||||||
"github.com/xtls/xray-core/features/policy"
|
"github.com/xtls/xray-core/features/policy"
|
||||||
@@ -35,18 +31,17 @@ type relayDest struct {
|
|||||||
destination net.Destination
|
destination net.Destination
|
||||||
email string
|
email string
|
||||||
level uint32
|
level uint32
|
||||||
key []byte
|
|
||||||
blockCipher cipher.Block
|
blockCipher cipher.Block
|
||||||
}
|
}
|
||||||
|
|
||||||
type RelayInbound struct {
|
type RelayInbound struct {
|
||||||
networks []net.Network
|
networks []net.Network
|
||||||
method *CipherMethod
|
method *CipherMethod
|
||||||
relayPSK []byte
|
relayPSK []byte
|
||||||
relayBlock cipher.Block
|
relayBlock cipher.Block
|
||||||
destinations map[[AESBlockSize]byte]*relayDest
|
destinations map[[AESBlockSize]byte]*relayDest
|
||||||
rawDestinations []*RelayDestination
|
udpSessions *UDPSessionManager
|
||||||
policyManager policy.Manager
|
policyManager policy.Manager
|
||||||
}
|
}
|
||||||
|
|
||||||
func NewRelayServer(ctx context.Context, config *RelayServerConfig) (*RelayInbound, error) {
|
func NewRelayServer(ctx context.Context, config *RelayServerConfig) (*RelayInbound, error) {
|
||||||
@@ -78,13 +73,13 @@ func NewRelayServer(ctx context.Context, config *RelayServerConfig) (*RelayInbou
|
|||||||
|
|
||||||
v := core.MustFromContext(ctx)
|
v := core.MustFromContext(ctx)
|
||||||
i := &RelayInbound{
|
i := &RelayInbound{
|
||||||
networks: networks,
|
networks: networks,
|
||||||
method: method,
|
method: method,
|
||||||
relayPSK: relayPSK,
|
relayPSK: relayPSK,
|
||||||
relayBlock: relayBlock,
|
relayBlock: relayBlock,
|
||||||
destinations: make(map[[AESBlockSize]byte]*relayDest),
|
destinations: make(map[[AESBlockSize]byte]*relayDest),
|
||||||
rawDestinations: config.Destinations,
|
udpSessions: NewUDPSessionManager(500 * time.Second),
|
||||||
policyManager: v.GetFeature(policy.ManagerType()).(policy.Manager),
|
policyManager: v.GetFeature(policy.ManagerType()).(policy.Manager),
|
||||||
}
|
}
|
||||||
|
|
||||||
for idx, d := range config.Destinations {
|
for idx, d := range config.Destinations {
|
||||||
@@ -108,7 +103,6 @@ func NewRelayServer(ctx context.Context, config *RelayServerConfig) (*RelayInbou
|
|||||||
destination: net.TCPDestination(d.Address.AsAddress(), net.Port(d.Port)),
|
destination: net.TCPDestination(d.Address.AsAddress(), net.Port(d.Port)),
|
||||||
email: d.Email,
|
email: d.Email,
|
||||||
level: uint32(d.Level),
|
level: uint32(d.Level),
|
||||||
key: destKey,
|
|
||||||
blockCipher: destBlock,
|
blockCipher: destBlock,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -139,28 +133,36 @@ func (i *RelayInbound) processTCP(ctx context.Context, conn net.Conn, dispatcher
|
|||||||
return errors.New("unable to set read deadline").Base(err)
|
return errors.New("unable to set read deadline").Base(err)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Read Salt + Outer EIH
|
// Read initial handshake in a single read call per SIP022 §3.1.3 & §3.1.4
|
||||||
needed := i.method.KeySaltLength + AESBlockSize
|
needed := i.method.KeySaltLength + AESBlockSize
|
||||||
var headerBuf [48]byte
|
requestHeader := buf.New()
|
||||||
headerSlice := headerBuf[:needed]
|
n, err := requestHeader.ReadFrom(conn)
|
||||||
if _, err := io.ReadFull(conn, headerSlice); err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
|
|
||||||
salt := headerSlice[:i.method.KeySaltLength]
|
|
||||||
eih := headerSlice[i.method.KeySaltLength:]
|
|
||||||
|
|
||||||
identitySubkey := DeriveIdentitySubKey(i.relayPSK, salt, i.method.KeySaltLength)
|
|
||||||
block, err := i.method.NewBlock(identitySubkey)
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
requestHeader.Release()
|
||||||
|
ResetTCPConn(conn)
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
if int(n) < needed {
|
||||||
|
requestHeader.Release()
|
||||||
|
ResetTCPConn(conn)
|
||||||
|
return ErrInvalidRequest
|
||||||
|
}
|
||||||
|
|
||||||
var decryptedHash [AESBlockSize]byte
|
headerSlice := requestHeader.Bytes()
|
||||||
block.Decrypt(decryptedHash[:], eih)
|
salt := headerSlice[:i.method.KeySaltLength]
|
||||||
|
eih := headerSlice[i.method.KeySaltLength:needed]
|
||||||
|
|
||||||
|
decryptedHash, err := DecryptEIH(i.method, i.relayPSK, salt, eih)
|
||||||
|
if err != nil {
|
||||||
|
requestHeader.Release()
|
||||||
|
ResetTCPConn(conn)
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
targetDest, ok := i.destinations[decryptedHash]
|
targetDest, ok := i.destinations[decryptedHash]
|
||||||
if !ok {
|
if !ok {
|
||||||
|
requestHeader.Release()
|
||||||
|
ResetTCPConn(conn)
|
||||||
return ErrInvalidRequest
|
return ErrInvalidRequest
|
||||||
}
|
}
|
||||||
conn.SetReadDeadline(time.Time{})
|
conn.SetReadDeadline(time.Time{})
|
||||||
@@ -182,45 +184,26 @@ func (i *RelayInbound) processTCP(ctx context.Context, conn net.Conn, dispatcher
|
|||||||
|
|
||||||
link, err := dispatcher.Dispatch(ctx, targetDest.destination)
|
link, err := dispatcher.Dispatch(ctx, targetDest.destination)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
requestHeader.Release()
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
// Unwrap outer EIH: send client salt to next hop, stripping this hop's EIH
|
// Unwrap outer EIH: send client salt and remaining handshake bytes to next hop
|
||||||
saltBuf := buf.New()
|
// in a single write call, satisfying downstream server's single-read handshake expectation (SIP022 §3.1.3).
|
||||||
saltBuf.Write(salt)
|
var saltCopy [32]byte
|
||||||
if err := link.Writer.WriteMultiBuffer(buf.MultiBuffer{saltBuf}); err != nil {
|
copy(saltCopy[:i.method.KeySaltLength], salt)
|
||||||
|
copy(requestHeader.Bytes()[AESBlockSize:AESBlockSize+i.method.KeySaltLength], saltCopy[:i.method.KeySaltLength])
|
||||||
|
requestHeader.Advance(AESBlockSize)
|
||||||
|
|
||||||
|
if err := link.Writer.WriteMultiBuffer(buf.MultiBuffer{requestHeader}); err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
sessionPolicy = i.policyManager.ForLevel(targetDest.level)
|
return TransportTCP(ctx, i.policyManager.ForLevel(targetDest.level), buf.NewReader(conn), buf.NewWriter(conn), link)
|
||||||
ctx, cancel := context.WithCancel(ctx)
|
|
||||||
timer := signal.CancelAfterInactivity(ctx, cancel, sessionPolicy.Timeouts.ConnectionIdle)
|
|
||||||
ctx = policy.ContextWithBufferPolicy(ctx, sessionPolicy.Buffer)
|
|
||||||
|
|
||||||
requestDone := func() error {
|
|
||||||
defer timer.SetTimeout(sessionPolicy.Timeouts.DownlinkOnly)
|
|
||||||
return buf.Copy(buf.NewReader(conn), link.Writer, buf.UpdateActivity(timer))
|
|
||||||
}
|
|
||||||
|
|
||||||
responseDone := func() error {
|
|
||||||
defer timer.SetTimeout(sessionPolicy.Timeouts.UplinkOnly)
|
|
||||||
return buf.Copy(link.Reader, buf.NewWriter(conn), buf.UpdateActivity(timer))
|
|
||||||
}
|
|
||||||
|
|
||||||
responseDoneAndCloseWriter := task.OnSuccess(responseDone, task.Close(link.Writer))
|
|
||||||
return task.Run(ctx, requestDone, responseDoneAndCloseWriter)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (i *RelayInbound) processUDP(ctx context.Context, conn stat.Connection, dispatcher routing.Dispatcher) error {
|
func (i *RelayInbound) processUDP(ctx context.Context, conn stat.Connection, dispatcher routing.Dispatcher) error {
|
||||||
udpConns := utils.NewTypedSyncMap[uint64, *udpConnEntry]()
|
reader := buf.NewPacketReader(conn)
|
||||||
defer func() {
|
|
||||||
udpConns.Range(func(key uint64, entry *udpConnEntry) bool {
|
|
||||||
entry.timer.SetTimeout(0)
|
|
||||||
return true
|
|
||||||
})
|
|
||||||
}()
|
|
||||||
|
|
||||||
reader := buf.NewReader(conn)
|
|
||||||
for {
|
for {
|
||||||
mb, err := reader.ReadMultiBuffer()
|
mb, err := reader.ReadMultiBuffer()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -238,11 +221,7 @@ func (i *RelayInbound) processUDP(ctx context.Context, conn stat.Connection, dis
|
|||||||
var packetHeader [AESBlockSize]byte
|
var packetHeader [AESBlockSize]byte
|
||||||
i.relayBlock.Decrypt(packetHeader[:], data[:AESBlockSize])
|
i.relayBlock.Decrypt(packetHeader[:], data[:AESBlockSize])
|
||||||
|
|
||||||
var eiHeader [AESBlockSize]byte
|
eiHeader := DecryptUDPEIH(i.relayBlock, packetHeader[:], data[AESBlockSize:2*AESBlockSize])
|
||||||
i.relayBlock.Decrypt(eiHeader[:], data[AESBlockSize:2*AESBlockSize])
|
|
||||||
for idx := 0; idx < AESBlockSize; idx++ {
|
|
||||||
eiHeader[idx] ^= packetHeader[idx]
|
|
||||||
}
|
|
||||||
|
|
||||||
targetDest, ok := i.destinations[eiHeader]
|
targetDest, ok := i.destinations[eiHeader]
|
||||||
if !ok {
|
if !ok {
|
||||||
@@ -263,68 +242,24 @@ func (i *RelayInbound) processUDP(ctx context.Context, conn stat.Connection, dis
|
|||||||
dest := targetDest.destination
|
dest := targetDest.destination
|
||||||
dest.Network = net.Network_UDP
|
dest.Network = net.Network_UDP
|
||||||
|
|
||||||
entry, ok := udpConns.Load(sessionID)
|
sessionItem := i.udpSessions.GetOrCreate(sessionID)
|
||||||
if !ok {
|
if sessionItem.User == nil {
|
||||||
sessCtx, cancel := context.WithCancel(ctx)
|
sessionItem.Lock()
|
||||||
inbound := session.InboundFromContext(sessCtx)
|
if sessionItem.User == nil {
|
||||||
inbound.User = &protocol.MemoryUser{
|
sessionItem.User = &protocol.MemoryUser{
|
||||||
Email: targetDest.email,
|
Email: targetDest.email,
|
||||||
Level: targetDest.level,
|
Level: targetDest.level,
|
||||||
}
|
}
|
||||||
|
|
||||||
sessCtx = log.ContextWithAccessMessage(sessCtx, &log.AccessMessage{
|
|
||||||
From: conn.RemoteAddr(),
|
|
||||||
To: dest,
|
|
||||||
Status: log.AccessAccepted,
|
|
||||||
Email: targetDest.email,
|
|
||||||
})
|
|
||||||
|
|
||||||
link, err := dispatcher.Dispatch(sessCtx, dest)
|
|
||||||
if err != nil {
|
|
||||||
cancel()
|
|
||||||
b.Release()
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
|
|
||||||
newEntry := &udpConnEntry{
|
|
||||||
link: link,
|
|
||||||
cancel: cancel,
|
|
||||||
}
|
|
||||||
sessionPolicy := i.policyManager.ForLevel(targetDest.level)
|
|
||||||
newEntry.timer = signal.CancelAfterInactivity(sessCtx, func() {
|
|
||||||
udpConns.Delete(sessionID)
|
|
||||||
common.Interrupt(link.Reader)
|
|
||||||
common.Interrupt(link.Writer)
|
|
||||||
cancel()
|
|
||||||
}, sessionPolicy.Timeouts.ConnectionIdle)
|
|
||||||
|
|
||||||
actual, loaded := udpConns.LoadOrStore(sessionID, newEntry)
|
|
||||||
if loaded {
|
|
||||||
newEntry.timer.SetTimeout(0)
|
|
||||||
entry = actual
|
|
||||||
} else {
|
|
||||||
entry = newEntry
|
|
||||||
go func(cEntry *udpConnEntry) {
|
|
||||||
defer func() {
|
|
||||||
cEntry.timer.SetTimeout(0)
|
|
||||||
}()
|
|
||||||
for {
|
|
||||||
resMb, err := cEntry.link.Reader.ReadMultiBuffer()
|
|
||||||
if err != nil {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
cEntry.timer.Update()
|
|
||||||
for _, rb := range resMb {
|
|
||||||
_, _ = conn.Write(rb.Bytes())
|
|
||||||
rb.Release()
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}(entry)
|
|
||||||
}
|
}
|
||||||
|
sessionItem.Unlock()
|
||||||
|
}
|
||||||
|
link, err := sessionItem.EnsureLink(ctx, conn, dest, dispatcher, i.policyManager, nil)
|
||||||
|
if err != nil {
|
||||||
|
b.Release()
|
||||||
|
continue
|
||||||
}
|
}
|
||||||
|
|
||||||
entry.timer.Update()
|
_ = link.Writer.WriteMultiBuffer(buf.MultiBuffer{b})
|
||||||
_ = entry.link.Writer.WriteMultiBuffer(buf.MultiBuffer{b})
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -61,3 +61,14 @@ func DeriveUserPSKHash(userPSK []byte) [AESBlockSize]byte {
|
|||||||
copy(out[:], h[:AESBlockSize])
|
copy(out[:], h[:AESBlockSize])
|
||||||
return out
|
return out
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func DecryptEIH(method *CipherMethod, key, salt, eih []byte) ([AESBlockSize]byte, error) {
|
||||||
|
identitySubkey := DeriveIdentitySubKey(key, salt, method.KeySaltLength)
|
||||||
|
block, err := method.NewBlock(identitySubkey)
|
||||||
|
if err != nil {
|
||||||
|
return [AESBlockSize]byte{}, err
|
||||||
|
}
|
||||||
|
var decryptedHash [AESBlockSize]byte
|
||||||
|
block.Decrypt(decryptedHash[:], eih)
|
||||||
|
return decryptedHash, nil
|
||||||
|
}
|
||||||
|
|||||||
@@ -4,7 +4,6 @@ import (
|
|||||||
"context"
|
"context"
|
||||||
"crypto/rand"
|
"crypto/rand"
|
||||||
"io"
|
"io"
|
||||||
"time"
|
|
||||||
|
|
||||||
"github.com/xtls/xray-core/common"
|
"github.com/xtls/xray-core/common"
|
||||||
"github.com/xtls/xray-core/common/buf"
|
"github.com/xtls/xray-core/common/buf"
|
||||||
@@ -46,8 +45,12 @@ func NewClient(ctx context.Context, config *ClientConfig) (*Outbound, error) {
|
|||||||
return nil, errors.New("invalid key: ", config.Key).Base(err)
|
return nil, errors.New("invalid key: ", config.Key).Base(err)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
if method.IsChaCha && len(pskList) > 1 {
|
||||||
|
return nil, errors.New("multi-key is not supported for chacha20-poly1305")
|
||||||
|
}
|
||||||
|
|
||||||
finalPSK := pskList[len(pskList)-1]
|
finalPSK := pskList[len(pskList)-1]
|
||||||
udpCodec, err := NewUDPPacketCodec(method, finalPSK)
|
udpCodec, err := NewUDPPacketCodec(method, pskList)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, errors.New("failed to create udp packet codec").Base(err)
|
return nil, errors.New("failed to create udp packet codec").Base(err)
|
||||||
}
|
}
|
||||||
@@ -126,18 +129,30 @@ func (o *Outbound) Process(ctx context.Context, link *transport.Link, dialer int
|
|||||||
|
|
||||||
requestDone := func() error {
|
requestDone := func() error {
|
||||||
defer timer.SetTimeout(sessionPolicy.Timeouts.DownlinkOnly)
|
defer timer.SetTimeout(sessionPolicy.Timeouts.DownlinkOnly)
|
||||||
bufferedWriter := buf.NewBufferedWriter(buf.NewWriter(conn))
|
|
||||||
bodyWriter, err := WriteTCPRequest(bufferedWriter, o.method, o.pskList, destination, clientSaltSlice, nil)
|
var initialPayload []byte
|
||||||
|
var firstBuf *buf.Buffer
|
||||||
|
var remainingMB buf.MultiBuffer
|
||||||
|
if timeoutReader, ok := link.Reader.(buf.TimeoutReader); ok {
|
||||||
|
if mb, err := timeoutReader.ReadMultiBufferTimeout(0); err == nil && !mb.IsEmpty() {
|
||||||
|
remainingMB, firstBuf = buf.SplitFirst(mb)
|
||||||
|
initialPayload = firstBuf.Bytes()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
bodyWriter, err := WriteTCPRequest(conn, o.method, o.pskList, destination, clientSaltSlice, initialPayload)
|
||||||
|
if firstBuf != nil {
|
||||||
|
firstBuf.Release()
|
||||||
|
}
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
buf.ReleaseMulti(remainingMB)
|
||||||
return errors.New("failed to write request").Base(err)
|
return errors.New("failed to write request").Base(err)
|
||||||
}
|
}
|
||||||
|
|
||||||
if err = buf.CopyOnceTimeout(link.Reader, bodyWriter, time.Millisecond*100); err != nil && err != buf.ErrNotTimeoutReader && err != buf.ErrReadTimeout {
|
if !remainingMB.IsEmpty() {
|
||||||
return errors.New("failed to write A request payload").Base(err)
|
if err := bodyWriter.WriteMultiBuffer(remainingMB); err != nil {
|
||||||
}
|
return err
|
||||||
|
}
|
||||||
if err := bufferedWriter.SetBuffered(false); err != nil {
|
|
||||||
return err
|
|
||||||
}
|
}
|
||||||
|
|
||||||
return buf.Copy(link.Reader, bodyWriter, buf.UpdateActivity(timer))
|
return buf.Copy(link.Reader, bodyWriter, buf.UpdateActivity(timer))
|
||||||
@@ -163,13 +178,18 @@ func (o *Outbound) Process(ctx context.Context, link *transport.Link, dialer int
|
|||||||
}
|
}
|
||||||
|
|
||||||
if network == net.Network_UDP {
|
if network == net.Network_UDP {
|
||||||
|
session, err := o.udpCodec.NewClientSession()
|
||||||
|
if err != nil {
|
||||||
|
return errors.New("failed to create client udp session").Base(err)
|
||||||
|
}
|
||||||
|
|
||||||
requestDone := func() error {
|
requestDone := func() error {
|
||||||
defer timer.SetTimeout(sessionPolicy.Timeouts.DownlinkOnly)
|
defer timer.SetTimeout(sessionPolicy.Timeouts.DownlinkOnly)
|
||||||
|
|
||||||
writer := &UDPWriter{
|
writer := &UDPWriter{
|
||||||
Writer: conn,
|
Writer: conn,
|
||||||
Destination: destination,
|
Destination: destination,
|
||||||
Codec: o.udpCodec,
|
Session: session,
|
||||||
}
|
}
|
||||||
|
|
||||||
if err := buf.Copy(link.Reader, writer, buf.UpdateActivity(timer)); err != nil {
|
if err := buf.Copy(link.Reader, writer, buf.UpdateActivity(timer)); err != nil {
|
||||||
@@ -182,8 +202,8 @@ func (o *Outbound) Process(ctx context.Context, link *transport.Link, dialer int
|
|||||||
defer timer.SetTimeout(sessionPolicy.Timeouts.UplinkOnly)
|
defer timer.SetTimeout(sessionPolicy.Timeouts.UplinkOnly)
|
||||||
|
|
||||||
reader := &UDPReader{
|
reader := &UDPReader{
|
||||||
Reader: conn,
|
Reader: conn,
|
||||||
Codec: o.udpCodec,
|
Session: session,
|
||||||
}
|
}
|
||||||
|
|
||||||
if err := buf.Copy(reader, link.Writer, buf.UpdateActivity(timer)); err != nil {
|
if err := buf.Copy(reader, link.Writer, buf.UpdateActivity(timer)); err != nil {
|
||||||
|
|||||||
+441
-187
@@ -16,14 +16,13 @@ import (
|
|||||||
)
|
)
|
||||||
|
|
||||||
type UDPCodec struct {
|
type UDPCodec struct {
|
||||||
method *CipherMethod
|
method *CipherMethod
|
||||||
psk []byte
|
pskList [][]byte
|
||||||
blockCipher cipher.Block
|
psk []byte
|
||||||
chachaCipher cipher.AEAD
|
blockCipher cipher.Block
|
||||||
clientBodyCipher cipher.AEAD
|
blockCiphers []cipher.Block
|
||||||
clientSessionID uint64
|
chachaCipher cipher.AEAD
|
||||||
nextPacketID atomic.Uint64
|
sessions *UDPSessionManager
|
||||||
sessions *UDPSessionManager
|
|
||||||
}
|
}
|
||||||
|
|
||||||
type (
|
type (
|
||||||
@@ -48,22 +47,23 @@ func newUDPCodec(method *CipherMethod, psk []byte) (*UDPCodec, error) {
|
|||||||
return c, nil
|
return c, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func NewUDPPacketCodec(method *CipherMethod, psk []byte) (*UDPCodec, error) {
|
func NewUDPPacketCodec(method *CipherMethod, pskList [][]byte) (*UDPCodec, error) {
|
||||||
c, err := newUDPCodec(method, psk)
|
if method.IsChaCha && len(pskList) > 1 {
|
||||||
|
return nil, errors.New("multi-key is not supported for chacha20-poly1305")
|
||||||
|
}
|
||||||
|
finalPSK := pskList[len(pskList)-1]
|
||||||
|
c, err := newUDPCodec(method, finalPSK)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
var sessID [8]byte
|
c.pskList = pskList
|
||||||
if _, err := io.ReadFull(rand.Reader, sessID[:]); err != nil {
|
if len(pskList) > 1 {
|
||||||
return nil, err
|
c.blockCiphers = make([]cipher.Block, len(pskList))
|
||||||
}
|
for i, psk := range pskList {
|
||||||
c.clientSessionID = binary.BigEndian.Uint64(sessID[:])
|
c.blockCiphers[i], err = method.NewBlock(psk)
|
||||||
|
if err != nil {
|
||||||
if !method.IsChaCha {
|
return nil, err
|
||||||
clientBodyKey := DeriveSessionSubKey(psk, sessID[:], method.KeySaltLength)
|
}
|
||||||
c.clientBodyCipher, err = method.NewAEAD(clientBodyKey)
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
return c, nil
|
return c, nil
|
||||||
@@ -78,108 +78,37 @@ func NewUDPServerCodec(method *CipherMethod, psk []byte, sessionTimeout time.Dur
|
|||||||
return c, nil
|
return c, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *UDPCodec) EncodeClientPacket(dest net.Destination, payload []byte) (*buf.Buffer, error) {
|
func (c *UDPCodec) Sessions() *UDPSessionManager {
|
||||||
packetID := c.nextPacketID.Add(1)
|
return c.sessions
|
||||||
sessID := c.clientSessionID
|
}
|
||||||
|
|
||||||
// Padding determination (e.g. DNS port 53 disguise)
|
func (c *UDPCodec) GetSession(sessionID uint64) *ServerUDPSession {
|
||||||
var paddingLen int
|
if c.sessions == nil {
|
||||||
if dest.Port == 53 && len(payload) < MaxPaddingLength {
|
return nil
|
||||||
paddingLen = mrand.IntN(MaxPaddingLength-len(payload)) + 1
|
|
||||||
}
|
}
|
||||||
|
return c.sessions.GetOrCreate(sessionID)
|
||||||
addrPortLen := AddrPortLength(dest)
|
|
||||||
|
|
||||||
if c.method.IsChaCha {
|
|
||||||
// ChaCha20 mode: 24-byte nonce + plaintext header (27B) + padding + dest + payload + AEAD tag (16B)
|
|
||||||
totalLen := PacketNonceSize + 27 + paddingLen + addrPortLen + len(payload) + AEADTagSize
|
|
||||||
if totalLen > buf.Size {
|
|
||||||
return nil, ErrPacketTooLarge
|
|
||||||
}
|
|
||||||
|
|
||||||
outBuf := buf.New()
|
|
||||||
|
|
||||||
var nonce [PacketNonceSize]byte
|
|
||||||
if _, err := io.ReadFull(rand.Reader, nonce[:]); err != nil {
|
|
||||||
outBuf.Release()
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
outBuf.Write(nonce[:])
|
|
||||||
|
|
||||||
var hdr [16 + 1 + 8 + 2]byte
|
|
||||||
binary.BigEndian.PutUint64(hdr[0:8], sessID)
|
|
||||||
binary.BigEndian.PutUint64(hdr[8:16], packetID)
|
|
||||||
hdr[16] = HeaderTypeClient
|
|
||||||
binary.BigEndian.PutUint64(hdr[17:25], uint64(time.Now().Unix()))
|
|
||||||
binary.BigEndian.PutUint16(hdr[25:27], uint16(paddingLen))
|
|
||||||
outBuf.Write(hdr[:])
|
|
||||||
if paddingLen > 0 {
|
|
||||||
outBuf.Write(zeroPadding[:paddingLen])
|
|
||||||
}
|
|
||||||
|
|
||||||
if err := WriteAddressPort(outBuf, dest); err != nil {
|
|
||||||
outBuf.Release()
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
outBuf.Write(payload)
|
|
||||||
|
|
||||||
plainBytes := outBuf.Bytes()[PacketNonceSize:]
|
|
||||||
outBuf.Extend(int32(c.chachaCipher.Overhead()))
|
|
||||||
c.chachaCipher.Seal(plainBytes[:0], nonce[:], plainBytes, nil)
|
|
||||||
return outBuf, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// AES mode:
|
|
||||||
// 16B Encrypted Header + (11B header + padding + dest + payload + 16B AEAD tag)
|
|
||||||
totalLen := 16 + 11 + paddingLen + addrPortLen + len(payload) + AEADTagSize
|
|
||||||
if totalLen > buf.Size {
|
|
||||||
return nil, ErrPacketTooLarge
|
|
||||||
}
|
|
||||||
|
|
||||||
outBuf := buf.New()
|
|
||||||
|
|
||||||
var rawHeader [16]byte
|
|
||||||
binary.BigEndian.PutUint64(rawHeader[:8], sessID)
|
|
||||||
binary.BigEndian.PutUint64(rawHeader[8:16], packetID)
|
|
||||||
|
|
||||||
var encryptedHeader [16]byte
|
|
||||||
c.blockCipher.Encrypt(encryptedHeader[:], rawHeader[:])
|
|
||||||
outBuf.Write(encryptedHeader[:])
|
|
||||||
|
|
||||||
bodyAead := c.clientBodyCipher
|
|
||||||
|
|
||||||
var hdr [1 + 8 + 2]byte
|
|
||||||
hdr[0] = HeaderTypeClient
|
|
||||||
binary.BigEndian.PutUint64(hdr[1:9], uint64(time.Now().Unix()))
|
|
||||||
binary.BigEndian.PutUint16(hdr[9:11], uint16(paddingLen))
|
|
||||||
outBuf.Write(hdr[:])
|
|
||||||
if paddingLen > 0 {
|
|
||||||
outBuf.Write(zeroPadding[:paddingLen])
|
|
||||||
}
|
|
||||||
|
|
||||||
if err := WriteAddressPort(outBuf, dest); err != nil {
|
|
||||||
outBuf.Release()
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
outBuf.Write(payload)
|
|
||||||
|
|
||||||
plainBytes := outBuf.Bytes()[16:]
|
|
||||||
bodyNonce := rawHeader[4:16]
|
|
||||||
outBuf.Extend(int32(bodyAead.Overhead()))
|
|
||||||
bodyAead.Seal(plainBytes[:0], bodyNonce, plainBytes, nil)
|
|
||||||
return outBuf, nil
|
|
||||||
}
|
}
|
||||||
|
|
||||||
type DecodedUDPPacket struct {
|
type DecodedUDPPacket struct {
|
||||||
SessionID uint64
|
SessionID uint64
|
||||||
PacketID uint64
|
PacketID uint64
|
||||||
HeaderType byte
|
HeaderType byte
|
||||||
Timestamp uint64
|
Timestamp uint64
|
||||||
Destination net.Destination
|
ClientSessionID uint64
|
||||||
Payload []byte
|
Destination net.Destination
|
||||||
|
Payload []byte
|
||||||
}
|
}
|
||||||
|
|
||||||
func parseAddressPort(data []byte) (net.Destination, int, error) {
|
func DecryptUDPEIH(block cipher.Block, rawHeader, eih []byte) [AESBlockSize]byte {
|
||||||
|
var decryptedHash [AESBlockSize]byte
|
||||||
|
block.Decrypt(decryptedHash[:], eih)
|
||||||
|
for k := 0; k < AESBlockSize; k++ {
|
||||||
|
decryptedHash[k] ^= rawHeader[k]
|
||||||
|
}
|
||||||
|
return decryptedHash
|
||||||
|
}
|
||||||
|
|
||||||
|
func ParseAddressPort(data []byte) (net.Destination, int, error) {
|
||||||
if len(data) < 1 {
|
if len(data) < 1 {
|
||||||
return net.Destination{}, 0, ErrPacketTooShort
|
return net.Destination{}, 0, ErrPacketTooShort
|
||||||
}
|
}
|
||||||
@@ -220,6 +149,9 @@ func parsePlainUDPPacket(sessionID, packetID uint64, bodyPlain []byte) (DecodedU
|
|||||||
}
|
}
|
||||||
|
|
||||||
headerType := bodyPlain[0]
|
headerType := bodyPlain[0]
|
||||||
|
if headerType != HeaderTypeClient && headerType != HeaderTypeServer {
|
||||||
|
return DecodedUDPPacket{}, ErrBadHeaderType
|
||||||
|
}
|
||||||
epoch := binary.BigEndian.Uint64(bodyPlain[1:9])
|
epoch := binary.BigEndian.Uint64(bodyPlain[1:9])
|
||||||
diff := int(math.Abs(float64(time.Now().Unix() - int64(epoch))))
|
diff := int(math.Abs(float64(time.Now().Unix() - int64(epoch))))
|
||||||
if diff > 30 {
|
if diff > 30 {
|
||||||
@@ -227,11 +159,13 @@ func parsePlainUDPPacket(sessionID, packetID uint64, bodyPlain []byte) (DecodedU
|
|||||||
}
|
}
|
||||||
|
|
||||||
offset := 9
|
offset := 9
|
||||||
|
var clientSessionID uint64
|
||||||
if headerType == HeaderTypeServer {
|
if headerType == HeaderTypeServer {
|
||||||
if len(bodyPlain) < offset+8+2 {
|
if len(bodyPlain) < offset+8+2 {
|
||||||
return DecodedUDPPacket{}, ErrPacketTooShort
|
return DecodedUDPPacket{}, ErrPacketTooShort
|
||||||
}
|
}
|
||||||
offset += 8 // skip clientSessionID
|
clientSessionID = binary.BigEndian.Uint64(bodyPlain[offset : offset+8])
|
||||||
|
offset += 8
|
||||||
}
|
}
|
||||||
|
|
||||||
paddingLen := int(binary.BigEndian.Uint16(bodyPlain[offset : offset+2]))
|
paddingLen := int(binary.BigEndian.Uint16(bodyPlain[offset : offset+2]))
|
||||||
@@ -242,19 +176,20 @@ func parsePlainUDPPacket(sessionID, packetID uint64, bodyPlain []byte) (DecodedU
|
|||||||
}
|
}
|
||||||
offset += paddingLen
|
offset += paddingLen
|
||||||
|
|
||||||
dest, addrLen, err := parseAddressPort(bodyPlain[offset:])
|
dest, addrLen, err := ParseAddressPort(bodyPlain[offset:])
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return DecodedUDPPacket{}, err
|
return DecodedUDPPacket{}, err
|
||||||
}
|
}
|
||||||
payload := bodyPlain[offset+addrLen:]
|
payload := bodyPlain[offset+addrLen:]
|
||||||
|
|
||||||
return DecodedUDPPacket{
|
return DecodedUDPPacket{
|
||||||
SessionID: sessionID,
|
SessionID: sessionID,
|
||||||
PacketID: packetID,
|
PacketID: packetID,
|
||||||
HeaderType: headerType,
|
HeaderType: headerType,
|
||||||
Timestamp: epoch,
|
Timestamp: epoch,
|
||||||
Destination: dest,
|
ClientSessionID: clientSessionID,
|
||||||
Payload: payload,
|
Destination: dest,
|
||||||
|
Payload: payload,
|
||||||
}, nil
|
}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -269,7 +204,7 @@ func (c *UDPCodec) DecodePacket(data []byte) (DecodedUDPPacket, error) {
|
|||||||
}
|
}
|
||||||
nonce := data[:PacketNonceSize]
|
nonce := data[:PacketNonceSize]
|
||||||
ciphertext := data[PacketNonceSize:]
|
ciphertext := data[PacketNonceSize:]
|
||||||
plain, err := c.chachaCipher.Open(ciphertext[:0], nonce, ciphertext, nil)
|
plain, err := c.chachaCipher.Open(nil, nonce, ciphertext, nil)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return DecodedUDPPacket{}, errors.New("failed to decrypt chacha udp packet").Base(err)
|
return DecodedUDPPacket{}, errors.New("failed to decrypt chacha udp packet").Base(err)
|
||||||
}
|
}
|
||||||
@@ -280,17 +215,22 @@ func (c *UDPCodec) DecodePacket(data []byte) (DecodedUDPPacket, error) {
|
|||||||
sessionID := binary.BigEndian.Uint64(plain[:8])
|
sessionID := binary.BigEndian.Uint64(plain[:8])
|
||||||
packetID := binary.BigEndian.Uint64(plain[8:16])
|
packetID := binary.BigEndian.Uint64(plain[8:16])
|
||||||
|
|
||||||
if c.sessions != nil {
|
sessionItem := c.sessions.GetOrCreate(sessionID)
|
||||||
sessionItem := c.sessions.GetOrCreate(sessionID)
|
if !sessionItem.CheckPacketID(packetID) {
|
||||||
sessionItem.Lock()
|
return DecodedUDPPacket{}, ErrPacketIdNotUnique
|
||||||
if !sessionItem.Window.CheckAndAdd(packetID) {
|
|
||||||
sessionItem.Unlock()
|
|
||||||
return DecodedUDPPacket{}, ErrPacketIdNotUnique
|
|
||||||
}
|
|
||||||
sessionItem.Unlock()
|
|
||||||
}
|
}
|
||||||
|
|
||||||
return parsePlainUDPPacket(sessionID, packetID, plain[16:])
|
decoded, err := parsePlainUDPPacket(sessionID, packetID, plain[16:])
|
||||||
|
if err != nil {
|
||||||
|
return DecodedUDPPacket{}, err
|
||||||
|
}
|
||||||
|
|
||||||
|
if decoded.HeaderType != HeaderTypeClient {
|
||||||
|
return DecodedUDPPacket{}, ErrBadHeaderType
|
||||||
|
}
|
||||||
|
|
||||||
|
sessionItem.AddPacketID(packetID)
|
||||||
|
return decoded, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// AES mode
|
// AES mode
|
||||||
@@ -299,54 +239,52 @@ func (c *UDPCodec) DecodePacket(data []byte) (DecodedUDPPacket, error) {
|
|||||||
sessionID := binary.BigEndian.Uint64(rawHeader[:8])
|
sessionID := binary.BigEndian.Uint64(rawHeader[:8])
|
||||||
packetID := binary.BigEndian.Uint64(rawHeader[8:16])
|
packetID := binary.BigEndian.Uint64(rawHeader[8:16])
|
||||||
|
|
||||||
var bodyAead cipher.AEAD
|
sessionItem := c.sessions.GetOrCreate(sessionID)
|
||||||
var sessionItem *ServerUDPSession
|
if !sessionItem.CheckPacketID(packetID) {
|
||||||
|
return DecodedUDPPacket{}, ErrPacketIdNotUnique
|
||||||
|
}
|
||||||
|
|
||||||
if c.sessions != nil {
|
return sessionItem.DecryptAESPayload(c.method, c.psk, sessionID, packetID, rawHeader[:], data[16:])
|
||||||
sessionItem = c.sessions.GetOrCreate(sessionID)
|
}
|
||||||
sessionItem.Lock()
|
|
||||||
if !sessionItem.Window.Check(packetID) {
|
|
||||||
sessionItem.Unlock()
|
|
||||||
return DecodedUDPPacket{}, ErrPacketIdNotUnique
|
|
||||||
}
|
|
||||||
sessionItem.Unlock()
|
|
||||||
|
|
||||||
bodyAead = sessionItem.GetRemoteCipher()
|
func (s *ServerUDPSession) DecryptAESPayload(method *CipherMethod, psk []byte, sessionID, packetID uint64, rawHeader, bodyCipher []byte) (DecodedUDPPacket, error) {
|
||||||
if bodyAead == nil {
|
bodyAead := s.clientBodyCipher
|
||||||
bodyKey := DeriveSessionSubKey(c.psk, rawHeader[:8], c.method.KeySaltLength)
|
isNewCipher := false
|
||||||
var err error
|
if bodyAead == nil {
|
||||||
bodyAead, err = c.method.NewAEAD(bodyKey)
|
bodyKey := DeriveSessionSubKey(psk, rawHeader[:8], method.KeySaltLength)
|
||||||
if err != nil {
|
|
||||||
return DecodedUDPPacket{}, err
|
|
||||||
}
|
|
||||||
sessionItem.SetRemoteCipher(bodyAead)
|
|
||||||
}
|
|
||||||
} else {
|
|
||||||
bodyKey := DeriveSessionSubKey(c.psk, rawHeader[:8], c.method.KeySaltLength)
|
|
||||||
var err error
|
var err error
|
||||||
bodyAead, err = c.method.NewAEAD(bodyKey)
|
bodyAead, err = method.NewAEAD(bodyKey)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return DecodedUDPPacket{}, err
|
return DecodedUDPPacket{}, err
|
||||||
}
|
}
|
||||||
|
isNewCipher = true
|
||||||
}
|
}
|
||||||
|
|
||||||
bodyNonce := rawHeader[4:16]
|
bodyNonce := rawHeader[4:16]
|
||||||
bodyCipher := data[16:]
|
bodyPlain, err := bodyAead.Open(nil, bodyNonce, bodyCipher, nil)
|
||||||
bodyPlain, err := bodyAead.Open(bodyCipher[:0], bodyNonce, bodyCipher, nil)
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return DecodedUDPPacket{}, errors.New("failed to decrypt aes udp body").Base(err)
|
return DecodedUDPPacket{}, errors.New("failed to decrypt aes udp body").Base(err)
|
||||||
}
|
}
|
||||||
|
|
||||||
if sessionItem != nil {
|
decoded, err := parsePlainUDPPacket(sessionID, packetID, bodyPlain)
|
||||||
sessionItem.Lock()
|
if err != nil {
|
||||||
sessionItem.Window.Add(packetID)
|
return DecodedUDPPacket{}, err
|
||||||
sessionItem.Unlock()
|
|
||||||
}
|
}
|
||||||
|
|
||||||
return parsePlainUDPPacket(sessionID, packetID, bodyPlain)
|
if decoded.HeaderType != HeaderTypeClient {
|
||||||
|
return DecodedUDPPacket{}, ErrBadHeaderType
|
||||||
|
}
|
||||||
|
|
||||||
|
s.AddPacketID(packetID)
|
||||||
|
|
||||||
|
if isNewCipher {
|
||||||
|
s.clientBodyCipher = bodyAead
|
||||||
|
}
|
||||||
|
|
||||||
|
return decoded, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (s *ServerUDPSession) EnsureServerState(method *CipherMethod, headerBlock cipher.Block, chachaCipher cipher.AEAD, psk []byte) error {
|
func (s *ServerUDPSession) EnsureServerState(method *CipherMethod, psk []byte) error {
|
||||||
s.Lock()
|
s.Lock()
|
||||||
defer s.Unlock()
|
defer s.Unlock()
|
||||||
if s.ServerSessionID != 0 {
|
if s.ServerSessionID != 0 {
|
||||||
@@ -363,23 +301,29 @@ func (s *ServerUDPSession) EnsureServerState(method *CipherMethod, headerBlock c
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
if method.IsChaCha {
|
if method.IsChaCha {
|
||||||
s.ServerChaCha = chachaCipher
|
var err error
|
||||||
} else {
|
s.serverChaCha, err = method.NewUDPCipher(psk)
|
||||||
s.ServerBlockCipher = headerBlock
|
return err
|
||||||
bodyKey := DeriveSessionSubKey(psk, sidBuf[:], method.KeySaltLength)
|
}
|
||||||
bodyAead, err := method.NewAEAD(bodyKey)
|
|
||||||
if err != nil {
|
var err error
|
||||||
s.ServerSessionID = 0
|
s.serverHeaderBlock, err = method.NewBlock(psk)
|
||||||
return err
|
if err != nil {
|
||||||
}
|
s.ServerSessionID = 0
|
||||||
s.ServerCipher = bodyAead
|
return err
|
||||||
|
}
|
||||||
|
bodyKey := DeriveSessionSubKey(psk, sidBuf[:], method.KeySaltLength)
|
||||||
|
s.serverBodyCipher, err = method.NewAEAD(bodyKey)
|
||||||
|
if err != nil {
|
||||||
|
s.ServerSessionID = 0
|
||||||
|
return err
|
||||||
}
|
}
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (s *ServerUDPSession) EncodeServerPacket(method *CipherMethod, clientSessionID uint64, dest net.Destination, payload []byte) ([]byte, error) {
|
func (s *ServerUDPSession) EncodeServerPacket(method *CipherMethod, clientSessionID uint64, dest net.Destination, payload []byte) ([]byte, error) {
|
||||||
serverSessionID := s.ServerSessionID
|
serverSessionID := s.ServerSessionID
|
||||||
serverPacketID := s.ServerPacketID.Add(1)
|
serverPacketID := s.ServerPacketID.Add(1) - 1
|
||||||
|
|
||||||
if method.IsChaCha {
|
if method.IsChaCha {
|
||||||
var nonce [PacketNonceSize]byte
|
var nonce [PacketNonceSize]byte
|
||||||
@@ -404,7 +348,7 @@ func (s *ServerUDPSession) EncodeServerPacket(method *CipherMethod, clientSessio
|
|||||||
}
|
}
|
||||||
plainBuf.Write(payload)
|
plainBuf.Write(payload)
|
||||||
|
|
||||||
sealed := s.ServerChaCha.Seal(nil, nonce[:], plainBuf.Bytes(), nil)
|
sealed := s.serverChaCha.Seal(nil, nonce[:], plainBuf.Bytes(), nil)
|
||||||
res := make([]byte, PacketNonceSize+len(sealed))
|
res := make([]byte, PacketNonceSize+len(sealed))
|
||||||
copy(res[:PacketNonceSize], nonce[:])
|
copy(res[:PacketNonceSize], nonce[:])
|
||||||
copy(res[PacketNonceSize:], sealed)
|
copy(res[PacketNonceSize:], sealed)
|
||||||
@@ -417,7 +361,7 @@ func (s *ServerUDPSession) EncodeServerPacket(method *CipherMethod, clientSessio
|
|||||||
binary.BigEndian.PutUint64(rawHeader[8:16], serverPacketID)
|
binary.BigEndian.PutUint64(rawHeader[8:16], serverPacketID)
|
||||||
|
|
||||||
var encryptedHeader [16]byte
|
var encryptedHeader [16]byte
|
||||||
s.ServerBlockCipher.Encrypt(encryptedHeader[:], rawHeader[:])
|
s.serverHeaderBlock.Encrypt(encryptedHeader[:], rawHeader[:])
|
||||||
|
|
||||||
bodyBuf := buf.New()
|
bodyBuf := buf.New()
|
||||||
defer bodyBuf.Release()
|
defer bodyBuf.Release()
|
||||||
@@ -435,7 +379,7 @@ func (s *ServerUDPSession) EncodeServerPacket(method *CipherMethod, clientSessio
|
|||||||
bodyBuf.Write(payload)
|
bodyBuf.Write(payload)
|
||||||
|
|
||||||
bodyNonce := rawHeader[4:16]
|
bodyNonce := rawHeader[4:16]
|
||||||
sealedBody := s.ServerCipher.Seal(nil, bodyNonce, bodyBuf.Bytes(), nil)
|
sealedBody := s.serverBodyCipher.Seal(nil, bodyNonce, bodyBuf.Bytes(), nil)
|
||||||
|
|
||||||
res := make([]byte, 16+len(sealedBody))
|
res := make([]byte, 16+len(sealedBody))
|
||||||
copy(res[:16], encryptedHeader[:])
|
copy(res[:16], encryptedHeader[:])
|
||||||
@@ -444,17 +388,327 @@ func (s *ServerUDPSession) EncodeServerPacket(method *CipherMethod, clientSessio
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (c *UDPCodec) EncodeServerPacket(clientSessionID uint64, dest net.Destination, payload []byte) ([]byte, error) {
|
func (c *UDPCodec) EncodeServerPacket(clientSessionID uint64, dest net.Destination, payload []byte) ([]byte, error) {
|
||||||
sessionItem := c.sessions.GetOrCreate(clientSessionID)
|
return c.sessions.EncodeServerPacket(c.method, c.psk, clientSessionID, dest, payload)
|
||||||
if err := sessionItem.EnsureServerState(c.method, c.blockCipher, c.chachaCipher, c.psk); err != nil {
|
}
|
||||||
|
|
||||||
|
type serverSessionState struct {
|
||||||
|
sessionID uint64
|
||||||
|
window *SlidingWindow
|
||||||
|
cipher cipher.AEAD
|
||||||
|
lastSeen atomic.Int64
|
||||||
|
}
|
||||||
|
|
||||||
|
func (st *serverSessionState) check(packetID uint64) bool {
|
||||||
|
if st.window == nil {
|
||||||
|
st.window = new(SlidingWindow)
|
||||||
|
}
|
||||||
|
return st.window.Check(packetID)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (st *serverSessionState) add(packetID uint64) {
|
||||||
|
if st.window == nil {
|
||||||
|
st.window = new(SlidingWindow)
|
||||||
|
}
|
||||||
|
st.window.Add(packetID)
|
||||||
|
}
|
||||||
|
|
||||||
|
type ClientUDPSession struct {
|
||||||
|
codec *UDPCodec
|
||||||
|
clientSessionID uint64
|
||||||
|
nextPacketID atomic.Uint64
|
||||||
|
clientBodyCipher cipher.AEAD
|
||||||
|
current atomic.Pointer[serverSessionState]
|
||||||
|
old atomic.Pointer[serverSessionState]
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *UDPCodec) NewClientSession() (*ClientUDPSession, error) {
|
||||||
|
var sessID [8]byte
|
||||||
|
if _, err := io.ReadFull(rand.Reader, sessID[:]); err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
return sessionItem.EncodeServerPacket(c.method, clientSessionID, dest, payload)
|
clientSessionID := binary.BigEndian.Uint64(sessID[:])
|
||||||
|
|
||||||
|
var clientBodyCipher cipher.AEAD
|
||||||
|
var err error
|
||||||
|
if !c.method.IsChaCha {
|
||||||
|
finalPSK := c.psk
|
||||||
|
clientBodyKey := DeriveSessionSubKey(finalPSK, sessID[:], c.method.KeySaltLength)
|
||||||
|
clientBodyCipher, err = c.method.NewAEAD(clientBodyKey)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return &ClientUDPSession{
|
||||||
|
codec: c,
|
||||||
|
clientSessionID: clientSessionID,
|
||||||
|
clientBodyCipher: clientBodyCipher,
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *ClientUDPSession) getServerSession(sessionID uint64, now int64) (*serverSessionState, error) {
|
||||||
|
cur := s.current.Load()
|
||||||
|
if cur != nil && cur.sessionID == sessionID {
|
||||||
|
return cur, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
old := s.old.Load()
|
||||||
|
if old != nil && old.sessionID == sessionID {
|
||||||
|
if now-old.lastSeen.Load() > 60 {
|
||||||
|
s.old.CompareAndSwap(old, nil)
|
||||||
|
return nil, errors.New("old server session expired")
|
||||||
|
}
|
||||||
|
return old, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// New server session:
|
||||||
|
// Spec §3.2.4: reject newer server sessions when the last packet received from the old session is less than 1 minute old.
|
||||||
|
if old != nil && now-old.lastSeen.Load() < 60 {
|
||||||
|
return nil, errors.New("newer server session rejected: old session is less than 1 minute old")
|
||||||
|
}
|
||||||
|
|
||||||
|
var bodyAead cipher.AEAD
|
||||||
|
if !s.codec.method.IsChaCha {
|
||||||
|
var sessBytes [8]byte
|
||||||
|
binary.BigEndian.PutUint64(sessBytes[:], sessionID)
|
||||||
|
bodyKey := DeriveSessionSubKey(s.codec.psk, sessBytes[:], s.codec.method.KeySaltLength)
|
||||||
|
var err error
|
||||||
|
bodyAead, err = s.codec.method.NewAEAD(bodyKey)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
newState := &serverSessionState{
|
||||||
|
sessionID: sessionID,
|
||||||
|
cipher: bodyAead,
|
||||||
|
}
|
||||||
|
newState.lastSeen.Store(now)
|
||||||
|
|
||||||
|
if cur == nil {
|
||||||
|
s.current.CompareAndSwap(nil, newState)
|
||||||
|
return s.current.Load(), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
s.old.Store(cur)
|
||||||
|
s.current.Store(newState)
|
||||||
|
return newState, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *ClientUDPSession) ClientSessionID() uint64 {
|
||||||
|
return s.clientSessionID
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *ClientUDPSession) EncodePacket(dest net.Destination, payload []byte) (*buf.Buffer, error) {
|
||||||
|
packetID := s.nextPacketID.Add(1) - 1
|
||||||
|
sessID := s.clientSessionID
|
||||||
|
|
||||||
|
var paddingLen int
|
||||||
|
if dest.Port == 53 && len(payload) < MaxPaddingLength {
|
||||||
|
paddingLen = mrand.IntN(MaxPaddingLength) + 1
|
||||||
|
}
|
||||||
|
|
||||||
|
addrPortLen := AddrPortLength(dest)
|
||||||
|
|
||||||
|
if s.codec.method.IsChaCha {
|
||||||
|
totalLen := PacketNonceSize + 27 + paddingLen + addrPortLen + len(payload) + AEADTagSize
|
||||||
|
if totalLen > buf.Size {
|
||||||
|
return nil, ErrPacketTooLarge
|
||||||
|
}
|
||||||
|
|
||||||
|
outBuf := buf.New()
|
||||||
|
|
||||||
|
var nonce [PacketNonceSize]byte
|
||||||
|
if _, err := io.ReadFull(rand.Reader, nonce[:]); err != nil {
|
||||||
|
outBuf.Release()
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
outBuf.Write(nonce[:])
|
||||||
|
|
||||||
|
var hdr [16 + 1 + 8 + 2]byte
|
||||||
|
binary.BigEndian.PutUint64(hdr[0:8], sessID)
|
||||||
|
binary.BigEndian.PutUint64(hdr[8:16], packetID)
|
||||||
|
hdr[16] = HeaderTypeClient
|
||||||
|
binary.BigEndian.PutUint64(hdr[17:25], uint64(time.Now().Unix()))
|
||||||
|
binary.BigEndian.PutUint16(hdr[25:27], uint16(paddingLen))
|
||||||
|
outBuf.Write(hdr[:])
|
||||||
|
if paddingLen > 0 {
|
||||||
|
outBuf.Write(zeroPadding[:paddingLen])
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := WriteAddressPort(outBuf, dest); err != nil {
|
||||||
|
outBuf.Release()
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
outBuf.Write(payload)
|
||||||
|
|
||||||
|
plainBytes := outBuf.Bytes()[PacketNonceSize:]
|
||||||
|
outBuf.Extend(int32(s.codec.chachaCipher.Overhead()))
|
||||||
|
s.codec.chachaCipher.Seal(plainBytes[:0], nonce[:], plainBytes, nil)
|
||||||
|
return outBuf, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// AES mode
|
||||||
|
var sessBytes [8]byte
|
||||||
|
binary.BigEndian.PutUint64(sessBytes[:], sessID)
|
||||||
|
|
||||||
|
var rawHeader [16]byte
|
||||||
|
copy(rawHeader[:8], sessBytes[:])
|
||||||
|
binary.BigEndian.PutUint64(rawHeader[8:16], packetID)
|
||||||
|
|
||||||
|
eihCount := 0
|
||||||
|
if len(s.codec.pskList) > 1 {
|
||||||
|
eihCount = len(s.codec.pskList) - 1
|
||||||
|
}
|
||||||
|
|
||||||
|
totalLen := 16 + eihCount*16 + 11 + paddingLen + addrPortLen + len(payload) + AEADTagSize
|
||||||
|
if totalLen > buf.Size {
|
||||||
|
return nil, ErrPacketTooLarge
|
||||||
|
}
|
||||||
|
|
||||||
|
outBuf := buf.New()
|
||||||
|
|
||||||
|
if len(s.codec.pskList) > 1 {
|
||||||
|
var encryptedHeader [16]byte
|
||||||
|
s.codec.blockCiphers[0].Encrypt(encryptedHeader[:], rawHeader[:])
|
||||||
|
outBuf.Write(encryptedHeader[:])
|
||||||
|
|
||||||
|
for i := 0; i < len(s.codec.pskList)-1; i++ {
|
||||||
|
nextPSK := s.codec.pskList[i+1]
|
||||||
|
pskHash := DeriveUserPSKHash(nextPSK)
|
||||||
|
var eihPlain [16]byte
|
||||||
|
for k := 0; k < 16; k++ {
|
||||||
|
eihPlain[k] = pskHash[k] ^ rawHeader[k]
|
||||||
|
}
|
||||||
|
var encryptedEIH [16]byte
|
||||||
|
s.codec.blockCiphers[i].Encrypt(encryptedEIH[:], eihPlain[:])
|
||||||
|
outBuf.Write(encryptedEIH[:])
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
var encryptedHeader [16]byte
|
||||||
|
s.codec.blockCipher.Encrypt(encryptedHeader[:], rawHeader[:])
|
||||||
|
outBuf.Write(encryptedHeader[:])
|
||||||
|
}
|
||||||
|
|
||||||
|
bodyAead := s.clientBodyCipher
|
||||||
|
|
||||||
|
var hdr [1 + 8 + 2]byte
|
||||||
|
hdr[0] = HeaderTypeClient
|
||||||
|
binary.BigEndian.PutUint64(hdr[1:9], uint64(time.Now().Unix()))
|
||||||
|
binary.BigEndian.PutUint16(hdr[9:11], uint16(paddingLen))
|
||||||
|
outBuf.Write(hdr[:])
|
||||||
|
if paddingLen > 0 {
|
||||||
|
outBuf.Write(zeroPadding[:paddingLen])
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := WriteAddressPort(outBuf, dest); err != nil {
|
||||||
|
outBuf.Release()
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
outBuf.Write(payload)
|
||||||
|
|
||||||
|
headerOffset := 16 + eihCount*16
|
||||||
|
plainBytes := outBuf.Bytes()[headerOffset:]
|
||||||
|
bodyNonce := rawHeader[4:16]
|
||||||
|
outBuf.Extend(int32(bodyAead.Overhead()))
|
||||||
|
bodyAead.Seal(plainBytes[:0], bodyNonce, plainBytes, nil)
|
||||||
|
return outBuf, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *ClientUDPSession) DecodePacket(data []byte) (DecodedUDPPacket, error) {
|
||||||
|
if len(data) < PacketMinimalHeaderSize {
|
||||||
|
return DecodedUDPPacket{}, ErrPacketTooShort
|
||||||
|
}
|
||||||
|
|
||||||
|
if s.codec.method.IsChaCha {
|
||||||
|
if len(data) < PacketNonceSize+AEADTagSize {
|
||||||
|
return DecodedUDPPacket{}, ErrPacketTooShort
|
||||||
|
}
|
||||||
|
nonce := data[:PacketNonceSize]
|
||||||
|
ciphertext := data[PacketNonceSize:]
|
||||||
|
plain, err := s.codec.chachaCipher.Open(nil, nonce, ciphertext, nil)
|
||||||
|
if err != nil {
|
||||||
|
return DecodedUDPPacket{}, errors.New("failed to decrypt chacha udp packet").Base(err)
|
||||||
|
}
|
||||||
|
if len(plain) < 16+1+8+2 {
|
||||||
|
return DecodedUDPPacket{}, ErrPacketTooShort
|
||||||
|
}
|
||||||
|
|
||||||
|
sessionID := binary.BigEndian.Uint64(plain[:8])
|
||||||
|
packetID := binary.BigEndian.Uint64(plain[8:16])
|
||||||
|
|
||||||
|
now := time.Now().Unix()
|
||||||
|
st, err := s.getServerSession(sessionID, now)
|
||||||
|
if err != nil {
|
||||||
|
return DecodedUDPPacket{}, err
|
||||||
|
}
|
||||||
|
if !st.check(packetID) {
|
||||||
|
return DecodedUDPPacket{}, ErrPacketIdNotUnique
|
||||||
|
}
|
||||||
|
|
||||||
|
decoded, err := parsePlainUDPPacket(sessionID, packetID, plain[16:])
|
||||||
|
if err != nil {
|
||||||
|
return DecodedUDPPacket{}, err
|
||||||
|
}
|
||||||
|
|
||||||
|
if decoded.HeaderType != HeaderTypeServer {
|
||||||
|
return DecodedUDPPacket{}, ErrBadHeaderType
|
||||||
|
}
|
||||||
|
if decoded.ClientSessionID != s.clientSessionID {
|
||||||
|
return DecodedUDPPacket{}, errors.New("client session ID mismatch")
|
||||||
|
}
|
||||||
|
|
||||||
|
st.add(packetID)
|
||||||
|
st.lastSeen.Store(now)
|
||||||
|
|
||||||
|
return decoded, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// AES mode
|
||||||
|
var rawHeader [16]byte
|
||||||
|
s.codec.blockCipher.Decrypt(rawHeader[:], data[:16])
|
||||||
|
sessionID := binary.BigEndian.Uint64(rawHeader[:8])
|
||||||
|
packetID := binary.BigEndian.Uint64(rawHeader[8:16])
|
||||||
|
|
||||||
|
now := time.Now().Unix()
|
||||||
|
st, err := s.getServerSession(sessionID, now)
|
||||||
|
if err != nil {
|
||||||
|
return DecodedUDPPacket{}, err
|
||||||
|
}
|
||||||
|
if !st.check(packetID) {
|
||||||
|
return DecodedUDPPacket{}, ErrPacketIdNotUnique
|
||||||
|
}
|
||||||
|
bodyAead := st.cipher
|
||||||
|
|
||||||
|
bodyNonce := rawHeader[4:16]
|
||||||
|
bodyCipher := data[16:]
|
||||||
|
bodyPlain, err := bodyAead.Open(nil, bodyNonce, bodyCipher, nil)
|
||||||
|
if err != nil {
|
||||||
|
return DecodedUDPPacket{}, errors.New("failed to decrypt aes udp body").Base(err)
|
||||||
|
}
|
||||||
|
|
||||||
|
decoded, err := parsePlainUDPPacket(sessionID, packetID, bodyPlain)
|
||||||
|
if err != nil {
|
||||||
|
return DecodedUDPPacket{}, err
|
||||||
|
}
|
||||||
|
|
||||||
|
if decoded.HeaderType != HeaderTypeServer {
|
||||||
|
return DecodedUDPPacket{}, ErrBadHeaderType
|
||||||
|
}
|
||||||
|
if decoded.ClientSessionID != s.clientSessionID {
|
||||||
|
return DecodedUDPPacket{}, errors.New("client session ID mismatch")
|
||||||
|
}
|
||||||
|
|
||||||
|
st.add(packetID)
|
||||||
|
st.lastSeen.Store(now)
|
||||||
|
|
||||||
|
return decoded, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
type UDPWriter struct {
|
type UDPWriter struct {
|
||||||
Writer io.Writer
|
Writer io.Writer
|
||||||
Destination net.Destination
|
Destination net.Destination
|
||||||
Codec *UDPPacketCodec
|
Session *ClientUDPSession
|
||||||
}
|
}
|
||||||
|
|
||||||
func (w *UDPWriter) WriteMultiBuffer(mb buf.MultiBuffer) error {
|
func (w *UDPWriter) WriteMultiBuffer(mb buf.MultiBuffer) error {
|
||||||
@@ -468,7 +722,7 @@ func (w *UDPWriter) WriteMultiBuffer(mb buf.MultiBuffer) error {
|
|||||||
if b.UDP != nil {
|
if b.UDP != nil {
|
||||||
dest = *b.UDP
|
dest = *b.UDP
|
||||||
}
|
}
|
||||||
pktBuf, err := w.Codec.EncodeClientPacket(dest, b.Bytes())
|
pktBuf, err := w.Session.EncodePacket(dest, b.Bytes())
|
||||||
b.Release()
|
b.Release()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
buf.ReleaseMulti(mb)
|
buf.ReleaseMulti(mb)
|
||||||
@@ -485,8 +739,8 @@ func (w *UDPWriter) WriteMultiBuffer(mb buf.MultiBuffer) error {
|
|||||||
}
|
}
|
||||||
|
|
||||||
type UDPReader struct {
|
type UDPReader struct {
|
||||||
Reader io.Reader
|
Reader io.Reader
|
||||||
Codec *UDPPacketCodec
|
Session *ClientUDPSession
|
||||||
}
|
}
|
||||||
|
|
||||||
func (r *UDPReader) ReadMultiBuffer() (buf.MultiBuffer, error) {
|
func (r *UDPReader) ReadMultiBuffer() (buf.MultiBuffer, error) {
|
||||||
@@ -498,7 +752,7 @@ func (r *UDPReader) ReadMultiBuffer() (buf.MultiBuffer, error) {
|
|||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
decoded, err := r.Codec.DecodePacket(buffer.Bytes())
|
decoded, err := r.Session.DecodePacket(buffer.Bytes())
|
||||||
if err != nil {
|
if err != nil {
|
||||||
buffer.Release()
|
buffer.Release()
|
||||||
continue
|
continue
|
||||||
|
|||||||
@@ -2,9 +2,11 @@ package shadowsocks_2022_test
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
|
"crypto/rand"
|
||||||
"encoding/base64"
|
"encoding/base64"
|
||||||
"encoding/binary"
|
"encoding/binary"
|
||||||
"errors"
|
"errors"
|
||||||
|
"io"
|
||||||
gonet "net"
|
gonet "net"
|
||||||
"sync"
|
"sync"
|
||||||
"sync/atomic"
|
"sync/atomic"
|
||||||
@@ -269,3 +271,106 @@ func (c *dummyStatConn) WriteMultiBuffer(mb buf.MultiBuffer) error {
|
|||||||
}
|
}
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestRelayTCPHandshakeForwarding(t *testing.T) {
|
||||||
|
methods := []string{MethodAES128GCM, MethodAES256GCM}
|
||||||
|
for _, methodName := range methods {
|
||||||
|
t.Run(methodName, func(t *testing.T) {
|
||||||
|
method, err := GetCipherMethod(methodName)
|
||||||
|
common.Must(err)
|
||||||
|
|
||||||
|
relayKey := make([]byte, method.KeySaltLength)
|
||||||
|
destKey := make([]byte, method.KeySaltLength)
|
||||||
|
_, _ = io.ReadFull(rand.Reader, relayKey)
|
||||||
|
_, _ = io.ReadFull(rand.Reader, destKey)
|
||||||
|
|
||||||
|
targetPort := uint32(54321)
|
||||||
|
relayConfig := &RelayServerConfig{
|
||||||
|
Method: methodName,
|
||||||
|
Key: base64.StdEncoding.EncodeToString(relayKey),
|
||||||
|
Destinations: []*RelayDestination{
|
||||||
|
{
|
||||||
|
Key: base64.StdEncoding.EncodeToString(destKey),
|
||||||
|
Address: net.NewIPOrDomain(net.LocalHostIP),
|
||||||
|
Port: targetPort,
|
||||||
|
Email: "test@xray.com",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
testCtx := newTestContext()
|
||||||
|
inbound, err := NewRelayServer(testCtx, relayConfig)
|
||||||
|
common.Must(err)
|
||||||
|
|
||||||
|
targetDest := net.TCPDestination(net.LocalHostIP, net.Port(targetPort))
|
||||||
|
|
||||||
|
downstreamR, downstreamW := gonet.Pipe()
|
||||||
|
defer downstreamR.Close()
|
||||||
|
defer downstreamW.Close()
|
||||||
|
|
||||||
|
disp := &dummyDispatcher{
|
||||||
|
onDispatch: func(ctx context.Context, dest net.Destination) (*transport.Link, error) {
|
||||||
|
inLink := &transport.Link{
|
||||||
|
Reader: buf.NewReader(downstreamR),
|
||||||
|
Writer: &customWriter{
|
||||||
|
write: func(mb buf.MultiBuffer) error {
|
||||||
|
defer buf.ReleaseMulti(mb)
|
||||||
|
for _, b := range mb {
|
||||||
|
if _, err := downstreamW.Write(b.Bytes()); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
return inLink, nil
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
clientConn, relayConn := gonet.Pipe()
|
||||||
|
defer clientConn.Close()
|
||||||
|
defer relayConn.Close()
|
||||||
|
|
||||||
|
go func() {
|
||||||
|
_ = inbound.Process(testCtx, net.Network_TCP, &dummyStatConn{Conn: relayConn}, disp)
|
||||||
|
}()
|
||||||
|
|
||||||
|
clientSalt := make([]byte, method.KeySaltLength)
|
||||||
|
_, _ = io.ReadFull(rand.Reader, clientSalt)
|
||||||
|
pskList := [][]byte{relayKey, destKey}
|
||||||
|
|
||||||
|
go func() {
|
||||||
|
_, err := WriteTCPRequest(clientConn, method, pskList, targetDest, clientSalt, []byte("relay payload"))
|
||||||
|
if err != nil {
|
||||||
|
t.Errorf("WriteTCPRequest failed: %v", err)
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
|
||||||
|
// Downstream server must be able to read Salt + Fixed chunk in a single Read call!
|
||||||
|
headerLen := method.KeySaltLength + RequestHeaderFixedChunkLength + AEADTagSize
|
||||||
|
headerBuf := make([]byte, headerLen)
|
||||||
|
n, err := downstreamR.Read(headerBuf)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("downstream failed to read handshake: %v", err)
|
||||||
|
}
|
||||||
|
if n < headerLen {
|
||||||
|
t.Fatalf("downstream expected single read >= %d bytes, got %d", headerLen, n)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Verify downstream can decode the fixed chunk and subsequent payload
|
||||||
|
sessionKey := DeriveSessionSubKey(destKey, headerBuf[:method.KeySaltLength], method.KeySaltLength)
|
||||||
|
aead, err := method.NewAEAD(sessionKey)
|
||||||
|
common.Must(err)
|
||||||
|
|
||||||
|
reader := NewStreamReader(downstreamR, aead)
|
||||||
|
reqHeader, err := ReadClientRequestHeaderWithFixed(reader, headerBuf[method.KeySaltLength:])
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("downstream failed to parse client request header: %v", err)
|
||||||
|
}
|
||||||
|
if string(reqHeader.EarlyData) != "relay payload" {
|
||||||
|
t.Fatalf("payload mismatch: expected 'relay payload', got '%s'", string(reqHeader.EarlyData))
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
@@ -6,8 +6,11 @@ import (
|
|||||||
"sync/atomic"
|
"sync/atomic"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
|
"github.com/xtls/xray-core/common/net"
|
||||||
"github.com/xtls/xray-core/common/protocol"
|
"github.com/xtls/xray-core/common/protocol"
|
||||||
|
"github.com/xtls/xray-core/common/signal"
|
||||||
"github.com/xtls/xray-core/common/utils"
|
"github.com/xtls/xray-core/common/utils"
|
||||||
|
"github.com/xtls/xray-core/transport"
|
||||||
)
|
)
|
||||||
|
|
||||||
const (
|
const (
|
||||||
@@ -74,30 +77,42 @@ func (f *SlidingWindow) CheckAndAdd(counter uint64) bool {
|
|||||||
|
|
||||||
type ServerUDPSession struct {
|
type ServerUDPSession struct {
|
||||||
sync.Mutex
|
sync.Mutex
|
||||||
SessionID uint64
|
SessionID uint64
|
||||||
RemoteCipher atomic.Pointer[cipher.AEAD]
|
Window *SlidingWindow
|
||||||
Window SlidingWindow
|
User *protocol.MemoryUser
|
||||||
User *protocol.MemoryUser
|
UserPSK []byte
|
||||||
UserPSK []byte
|
LastActive atomic.Int64 // Unix timestamp in seconds
|
||||||
LastActive atomic.Int64 // Unix timestamp in seconds
|
|
||||||
|
clientBodyCipher cipher.AEAD
|
||||||
|
|
||||||
ServerSessionID uint64
|
ServerSessionID uint64
|
||||||
ServerPacketID atomic.Uint64
|
ServerPacketID atomic.Uint64
|
||||||
ServerCipher cipher.AEAD
|
serverBodyCipher cipher.AEAD
|
||||||
ServerBlockCipher cipher.Block
|
serverHeaderBlock cipher.Block
|
||||||
ServerChaCha cipher.AEAD
|
serverChaCha cipher.AEAD
|
||||||
|
|
||||||
|
manager *UDPSessionManager
|
||||||
|
link atomic.Pointer[transport.Link]
|
||||||
|
timer *signal.ActivityTimer
|
||||||
|
currentConn atomic.Value // stores stat.Connection
|
||||||
}
|
}
|
||||||
|
|
||||||
func (s *ServerUDPSession) GetRemoteCipher() cipher.AEAD {
|
func (s *ServerUDPSession) CheckPacketID(packetID uint64) bool {
|
||||||
ptr := s.RemoteCipher.Load()
|
s.Lock()
|
||||||
if ptr == nil {
|
defer s.Unlock()
|
||||||
return nil
|
if s.Window == nil {
|
||||||
|
s.Window = new(SlidingWindow)
|
||||||
}
|
}
|
||||||
return *ptr
|
return s.Window.Check(packetID)
|
||||||
}
|
}
|
||||||
|
|
||||||
func (s *ServerUDPSession) SetRemoteCipher(c cipher.AEAD) {
|
func (s *ServerUDPSession) AddPacketID(packetID uint64) {
|
||||||
s.RemoteCipher.Store(&c)
|
s.Lock()
|
||||||
|
defer s.Unlock()
|
||||||
|
if s.Window == nil {
|
||||||
|
s.Window = new(SlidingWindow)
|
||||||
|
}
|
||||||
|
s.Window.Add(packetID)
|
||||||
}
|
}
|
||||||
|
|
||||||
type UDPSessionManager struct {
|
type UDPSessionManager struct {
|
||||||
@@ -122,6 +137,7 @@ func (m *UDPSessionManager) GetOrCreate(sessionID uint64) *ServerUDPSession {
|
|||||||
|
|
||||||
s := &ServerUDPSession{
|
s := &ServerUDPSession{
|
||||||
SessionID: sessionID,
|
SessionID: sessionID,
|
||||||
|
manager: m,
|
||||||
}
|
}
|
||||||
s.LastActive.Store(now)
|
s.LastActive.Store(now)
|
||||||
|
|
||||||
@@ -148,6 +164,7 @@ func (m *UDPSessionManager) cleanup(now int64) {
|
|||||||
m.sessions.Range(func(k uint64, v *ServerUDPSession) bool {
|
m.sessions.Range(func(k uint64, v *ServerUDPSession) bool {
|
||||||
if now-v.LastActive.Load() > timeoutSec {
|
if now-v.LastActive.Load() > timeoutSec {
|
||||||
m.sessions.Delete(k)
|
m.sessions.Delete(k)
|
||||||
|
v.Close()
|
||||||
}
|
}
|
||||||
return true
|
return true
|
||||||
})
|
})
|
||||||
@@ -156,3 +173,11 @@ func (m *UDPSessionManager) cleanup(now int64) {
|
|||||||
func (m *UDPSessionManager) Delete(sessionID uint64) {
|
func (m *UDPSessionManager) Delete(sessionID uint64) {
|
||||||
m.sessions.Delete(sessionID)
|
m.sessions.Delete(sessionID)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (m *UDPSessionManager) EncodeServerPacket(method *CipherMethod, psk []byte, clientSessionID uint64, dest net.Destination, payload []byte) ([]byte, error) {
|
||||||
|
sessionItem := m.GetOrCreate(clientSessionID)
|
||||||
|
if err := sessionItem.EnsureServerState(method, psk); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
return sessionItem.EncodeServerPacket(method, clientSessionID, dest, payload)
|
||||||
|
}
|
||||||
|
|||||||
@@ -2,18 +2,161 @@ package shadowsocks_2022
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"sync"
|
|
||||||
|
|
||||||
|
"github.com/xtls/xray-core/common"
|
||||||
|
"github.com/xtls/xray-core/common/buf"
|
||||||
"github.com/xtls/xray-core/common/errors"
|
"github.com/xtls/xray-core/common/errors"
|
||||||
|
"github.com/xtls/xray-core/common/log"
|
||||||
|
"github.com/xtls/xray-core/common/net"
|
||||||
|
"github.com/xtls/xray-core/common/session"
|
||||||
"github.com/xtls/xray-core/common/signal"
|
"github.com/xtls/xray-core/common/signal"
|
||||||
|
"github.com/xtls/xray-core/features/policy"
|
||||||
|
"github.com/xtls/xray-core/features/routing"
|
||||||
|
"github.com/xtls/xray-core/proxy"
|
||||||
"github.com/xtls/xray-core/transport"
|
"github.com/xtls/xray-core/transport"
|
||||||
|
"github.com/xtls/xray-core/transport/internet/stat"
|
||||||
)
|
)
|
||||||
|
|
||||||
type udpConnEntry struct {
|
func (s *ServerUDPSession) UpdateConn(conn stat.Connection) {
|
||||||
sync.Mutex
|
if s.currentConn.Load() == nil {
|
||||||
link *transport.Link
|
s.currentConn.Store(conn)
|
||||||
timer *signal.ActivityTimer
|
}
|
||||||
cancel context.CancelFunc
|
if s.timer != nil {
|
||||||
|
s.timer.Update()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *ServerUDPSession) WriteToClient(b []byte) error {
|
||||||
|
connVal := s.currentConn.Load()
|
||||||
|
if connVal == nil {
|
||||||
|
return errors.New("client connection closed")
|
||||||
|
}
|
||||||
|
conn, ok := connVal.(stat.Connection)
|
||||||
|
if !ok || conn == nil {
|
||||||
|
return errors.New("client connection closed")
|
||||||
|
}
|
||||||
|
_, err := conn.Write(b)
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *ServerUDPSession) Close() {
|
||||||
|
if s.timer != nil {
|
||||||
|
s.timer.SetTimeout(0)
|
||||||
|
}
|
||||||
|
if link := s.link.Load(); link != nil {
|
||||||
|
common.Interrupt(link.Reader)
|
||||||
|
common.Interrupt(link.Writer)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *ServerUDPSession) EnsureLink(
|
||||||
|
ctx context.Context,
|
||||||
|
conn stat.Connection,
|
||||||
|
dest net.Destination,
|
||||||
|
dispatcher routing.Dispatcher,
|
||||||
|
policyManager policy.Manager,
|
||||||
|
responseEncoder func(dest net.Destination, payload []byte) ([]byte, error),
|
||||||
|
) (*transport.Link, error) {
|
||||||
|
s.UpdateConn(conn)
|
||||||
|
|
||||||
|
if link := s.link.Load(); link != nil {
|
||||||
|
return link, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
s.Lock()
|
||||||
|
defer s.Unlock()
|
||||||
|
|
||||||
|
if link := s.link.Load(); link != nil {
|
||||||
|
return link, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
sessCtx, cancel := context.WithCancel(ctx)
|
||||||
|
inbound := session.InboundFromContext(sessCtx)
|
||||||
|
if inbound != nil && s.User != nil {
|
||||||
|
inbound.User = s.User
|
||||||
|
}
|
||||||
|
var email string
|
||||||
|
var level uint32
|
||||||
|
if s.User != nil {
|
||||||
|
email = s.User.Email
|
||||||
|
level = s.User.Level
|
||||||
|
}
|
||||||
|
sessCtx = log.ContextWithAccessMessage(sessCtx, &log.AccessMessage{
|
||||||
|
From: conn.RemoteAddr(),
|
||||||
|
To: dest,
|
||||||
|
Status: log.AccessAccepted,
|
||||||
|
Email: email,
|
||||||
|
})
|
||||||
|
|
||||||
|
link, err := dispatcher.Dispatch(sessCtx, dest)
|
||||||
|
if err != nil {
|
||||||
|
cancel()
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
s.link.Store(link)
|
||||||
|
sessionPolicy := policyManager.ForLevel(level)
|
||||||
|
s.timer = signal.CancelAfterInactivity(sessCtx, func() {
|
||||||
|
if s.manager != nil {
|
||||||
|
s.manager.Delete(s.SessionID)
|
||||||
|
}
|
||||||
|
s.Close()
|
||||||
|
cancel()
|
||||||
|
}, sessionPolicy.Timeouts.ConnectionIdle)
|
||||||
|
|
||||||
|
go handleUDPResponse(s, link, dest, responseEncoder)
|
||||||
|
return link, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// ResetTCPConn sets SO_LINGER to 0 per SIP022 §3.1.4 to consistently send RST on close
|
||||||
|
// when handshake or header validation fails.
|
||||||
|
func ResetTCPConn(conn net.Conn) {
|
||||||
|
rawConn, _, _ := proxy.UnwrapRawConn(conn)
|
||||||
|
if tcpConn, ok := rawConn.(*net.TCPConn); ok {
|
||||||
|
_ = tcpConn.SetLinger(0)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func handleUDPResponse(s *ServerUDPSession, link *transport.Link, fallbackDest net.Destination, encode func(dest net.Destination, payload []byte) ([]byte, error)) {
|
||||||
|
defer func() {
|
||||||
|
if s.timer != nil {
|
||||||
|
s.timer.SetTimeout(0)
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
for {
|
||||||
|
resMb, err := link.Reader.ReadMultiBuffer()
|
||||||
|
if err != nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if s.timer != nil {
|
||||||
|
s.timer.Update()
|
||||||
|
}
|
||||||
|
for i, rb := range resMb {
|
||||||
|
b := rb.Bytes()
|
||||||
|
if encode != nil {
|
||||||
|
replyDest := fallbackDest
|
||||||
|
if rb.UDP != nil {
|
||||||
|
replyDest = *rb.UDP
|
||||||
|
}
|
||||||
|
encPacket, err := encode(replyDest, b)
|
||||||
|
rb.Release()
|
||||||
|
if err != nil {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if err := s.WriteToClient(encPacket); err != nil {
|
||||||
|
buf.ReleaseMulti(resMb[i+1:])
|
||||||
|
return
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
err := s.WriteToClient(b)
|
||||||
|
rb.Release()
|
||||||
|
if err != nil {
|
||||||
|
buf.ReleaseMulti(resMb[i+1:])
|
||||||
|
return
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
const (
|
const (
|
||||||
|
|||||||
@@ -182,57 +182,48 @@ func TestTCPStream(t *testing.T) {
|
|||||||
common.Must(err)
|
common.Must(err)
|
||||||
IncreaseNonce(reader.Nonce())
|
IncreaseNonce(reader.Nonce())
|
||||||
|
|
||||||
vBuf := buf.New()
|
dest, addrLen, err := ParseAddressPort(plainVar)
|
||||||
vBuf.Write(plainVar)
|
|
||||||
receivedDest, err = ReadAddressPort(vBuf)
|
|
||||||
common.Must(err)
|
common.Must(err)
|
||||||
|
receivedDest = net.TCPDestination(dest.Address, dest.Port)
|
||||||
|
plainVar = plainVar[addrLen:]
|
||||||
|
padLen := int(binary.BigEndian.Uint16(plainVar[:2]))
|
||||||
|
receivedPayload = plainVar[2+padLen:]
|
||||||
|
|
||||||
// Skip padding
|
// Server sends response stream with receivedPayload as first payload
|
||||||
var padBytes [2]byte
|
writer := NewServerStreamWriter(serverConn, method, rawKey, salt)
|
||||||
_, _ = vBuf.Read(padBytes[:])
|
pBuf := buf.New()
|
||||||
padLen := int(padBytes[0])<<8 | int(padBytes[1])
|
pBuf.Write(receivedPayload)
|
||||||
vBuf.Advance(int32(padLen))
|
_ = writer.WriteMultiBuffer(buf.MultiBuffer{pBuf})
|
||||||
|
|
||||||
receivedPayload = make([]byte, vBuf.Len())
|
// Read and echo additional stream data
|
||||||
copy(receivedPayload, vBuf.Bytes())
|
|
||||||
vBuf.Release()
|
|
||||||
|
|
||||||
// Server sends response handshake
|
|
||||||
serverSalt := make([]byte, method.KeySaltLength)
|
|
||||||
_, _ = rand.Read(serverSalt)
|
|
||||||
respKey := DeriveSessionSubKey(rawKey, serverSalt, method.KeySaltLength)
|
|
||||||
respAead, err := method.NewAEAD(respKey)
|
|
||||||
writer := NewStreamWriter(serverConn, respAead)
|
|
||||||
_, _ = serverConn.Write(serverSalt)
|
|
||||||
|
|
||||||
fixedResp := make([]byte, 1+8+method.KeySaltLength+2)
|
|
||||||
fixedResp[0] = HeaderTypeServer
|
|
||||||
binary.BigEndian.PutUint64(fixedResp[1:9], uint64(time.Now().Unix()))
|
|
||||||
copy(fixedResp[9:9+method.KeySaltLength], salt)
|
|
||||||
binary.BigEndian.PutUint16(fixedResp[9+method.KeySaltLength:11+method.KeySaltLength], 0)
|
|
||||||
|
|
||||||
fixedChunk := respAead.Seal(nil, writer.Nonce(), fixedResp, nil)
|
|
||||||
IncreaseNonce(writer.Nonce())
|
|
||||||
_, _ = serverConn.Write(fixedChunk)
|
|
||||||
|
|
||||||
// Echo stream data
|
|
||||||
mb, err := reader.ReadMultiBuffer()
|
mb, err := reader.ReadMultiBuffer()
|
||||||
common.Must(err)
|
common.Must(err)
|
||||||
_ = writer.WriteMultiBuffer(mb)
|
_ = writer.WriteMultiBuffer(mb)
|
||||||
|
_ = writer.Close()
|
||||||
}()
|
}()
|
||||||
|
|
||||||
// Client goroutine
|
// Client goroutine
|
||||||
go func() {
|
go func() {
|
||||||
defer wg.Done()
|
defer wg.Done()
|
||||||
clientSalt, writer, err := ClientHandshake(clientConn, method, [][]byte{rawKey}, dest, testPayload)
|
clientSalt := make([]byte, method.KeySaltLength)
|
||||||
|
common.Must2(io.ReadFull(rand.Reader, clientSalt))
|
||||||
|
writer, err := WriteTCPRequest(clientConn, method, [][]byte{rawKey}, dest, clientSalt, testPayload)
|
||||||
common.Must(err)
|
common.Must(err)
|
||||||
|
|
||||||
reader, _, err := ClientVerifyServerResponse(clientConn, method, rawKey, clientSalt)
|
reader, err := ReadTCPResponse(clientConn, method, rawKey, clientSalt)
|
||||||
common.Must(err)
|
common.Must(err)
|
||||||
|
|
||||||
|
// The first ReadMultiBuffer drains initialPayload from reader cache
|
||||||
|
mbInit, err := reader.ReadMultiBuffer()
|
||||||
|
common.Must(err)
|
||||||
|
if !bytes.Equal(mbInit[0].Bytes(), testPayload) {
|
||||||
|
t.Errorf("drained initial payload mismatch: got %s, want %s", mbInit[0].Bytes(), testPayload)
|
||||||
|
}
|
||||||
|
buf.ReleaseMulti(mbInit)
|
||||||
|
|
||||||
// Send additional stream data
|
// Send additional stream data
|
||||||
streamData := []byte("stream chunk test")
|
streamData := []byte("stream chunk test")
|
||||||
_ = writer.WriteChunk(streamData)
|
_ = writer.WriteMultiBuffer(buf.MultiBuffer{buf.FromBytes(streamData)})
|
||||||
|
|
||||||
mb, err := reader.ReadMultiBuffer()
|
mb, err := reader.ReadMultiBuffer()
|
||||||
common.Must(err)
|
common.Must(err)
|
||||||
@@ -272,12 +263,14 @@ func TestUDPCodec(t *testing.T) {
|
|||||||
psk := make([]byte, method.KeySaltLength)
|
psk := make([]byte, method.KeySaltLength)
|
||||||
_, _ = rand.Read(psk)
|
_, _ = rand.Read(psk)
|
||||||
|
|
||||||
clientCodec, err := NewUDPPacketCodec(method, psk)
|
clientCodec, err := NewUDPPacketCodec(method, [][]byte{psk})
|
||||||
common.Must(err)
|
common.Must(err)
|
||||||
serverCodec, err := NewUDPServerCodec(method, psk, time.Minute)
|
serverCodec, err := NewUDPServerCodec(method, psk, time.Minute)
|
||||||
common.Must(err)
|
common.Must(err)
|
||||||
|
|
||||||
pktBuf, err := clientCodec.EncodeClientPacket(dest, payload)
|
session, err := clientCodec.NewClientSession()
|
||||||
|
common.Must(err)
|
||||||
|
pktBuf, err := session.EncodePacket(dest, payload)
|
||||||
common.Must(err)
|
common.Must(err)
|
||||||
defer pktBuf.Release()
|
defer pktBuf.Release()
|
||||||
|
|
||||||
@@ -360,3 +353,146 @@ func TestMultiUserManager(t *testing.T) {
|
|||||||
t.Fatal("user1 should have been removed")
|
t.Fatal("user1 should have been removed")
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestLargeStreamTransfer(t *testing.T) {
|
||||||
|
method, err := GetCipherMethod(MethodAES128GCM)
|
||||||
|
common.Must(err)
|
||||||
|
sessionKey := make([]byte, 16)
|
||||||
|
_, _ = rand.Read(sessionKey)
|
||||||
|
|
||||||
|
clientAead, err := method.NewAEAD(sessionKey)
|
||||||
|
common.Must(err)
|
||||||
|
serverAead, err := method.NewAEAD(sessionKey)
|
||||||
|
common.Must(err)
|
||||||
|
|
||||||
|
r, w := io.Pipe()
|
||||||
|
defer r.Close()
|
||||||
|
defer w.Close()
|
||||||
|
|
||||||
|
writer := NewStreamWriter(w, clientAead)
|
||||||
|
reader := NewStreamReader(r, serverAead)
|
||||||
|
|
||||||
|
const totalSize = 100 * 1024 // 100 KB
|
||||||
|
data := make([]byte, totalSize)
|
||||||
|
_, _ = rand.Read(data)
|
||||||
|
|
||||||
|
errCh := make(chan error, 1)
|
||||||
|
go func() {
|
||||||
|
// Write using Write (which splits by MaxPacketSize = 65535)
|
||||||
|
_, werr := writer.Write(data)
|
||||||
|
if werr != nil {
|
||||||
|
errCh <- werr
|
||||||
|
return
|
||||||
|
}
|
||||||
|
_ = w.Close()
|
||||||
|
errCh <- nil
|
||||||
|
}()
|
||||||
|
|
||||||
|
var received []byte
|
||||||
|
for {
|
||||||
|
mb, rerr := reader.ReadMultiBuffer()
|
||||||
|
if !mb.IsEmpty() {
|
||||||
|
for _, b := range mb {
|
||||||
|
received = append(received, b.Bytes()...)
|
||||||
|
}
|
||||||
|
buf.ReleaseMulti(mb)
|
||||||
|
}
|
||||||
|
if rerr != nil {
|
||||||
|
if rerr == io.EOF {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
t.Fatalf("ReadMultiBuffer error: %v", rerr)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if werr := <-errCh; werr != nil {
|
||||||
|
t.Fatalf("writer error: %v", werr)
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(received) != totalSize {
|
||||||
|
t.Fatalf("received size mismatch: got %d, want %d", len(received), totalSize)
|
||||||
|
}
|
||||||
|
if !bytes.Equal(received, data) {
|
||||||
|
t.Fatal("received data does not match sent data")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestClientUDPSessionMultiDestination(t *testing.T) {
|
||||||
|
for _, methodName := range []string{MethodAES128GCM, MethodAES256GCM, MethodChaCha20Poly1305} {
|
||||||
|
t.Run(methodName, func(t *testing.T) {
|
||||||
|
method, err := GetCipherMethod(methodName)
|
||||||
|
common.Must(err)
|
||||||
|
rawKey := make([]byte, method.KeySaltLength)
|
||||||
|
_, _ = rand.Read(rawKey)
|
||||||
|
|
||||||
|
clientCodec, err := NewUDPPacketCodec(method, [][]byte{rawKey})
|
||||||
|
common.Must(err)
|
||||||
|
serverCodec, err := NewUDPServerCodec(method, rawKey, time.Minute)
|
||||||
|
common.Must(err)
|
||||||
|
|
||||||
|
session, err := clientCodec.NewClientSession()
|
||||||
|
common.Must(err)
|
||||||
|
|
||||||
|
dest1 := net.UDPDestination(net.LocalHostIP, net.Port(53))
|
||||||
|
dest2 := net.UDPDestination(net.IPAddress([]byte{127, 0, 0, 2}), net.Port(53))
|
||||||
|
|
||||||
|
payload1 := []byte("query-google-dns")
|
||||||
|
payload2 := []byte("query-cloudflare-dns")
|
||||||
|
|
||||||
|
// Client sends to dest1 and dest2 using SAME session
|
||||||
|
pkt1, err := session.EncodePacket(dest1, payload1)
|
||||||
|
common.Must(err)
|
||||||
|
defer pkt1.Release()
|
||||||
|
pkt2, err := session.EncodePacket(dest2, payload2)
|
||||||
|
common.Must(err)
|
||||||
|
defer pkt2.Release()
|
||||||
|
|
||||||
|
// Server decodes both
|
||||||
|
dec1, err := serverCodec.DecodePacket(pkt1.Bytes())
|
||||||
|
common.Must(err)
|
||||||
|
dec2, err := serverCodec.DecodePacket(pkt2.Bytes())
|
||||||
|
common.Must(err)
|
||||||
|
|
||||||
|
if dec1.SessionID != session.ClientSessionID() || dec2.SessionID != session.ClientSessionID() {
|
||||||
|
t.Fatalf("both packets must share client session ID %d, got %d and %d", session.ClientSessionID(), dec1.SessionID, dec2.SessionID)
|
||||||
|
}
|
||||||
|
if dec1.Destination.String() != dest1.String() {
|
||||||
|
t.Fatalf("expected dest1 %s, got %s", dest1, dec1.Destination)
|
||||||
|
}
|
||||||
|
if dec2.Destination.String() != dest2.String() {
|
||||||
|
t.Fatalf("expected dest2 %s, got %s", dest2, dec2.Destination)
|
||||||
|
}
|
||||||
|
if !bytes.Equal(dec1.Payload, payload1) || !bytes.Equal(dec2.Payload, payload2) {
|
||||||
|
t.Fatal("payload mismatch")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Server replies to dest1 and dest2
|
||||||
|
respPayload1 := []byte("reply-google-dns")
|
||||||
|
respPayload2 := []byte("reply-cloudflare-dns")
|
||||||
|
|
||||||
|
respPkt1, err := serverCodec.EncodeServerPacket(dec1.SessionID, dest1, respPayload1)
|
||||||
|
common.Must(err)
|
||||||
|
respPkt2, err := serverCodec.EncodeServerPacket(dec2.SessionID, dest2, respPayload2)
|
||||||
|
common.Must(err)
|
||||||
|
|
||||||
|
// Client decodes replies
|
||||||
|
clientDec1, err := session.DecodePacket(respPkt1)
|
||||||
|
common.Must(err)
|
||||||
|
if clientDec1.Destination.String() != dest1.String() {
|
||||||
|
t.Fatalf("expected client dec1 dest %s, got %s", dest1, clientDec1.Destination)
|
||||||
|
}
|
||||||
|
if !bytes.Equal(clientDec1.Payload, respPayload1) {
|
||||||
|
t.Fatal("reply payload 1 mismatch")
|
||||||
|
}
|
||||||
|
|
||||||
|
clientDec2, err := session.DecodePacket(respPkt2)
|
||||||
|
common.Must(err)
|
||||||
|
if clientDec2.Destination.String() != dest2.String() {
|
||||||
|
t.Fatalf("expected client dec2 dest %s, got %s", dest2, clientDec2.Destination)
|
||||||
|
}
|
||||||
|
if !bytes.Equal(clientDec2.Payload, respPayload2) {
|
||||||
|
t.Fatal("reply payload 2 mismatch")
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
+234
-114
@@ -1,18 +1,25 @@
|
|||||||
package shadowsocks_2022
|
package shadowsocks_2022
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"context"
|
||||||
"crypto/cipher"
|
"crypto/cipher"
|
||||||
"crypto/rand"
|
"crypto/rand"
|
||||||
"encoding/binary"
|
"encoding/binary"
|
||||||
"io"
|
"io"
|
||||||
"math"
|
"math"
|
||||||
mrand "math/rand/v2"
|
mrand "math/rand/v2"
|
||||||
|
"sync"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
|
"github.com/xtls/xray-core/common/antireplay"
|
||||||
"github.com/xtls/xray-core/common/buf"
|
"github.com/xtls/xray-core/common/buf"
|
||||||
"github.com/xtls/xray-core/common/errors"
|
"github.com/xtls/xray-core/common/errors"
|
||||||
"github.com/xtls/xray-core/common/net"
|
"github.com/xtls/xray-core/common/net"
|
||||||
"github.com/xtls/xray-core/common/protocol"
|
"github.com/xtls/xray-core/common/protocol"
|
||||||
|
"github.com/xtls/xray-core/common/signal"
|
||||||
|
"github.com/xtls/xray-core/common/task"
|
||||||
|
"github.com/xtls/xray-core/features/policy"
|
||||||
|
"github.com/xtls/xray-core/transport"
|
||||||
)
|
)
|
||||||
|
|
||||||
var addrParser = protocol.NewAddressParser(
|
var addrParser = protocol.NewAddressParser(
|
||||||
@@ -38,15 +45,6 @@ func WriteAddressPort(w io.Writer, dest net.Destination) error {
|
|||||||
return addrParser.WriteAddressPort(w, dest.Address, dest.Port)
|
return addrParser.WriteAddressPort(w, dest.Address, dest.Port)
|
||||||
}
|
}
|
||||||
|
|
||||||
// ReadAddressPort reads a destination address and port in SOCKS5 format
|
|
||||||
func ReadAddressPort(r io.Reader) (net.Destination, error) {
|
|
||||||
addr, port, err := addrParser.ReadAddressPort(nil, r)
|
|
||||||
if err != nil {
|
|
||||||
return net.Destination{}, err
|
|
||||||
}
|
|
||||||
return net.TCPDestination(addr, port), nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// AddrPortLength returns the serialized length of a destination in SOCKS5 format
|
// AddrPortLength returns the serialized length of a destination in SOCKS5 format
|
||||||
func AddrPortLength(dest net.Destination) int {
|
func AddrPortLength(dest net.Destination) int {
|
||||||
switch dest.Address.Family() {
|
switch dest.Address.Family() {
|
||||||
@@ -119,8 +117,16 @@ func (w *StreamWriter) Write(p []byte) (int, error) {
|
|||||||
func (w *StreamWriter) WriteMultiBuffer(mb buf.MultiBuffer) error {
|
func (w *StreamWriter) WriteMultiBuffer(mb buf.MultiBuffer) error {
|
||||||
defer buf.ReleaseMulti(mb)
|
defer buf.ReleaseMulti(mb)
|
||||||
for _, b := range mb {
|
for _, b := range mb {
|
||||||
if err := w.WriteChunk(b.Bytes()); err != nil {
|
p := b.Bytes()
|
||||||
return err
|
for len(p) > 0 {
|
||||||
|
chunkSize := len(p)
|
||||||
|
if chunkSize > MaxPacketSize {
|
||||||
|
chunkSize = MaxPacketSize
|
||||||
|
}
|
||||||
|
if err := w.WriteChunk(p[:chunkSize]); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
p = p[chunkSize:]
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
return nil
|
return nil
|
||||||
@@ -168,7 +174,7 @@ func (r *StreamReader) Read(p []byte) (int, error) {
|
|||||||
IncreaseNonce(r.nonce[:])
|
IncreaseNonce(r.nonce[:])
|
||||||
|
|
||||||
payloadLen := int(binary.BigEndian.Uint16(decryptedLen))
|
payloadLen := int(binary.BigEndian.Uint16(decryptedLen))
|
||||||
if payloadLen == 0 {
|
if payloadLen == 0 || payloadLen > MaxPacketSize {
|
||||||
return 0, ErrInvalidRequest
|
return 0, ErrInvalidRequest
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -194,11 +200,10 @@ func (r *StreamReader) Read(p []byte) (int, error) {
|
|||||||
|
|
||||||
func (r *StreamReader) ReadMultiBuffer() (buf.MultiBuffer, error) {
|
func (r *StreamReader) ReadMultiBuffer() (buf.MultiBuffer, error) {
|
||||||
if r.cached > 0 {
|
if r.cached > 0 {
|
||||||
b := buf.New()
|
mb := buf.MergeBytes(nil, r.buffer[r.offset:r.offset+r.cached])
|
||||||
b.Write(r.buffer[r.offset : r.offset+r.cached])
|
|
||||||
r.cached = 0
|
r.cached = 0
|
||||||
r.offset = 0
|
r.offset = 0
|
||||||
return buf.MultiBuffer{b}, nil
|
return mb, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
if _, err := io.ReadFull(r.reader, r.lenBuf[:]); err != nil {
|
if _, err := io.ReadFull(r.reader, r.lenBuf[:]); err != nil {
|
||||||
@@ -212,7 +217,7 @@ func (r *StreamReader) ReadMultiBuffer() (buf.MultiBuffer, error) {
|
|||||||
IncreaseNonce(r.nonce[:])
|
IncreaseNonce(r.nonce[:])
|
||||||
|
|
||||||
payloadLen := int(binary.BigEndian.Uint16(decryptedLen))
|
payloadLen := int(binary.BigEndian.Uint16(decryptedLen))
|
||||||
if payloadLen == 0 {
|
if payloadLen == 0 || payloadLen > MaxPacketSize {
|
||||||
return nil, ErrInvalidRequest
|
return nil, ErrInvalidRequest
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -227,9 +232,8 @@ func (r *StreamReader) ReadMultiBuffer() (buf.MultiBuffer, error) {
|
|||||||
}
|
}
|
||||||
IncreaseNonce(r.nonce[:])
|
IncreaseNonce(r.nonce[:])
|
||||||
|
|
||||||
b := buf.New()
|
mb := buf.MergeBytes(nil, decryptedPayload)
|
||||||
b.Write(decryptedPayload)
|
return mb, nil
|
||||||
return buf.MultiBuffer{b}, nil
|
|
||||||
}
|
}
|
||||||
|
|
||||||
type ClientRequestHeader struct {
|
type ClientRequestHeader struct {
|
||||||
@@ -237,13 +241,8 @@ type ClientRequestHeader struct {
|
|||||||
EarlyData []byte
|
EarlyData []byte
|
||||||
}
|
}
|
||||||
|
|
||||||
func ReadClientRequestHeader(conn io.Reader, reader *StreamReader) (*ClientRequestHeader, error) {
|
func ReadClientRequestHeaderWithFixed(reader *StreamReader, fixedChunk []byte) (*ClientRequestHeader, error) {
|
||||||
var fixedBuf [RequestHeaderFixedChunkLength + AEADTagSize]byte
|
plainFixed, err := reader.cipher.Open(fixedChunk[:0], reader.Nonce(), fixedChunk, nil)
|
||||||
if _, err := io.ReadFull(conn, fixedBuf[:]); err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
plainFixed, err := reader.cipher.Open(fixedBuf[:0], reader.Nonce(), fixedBuf[:], nil)
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, errors.New("failed to decrypt client request header").Base(err)
|
return nil, errors.New("failed to decrypt client request header").Base(err)
|
||||||
}
|
}
|
||||||
@@ -272,7 +271,7 @@ func ReadClientRequestHeader(conn io.Reader, reader *StreamReader) (*ClientReque
|
|||||||
} else {
|
} else {
|
||||||
varChunkCipher = make([]byte, needed)
|
varChunkCipher = make([]byte, needed)
|
||||||
}
|
}
|
||||||
if _, err := io.ReadFull(conn, varChunkCipher); err != nil {
|
if _, err := io.ReadFull(reader.reader, varChunkCipher); err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -282,31 +281,34 @@ func ReadClientRequestHeader(conn io.Reader, reader *StreamReader) (*ClientReque
|
|||||||
}
|
}
|
||||||
IncreaseNonce(reader.Nonce())
|
IncreaseNonce(reader.Nonce())
|
||||||
|
|
||||||
b := buf.New()
|
dest, addrLen, err := ParseAddressPort(plainVar)
|
||||||
b.Write(plainVar)
|
|
||||||
defer b.Release()
|
|
||||||
|
|
||||||
dest, err := ReadAddressPort(b)
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
dest.Network = net.Network_TCP
|
||||||
|
|
||||||
var padLenBytes [2]byte
|
offset := addrLen
|
||||||
if _, err := b.Read(padLenBytes[:]); err != nil {
|
if len(plainVar) < offset+2 {
|
||||||
return nil, err
|
return nil, ErrPacketTooShort
|
||||||
}
|
}
|
||||||
paddingLen := int(binary.BigEndian.Uint16(padLenBytes[:]))
|
paddingLen := int(binary.BigEndian.Uint16(plainVar[offset : offset+2]))
|
||||||
if int(b.Len()) < paddingLen {
|
offset += 2
|
||||||
|
|
||||||
|
if len(plainVar) < offset+paddingLen {
|
||||||
return nil, ErrNoPadding
|
return nil, ErrNoPadding
|
||||||
}
|
}
|
||||||
if paddingLen > 0 {
|
offset += paddingLen
|
||||||
b.Advance(int32(paddingLen))
|
|
||||||
}
|
|
||||||
|
|
||||||
var earlyData []byte
|
var earlyData []byte
|
||||||
if b.Len() > 0 {
|
var payloadLen int
|
||||||
earlyData = make([]byte, b.Len())
|
if len(plainVar) > offset {
|
||||||
copy(earlyData, b.Bytes())
|
earlyData = plainVar[offset:]
|
||||||
|
payloadLen = len(earlyData)
|
||||||
|
}
|
||||||
|
|
||||||
|
// SIP022 §3.1.4: Servers MUST reject the request if the variable-length header chunk does not contain payload and the padding length is 0.
|
||||||
|
if paddingLen == 0 && payloadLen == 0 {
|
||||||
|
return nil, errors.New("request without payload and padding is not allowed")
|
||||||
}
|
}
|
||||||
|
|
||||||
return &ClientRequestHeader{
|
return &ClientRequestHeader{
|
||||||
@@ -315,34 +317,6 @@ func ReadClientRequestHeader(conn io.Reader, reader *StreamReader) (*ClientReque
|
|||||||
}, nil
|
}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// ClientHandshake writes the full client request header to w
|
|
||||||
func ClientHandshake(w io.Writer, method *CipherMethod, pskList [][]byte, dest net.Destination, payload []byte) ([]byte, *StreamWriter, error) {
|
|
||||||
salt := make([]byte, method.KeySaltLength)
|
|
||||||
if _, err := io.ReadFull(rand.Reader, salt); err != nil {
|
|
||||||
return nil, nil, err
|
|
||||||
}
|
|
||||||
writer, err := WriteTCPRequest(w, method, pskList, dest, salt, payload)
|
|
||||||
if err != nil {
|
|
||||||
return nil, nil, err
|
|
||||||
}
|
|
||||||
return salt, writer.(*StreamWriter), nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// ClientVerifyServerResponse reads and verifies the server's handshake response
|
|
||||||
func ClientVerifyServerResponse(r io.Reader, method *CipherMethod, psk []byte, clientSalt []byte) (*StreamReader, []byte, error) {
|
|
||||||
reader, err := ReadTCPResponse(r, method, psk, clientSalt)
|
|
||||||
if err != nil {
|
|
||||||
return nil, nil, err
|
|
||||||
}
|
|
||||||
sr := reader.(*StreamReader)
|
|
||||||
var initialPayload []byte
|
|
||||||
if sr.cached > 0 {
|
|
||||||
initialPayload = make([]byte, sr.cached)
|
|
||||||
copy(initialPayload, sr.buffer[sr.offset:sr.offset+sr.cached])
|
|
||||||
}
|
|
||||||
return sr, initialPayload, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// WriteTCPRequest writes the Shadowsocks 2022 request header into w and returns a body writer.
|
// WriteTCPRequest writes the Shadowsocks 2022 request header into w and returns a body writer.
|
||||||
func WriteTCPRequest(w io.Writer, method *CipherMethod, pskList [][]byte, dest net.Destination, clientSalt []byte, payload []byte) (buf.Writer, error) {
|
func WriteTCPRequest(w io.Writer, method *CipherMethod, pskList [][]byte, dest net.Destination, clientSalt []byte, payload []byte) (buf.Writer, error) {
|
||||||
finalPSK := pskList[len(pskList)-1]
|
finalPSK := pskList[len(pskList)-1]
|
||||||
@@ -354,7 +328,16 @@ func WriteTCPRequest(w io.Writer, method *CipherMethod, pskList [][]byte, dest n
|
|||||||
|
|
||||||
writer := NewStreamWriter(w, aead)
|
writer := NewStreamWriter(w, aead)
|
||||||
|
|
||||||
handshakeBuf := buf.New()
|
payloadLen := len(payload)
|
||||||
|
var paddingLen int
|
||||||
|
if payloadLen < MaxPaddingLength {
|
||||||
|
paddingLen = mrand.IntN(MaxPaddingLength) + 1
|
||||||
|
}
|
||||||
|
addrPortLen := AddrPortLength(dest)
|
||||||
|
varHeaderLen := addrPortLen + 2 + paddingLen + payloadLen
|
||||||
|
|
||||||
|
totalHandshakeLen := int32(method.KeySaltLength + len(pskList)*AESBlockSize + RequestHeaderFixedChunkLength + AEADTagSize + varHeaderLen + AEADTagSize)
|
||||||
|
handshakeBuf := buf.NewWithSize(totalHandshakeLen)
|
||||||
defer handshakeBuf.Release()
|
defer handshakeBuf.Release()
|
||||||
|
|
||||||
handshakeBuf.Write(clientSalt)
|
handshakeBuf.Write(clientSalt)
|
||||||
@@ -372,14 +355,6 @@ func WriteTCPRequest(w io.Writer, method *CipherMethod, pskList [][]byte, dest n
|
|||||||
handshakeBuf.Write(encryptedEIH[:])
|
handshakeBuf.Write(encryptedEIH[:])
|
||||||
}
|
}
|
||||||
|
|
||||||
payloadLen := len(payload)
|
|
||||||
var paddingLen int
|
|
||||||
if payloadLen < MaxPaddingLength {
|
|
||||||
paddingLen = mrand.IntN(MaxPaddingLength-payloadLen) + 1
|
|
||||||
}
|
|
||||||
addrPortLen := AddrPortLength(dest)
|
|
||||||
varHeaderLen := addrPortLen + 2 + paddingLen + payloadLen
|
|
||||||
|
|
||||||
var fixedHeaderPlaintext [RequestHeaderFixedChunkLength]byte
|
var fixedHeaderPlaintext [RequestHeaderFixedChunkLength]byte
|
||||||
fixedHeaderPlaintext[0] = HeaderTypeClient
|
fixedHeaderPlaintext[0] = HeaderTypeClient
|
||||||
binary.BigEndian.PutUint64(fixedHeaderPlaintext[1:9], uint64(time.Now().Unix()))
|
binary.BigEndian.PutUint64(fixedHeaderPlaintext[1:9], uint64(time.Now().Unix()))
|
||||||
@@ -389,7 +364,7 @@ func WriteTCPRequest(w io.Writer, method *CipherMethod, pskList [][]byte, dest n
|
|||||||
IncreaseNonce(writer.nonce[:])
|
IncreaseNonce(writer.nonce[:])
|
||||||
handshakeBuf.Write(fixedChunk)
|
handshakeBuf.Write(fixedChunk)
|
||||||
|
|
||||||
varHeaderBuf := buf.New()
|
varHeaderBuf := buf.NewWithSize(int32(varHeaderLen))
|
||||||
defer varHeaderBuf.Release()
|
defer varHeaderBuf.Release()
|
||||||
|
|
||||||
if err := WriteAddressPort(varHeaderBuf, dest); err != nil {
|
if err := WriteAddressPort(varHeaderBuf, dest); err != nil {
|
||||||
@@ -421,12 +396,21 @@ func WriteTCPRequest(w io.Writer, method *CipherMethod, pskList [][]byte, dest n
|
|||||||
|
|
||||||
// ReadTCPResponse reads and verifies the server's handshake response and returns a reader for the stream.
|
// ReadTCPResponse reads and verifies the server's handshake response and returns a reader for the stream.
|
||||||
func ReadTCPResponse(r io.Reader, method *CipherMethod, psk []byte, clientSalt []byte) (buf.Reader, error) {
|
func ReadTCPResponse(r io.Reader, method *CipherMethod, psk []byte, clientSalt []byte) (buf.Reader, error) {
|
||||||
var serverSalt [32]byte
|
fixedPlainLen := 1 + 8 + method.KeySaltLength + 2
|
||||||
serverSaltSlice := serverSalt[:method.KeySaltLength]
|
chunkCipherLen := fixedPlainLen + AEADTagSize
|
||||||
if _, err := io.ReadFull(r, serverSaltSlice); err != nil {
|
headerLen := method.KeySaltLength + chunkCipherLen
|
||||||
return nil, err
|
|
||||||
|
// Single read call for Salt + Fixed-length response header chunk per SIP022 §3.1.4
|
||||||
|
var headerBuf [128]byte
|
||||||
|
headerSlice := headerBuf[:headerLen]
|
||||||
|
n, err := r.Read(headerSlice)
|
||||||
|
if err != nil || n < headerLen {
|
||||||
|
return nil, errors.New("failed to read complete server response header")
|
||||||
}
|
}
|
||||||
|
|
||||||
|
serverSaltSlice := headerSlice[:method.KeySaltLength]
|
||||||
|
chunkSlice := headerSlice[method.KeySaltLength:headerLen]
|
||||||
|
|
||||||
sessionKey := DeriveSessionSubKey(psk, serverSaltSlice, method.KeySaltLength)
|
sessionKey := DeriveSessionSubKey(psk, serverSaltSlice, method.KeySaltLength)
|
||||||
aead, err := method.NewAEAD(sessionKey)
|
aead, err := method.NewAEAD(sessionKey)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -435,14 +419,6 @@ func ReadTCPResponse(r io.Reader, method *CipherMethod, psk []byte, clientSalt [
|
|||||||
|
|
||||||
reader := NewStreamReader(r, aead)
|
reader := NewStreamReader(r, aead)
|
||||||
|
|
||||||
fixedPlainLen := 1 + 8 + method.KeySaltLength + 2
|
|
||||||
chunkCipherLen := fixedPlainLen + AEADTagSize
|
|
||||||
var chunkBuf [64]byte
|
|
||||||
chunkSlice := chunkBuf[:chunkCipherLen]
|
|
||||||
if _, err := io.ReadFull(r, chunkSlice); err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
decryptedFixed, err := reader.cipher.Open(chunkSlice[:0], reader.nonce[:], chunkSlice, nil)
|
decryptedFixed, err := reader.cipher.Open(chunkSlice[:0], reader.nonce[:], chunkSlice, nil)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, errors.New("failed to decrypt server response header").Base(err)
|
return nil, errors.New("failed to decrypt server response header").Base(err)
|
||||||
@@ -484,46 +460,190 @@ func ReadTCPResponse(r io.Reader, method *CipherMethod, psk []byte, clientSalt [
|
|||||||
return reader, nil
|
return reader, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// WriteTCPResponse writes the server handshake response and returns a body writer for server stream.
|
// ServerStreamWriter lazily sends the response header along with the first payload chunk per SIP022 §3.1.2 & §3.1.4.
|
||||||
func WriteTCPResponse(w io.Writer, method *CipherMethod, psk []byte, clientSalt []byte, initialPayload []byte) (buf.Writer, error) {
|
type ServerStreamWriter struct {
|
||||||
|
mu sync.Mutex
|
||||||
|
w io.Writer
|
||||||
|
method *CipherMethod
|
||||||
|
psk []byte
|
||||||
|
clientSalt []byte
|
||||||
|
streamWriter *StreamWriter
|
||||||
|
}
|
||||||
|
|
||||||
|
func NewServerStreamWriter(w io.Writer, method *CipherMethod, psk []byte, clientSalt []byte) *ServerStreamWriter {
|
||||||
|
return &ServerStreamWriter{
|
||||||
|
w: w,
|
||||||
|
method: method,
|
||||||
|
psk: psk,
|
||||||
|
clientSalt: clientSalt,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *ServerStreamWriter) sendHeaderWithFirstPayload(payload []byte) (*StreamWriter, error) {
|
||||||
var serverSalt [32]byte
|
var serverSalt [32]byte
|
||||||
serverSaltSlice := serverSalt[:method.KeySaltLength]
|
serverSaltSlice := serverSalt[:s.method.KeySaltLength]
|
||||||
if _, err := io.ReadFull(rand.Reader, serverSaltSlice); err != nil {
|
if _, err := io.ReadFull(rand.Reader, serverSaltSlice); err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
respKey := DeriveSessionSubKey(psk, serverSaltSlice, method.KeySaltLength)
|
respKey := DeriveSessionSubKey(s.psk, serverSaltSlice, s.method.KeySaltLength)
|
||||||
respAead, err := method.NewAEAD(respKey)
|
respAead, err := s.method.NewAEAD(respKey)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
writer := NewStreamWriter(w, respAead)
|
sw := NewStreamWriter(s.w, respAead)
|
||||||
|
|
||||||
respBuf := buf.New()
|
totalHeaderLen := int32(s.method.KeySaltLength + 1 + 8 + s.method.KeySaltLength + 2 + AEADTagSize + len(payload) + AEADTagSize)
|
||||||
defer respBuf.Release()
|
outBuf := buf.NewWithSize(totalHeaderLen)
|
||||||
|
defer outBuf.Release()
|
||||||
|
|
||||||
respBuf.Write(serverSaltSlice)
|
outBuf.Write(serverSaltSlice)
|
||||||
|
|
||||||
var fixedRespPlain [1 + 8 + 32 + 2]byte
|
var fixedRespPlain [1 + 8 + 32 + 2]byte
|
||||||
fixedRespSlice := fixedRespPlain[:1+8+method.KeySaltLength+2]
|
fixedRespSlice := fixedRespPlain[:1+8+s.method.KeySaltLength+2]
|
||||||
fixedRespSlice[0] = HeaderTypeServer
|
fixedRespSlice[0] = HeaderTypeServer
|
||||||
binary.BigEndian.PutUint64(fixedRespSlice[1:9], uint64(time.Now().Unix()))
|
binary.BigEndian.PutUint64(fixedRespSlice[1:9], uint64(time.Now().Unix()))
|
||||||
copy(fixedRespSlice[9:9+method.KeySaltLength], clientSalt)
|
copy(fixedRespSlice[9:9+s.method.KeySaltLength], s.clientSalt)
|
||||||
binary.BigEndian.PutUint16(fixedRespSlice[9+method.KeySaltLength:11+method.KeySaltLength], uint16(len(initialPayload)))
|
binary.BigEndian.PutUint16(fixedRespSlice[9+s.method.KeySaltLength:11+s.method.KeySaltLength], uint16(len(payload)))
|
||||||
|
|
||||||
fixedRespChunk := writer.cipher.Seal(nil, writer.nonce[:], fixedRespSlice, nil)
|
fixedRespChunk := sw.cipher.Seal(nil, sw.nonce[:], fixedRespSlice, nil)
|
||||||
IncreaseNonce(writer.nonce[:])
|
IncreaseNonce(sw.nonce[:])
|
||||||
respBuf.Write(fixedRespChunk)
|
outBuf.Write(fixedRespChunk)
|
||||||
|
|
||||||
if len(initialPayload) > 0 {
|
if len(payload) > 0 {
|
||||||
initialChunk := writer.cipher.Seal(nil, writer.nonce[:], initialPayload, nil)
|
payloadChunk := sw.cipher.Seal(nil, sw.nonce[:], payload, nil)
|
||||||
IncreaseNonce(writer.nonce[:])
|
IncreaseNonce(sw.nonce[:])
|
||||||
respBuf.Write(initialChunk)
|
outBuf.Write(payloadChunk)
|
||||||
}
|
}
|
||||||
|
|
||||||
if _, err := w.Write(respBuf.Bytes()); err != nil {
|
if _, err := s.w.Write(outBuf.Bytes()); err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
return sw, nil
|
||||||
|
}
|
||||||
|
|
||||||
return writer, nil
|
func (s *ServerStreamWriter) WriteMultiBuffer(mb buf.MultiBuffer) error {
|
||||||
|
if mb.IsEmpty() {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
if s.streamWriter == nil {
|
||||||
|
s.mu.Lock()
|
||||||
|
if s.streamWriter == nil {
|
||||||
|
firstBuf := mb[0]
|
||||||
|
firstBytes := firstBuf.Bytes()
|
||||||
|
chunkSize := len(firstBytes)
|
||||||
|
if chunkSize > MaxPacketSize {
|
||||||
|
chunkSize = MaxPacketSize
|
||||||
|
}
|
||||||
|
firstPayload := firstBytes[:chunkSize]
|
||||||
|
sw, err := s.sendHeaderWithFirstPayload(firstPayload)
|
||||||
|
if err != nil {
|
||||||
|
s.mu.Unlock()
|
||||||
|
buf.ReleaseMulti(mb)
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
s.streamWriter = sw
|
||||||
|
|
||||||
|
firstBuf.Advance(int32(chunkSize))
|
||||||
|
if firstBuf.IsEmpty() {
|
||||||
|
firstBuf.Release()
|
||||||
|
mb = mb[1:]
|
||||||
|
}
|
||||||
|
}
|
||||||
|
s.mu.Unlock()
|
||||||
|
if len(mb) == 0 {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return s.streamWriter.WriteMultiBuffer(mb)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *ServerStreamWriter) Write(p []byte) (int, error) {
|
||||||
|
n := len(p)
|
||||||
|
if s.streamWriter == nil {
|
||||||
|
s.mu.Lock()
|
||||||
|
if s.streamWriter == nil {
|
||||||
|
chunkSize := len(p)
|
||||||
|
if chunkSize > MaxPacketSize {
|
||||||
|
chunkSize = MaxPacketSize
|
||||||
|
}
|
||||||
|
firstPayload := p[:chunkSize]
|
||||||
|
sw, err := s.sendHeaderWithFirstPayload(firstPayload)
|
||||||
|
if err != nil {
|
||||||
|
s.mu.Unlock()
|
||||||
|
return 0, err
|
||||||
|
}
|
||||||
|
s.streamWriter = sw
|
||||||
|
p = p[chunkSize:]
|
||||||
|
}
|
||||||
|
s.mu.Unlock()
|
||||||
|
if len(p) == 0 {
|
||||||
|
return n, nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
_, err := s.streamWriter.Write(p)
|
||||||
|
return n, err
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *ServerStreamWriter) Close() error {
|
||||||
|
if s.streamWriter == nil {
|
||||||
|
s.mu.Lock()
|
||||||
|
defer s.mu.Unlock()
|
||||||
|
if s.streamWriter == nil {
|
||||||
|
sw, err := s.sendHeaderWithFirstPayload(nil)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
s.streamWriter = sw
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// InitServerStream decrypts the client request header, verifies the timestamp and replay filter,
|
||||||
|
// and returns a StreamReader for subsequent stream chunks.
|
||||||
|
func InitServerStream(conn net.Conn, method *CipherMethod, psk, saltSlice []byte, salt [32]byte, fixedChunk []byte, saltFilter *antireplay.ReplayFilter[[32]byte]) (*StreamReader, *ClientRequestHeader, error) {
|
||||||
|
sessionKey := DeriveSessionSubKey(psk, saltSlice, method.KeySaltLength)
|
||||||
|
aead, err := method.NewAEAD(sessionKey)
|
||||||
|
if err != nil {
|
||||||
|
return nil, nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
reader := NewStreamReader(conn, aead)
|
||||||
|
|
||||||
|
reqHeader, err := ReadClientRequestHeaderWithFixed(reader, fixedChunk)
|
||||||
|
if err != nil {
|
||||||
|
return nil, nil, err
|
||||||
|
}
|
||||||
|
_ = conn.SetReadDeadline(time.Time{})
|
||||||
|
|
||||||
|
if !saltFilter.Check(salt) {
|
||||||
|
return nil, nil, ErrSaltNotUnique
|
||||||
|
}
|
||||||
|
return reader, reqHeader, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func TransportTCP(ctx context.Context, sessionPolicy policy.Session, reader buf.Reader, writer buf.Writer, link *transport.Link) error {
|
||||||
|
ctx, cancel := context.WithCancel(ctx)
|
||||||
|
timer := signal.CancelAfterInactivity(ctx, cancel, sessionPolicy.Timeouts.ConnectionIdle)
|
||||||
|
ctx = policy.ContextWithBufferPolicy(ctx, sessionPolicy.Buffer)
|
||||||
|
|
||||||
|
requestDone := func() error {
|
||||||
|
defer timer.SetTimeout(sessionPolicy.Timeouts.DownlinkOnly)
|
||||||
|
return buf.Copy(reader, link.Writer, buf.UpdateActivity(timer))
|
||||||
|
}
|
||||||
|
|
||||||
|
responseDone := func() error {
|
||||||
|
defer timer.SetTimeout(sessionPolicy.Timeouts.UplinkOnly)
|
||||||
|
if c, ok := writer.(io.Closer); ok {
|
||||||
|
defer c.Close()
|
||||||
|
}
|
||||||
|
return buf.Copy(link.Reader, writer, buf.UpdateActivity(timer))
|
||||||
|
}
|
||||||
|
|
||||||
|
responseDoneAndCloseWriter := task.OnSuccess(responseDone, task.Close(link.Writer))
|
||||||
|
return task.Run(ctx, requestDone, responseDoneAndCloseWriter)
|
||||||
}
|
}
|
||||||
|
|||||||
+28
-13
@@ -15,27 +15,28 @@ Plainly enabling it in the config probably will result nothing, or lock your rou
|
|||||||
## DETAILS
|
## DETAILS
|
||||||
|
|
||||||
By default, enabling the feature will only bring the tun interface up. \
|
By default, enabling the feature will only bring the tun interface up. \
|
||||||
When configured explicitly, Windows and Linux can apply interface addresses from `gateway`, while macOS uses the first IPv4 prefix from `gateway` to configure the utun point-to-point address. \
|
When configured explicitly, Windows and Linux can apply interface addresses from `gateway`, while macOS and FreeBSD use the first IPv4 prefix from `gateway` for the point-to-point address. \
|
||||||
|
Without `gateway`, the systems differ: Xray assigns no address on Linux, Windows gives the interface link-local addresses itself (an IPv6 one at once, an IPv4 one from `169.254.0.0/16` after a few seconds), and macOS and FreeBSD use `169.254.10.1/30`. \
|
||||||
Windows, Linux and macOS can also apply system routes from `autoSystemRoutingTable`.
|
Windows, Linux and macOS can also apply system routes from `autoSystemRoutingTable`.
|
||||||
macOS does not configure system DNS from the `dns` field, and neither does Linux by default; system DNS remains managed by the OS or distribution-specific network services. \
|
macOS does not configure system DNS from the `dns` field, and neither does Linux by default; system DNS remains managed by the OS or distribution-specific network services. \
|
||||||
For more advanced routing policies or rules, OS level configuration can still manage the named interface (e.g. xray0) when it appears.
|
For more advanced routing policies or rules, OS level configuration can still manage the named interface (e.g. xray0) when it appears.
|
||||||
This keeps complex system level routing and rules in a single place of responsibility - the OS itself. \
|
This keeps complex system level routing and rules in a single place of responsibility - the OS itself. \
|
||||||
Examples of how to achieve this on a simple Linux system (Ubuntu with systemd-networkd) can be found at the end of this README.
|
Examples of how to achieve this on a simple Linux system (Ubuntu with systemd-networkd) can be found at the end of this README.
|
||||||
|
|
||||||
### SYSTEM DNS ON LINUX (`autoSystemDNS`)
|
### SYSTEM DNS ON LINUX (`autoSystemDnsToGateway`)
|
||||||
|
|
||||||
On Linux, setting `autoSystemDNS` to `true` lets the inbound point the system resolver at the tun interface, so name lookups resolve through Xray instead of going out over the physical link. It is off by default, and it is Linux-only.
|
On Linux, setting `autoSystemDnsToGateway` to `true` lets the inbound point the system resolver at the tun interface, so name lookups resolve through Xray instead of going out over the physical link. It is off by default, and it is Linux-only.
|
||||||
|
|
||||||
It uses `resolvectl`, which means it applies only when all of these hold:
|
It uses `resolvectl`, which means it only works when all of these hold. Where Xray can tell that one does not, it does not start:
|
||||||
|
|
||||||
- the system runs systemd and `resolvectl` is on `PATH`
|
- the system runs systemd and `resolvectl` is on `PATH`
|
||||||
- `systemd-resolved` is enabled and actually managing DNS (installed but not running has no effect)
|
- `systemd-resolved` is enabled and actually managing DNS (installed but not running is not enough)
|
||||||
- systemd-resolved is version 240 or newer, where `default-route` exists
|
- systemd-resolved is version 240 or newer, where `default-route` exists
|
||||||
- no `dns` upstream resolves through the system resolver, directly or through its own bootstrap (see below)
|
- no `dns` upstream resolves through the system resolver, directly or through its own bootstrap (see below)
|
||||||
|
|
||||||
The address handed over is the first IPv4 `gateway` incremented by one (e.g. `192.168.100.1/30` -> `192.168.100.2`). It is not taken from `dns`: handing `1.1.1.1` to `resolvectl dns` would make systemd-resolved query that server directly over the physical link, which is the leak this option exists to close.
|
The address handed over is the first IPv4 `gateway`, or without one the first IPv6 `gateway`, incremented by one (e.g. `192.168.100.1/30` -> `192.168.100.2`, `fc00::1/64` -> `fc00::2`). Without any `gateway`, the config is rejected. It is not taken from `dns`: handing `1.1.1.1` to `resolvectl dns` would make systemd-resolved query that server directly over the physical link, which is the leak this option exists to close.
|
||||||
|
|
||||||
Because that address has to actually answer, the takeover is checked before it happens. A query from the interface address to that address is routed through the configured rules, and host-wide DNS is only changed when the result is a DNS-capable outbound. Otherwise the option does nothing and DNS is left to the OS. In practice this means you also need a routing rule sending the interface's port 53 to a `dns` outbound, for example:
|
Because that address has to actually answer, the takeover is checked before it happens. A query from the interface address to that address is routed through the configured rules, and host-wide DNS is only changed when the result is a DNS-capable outbound. Otherwise DNS is left alone and Xray does not start. In practice this means you also need a routing rule sending the interface's port 53 to a `dns` outbound, for example:
|
||||||
|
|
||||||
```json
|
```json
|
||||||
"routing": {
|
"routing": {
|
||||||
@@ -49,19 +50,19 @@ The check is a preflight, not a proof for arbitrary rules. It sends its query fr
|
|||||||
|
|
||||||
It is also a check for the dependencies it knows about, not a proof that no indirect one exists. A hostname-based upstream that bootstraps through system DNS is the case in point: `https+local://dns.google/dns-query` resolves its own hostname with `DialSystem`, so once the takeover is in place that bootstrap goes `resolved -> TUN -> DNS outbound -> bootstrap -> resolved` and the query times out. The preflight does not see it, because the dependency sits in the upstream's bootstrap rather than in the clients it inspects. Upstream resolution, bootstrap included, therefore has to stay independent of the resolver path being redirected; configuring the address instead of the hostname, or resolving the hostname beforehand, avoids it.
|
It is also a check for the dependencies it knows about, not a proof that no indirect one exists. A hostname-based upstream that bootstraps through system DNS is the case in point: `https+local://dns.google/dns-query` resolves its own hostname with `DialSystem`, so once the takeover is in place that bootstrap goes `resolved -> TUN -> DNS outbound -> bootstrap -> resolved` and the query times out. The preflight does not see it, because the dependency sits in the upstream's bootstrap rather than in the clients it inspects. Upstream resolution, bootstrap included, therefore has to stay independent of the resolver path being redirected; configuring the address instead of the hostname, or resolving the hostname beforehand, avoids it.
|
||||||
|
|
||||||
The upstream requirement in the list above matters as much as the routing rule. With no name servers configured, Core resolves through a client that forwards to the system resolver; pointing the system resolver at the TUN would then close a loop through the DNS outbound, `resolved -> TUN -> DNS outbound -> system resolver -> resolved`, and resolution stops. The takeover is refused in that case.
|
The upstream requirement in the list above matters as much as the routing rule. With no name servers configured, Core resolves through a client that forwards to the system resolver; pointing the system resolver at the TUN would then close a loop through the DNS outbound, `resolved -> TUN -> DNS outbound -> system resolver -> resolved`, and resolution stops. The takeover is refused in that case, and Xray does not start.
|
||||||
|
|
||||||
The same applies to a name server pointed at `localhost`, and to a `dns` section that is present but lists no name servers. One such upstream is enough to refuse the takeover even when independent upstreams are configured alongside it: name servers are selected per domain, so a domain-specific rule can still choose the local one, and the loop then affects whichever domains reach it. The check is deliberately broader than the loop it observed, because the alternative would be to drop a name server the user configured.
|
The same applies to a name server pointed at `localhost`, and to a `dns` section that is present but lists no name servers. One such upstream is enough to refuse the takeover even when independent upstreams are configured alongside it: name servers are selected per domain, so a domain-specific rule can still choose the local one, and the loop then affects whichever domains reach it. The check is deliberately broader than the loop it observed, because the alternative would be to drop a name server the user configured.
|
||||||
|
|
||||||
Where it does not apply, DNS is left alone and the leak described in XTLS/Xray-core#6454 remains:
|
Where it cannot apply, Xray does not start, rather than run with the leak described in XTLS/Xray-core#6454, so leave the option off there:
|
||||||
|
|
||||||
| Environment | Behaviour |
|
| Environment | Behaviour |
|
||||||
|---|---|
|
|---|---|
|
||||||
| systemd distribution with systemd-resolved enabled | applies |
|
| systemd distribution with systemd-resolved enabled | applies |
|
||||||
| Alpine, Void, Devuan, OpenRC-based, OpenWrt | no `resolvectl`, skipped |
|
| Alpine, Void, Devuan, OpenRC-based, OpenWrt | no `resolvectl`, does not start |
|
||||||
| DNS managed by dnsmasq / unbound / BIND / static `resolv.conf` | unreachable by `resolvectl`, skipped |
|
| DNS managed by dnsmasq / unbound / BIND / static `resolv.conf` | unreachable by `resolvectl`, does not start |
|
||||||
| Containers without a systemd-resolved daemon | skipped |
|
| Containers without a systemd-resolved daemon | does not start |
|
||||||
| systemd older than 240 | `default-route` unavailable, skipped |
|
| systemd older than 240 | `default-route` unavailable, does not start |
|
||||||
|
|
||||||
On `Close()` the setting is reverted. It is **not** reverted if the process is killed with `SIGKILL`, since a process cannot handle that signal; run `resolvectl revert <iface>` to clean up by hand. An application that brings its own DNS endpoint is unaffected either way — this only covers the system resolver.
|
On `Close()` the setting is reverted. It is **not** reverted if the process is killed with `SIGKILL`, since a process cannot handle that signal; run `resolvectl revert <iface>` to clean up by hand. An application that brings its own DNS endpoint is unaffected either way — this only covers the system resolver.
|
||||||
|
|
||||||
@@ -198,6 +199,20 @@ To make it start, wintun.dll specific for your Windows/arch must be present next
|
|||||||
|
|
||||||
After the start network adapter with the name you chose in the config will be created in the system, and exist while Xray is running.
|
After the start network adapter with the name you chose in the config will be created in the system, and exist while Xray is running.
|
||||||
|
|
||||||
|
When `dns` is set, those servers are applied to the adapter. Windows is kept from registering the TUN's addresses in DNS, and its DNS cache is flushed when the TUN starts and stops.
|
||||||
|
|
||||||
|
With `autoSystemWfpBlockLeak`, which needs `autoSystemRoutingTable` (the config is rejected otherwise), Xray also adds Windows Filtering Platform filters that keep two kinds of traffic of every program but Xray itself from leaving outside the TUN, each chosen by a value in the list, e.g. `"autoSystemWfpBlockLeak": ["dns", "misconfigtun"]`:
|
||||||
|
- `"dns"` (needs `dns`, the config is rejected otherwise): DNS (port 53) only goes through the TUN. Windows keeps sending name queries to the DNS servers of the other interfaces as well, out through those interfaces whatever the routes say, and other programs reach a resolver on the local network (e.g. `192.168.1.1` handed out by DHCP) through its more specific LAN route instead of the TUN. On Windows 11 and Server 2022 and later, where those queries may also go over HTTPS or TLS, Windows' DNS Client service cannot connect outside the TUN at all, except for name resolution on the local network (LLMNR, mDNS). The `dns` servers therefore have to lie within `gateway` or `autoSystemRoutingTable` (a warning is logged otherwise), and DNS servers that should be reached directly belong in Xray's own `dns` settings.
|
||||||
|
- `"misconfigtun"`: an IP version without routes in `autoSystemRoutingTable`, IPv4 or IPv6, is blocked entirely, in both directions, as it would bypass the TUN. Only loopback and what Windows itself needs on the local link (DHCP, and for IPv6 neighbor and multicast listener discovery) remain allowed. An address of that version in `gateway` is not needed: without one, Windows gives the TUN link-local addresses itself, an IPv6 one at once and an IPv4 one from `169.254.0.0/16` after some seconds (until then, IPv4 routed to the TUN is unreachable), and what is routed to the TUN goes through it with those.
|
||||||
|
|
||||||
|
With the filters in place, Xray's own connections out also get past Windows Firewall's block rules (other firewalls may still block them), while connections to Xray's inbounds stay subject to them.
|
||||||
|
|
||||||
|
Names that Xray resolves through the system resolver, such as an outbound's server address given as a domain with the default `AsIs` domain strategy, would be looked up by Windows on Xray's behalf, and those queries would then go into the TUN too. While DNS is restricted this way and `autoOutboundsInterface` is in use (the default with `autoSystemRoutingTable`), Xray therefore resolves them itself, with its own queries to the DNS servers of the other interfaces. That bypasses Windows' DNS cache, and its name resolution on the local network (LLMNR, mDNS): a server address given as a domain is looked up again for every connection, and a DNS server that does not answer delays each lookup. Having Xray's own `dns` resolve it, through the outbound's `sockopt.domainStrategy`, avoids that. The `localhost` DNS server queries the same servers whenever `autoOutboundsInterface` is in use. Both skip the TUN's own DNS servers, unless another interface uses them as well: queried from Xray itself, they would lead back into it, or nowhere.
|
||||||
|
|
||||||
|
If the filters cannot be added, Xray does not start. They are removed when Xray exits. Not covered is name resolution on the local network (LLMNR, mDNS, NetBIOS), except over an IP version that is blocked.
|
||||||
|
|
||||||
|
`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"]`.
|
||||||
|
|
||||||
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. \
|
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.
|
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.
|
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.
|
||||||
|
|||||||
+16
-6
@@ -32,7 +32,8 @@ type Config struct {
|
|||||||
AutoSystemRoutingTable []string `protobuf:"bytes,6,rep,name=auto_system_routing_table,json=autoSystemRoutingTable,proto3" json:"auto_system_routing_table,omitempty"`
|
AutoSystemRoutingTable []string `protobuf:"bytes,6,rep,name=auto_system_routing_table,json=autoSystemRoutingTable,proto3" json:"auto_system_routing_table,omitempty"`
|
||||||
AutoOutboundsInterface string `protobuf:"bytes,7,opt,name=auto_outbounds_interface,json=autoOutboundsInterface,proto3" json:"auto_outbounds_interface,omitempty"`
|
AutoOutboundsInterface string `protobuf:"bytes,7,opt,name=auto_outbounds_interface,json=autoOutboundsInterface,proto3" json:"auto_outbounds_interface,omitempty"`
|
||||||
Desc string `protobuf:"bytes,8,opt,name=desc,proto3" json:"desc,omitempty"`
|
Desc string `protobuf:"bytes,8,opt,name=desc,proto3" json:"desc,omitempty"`
|
||||||
AutoSystemDns bool `protobuf:"varint,9,opt,name=auto_system_dns,json=autoSystemDns,proto3" json:"auto_system_dns,omitempty"`
|
AutoSystemDnsToGateway bool `protobuf:"varint,9,opt,name=auto_system_dns_to_gateway,json=autoSystemDnsToGateway,proto3" json:"auto_system_dns_to_gateway,omitempty"`
|
||||||
|
AutoSystemWfpBlockLeak []string `protobuf:"bytes,10,rep,name=auto_system_wfp_block_leak,json=autoSystemWfpBlockLeak,proto3" json:"auto_system_wfp_block_leak,omitempty"`
|
||||||
unknownFields protoimpl.UnknownFields
|
unknownFields protoimpl.UnknownFields
|
||||||
sizeCache protoimpl.SizeCache
|
sizeCache protoimpl.SizeCache
|
||||||
}
|
}
|
||||||
@@ -123,18 +124,25 @@ func (x *Config) GetDesc() string {
|
|||||||
return ""
|
return ""
|
||||||
}
|
}
|
||||||
|
|
||||||
func (x *Config) GetAutoSystemDns() bool {
|
func (x *Config) GetAutoSystemDnsToGateway() bool {
|
||||||
if x != nil {
|
if x != nil {
|
||||||
return x.AutoSystemDns
|
return x.AutoSystemDnsToGateway
|
||||||
}
|
}
|
||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (x *Config) GetAutoSystemWfpBlockLeak() []string {
|
||||||
|
if x != nil {
|
||||||
|
return x.AutoSystemWfpBlockLeak
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
var File_proxy_tun_config_proto protoreflect.FileDescriptor
|
var File_proxy_tun_config_proto protoreflect.FileDescriptor
|
||||||
|
|
||||||
const file_proxy_tun_config_proto_rawDesc = "" +
|
const file_proxy_tun_config_proto_rawDesc = "" +
|
||||||
"\n" +
|
"\n" +
|
||||||
"\x16proxy/tun/config.proto\x12\x0exray.proxy.tun\"\xaa\x02\n" +
|
"\x16proxy/tun/config.proto\x12\x0exray.proxy.tun\"\xfa\x02\n" +
|
||||||
"\x06Config\x12\x12\n" +
|
"\x06Config\x12\x12\n" +
|
||||||
"\x04name\x18\x01 \x01(\tR\x04name\x12\x10\n" +
|
"\x04name\x18\x01 \x01(\tR\x04name\x12\x10\n" +
|
||||||
"\x03MTU\x18\x02 \x01(\rR\x03MTU\x12\x18\n" +
|
"\x03MTU\x18\x02 \x01(\rR\x03MTU\x12\x18\n" +
|
||||||
@@ -144,8 +152,10 @@ const file_proxy_tun_config_proto_rawDesc = "" +
|
|||||||
"user_level\x18\x05 \x01(\rR\tuserLevel\x129\n" +
|
"user_level\x18\x05 \x01(\rR\tuserLevel\x129\n" +
|
||||||
"\x19auto_system_routing_table\x18\x06 \x03(\tR\x16autoSystemRoutingTable\x128\n" +
|
"\x19auto_system_routing_table\x18\x06 \x03(\tR\x16autoSystemRoutingTable\x128\n" +
|
||||||
"\x18auto_outbounds_interface\x18\a \x01(\tR\x16autoOutboundsInterface\x12\x12\n" +
|
"\x18auto_outbounds_interface\x18\a \x01(\tR\x16autoOutboundsInterface\x12\x12\n" +
|
||||||
"\x04desc\x18\b \x01(\tR\x04desc\x12&\n" +
|
"\x04desc\x18\b \x01(\tR\x04desc\x12:\n" +
|
||||||
"\x0fauto_system_dns\x18\t \x01(\bR\rautoSystemDnsBL\n" +
|
"\x1aauto_system_dns_to_gateway\x18\t \x01(\bR\x16autoSystemDnsToGateway\x12:\n" +
|
||||||
|
"\x1aauto_system_wfp_block_leak\x18\n" +
|
||||||
|
" \x03(\tR\x16autoSystemWfpBlockLeakBL\n" +
|
||||||
"\x12com.xray.proxy.tunP\x01Z#github.com/xtls/xray-core/proxy/tun\xaa\x02\x0eXray.Proxy.Tunb\x06proto3"
|
"\x12com.xray.proxy.tunP\x01Z#github.com/xtls/xray-core/proxy/tun\xaa\x02\x0eXray.Proxy.Tunb\x06proto3"
|
||||||
|
|
||||||
var (
|
var (
|
||||||
|
|||||||
@@ -15,5 +15,6 @@ message Config {
|
|||||||
repeated string auto_system_routing_table = 6;
|
repeated string auto_system_routing_table = 6;
|
||||||
string auto_outbounds_interface = 7;
|
string auto_outbounds_interface = 7;
|
||||||
string desc = 8;
|
string desc = 8;
|
||||||
bool auto_system_dns = 9;
|
bool auto_system_dns_to_gateway = 9;
|
||||||
|
repeated string auto_system_wfp_block_leak = 10;
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -166,12 +166,14 @@ func (t *Handler) Start() error {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// Platform-specific system DNS takeover, where the platform implements it.
|
// Platform-specific system DNS takeover, where the platform implements it.
|
||||||
// Non-fatal: a failure leaves DNS management with the OS.
|
// Rather no TUN than one that the system DNS bypasses.
|
||||||
if c, ok := tunInterface.(interface {
|
if c, ok := tunInterface.(interface {
|
||||||
ConfigureSystemDNS(context.Context, string) error
|
ConfigureSystemDNS(context.Context, string) error
|
||||||
}); ok {
|
}); ok {
|
||||||
if err := c.ConfigureSystemDNS(t.ctx, t.tag); err != nil {
|
if err := c.ConfigureSystemDNS(t.ctx, t.tag); err != nil {
|
||||||
errors.LogInfoInner(t.ctx, err, "[tun] system DNS not configured")
|
_ = tunStack.Close()
|
||||||
|
_ = tunInterface.Close()
|
||||||
|
return errors.New("unable to set the system DNS (remove autoSystemDnsToGateway to run without)").Base(err)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
+21
-15
@@ -53,23 +53,29 @@ var resolvectlRunner = func(name string, args ...string) ([]byte, error) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// systemDNSAddrs derives the addresses used for the system DNS takeover from the
|
// systemDNSAddrs derives the addresses used for the system DNS takeover from the
|
||||||
// first IPv4 gateway: the gateway address itself is what a query from this
|
// first IPv4 gateway, or without one, the first IPv6 gateway: the gateway
|
||||||
// interface appears to come from, and the next address is what the resolver is
|
// address itself is what a query from this interface appears to come from, and
|
||||||
// pointed at. The latter belongs to the TUN and is answered inside Xray;
|
// the next address is what the resolver is pointed at. The latter belongs to
|
||||||
// handing the configured public resolvers to resolvectl instead would leave the
|
// the TUN and is answered inside Xray; handing the configured public resolvers
|
||||||
// system querying them directly over the physical link, defeating the point of
|
// to resolvectl instead would leave the system querying them directly over the
|
||||||
// the TUN.
|
// physical link, defeating the point of the TUN.
|
||||||
func systemDNSAddrs(gateway []string) (source, dns netip.Addr, ok bool) {
|
func systemDNSAddrs(gateway []string) (source, dns netip.Addr, ok bool) {
|
||||||
|
var first6 netip.Addr
|
||||||
for _, address := range gateway {
|
for _, address := range gateway {
|
||||||
prefix, err := netip.ParsePrefix(address)
|
prefix, err := netip.ParsePrefix(address)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
addr := prefix.Addr()
|
addr := prefix.Addr()
|
||||||
if !addr.Is4() {
|
if addr.Is4() {
|
||||||
continue
|
return addr, addr.Next(), true
|
||||||
}
|
}
|
||||||
return addr, addr.Next(), true
|
if !first6.IsValid() {
|
||||||
|
first6 = addr
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if first6.IsValid() {
|
||||||
|
return first6, first6.Next(), true
|
||||||
}
|
}
|
||||||
return netip.Addr{}, netip.Addr{}, false
|
return netip.Addr{}, netip.Addr{}, false
|
||||||
}
|
}
|
||||||
@@ -115,11 +121,11 @@ const probeSourcePort = 49152
|
|||||||
// Overridable for tests.
|
// Overridable for tests.
|
||||||
var verifyDNSRouting = func(ctx context.Context, inboundTag, source, address string) error {
|
var verifyDNSRouting = func(ctx context.Context, inboundTag, source, address string) error {
|
||||||
ip, err := netip.ParseAddr(address)
|
ip, err := netip.ParseAddr(address)
|
||||||
if err != nil || !ip.Is4() {
|
if err != nil {
|
||||||
return errors.New("invalid DNS address ", address).Base(err)
|
return errors.New("invalid DNS address ", address).Base(err)
|
||||||
}
|
}
|
||||||
src, err := netip.ParseAddr(source)
|
src, err := netip.ParseAddr(source)
|
||||||
if err != nil || !src.Is4() {
|
if err != nil || src.Is4() != ip.Is4() {
|
||||||
return errors.New("invalid source address ", source).Base(err)
|
return errors.New("invalid source address ", source).Base(err)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -182,10 +188,10 @@ var verifyDNSRouting = func(ctx context.Context, inboundTag, source, address str
|
|||||||
//
|
//
|
||||||
// It acts only when the config opts in, and it verifies the data path first:
|
// It acts only when the config opts in, and it verifies the data path first:
|
||||||
// unless a query to the advertised address would actually be handled, host-wide
|
// unless a query to the advertised address would actually be handled, host-wide
|
||||||
// resolution is left to the OS, which is the documented default. Errors are
|
// resolution is left to the OS and an error returned. The caller does not start
|
||||||
// returned to the caller, which treats them as non-fatal.
|
// the TUN on an error, as the system DNS would bypass it.
|
||||||
func (t *LinuxTun) ConfigureSystemDNS(ctx context.Context, inboundTag string) error {
|
func (t *LinuxTun) ConfigureSystemDNS(ctx context.Context, inboundTag string) error {
|
||||||
if !t.options.AutoSystemDns {
|
if !t.options.AutoSystemDnsToGateway {
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
if t.systemDNSSet {
|
if t.systemDNSSet {
|
||||||
@@ -202,7 +208,7 @@ func (t *LinuxTun) ConfigureSystemDNS(ctx context.Context, inboundTag string) er
|
|||||||
|
|
||||||
source, address, ok := systemDNSAddrs(t.options.Gateway)
|
source, address, ok := systemDNSAddrs(t.options.Gateway)
|
||||||
if !ok {
|
if !ok {
|
||||||
return errors.New("no IPv4 gateway, cannot derive a system DNS address")
|
return errors.New("no gateway, cannot derive a system DNS address")
|
||||||
}
|
}
|
||||||
|
|
||||||
iface := t.ifaceName()
|
iface := t.ifaceName()
|
||||||
|
|||||||
@@ -191,3 +191,15 @@ func TestVerifyDNSRoutingDecisions(t *testing.T) {
|
|||||||
})
|
})
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Without an IPv4 gateway, the takeover uses the first IPv6 one, and the probe
|
||||||
|
// carries IPv6 addresses.
|
||||||
|
func TestVerifyDNSRoutingIPv6(t *testing.T) {
|
||||||
|
ctx := newRouteTestContext(t, true, udpNameServer([]byte{9, 9, 9, 9}), []*router.RoutingRule{port53Rule()})
|
||||||
|
if err := verifyDNSRouting(ctx, routeTestInboundTag, "fc00::1", "fc00::2"); err != nil {
|
||||||
|
t.Fatalf("expected the takeover to be accepted, got: %v", err)
|
||||||
|
}
|
||||||
|
if err := verifyDNSRouting(ctx, routeTestInboundTag, routeTestSource, "fc00::2"); err == nil {
|
||||||
|
t.Fatal("expected mixed IPv4 and IPv6 addresses to be refused")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
@@ -58,9 +58,9 @@ func recorder(t *testing.T, failOn string) *[][]string {
|
|||||||
func optedInTun() *LinuxTun {
|
func optedInTun() *LinuxTun {
|
||||||
return &LinuxTun{
|
return &LinuxTun{
|
||||||
options: &Config{
|
options: &Config{
|
||||||
Name: "xray_tun",
|
Name: "xray_tun",
|
||||||
Gateway: []string{"192.168.100.1/30"},
|
Gateway: []string{"192.168.100.1/30"},
|
||||||
AutoSystemDns: true,
|
AutoSystemDnsToGateway: true,
|
||||||
},
|
},
|
||||||
tunLink: testLink("xray_tun"),
|
tunLink: testLink("xray_tun"),
|
||||||
}
|
}
|
||||||
@@ -79,7 +79,7 @@ func TestConfigureSystemDNSDisabledByDefault(t *testing.T) {
|
|||||||
calls := recorder(t, "")
|
calls := recorder(t, "")
|
||||||
|
|
||||||
t1 := optedInTun()
|
t1 := optedInTun()
|
||||||
t1.options.AutoSystemDns = false
|
t1.options.AutoSystemDnsToGateway = false
|
||||||
|
|
||||||
if err := t1.ConfigureSystemDNS(context.Background(), "tun"); err != nil {
|
if err := t1.ConfigureSystemDNS(context.Background(), "tun"); err != nil {
|
||||||
t.Fatalf("unexpected error: %v", err)
|
t.Fatalf("unexpected error: %v", err)
|
||||||
@@ -103,7 +103,7 @@ func TestConfigureSystemDNSNoGateway(t *testing.T) {
|
|||||||
t1.options.Gateway = nil
|
t1.options.Gateway = nil
|
||||||
|
|
||||||
if err := t1.ConfigureSystemDNS(context.Background(), "tun"); err == nil {
|
if err := t1.ConfigureSystemDNS(context.Background(), "tun"); err == nil {
|
||||||
t.Fatal("expected an error when no IPv4 gateway is configured")
|
t.Fatal("expected an error when no gateway is configured")
|
||||||
}
|
}
|
||||||
if len(*probes) != 0 {
|
if len(*probes) != 0 {
|
||||||
t.Errorf("routing probe must not run without a gateway, got %d calls", len(*probes))
|
t.Errorf("routing probe must not run without a gateway, got %d calls", len(*probes))
|
||||||
@@ -351,9 +351,18 @@ func TestSystemDNSAddrs(t *testing.T) {
|
|||||||
wantOK: false,
|
wantOK: false,
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
name: "ipv6 only",
|
name: "ipv6 only",
|
||||||
gateway: []string{"fc00::1/64"},
|
gateway: []string{"fc00::1/64"},
|
||||||
wantOK: false,
|
wantSource: "fc00::1",
|
||||||
|
wantDNS: "fc00::2",
|
||||||
|
wantOK: true,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "first ipv6 without ipv4",
|
||||||
|
gateway: []string{"fc00::1/64", "fd00::1/64"},
|
||||||
|
wantSource: "fc00::1",
|
||||||
|
wantDNS: "fc00::2",
|
||||||
|
wantOK: true,
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
+229
-2
@@ -3,17 +3,25 @@
|
|||||||
package tun
|
package tun
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"bytes"
|
||||||
"context"
|
"context"
|
||||||
"crypto/md5"
|
"crypto/md5"
|
||||||
"encoding/binary"
|
"encoding/binary"
|
||||||
go_errors "errors"
|
go_errors "errors"
|
||||||
"net"
|
"net"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
|
"os/exec"
|
||||||
|
"path/filepath"
|
||||||
|
"slices"
|
||||||
|
"strconv"
|
||||||
|
"strings"
|
||||||
"sync"
|
"sync"
|
||||||
|
"syscall"
|
||||||
"time"
|
"time"
|
||||||
"unsafe"
|
"unsafe"
|
||||||
|
|
||||||
"github.com/xtls/xray-core/common/errors"
|
"github.com/xtls/xray-core/common/errors"
|
||||||
|
"github.com/xtls/xray-core/transport/internet"
|
||||||
"golang.org/x/sys/windows"
|
"golang.org/x/sys/windows"
|
||||||
"golang.zx2c4.com/wintun"
|
"golang.zx2c4.com/wintun"
|
||||||
"golang.zx2c4.com/wireguard/windows/tunnel/winipcfg"
|
"golang.zx2c4.com/wireguard/windows/tunnel/winipcfg"
|
||||||
@@ -38,6 +46,10 @@ type WindowsTun struct {
|
|||||||
luid winipcfg.LUID
|
luid winipcfg.LUID
|
||||||
cbr winipcfg.ChangeCallback
|
cbr winipcfg.ChangeCallback
|
||||||
cbi winipcfg.ChangeCallback
|
cbi winipcfg.ChangeCallback
|
||||||
|
wfp windows.Handle
|
||||||
|
resolver *savedResolver
|
||||||
|
skipStop chan struct{}
|
||||||
|
skipDone chan struct{}
|
||||||
closed bool
|
closed bool
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -197,19 +209,105 @@ startOver:
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Windows lists the TUN's DNS servers among the system's ones, which Go's
|
||||||
|
// resolver queries for Xray's own lookups past the TUN, where they lead
|
||||||
|
// nowhere or back into Xray. Not skipped are those another interface uses
|
||||||
|
// as well, as that could leave no server at all. As those can change at
|
||||||
|
// any time, they are looked at again as often as Go rereads its servers.
|
||||||
|
if len(dns) > 0 {
|
||||||
|
skipped, err := tunOnlyDNS(t.luid, dns)
|
||||||
|
if err != nil {
|
||||||
|
skipped = dns
|
||||||
|
}
|
||||||
|
internet.SkipDNSServers(skipped)
|
||||||
|
t.skipStop, t.skipDone = make(chan struct{}), make(chan struct{})
|
||||||
|
go func() {
|
||||||
|
defer close(t.skipDone)
|
||||||
|
ticker := time.NewTicker(5 * time.Second)
|
||||||
|
defer ticker.Stop()
|
||||||
|
for {
|
||||||
|
select {
|
||||||
|
case <-ticker.C:
|
||||||
|
if skipped, err := tunOnlyDNS(t.luid, dns); err == nil {
|
||||||
|
internet.SkipDNSServers(skipped)
|
||||||
|
}
|
||||||
|
case <-t.skipStop:
|
||||||
|
return
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
}
|
||||||
|
|
||||||
|
// Keep Windows from registering the TUN's addresses, and the host name
|
||||||
|
// with them, through dynamic DNS updates. Best effort.
|
||||||
|
if address4 || address6 {
|
||||||
|
if err := disableDNSRegistration(t.luid, dns); err != nil {
|
||||||
|
errors.LogDebugInner(context.Background(), err, "[tun] unable to disable DNS registration")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// With autoSystemWfpBlockLeak, once the system routes lead to the TUN,
|
||||||
|
// keep DNS ("dns", if dns is set), and an IP version no route of which
|
||||||
|
// leads to the TUN ("misconfigtun"), from leaving through the other
|
||||||
|
// interfaces. Addresses do not matter: without one of a version in
|
||||||
|
// gateway, Windows gives the TUN a link-local one.
|
||||||
|
leaks := t.options.AutoSystemWfpBlockLeak
|
||||||
|
blockDNS := slices.Contains(leaks, "dns") && len(dns) > 0
|
||||||
|
blockIPv4 := slices.Contains(leaks, "misconfigtun") && !route4
|
||||||
|
blockIPv6 := slices.Contains(leaks, "misconfigtun") && !route6
|
||||||
|
if (route4 || route6) && (blockDNS || blockIPv4 || blockIPv6) {
|
||||||
|
if t.wfp, err = blockLeaks(t.luid, blockDNS, blockIPv4, blockIPv6); err != nil {
|
||||||
|
var blocked []string
|
||||||
|
for _, b := range []struct {
|
||||||
|
on bool
|
||||||
|
what string
|
||||||
|
}{{blockDNS, "DNS"}, {blockIPv4, "IPv4"}, {blockIPv6, "IPv6"}} {
|
||||||
|
if b.on {
|
||||||
|
blocked = append(blocked, b.what)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
// Rather no TUN than a leaking one.
|
||||||
|
return errors.New("unable to block ", strings.Join(blocked, " and "), " outside the TUN (remove autoSystemWfpBlockLeak to run without)").Base(err)
|
||||||
|
}
|
||||||
|
errors.LogInfo(context.Background(), "[tun] outside the TUN, blocked DNS: ", blockDNS, ", blocked IPv4: ", blockIPv4, ", blocked IPv6: ", blockIPv6)
|
||||||
|
if blockDNS {
|
||||||
|
covered := slices.Clone(addresses)
|
||||||
|
for _, route := range routesData {
|
||||||
|
covered = append(covered, route.Destination)
|
||||||
|
}
|
||||||
|
for _, server := range dnsOutsideTUN(dns, covered) {
|
||||||
|
errors.LogWarning(context.Background(), "[tun] DNS server ", server, " is in neither gateway nor autoSystemRoutingTable, so queries to it cannot go through the TUN and are blocked")
|
||||||
|
}
|
||||||
|
// With updater, the dialer controllers bind Xray's own sockets
|
||||||
|
// to the physical interface.
|
||||||
|
if updater != nil {
|
||||||
|
t.resolver = resolveOnOwn()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if len(dns) > 0 || route4 || route6 {
|
||||||
|
if err := flushDNSCache(); err != nil {
|
||||||
|
errors.LogInfoInner(context.Background(), err, "[tun] unable to flush DNS cache")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
if updater != nil {
|
if updater != nil {
|
||||||
t.cbr, err = winipcfg.RegisterRouteChangeCallback(func(notificationType winipcfg.MibNotificationType, route *winipcfg.MibIPforwardRow2) {
|
// 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()
|
updater.Update()
|
||||||
})
|
})
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
t.cbi, err = winipcfg.RegisterInterfaceChangeCallback(func(notificationType winipcfg.MibNotificationType, iface *winipcfg.MibIPInterfaceRow) {
|
t.cbr = cbr
|
||||||
|
cbi, err := winipcfg.RegisterInterfaceChangeCallback(func(notificationType winipcfg.MibNotificationType, iface *winipcfg.MibIPInterfaceRow) {
|
||||||
updater.Update()
|
updater.Update()
|
||||||
})
|
})
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
t.cbi = cbi
|
||||||
}
|
}
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
@@ -236,6 +334,20 @@ func (t *WindowsTun) Close() error {
|
|||||||
t.luid.FlushIPAddresses(windows.AF_INET6)
|
t.luid.FlushIPAddresses(windows.AF_INET6)
|
||||||
t.luid.FlushDNS(windows.AF_INET6)
|
t.luid.FlushDNS(windows.AF_INET6)
|
||||||
}
|
}
|
||||||
|
if t.wfp != 0 {
|
||||||
|
closeWFPEngine(t.wfp)
|
||||||
|
}
|
||||||
|
if t.resolver != nil {
|
||||||
|
t.resolver.restore()
|
||||||
|
}
|
||||||
|
if t.skipStop != nil {
|
||||||
|
close(t.skipStop)
|
||||||
|
<-t.skipDone
|
||||||
|
}
|
||||||
|
internet.SkipDNSServers(nil)
|
||||||
|
if len(t.options.DNS) > 0 || len(t.options.AutoSystemRoutingTable) > 0 {
|
||||||
|
flushDNSCache()
|
||||||
|
}
|
||||||
if t.session != (wintun.Session{}) {
|
if t.session != (wintun.Session{}) {
|
||||||
t.session.End()
|
t.session.End()
|
||||||
}
|
}
|
||||||
@@ -245,6 +357,121 @@ func (t *WindowsTun) Close() error {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
type savedResolver struct {
|
||||||
|
preferGo bool
|
||||||
|
dial func(ctx context.Context, network, address string) (net.Conn, error)
|
||||||
|
}
|
||||||
|
|
||||||
|
// resolveOnOwn has Go resolve the names Xray would otherwise ask Windows for,
|
||||||
|
// on Xray's own sockets, which the dialer controllers bind to the physical
|
||||||
|
// interface, and skipping the TUN's DNS servers, as localdns does. Windows'
|
||||||
|
// resolver runs in the DNS Client service, whose queries the DNS filter lets
|
||||||
|
// through the TUN only, so Xray's own lookups, like of an outbound's server
|
||||||
|
// domain, would go into Xray again and could end up waiting on themselves.
|
||||||
|
//
|
||||||
|
// It changes net.DefaultResolver for the whole process, which covers every
|
||||||
|
// lookup that would reach Windows' resolver; restore undoes it.
|
||||||
|
func resolveOnOwn() *savedResolver {
|
||||||
|
saved := &savedResolver{net.DefaultResolver.PreferGo, net.DefaultResolver.Dial}
|
||||||
|
dialer := &net.Dialer{Control: func(network, address string, c syscall.RawConn) error {
|
||||||
|
for _, ctl := range internet.Controllers {
|
||||||
|
if err := ctl(network, address, c); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}}
|
||||||
|
// Go's resolver moves on to the next server right away when a dial fails.
|
||||||
|
net.DefaultResolver.Dial = func(ctx context.Context, network, address string) (net.Conn, error) {
|
||||||
|
if internet.IsSkippedDNSServer(address) {
|
||||||
|
return nil, errors.New("skipped DNS server ", address)
|
||||||
|
}
|
||||||
|
return dialer.DialContext(ctx, network, address)
|
||||||
|
}
|
||||||
|
net.DefaultResolver.PreferGo = true
|
||||||
|
return saved
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *savedResolver) restore() {
|
||||||
|
net.DefaultResolver.PreferGo = s.preferGo
|
||||||
|
net.DefaultResolver.Dial = s.dial
|
||||||
|
}
|
||||||
|
|
||||||
|
// tunOnlyDNS returns those of servers, the TUN's DNS servers, that Go's
|
||||||
|
// resolver does not also get from another interface: one that is up and has
|
||||||
|
// a gateway, as it reads them.
|
||||||
|
func tunOnlyDNS(tun winipcfg.LUID, servers []netip.Addr) ([]netip.Addr, error) {
|
||||||
|
adapters, err := winipcfg.GetAdaptersAddresses(windows.AF_UNSPEC, winipcfg.GAAFlagIncludeGateways)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
var others []netip.Addr
|
||||||
|
for _, adapter := range adapters {
|
||||||
|
if adapter.LUID == tun || adapter.OperStatus != winipcfg.IfOperStatusUp || adapter.FirstGatewayAddress == nil {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
for server := adapter.FirstDNSServerAddress; server != nil; server = server.Next {
|
||||||
|
if addr, ok := netip.AddrFromSlice(server.Address.IP()); ok {
|
||||||
|
others = append(others, addr.Unmap())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return slices.DeleteFunc(slices.Clone(servers), func(server netip.Addr) bool {
|
||||||
|
return slices.Contains(others, server.Unmap())
|
||||||
|
}), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// disableDNSRegistration turns off the dynamic DNS registration of the
|
||||||
|
// interface's addresses. dns are its DNS servers.
|
||||||
|
func disableDNSRegistration(luid winipcfg.LUID, dns []netip.Addr) error {
|
||||||
|
guid, err := luid.GUID()
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
err = winipcfg.SetInterfaceDnsSettings(*guid, &winipcfg.DnsInterfaceSettings{
|
||||||
|
Version: winipcfg.DnsInterfaceSettingsVersion1,
|
||||||
|
Flags: winipcfg.DnsInterfaceSettingsFlagRegistrationEnabled,
|
||||||
|
})
|
||||||
|
if err == nil || !go_errors.Is(err, windows.ERROR_PROC_NOT_FOUND) {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
return disableDNSRegistrationByNetsh(luid, dns)
|
||||||
|
}
|
||||||
|
|
||||||
|
// disableDNSRegistrationByNetsh does it for Windows before 10 1809, which
|
||||||
|
// lacks SetInterfaceDnsSettings. The setting is the interface's, not the
|
||||||
|
// address family's, but netsh only applies it along with a DNS server, which
|
||||||
|
// replaces the IPv4 ones, so they are set again afterwards.
|
||||||
|
func disableDNSRegistrationByNetsh(luid winipcfg.LUID, dns []netip.Addr) error {
|
||||||
|
row, err := luid.Interface()
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
server := "127.0.0.1" // any will do when there is no IPv4 one
|
||||||
|
if i := slices.IndexFunc(dns, netip.Addr.Is4); i >= 0 {
|
||||||
|
server = dns[i].String()
|
||||||
|
}
|
||||||
|
err = runNetsh("interface", "ipv4", "set", "dnsservers", "name="+strconv.FormatUint(uint64(row.InterfaceIndex), 10), "source=static", "address="+server, "register=none", "validate=no")
|
||||||
|
return errors.Combine(err, luid.SetDNS(windows.AF_INET, dns, nil))
|
||||||
|
}
|
||||||
|
|
||||||
|
// runNetsh runs netsh.exe from the system directory. netsh reports some
|
||||||
|
// failures, like a syntax error, only in its output, even with exit code 0,
|
||||||
|
// so any output counts as a failure.
|
||||||
|
func runNetsh(args ...string) error {
|
||||||
|
system32, err := windows.GetSystemDirectory()
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
cmd := exec.Command(filepath.Join(system32, "netsh.exe"), args...)
|
||||||
|
cmd.SysProcAttr = &syscall.SysProcAttr{HideWindow: true}
|
||||||
|
output, err := cmd.CombinedOutput()
|
||||||
|
if output = bytes.TrimSpace(output); err != nil || len(output) > 0 {
|
||||||
|
return errors.New("netsh ", strings.Join(args, " "), ": ", string(output)).Base(err)
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
func (t *WindowsTun) Name() (string, error) {
|
func (t *WindowsTun) Name() (string, error) {
|
||||||
row, err := t.luid.Interface()
|
row, err := t.luid.Interface()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
|||||||
@@ -0,0 +1,471 @@
|
|||||||
|
//go:build windows
|
||||||
|
|
||||||
|
package tun
|
||||||
|
|
||||||
|
import (
|
||||||
|
"net/netip"
|
||||||
|
"os"
|
||||||
|
"runtime"
|
||||||
|
"slices"
|
||||||
|
"unsafe"
|
||||||
|
|
||||||
|
"github.com/xtls/xray-core/common/errors"
|
||||||
|
"golang.org/x/sys/windows"
|
||||||
|
"golang.zx2c4.com/wireguard/windows/tunnel/winipcfg"
|
||||||
|
)
|
||||||
|
|
||||||
|
var (
|
||||||
|
modfwpuclnt = windows.NewLazySystemDLL("fwpuclnt.dll")
|
||||||
|
moddnsapi = windows.NewLazySystemDLL("dnsapi.dll")
|
||||||
|
|
||||||
|
procFwpmEngineOpen0 = modfwpuclnt.NewProc("FwpmEngineOpen0")
|
||||||
|
procFwpmEngineClose0 = modfwpuclnt.NewProc("FwpmEngineClose0")
|
||||||
|
procFwpmTransactionBegin0 = modfwpuclnt.NewProc("FwpmTransactionBegin0")
|
||||||
|
procFwpmTransactionCommit0 = modfwpuclnt.NewProc("FwpmTransactionCommit0")
|
||||||
|
procFwpmTransactionAbort0 = modfwpuclnt.NewProc("FwpmTransactionAbort0")
|
||||||
|
procFwpmSubLayerAdd0 = modfwpuclnt.NewProc("FwpmSubLayerAdd0")
|
||||||
|
procFwpmFilterAdd0 = modfwpuclnt.NewProc("FwpmFilterAdd0")
|
||||||
|
procFwpmGetAppIdFromFileName0 = modfwpuclnt.NewProc("FwpmGetAppIdFromFileName0")
|
||||||
|
procFwpmFreeMemory0 = modfwpuclnt.NewProc("FwpmFreeMemory0")
|
||||||
|
procDnsFlushResolverCache = moddnsapi.NewProc("DnsFlushResolverCache")
|
||||||
|
)
|
||||||
|
|
||||||
|
// fwptypes.h and fwpmtypes.h
|
||||||
|
const (
|
||||||
|
rpcCAuthnWinNT = 10 // RPC_C_AUTHN_WINNT
|
||||||
|
fwpmSessionFlagDynamic = 1 // FWPM_SESSION_FLAG_DYNAMIC
|
||||||
|
fwpmFilterFlagClearActionRight = 8 // FWPM_FILTER_FLAG_CLEAR_ACTION_RIGHT
|
||||||
|
|
||||||
|
fwpUint8 = 1 // FWP_UINT8
|
||||||
|
fwpUint16 = 2 // FWP_UINT16
|
||||||
|
fwpUint32 = 3 // FWP_UINT32
|
||||||
|
fwpUint64 = 4 // FWP_UINT64
|
||||||
|
fwpByteArray16Type = 11 // FWP_BYTE_ARRAY16_TYPE
|
||||||
|
fwpByteBlobType = 12 // FWP_BYTE_BLOB_TYPE
|
||||||
|
fwpSecurityDescriptorType = 14 // FWP_SECURITY_DESCRIPTOR_TYPE
|
||||||
|
|
||||||
|
fwpMatchEqual = 0 // FWP_MATCH_EQUAL
|
||||||
|
fwpMatchFlagsAllSet = 6 // FWP_MATCH_FLAGS_ALL_SET
|
||||||
|
|
||||||
|
fwpConditionFlagIsLoopback = 1 // FWP_CONDITION_FLAG_IS_LOOPBACK
|
||||||
|
|
||||||
|
fwpActionBlock = 0x1001 // FWP_ACTION_BLOCK
|
||||||
|
fwpActionPermit = 0x1002 // FWP_ACTION_PERMIT
|
||||||
|
)
|
||||||
|
|
||||||
|
// fwpmu.h
|
||||||
|
var (
|
||||||
|
fwpmLayerALEAuthConnectV4 = windows.GUID{Data1: 0xc38d57d1, Data2: 0x05a7, Data3: 0x4c33, Data4: [8]byte{0x90, 0x4f, 0x7f, 0xbc, 0xee, 0xe6, 0x0e, 0x82}}
|
||||||
|
fwpmLayerALEAuthConnectV6 = windows.GUID{Data1: 0x4a72393b, Data2: 0x319f, Data3: 0x44bc, Data4: [8]byte{0x84, 0xc3, 0xba, 0x54, 0xdc, 0xb3, 0xb6, 0xb4}}
|
||||||
|
fwpmLayerALEAuthRecvAcceptV4 = windows.GUID{Data1: 0xe1cd9fe7, Data2: 0xf4b5, Data3: 0x4273, Data4: [8]byte{0x96, 0xc0, 0x59, 0x2e, 0x48, 0x7b, 0x86, 0x50}}
|
||||||
|
fwpmLayerALEAuthRecvAcceptV6 = windows.GUID{Data1: 0xa3b42c97, Data2: 0x9f04, Data3: 0x4672, Data4: [8]byte{0xb8, 0x7e, 0xce, 0xe9, 0xc4, 0x83, 0x25, 0x7f}}
|
||||||
|
|
||||||
|
fwpmConditionFlags = windows.GUID{Data1: 0x632ce23b, Data2: 0x5167, Data3: 0x435c, Data4: [8]byte{0x86, 0xd7, 0xe9, 0x03, 0x68, 0x4a, 0xa8, 0x0c}}
|
||||||
|
fwpmConditionIPArrivalInterface = windows.GUID{Data1: 0x618a9b6d, Data2: 0x386b, Data3: 0x4136, Data4: [8]byte{0xad, 0x6e, 0xb5, 0x15, 0x87, 0xcf, 0xb1, 0xcd}}
|
||||||
|
fwpmConditionIPLocalInterface = windows.GUID{Data1: 0x4cd62a49, Data2: 0x59c3, Data3: 0x4969, Data4: [8]byte{0xb7, 0xf3, 0xbd, 0xa5, 0xd3, 0x28, 0x90, 0xa4}}
|
||||||
|
fwpmConditionIPLocalPort = windows.GUID{Data1: 0x0c1ba1af, Data2: 0x5765, Data3: 0x453f, Data4: [8]byte{0xaf, 0x22, 0xa8, 0xf7, 0x91, 0xac, 0x77, 0x5b}} // also FWPM_CONDITION_ICMP_TYPE
|
||||||
|
fwpmConditionIPNexthopInterface = windows.GUID{Data1: 0x93ae8f5b, Data2: 0x7f6f, Data3: 0x4719, Data4: [8]byte{0x98, 0xc8, 0x14, 0xe9, 0x74, 0x29, 0xef, 0x04}}
|
||||||
|
fwpmConditionIPProtocol = windows.GUID{Data1: 0x3971ef2b, Data2: 0x623e, Data3: 0x4f9a, Data4: [8]byte{0x8c, 0xb1, 0x6e, 0x79, 0xb8, 0x06, 0xb9, 0xa7}}
|
||||||
|
fwpmConditionIPRemoteAddress = windows.GUID{Data1: 0xb235ae9a, Data2: 0x1d64, Data3: 0x49b8, Data4: [8]byte{0xa4, 0x4c, 0x5f, 0xf3, 0xd9, 0x09, 0x50, 0x45}}
|
||||||
|
fwpmConditionIPRemotePort = windows.GUID{Data1: 0xc35a604d, Data2: 0xd22b, Data3: 0x4e1a, Data4: [8]byte{0x91, 0xb4, 0x68, 0xf6, 0x74, 0xee, 0x67, 0x4b}} // also FWPM_CONDITION_ICMP_CODE
|
||||||
|
fwpmConditionALEAppID = windows.GUID{Data1: 0xd78e1e87, Data2: 0x8644, Data3: 0x4ea5, Data4: [8]byte{0x94, 0x37, 0xd8, 0x09, 0xec, 0xef, 0xc9, 0x71}}
|
||||||
|
fwpmConditionALEUserID = windows.GUID{Data1: 0xaf043a0a, Data2: 0xb34d, Data3: 0x4f86, Data4: [8]byte{0x97, 0x9c, 0xc9, 0x03, 0x71, 0xaf, 0x6e, 0x66}}
|
||||||
|
)
|
||||||
|
|
||||||
|
// dnsClientSID is the SID of Windows' DNS Client service, NT SERVICE\Dnscache.
|
||||||
|
// Service SIDs derive from the service name, so it is the same everywhere (sc
|
||||||
|
// showsid dnscache).
|
||||||
|
const dnsClientSID = "S-1-5-80-859482183-879914841-863379149-1145462774-2388618682"
|
||||||
|
|
||||||
|
// ff02::1:2, where DHCPv6 clients send to. A package-level variable never
|
||||||
|
// moves, so conditions may refer to it through uintptr.
|
||||||
|
var ipv6AllDHCPv6Servers = [16]byte{0xff, 0x02, 13: 0x01, 15: 0x02}
|
||||||
|
|
||||||
|
type fwpByteBlob struct {
|
||||||
|
size uint32
|
||||||
|
data *byte
|
||||||
|
}
|
||||||
|
|
||||||
|
// fwpValue0 is FWP_VALUE0 as well as FWP_CONDITION_VALUE0. Their union holds
|
||||||
|
// a scalar of at most 32 bits, or a pointer for the larger types.
|
||||||
|
type fwpValue0 struct {
|
||||||
|
typ uint32
|
||||||
|
value uintptr
|
||||||
|
}
|
||||||
|
|
||||||
|
type fwpmDisplayData0 struct {
|
||||||
|
name *uint16
|
||||||
|
description *uint16
|
||||||
|
}
|
||||||
|
|
||||||
|
type fwpmSession0 struct {
|
||||||
|
sessionKey windows.GUID
|
||||||
|
displayData fwpmDisplayData0
|
||||||
|
flags uint32
|
||||||
|
txnWaitTimeoutInMSec uint32
|
||||||
|
processID uint32
|
||||||
|
sid *windows.SID
|
||||||
|
username *uint16
|
||||||
|
kernelMode int32
|
||||||
|
}
|
||||||
|
|
||||||
|
type fwpmSublayer0 struct {
|
||||||
|
subLayerKey windows.GUID
|
||||||
|
displayData fwpmDisplayData0
|
||||||
|
flags uint32
|
||||||
|
providerKey *windows.GUID
|
||||||
|
providerData fwpByteBlob
|
||||||
|
weight uint16
|
||||||
|
}
|
||||||
|
|
||||||
|
type fwpmFilterCondition0 struct {
|
||||||
|
fieldKey windows.GUID
|
||||||
|
matchType uint32
|
||||||
|
conditionValue fwpValue0
|
||||||
|
}
|
||||||
|
|
||||||
|
type fwpmAction0 struct {
|
||||||
|
typ uint32
|
||||||
|
filterType windows.GUID
|
||||||
|
}
|
||||||
|
|
||||||
|
type fwpmFilter0 struct {
|
||||||
|
filterKey windows.GUID
|
||||||
|
displayData fwpmDisplayData0
|
||||||
|
flags uint32
|
||||||
|
providerKey *windows.GUID
|
||||||
|
providerData fwpByteBlob
|
||||||
|
layerKey windows.GUID
|
||||||
|
subLayerKey windows.GUID
|
||||||
|
weight fwpValue0
|
||||||
|
numFilterConditions uint32
|
||||||
|
filterCondition *fwpmFilterCondition0
|
||||||
|
action fwpmAction0
|
||||||
|
_ uint32 // C aligns the following union to 8 bytes, as it holds a UINT64
|
||||||
|
providerContextKey windows.GUID
|
||||||
|
reserved *windows.GUID
|
||||||
|
_ [8 - unsafe.Sizeof(uintptr(0))]byte // and filterId as well, also on 32-bit
|
||||||
|
filterID uint64
|
||||||
|
effectiveWeight fwpValue0
|
||||||
|
}
|
||||||
|
|
||||||
|
// fwpmResult converts the DWORD status the Fwpm functions return.
|
||||||
|
func fwpmResult(r1, _ uintptr, _ error) error {
|
||||||
|
if r1 != 0 {
|
||||||
|
return windows.Errno(r1)
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func utf16Ptr(s string) *uint16 {
|
||||||
|
p, _ := windows.UTF16PtrFromString(s)
|
||||||
|
return p
|
||||||
|
}
|
||||||
|
|
||||||
|
func condition(field *windows.GUID, typ uint32, value uintptr) fwpmFilterCondition0 {
|
||||||
|
return fwpmFilterCondition0{
|
||||||
|
fieldKey: *field,
|
||||||
|
matchType: fwpMatchEqual,
|
||||||
|
conditionValue: fwpValue0{typ: typ, value: value},
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// blockLeaks keeps traffic from leaving through interfaces other than tun,
|
||||||
|
// for every program but Xray itself, whose outbounds (DNS included) use the
|
||||||
|
// other interfaces on purpose:
|
||||||
|
//
|
||||||
|
// - dns: DNS (port 53) may only go through the TUN. Windows sends a name
|
||||||
|
// query to the DNS servers of all interfaces, not only to those of the TUN:
|
||||||
|
// to the first server of each interface, then to all of them when no answer
|
||||||
|
// arrives within a second or two. It sends the queries for the servers of
|
||||||
|
// an interface out through that interface, whatever the routes say, and
|
||||||
|
// other programs reach an on-link resolver, like 192.168.1.1 from DHCP,
|
||||||
|
// through its LAN route, which is more specific than the TUN's default
|
||||||
|
// route. Since Windows 11 and Server 2022, Windows may also send its
|
||||||
|
// queries over HTTPS or TLS, so there its DNS Client service may not
|
||||||
|
// connect outside the TUN at all, except for name resolution on the local
|
||||||
|
// link (mDNS, LLMNR).
|
||||||
|
// - ipv4, ipv6: no IPv4, or no IPv6, at all, in either direction, for a TUN
|
||||||
|
// that no route of it leads to, except loopback and what Windows itself
|
||||||
|
// needs on the local link (DHCP, and for IPv6 neighbor and multicast
|
||||||
|
// listener discovery), none of which can leave it. The TUN carries what
|
||||||
|
// is routed to it even without an address of that IP version in gateway:
|
||||||
|
// Windows gives it link-local ones itself, an IPv6 one at once, an IPv4
|
||||||
|
// one from 169.254.0.0/16 after some seconds (until then, IPv4 routed to
|
||||||
|
// the TUN is unreachable).
|
||||||
|
//
|
||||||
|
// The filters live in a dynamic WFP session: closing the returned engine handle
|
||||||
|
// with closeWFPEngine deletes them, and so does Windows when the process dies.
|
||||||
|
func blockLeaks(tun winipcfg.LUID, dns, ipv4, ipv6 bool) (windows.Handle, error) {
|
||||||
|
engine, err := openWFPEngine()
|
||||||
|
if err != nil {
|
||||||
|
return 0, err
|
||||||
|
}
|
||||||
|
if err := fwpmResult(procFwpmTransactionBegin0.Call(uintptr(engine), 0)); err != nil {
|
||||||
|
closeWFPEngine(engine)
|
||||||
|
return 0, errors.New("FwpmTransactionBegin0 failed").Base(err)
|
||||||
|
}
|
||||||
|
err = addLeakFilters(engine, tun, dns, ipv4, ipv6)
|
||||||
|
if err == nil {
|
||||||
|
if err = fwpmResult(procFwpmTransactionCommit0.Call(uintptr(engine))); err != nil {
|
||||||
|
err = errors.New("FwpmTransactionCommit0 failed").Base(err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if err != nil {
|
||||||
|
procFwpmTransactionAbort0.Call(uintptr(engine))
|
||||||
|
closeWFPEngine(engine)
|
||||||
|
return 0, err
|
||||||
|
}
|
||||||
|
return engine, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func openWFPEngine() (windows.Handle, error) {
|
||||||
|
if err := modfwpuclnt.Load(); err != nil {
|
||||||
|
return 0, err
|
||||||
|
}
|
||||||
|
// txnWaitTimeoutInMSec stays 0 for BFE's default, so that a transaction
|
||||||
|
// held by another program cannot hang the start forever.
|
||||||
|
session := fwpmSession0{
|
||||||
|
displayData: fwpmDisplayData0{name: utf16Ptr("Xray TUN")},
|
||||||
|
flags: fwpmSessionFlagDynamic,
|
||||||
|
}
|
||||||
|
var engine windows.Handle
|
||||||
|
if err := fwpmResult(procFwpmEngineOpen0.Call(0, rpcCAuthnWinNT, 0, uintptr(unsafe.Pointer(&session)), uintptr(unsafe.Pointer(&engine)))); err != nil {
|
||||||
|
return 0, errors.New("FwpmEngineOpen0 failed").Base(err)
|
||||||
|
}
|
||||||
|
return engine, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func closeWFPEngine(engine windows.Handle) {
|
||||||
|
procFwpmEngineClose0.Call(uintptr(engine))
|
||||||
|
}
|
||||||
|
|
||||||
|
// addLeakFilters adds the filters of blockLeaks in a sublayer of their own.
|
||||||
|
// blockLeaks runs it in a transaction, so that they take effect all at once.
|
||||||
|
func addLeakFilters(engine windows.Handle, tun winipcfg.LUID, dns, ipv4, ipv6 bool) error {
|
||||||
|
exe, err := os.Executable()
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
exePath, err := windows.UTF16PtrFromString(exe)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
var appID *fwpByteBlob
|
||||||
|
if err := fwpmResult(procFwpmGetAppIdFromFileName0.Call(uintptr(unsafe.Pointer(exePath)), uintptr(unsafe.Pointer(&appID)))); err != nil {
|
||||||
|
return errors.New("FwpmGetAppIdFromFileName0 failed for ", exe).Base(err)
|
||||||
|
}
|
||||||
|
defer func() { procFwpmFreeMemory0.Call(uintptr(unsafe.Pointer(&appID))) }()
|
||||||
|
|
||||||
|
sublayer := fwpmSublayer0{
|
||||||
|
displayData: fwpmDisplayData0{name: utf16Ptr("Xray TUN")},
|
||||||
|
weight: 0xffff,
|
||||||
|
}
|
||||||
|
if sublayer.subLayerKey, err = windows.GenerateGUID(); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if err := fwpmResult(procFwpmSubLayerAdd0.Call(uintptr(engine), uintptr(unsafe.Pointer(&sublayer)), 0)); err != nil {
|
||||||
|
return errors.New("FwpmSubLayerAdd0 failed").Base(err)
|
||||||
|
}
|
||||||
|
add := func(layer *windows.GUID, name string, flags, action uint32, weight uint8, conditions ...fwpmFilterCondition0) error {
|
||||||
|
return addFilter(engine, &sublayer.subLayerKey, layer, "Xray TUN: "+name, flags, action, weight, conditions...)
|
||||||
|
}
|
||||||
|
|
||||||
|
var pinner runtime.Pinner
|
||||||
|
defer pinner.Unpin()
|
||||||
|
tunLUID := new(uint64)
|
||||||
|
*tunLUID = uint64(tun)
|
||||||
|
pinner.Pin(tunLUID) // the condition only holds it as uintptr
|
||||||
|
|
||||||
|
// The heaviest matching filter of a sublayer decides. All sublayers have
|
||||||
|
// their say, though, and a block in any of them beats a permit, unless
|
||||||
|
// the permit is hard: it clears the action right, and then the blocks of
|
||||||
|
// lower sublayers, Windows Firewall rules among them, no longer override
|
||||||
|
// it, only a callout's veto does. Xray's own connections out get such a
|
||||||
|
// hard permit. Connections from outside to Xray get an ordinary one, so
|
||||||
|
// that firewalls keep guarding its inbounds.
|
||||||
|
self := condition(&fwpmConditionALEAppID, fwpByteBlobType, uintptr(unsafe.Pointer(appID)))
|
||||||
|
dns53 := condition(&fwpmConditionIPRemotePort, fwpUint16, 53)
|
||||||
|
// DNS goes through the TUN when its local address is the TUN's, and it
|
||||||
|
// also leaves, or arrives, through the TUN. The local address alone
|
||||||
|
// decides by default, but with weak host sending or receiving enabled,
|
||||||
|
// packets of the TUN's address can use other interfaces. (The next hop,
|
||||||
|
// the interface replies would leave by, is not known for arriving ones.)
|
||||||
|
onTUN := func(field *windows.GUID) fwpmFilterCondition0 {
|
||||||
|
return condition(field, fwpUint64, uintptr(unsafe.Pointer(tunLUID)))
|
||||||
|
}
|
||||||
|
out := []fwpmFilterCondition0{dns53, onTUN(&fwpmConditionIPLocalInterface), onTUN(&fwpmConditionIPNexthopInterface)}
|
||||||
|
in := []fwpmFilterCondition0{dns53, onTUN(&fwpmConditionIPLocalInterface), onTUN(&fwpmConditionIPArrivalInterface)}
|
||||||
|
for _, layer := range []struct {
|
||||||
|
key *windows.GUID
|
||||||
|
selfFlags uint32
|
||||||
|
throughTUN []fwpmFilterCondition0
|
||||||
|
}{
|
||||||
|
{&fwpmLayerALEAuthConnectV4, fwpmFilterFlagClearActionRight, out},
|
||||||
|
{&fwpmLayerALEAuthRecvAcceptV4, 0, in},
|
||||||
|
{&fwpmLayerALEAuthConnectV6, fwpmFilterFlagClearActionRight, out},
|
||||||
|
{&fwpmLayerALEAuthRecvAcceptV6, 0, in},
|
||||||
|
} {
|
||||||
|
if err := add(layer.key, "permit Xray", layer.selfFlags, fwpActionPermit, 4, self); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if dns {
|
||||||
|
if err := add(layer.key, "permit DNS through the TUN", 0, fwpActionPermit, 3, layer.throughTUN...); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if err := add(layer.key, "block DNS", 0, fwpActionBlock, 2, dns53); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Since Windows 11 and Server 2022 (build 20348), the DNS Client service
|
||||||
|
// may also send the queries for an interface's servers over HTTPS or TLS,
|
||||||
|
// out through that interface and to any port. So there it may only
|
||||||
|
// connect through the TUN, except for mDNS and LLMNR, which stay on the
|
||||||
|
// local link (over an IP version only while it is not blocked altogether).
|
||||||
|
// Earlier versions only query port 53, and may run the service in one
|
||||||
|
// process with others, which the filters would catch as well. Like
|
||||||
|
// Windows Firewall's rules for it, they recognize the service by its SID,
|
||||||
|
// which Windows puts in the token of its process: the security descriptor
|
||||||
|
// grants that SID the right to match (FWP_ACTRL_MATCH_FILTER, CC in SDDL).
|
||||||
|
if _, _, build := windows.RtlGetNtVersionNumbers(); dns && build >= 20348 {
|
||||||
|
sd, err := windows.SecurityDescriptorFromString("O:SYG:SYD:(A;;CCRC;;;" + dnsClientSID + ")")
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
sdBlob := &fwpByteBlob{size: sd.Length(), data: (*byte)(unsafe.Pointer(sd))}
|
||||||
|
pinner.Pin(sdBlob) // the condition only holds it as uintptr
|
||||||
|
dnsClient := condition(&fwpmConditionALEUserID, fwpSecurityDescriptorType, uintptr(unsafe.Pointer(sdBlob)))
|
||||||
|
// Conditions on the same field match when any of them does.
|
||||||
|
mdnsLLMNR := []fwpmFilterCondition0{dnsClient, condition(&fwpmConditionIPRemotePort, fwpUint16, 5353), condition(&fwpmConditionIPRemotePort, fwpUint16, 5355)}
|
||||||
|
for _, layer := range []struct {
|
||||||
|
key *windows.GUID
|
||||||
|
localLink bool
|
||||||
|
}{
|
||||||
|
{&fwpmLayerALEAuthConnectV4, !ipv4},
|
||||||
|
{&fwpmLayerALEAuthConnectV6, !ipv6},
|
||||||
|
} {
|
||||||
|
if err := add(layer.key, "permit the DNS Client service through the TUN", 0, fwpActionPermit, 3, dnsClient, onTUN(&fwpmConditionIPLocalInterface), onTUN(&fwpmConditionIPNexthopInterface)); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if layer.localLink {
|
||||||
|
if err := add(layer.key, "permit the DNS Client service's mDNS and LLMNR", 0, fwpActionPermit, 3, mdnsLLMNR...); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if err := add(layer.key, "block the DNS Client service", 0, fwpActionBlock, 2, dnsClient); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Both directions: replies to a connection accepted from outside would
|
||||||
|
// leave through the physical link as well.
|
||||||
|
loopback := fwpmFilterCondition0{
|
||||||
|
fieldKey: fwpmConditionFlags,
|
||||||
|
matchType: fwpMatchFlagsAllSet,
|
||||||
|
conditionValue: fwpValue0{typ: fwpUint32, value: fwpConditionFlagIsLoopback},
|
||||||
|
}
|
||||||
|
if ipv4 {
|
||||||
|
// DHCP keeps the addresses of the other interfaces, which Xray's own
|
||||||
|
// connections use.
|
||||||
|
dhcp := []fwpmFilterCondition0{
|
||||||
|
condition(&fwpmConditionIPProtocol, fwpUint8, windows.IPPROTO_UDP),
|
||||||
|
condition(&fwpmConditionIPLocalPort, fwpUint16, 68),
|
||||||
|
condition(&fwpmConditionIPRemotePort, fwpUint16, 67),
|
||||||
|
}
|
||||||
|
for _, layer := range []*windows.GUID{&fwpmLayerALEAuthConnectV4, &fwpmLayerALEAuthRecvAcceptV4} {
|
||||||
|
if err := add(layer, "permit IPv4 loopback", 0, fwpActionPermit, 1, loopback); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if err := add(layer, "permit DHCP", 0, fwpActionPermit, 1, dhcp...); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if err := add(layer, "block IPv4", 0, fwpActionBlock, 0); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if ipv6 {
|
||||||
|
// Neighbor and multicast listener discovery, ICMPv6 130-137 and 143,
|
||||||
|
// whose type and code sit where the local and remote port are.
|
||||||
|
discovery := []fwpmFilterCondition0{condition(&fwpmConditionIPProtocol, fwpUint8, windows.IPPROTO_ICMPV6)}
|
||||||
|
for _, typ := range []uintptr{130, 131, 132, 133, 134, 135, 136, 137, 143} {
|
||||||
|
discovery = append(discovery, condition(&fwpmConditionIPLocalPort, fwpUint16, typ))
|
||||||
|
}
|
||||||
|
discovery = append(discovery, condition(&fwpmConditionIPRemotePort, fwpUint16, 0))
|
||||||
|
dhcpv6 := []fwpmFilterCondition0{
|
||||||
|
condition(&fwpmConditionIPProtocol, fwpUint8, windows.IPPROTO_UDP),
|
||||||
|
condition(&fwpmConditionIPLocalPort, fwpUint16, 546),
|
||||||
|
condition(&fwpmConditionIPRemotePort, fwpUint16, 547),
|
||||||
|
}
|
||||||
|
for _, direction := range []struct {
|
||||||
|
layer *windows.GUID
|
||||||
|
dhcpv6 []fwpmFilterCondition0
|
||||||
|
}{
|
||||||
|
// The client sends to the servers' multicast address, and they
|
||||||
|
// answer from their own.
|
||||||
|
{&fwpmLayerALEAuthConnectV6, slices.Concat(dhcpv6, []fwpmFilterCondition0{condition(&fwpmConditionIPRemoteAddress, fwpByteArray16Type, uintptr(unsafe.Pointer(&ipv6AllDHCPv6Servers)))})},
|
||||||
|
{&fwpmLayerALEAuthRecvAcceptV6, dhcpv6},
|
||||||
|
} {
|
||||||
|
if err := add(direction.layer, "permit IPv6 loopback", 0, fwpActionPermit, 1, loopback); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if err := add(direction.layer, "permit IPv6 neighbor and multicast listener discovery", 0, fwpActionPermit, 1, discovery...); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if err := add(direction.layer, "permit DHCPv6", 0, fwpActionPermit, 1, direction.dhcpv6...); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if err := add(direction.layer, "block IPv6", 0, fwpActionBlock, 0); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func addFilter(engine windows.Handle, sublayer, layer *windows.GUID, name string, flags, action uint32, weight uint8, conditions ...fwpmFilterCondition0) error {
|
||||||
|
filter := fwpmFilter0{
|
||||||
|
displayData: fwpmDisplayData0{name: utf16Ptr(name)},
|
||||||
|
flags: flags,
|
||||||
|
layerKey: *layer,
|
||||||
|
subLayerKey: *sublayer,
|
||||||
|
weight: fwpValue0{typ: fwpUint8, value: uintptr(weight)},
|
||||||
|
numFilterConditions: uint32(len(conditions)),
|
||||||
|
action: fwpmAction0{typ: action},
|
||||||
|
}
|
||||||
|
if len(conditions) > 0 {
|
||||||
|
filter.filterCondition = &conditions[0]
|
||||||
|
}
|
||||||
|
if err := fwpmResult(procFwpmFilterAdd0.Call(uintptr(engine), uintptr(unsafe.Pointer(&filter)), 0, 0)); err != nil {
|
||||||
|
return errors.New("FwpmFilterAdd0 failed for ", name).Base(err)
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// dnsOutsideTUN returns the servers outside all of prefixes, the TUN's own
|
||||||
|
// subnets and routes: queries to them cannot go through the TUN.
|
||||||
|
func dnsOutsideTUN(servers []netip.Addr, prefixes []netip.Prefix) []netip.Addr {
|
||||||
|
var outside []netip.Addr
|
||||||
|
for _, server := range servers {
|
||||||
|
server = server.Unmap()
|
||||||
|
if !slices.ContainsFunc(prefixes, func(p netip.Prefix) bool { return p.Contains(server) }) {
|
||||||
|
outside = append(outside, server)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return outside
|
||||||
|
}
|
||||||
|
|
||||||
|
// flushDNSCache drops the answers Windows cached so far, like ipconfig
|
||||||
|
// /flushdns, so that names get resolved again with the current DNS setup.
|
||||||
|
func flushDNSCache() error {
|
||||||
|
if err := procDnsFlushResolverCache.Find(); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if r, _, err := procDnsFlushResolverCache.Call(); r == 0 {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
@@ -0,0 +1,206 @@
|
|||||||
|
//go:build windows
|
||||||
|
|
||||||
|
package tun
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
go_errors "errors"
|
||||||
|
"net"
|
||||||
|
"net/netip"
|
||||||
|
"slices"
|
||||||
|
"testing"
|
||||||
|
"unsafe"
|
||||||
|
|
||||||
|
"github.com/xtls/xray-core/transport/internet"
|
||||||
|
"golang.org/x/sys/windows"
|
||||||
|
"golang.zx2c4.com/wireguard/windows/tunnel/winipcfg"
|
||||||
|
)
|
||||||
|
|
||||||
|
// The WFP structures are handed to fwpuclnt.dll as they are, so their layout
|
||||||
|
// has to match what MSVC produces for 64-bit and for 32-bit Windows.
|
||||||
|
func TestWFPStructLayout(t *testing.T) {
|
||||||
|
check := func(name string, got, want64, want32 []uintptr) {
|
||||||
|
t.Helper()
|
||||||
|
want := want32
|
||||||
|
if unsafe.Sizeof(uintptr(0)) == 8 {
|
||||||
|
want = want64
|
||||||
|
}
|
||||||
|
if !slices.Equal(got, want) {
|
||||||
|
t.Errorf("%s: size and offsets are %v, want %v", name, got, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
var blob fwpByteBlob
|
||||||
|
check("FWP_BYTE_BLOB",
|
||||||
|
[]uintptr{unsafe.Sizeof(blob), unsafe.Offsetof(blob.data)},
|
||||||
|
[]uintptr{16, 8}, []uintptr{8, 4})
|
||||||
|
|
||||||
|
var value fwpValue0
|
||||||
|
check("FWP_VALUE0",
|
||||||
|
[]uintptr{unsafe.Sizeof(value), unsafe.Offsetof(value.value)},
|
||||||
|
[]uintptr{16, 8}, []uintptr{8, 4})
|
||||||
|
|
||||||
|
var display fwpmDisplayData0
|
||||||
|
check("FWPM_DISPLAY_DATA0",
|
||||||
|
[]uintptr{unsafe.Sizeof(display), unsafe.Offsetof(display.description)},
|
||||||
|
[]uintptr{16, 8}, []uintptr{8, 4})
|
||||||
|
|
||||||
|
var action fwpmAction0
|
||||||
|
check("FWPM_ACTION0",
|
||||||
|
[]uintptr{unsafe.Sizeof(action), unsafe.Offsetof(action.filterType)},
|
||||||
|
[]uintptr{20, 4}, []uintptr{20, 4})
|
||||||
|
|
||||||
|
var cond fwpmFilterCondition0
|
||||||
|
check("FWPM_FILTER_CONDITION0",
|
||||||
|
[]uintptr{unsafe.Sizeof(cond), unsafe.Offsetof(cond.matchType), unsafe.Offsetof(cond.conditionValue)},
|
||||||
|
[]uintptr{40, 16, 24}, []uintptr{28, 16, 20})
|
||||||
|
|
||||||
|
var session fwpmSession0
|
||||||
|
check("FWPM_SESSION0",
|
||||||
|
[]uintptr{
|
||||||
|
unsafe.Sizeof(session), unsafe.Offsetof(session.displayData), unsafe.Offsetof(session.flags),
|
||||||
|
unsafe.Offsetof(session.txnWaitTimeoutInMSec), unsafe.Offsetof(session.processID), unsafe.Offsetof(session.sid),
|
||||||
|
unsafe.Offsetof(session.username), unsafe.Offsetof(session.kernelMode),
|
||||||
|
},
|
||||||
|
[]uintptr{72, 16, 32, 36, 40, 48, 56, 64},
|
||||||
|
[]uintptr{48, 16, 24, 28, 32, 36, 40, 44})
|
||||||
|
|
||||||
|
var sublayer fwpmSublayer0
|
||||||
|
check("FWPM_SUBLAYER0",
|
||||||
|
[]uintptr{
|
||||||
|
unsafe.Sizeof(sublayer), unsafe.Offsetof(sublayer.displayData), unsafe.Offsetof(sublayer.flags),
|
||||||
|
unsafe.Offsetof(sublayer.providerKey), unsafe.Offsetof(sublayer.providerData), unsafe.Offsetof(sublayer.weight),
|
||||||
|
},
|
||||||
|
[]uintptr{72, 16, 32, 40, 48, 64},
|
||||||
|
[]uintptr{44, 16, 24, 28, 32, 40})
|
||||||
|
|
||||||
|
var filter fwpmFilter0
|
||||||
|
check("FWPM_FILTER0",
|
||||||
|
[]uintptr{
|
||||||
|
unsafe.Sizeof(filter), unsafe.Offsetof(filter.displayData), unsafe.Offsetof(filter.flags),
|
||||||
|
unsafe.Offsetof(filter.providerKey), unsafe.Offsetof(filter.providerData), unsafe.Offsetof(filter.layerKey),
|
||||||
|
unsafe.Offsetof(filter.subLayerKey), unsafe.Offsetof(filter.weight), unsafe.Offsetof(filter.numFilterConditions),
|
||||||
|
unsafe.Offsetof(filter.filterCondition), unsafe.Offsetof(filter.action), unsafe.Offsetof(filter.providerContextKey),
|
||||||
|
unsafe.Offsetof(filter.reserved), unsafe.Offsetof(filter.filterID), unsafe.Offsetof(filter.effectiveWeight),
|
||||||
|
},
|
||||||
|
[]uintptr{200, 16, 32, 40, 48, 64, 80, 96, 112, 120, 128, 152, 168, 176, 184},
|
||||||
|
[]uintptr{152, 16, 24, 28, 32, 40, 56, 72, 80, 84, 88, 112, 128, 136, 144})
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestLeakFiltersAccepted has WFP validate the filters by adding them inside a
|
||||||
|
// transaction that is then aborted, which leaves the system untouched. Adding
|
||||||
|
// filters requires an elevated process.
|
||||||
|
func TestLeakFiltersAccepted(t *testing.T) {
|
||||||
|
skipUnlessElevated := func(err error) {
|
||||||
|
t.Helper()
|
||||||
|
if go_errors.Is(err, windows.ERROR_ACCESS_DENIED) {
|
||||||
|
t.Skipf("WFP filters can only be added by an elevated process: %v", err)
|
||||||
|
}
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
|
||||||
|
engine, err := openWFPEngine()
|
||||||
|
if err != nil {
|
||||||
|
skipUnlessElevated(err)
|
||||||
|
}
|
||||||
|
defer closeWFPEngine(engine)
|
||||||
|
if err := fwpmResult(procFwpmTransactionBegin0.Call(uintptr(engine), 0)); err != nil {
|
||||||
|
skipUnlessElevated(err)
|
||||||
|
}
|
||||||
|
defer procFwpmTransactionAbort0.Call(uintptr(engine))
|
||||||
|
|
||||||
|
// Any interface stands in for the TUN; the loopback one always exists.
|
||||||
|
loopback, err := winipcfg.LUIDFromIndex(1)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if err := addLeakFilters(engine, loopback, true, true, true); err != nil {
|
||||||
|
skipUnlessElevated(err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestDNSClientSID(t *testing.T) {
|
||||||
|
sid, _, _, err := windows.LookupSID("", `NT SERVICE\Dnscache`)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if sid.String() != dnsClientSID {
|
||||||
|
t.Errorf(`NT SERVICE\Dnscache is %v, not %v`, sid, dnsClientSID)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestDNSOutsideTUN(t *testing.T) {
|
||||||
|
prefixes := []netip.Prefix{
|
||||||
|
netip.MustParsePrefix("198.51.100.1/30"), // gateway, not masked
|
||||||
|
netip.MustParsePrefix("203.0.113.0/24"), // route
|
||||||
|
}
|
||||||
|
servers := []netip.Addr{
|
||||||
|
netip.MustParseAddr("198.51.100.2"),
|
||||||
|
netip.MustParseAddr("203.0.113.53"),
|
||||||
|
netip.MustParseAddr("::ffff:203.0.113.54"),
|
||||||
|
netip.MustParseAddr("8.8.8.8"),
|
||||||
|
netip.MustParseAddr("2001:db8::53"),
|
||||||
|
}
|
||||||
|
want := []netip.Addr{netip.MustParseAddr("8.8.8.8"), netip.MustParseAddr("2001:db8::53")}
|
||||||
|
if got := dnsOutsideTUN(servers, prefixes); !slices.Equal(got, want) {
|
||||||
|
t.Errorf("got %v, want %v", got, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestResolveOnOwn(t *testing.T) {
|
||||||
|
internet.SkipDNSServers([]netip.Addr{netip.MustParseAddr("::ffff:203.0.113.53")})
|
||||||
|
t.Cleanup(func() { internet.SkipDNSServers(nil) })
|
||||||
|
preferGo, dial := net.DefaultResolver.PreferGo, net.DefaultResolver.Dial
|
||||||
|
saved := resolveOnOwn()
|
||||||
|
t.Cleanup(saved.restore)
|
||||||
|
if !net.DefaultResolver.PreferGo || net.DefaultResolver.Dial == nil {
|
||||||
|
t.Fatal("net.DefaultResolver is unchanged")
|
||||||
|
}
|
||||||
|
if _, err := net.DefaultResolver.Dial(context.Background(), "udp", "203.0.113.53:53"); err == nil {
|
||||||
|
t.Error("the TUN's DNS server was not skipped")
|
||||||
|
}
|
||||||
|
conn, err := net.DefaultResolver.Dial(context.Background(), "udp", "127.0.0.1:53")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
conn.Close()
|
||||||
|
saved.restore()
|
||||||
|
if net.DefaultResolver.PreferGo != preferGo || (net.DefaultResolver.Dial == nil) != (dial == nil) {
|
||||||
|
t.Error("net.DefaultResolver is not restored")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestTunOnlyDNS checks that a DNS server another interface uses as well is
|
||||||
|
// not skipped, while one of the TUN alone is.
|
||||||
|
func TestTunOnlyDNS(t *testing.T) {
|
||||||
|
adapters, err := winipcfg.GetAdaptersAddresses(windows.AF_UNSPEC, winipcfg.GAAFlagIncludeGateways)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
var other netip.Addr
|
||||||
|
for _, adapter := range adapters {
|
||||||
|
if adapter.OperStatus == winipcfg.IfOperStatusUp && adapter.FirstGatewayAddress != nil && adapter.FirstDNSServerAddress != nil {
|
||||||
|
other, _ = netip.AddrFromSlice(adapter.FirstDNSServerAddress.Address.IP())
|
||||||
|
other = other.Unmap()
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if !other.IsValid() {
|
||||||
|
t.Skip("no interface with a gateway and a DNS server")
|
||||||
|
}
|
||||||
|
tunOnly := netip.MustParseAddr("203.0.113.53")
|
||||||
|
// LUID 0 is no interface, so every one counts as another.
|
||||||
|
got, err := tunOnlyDNS(0, []netip.Addr{other, tunOnly})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if !slices.Equal(got, []netip.Addr{tunOnly}) {
|
||||||
|
t.Errorf("got %v, want [%v]", got, tunOnly)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestFlushDNSCache(t *testing.T) {
|
||||||
|
if err := flushDNSCache(); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
}
|
||||||
+12
-2
@@ -52,9 +52,12 @@ func (b *bind) Open(port uint16) (fns []conn.ReceiveFunc, actualPort uint16, err
|
|||||||
case <-ch:
|
case <-ch:
|
||||||
default:
|
default:
|
||||||
errors.LogErrorInner(context.Background(), err, "unexpected closed")
|
errors.LogErrorInner(context.Background(), err, "unexpected closed")
|
||||||
if b.downFunc != nil {
|
b.mu.Lock()
|
||||||
|
downFunc := b.downFunc
|
||||||
|
b.mu.Unlock()
|
||||||
|
if downFunc != nil {
|
||||||
go func() {
|
go func() {
|
||||||
common.Must(b.downFunc())
|
common.Must(downFunc())
|
||||||
}()
|
}()
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -76,6 +79,13 @@ func (b *bind) Open(port uint16) (fns []conn.ReceiveFunc, actualPort uint16, err
|
|||||||
}, uint16(c.LocalAddr().(*net.UDPAddr).Port), nil
|
}, uint16(c.LocalAddr().(*net.UDPAddr).Port), nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// setDownFunc sets downFunc after the device is created, since the device may already be using the bind.
|
||||||
|
func (b *bind) setDownFunc(f func() error) {
|
||||||
|
b.mu.Lock()
|
||||||
|
defer b.mu.Unlock()
|
||||||
|
b.downFunc = f
|
||||||
|
}
|
||||||
|
|
||||||
func (b *bind) Close() error {
|
func (b *bind) Close() error {
|
||||||
b.mu.Lock()
|
b.mu.Lock()
|
||||||
defer b.mu.Unlock()
|
defer b.mu.Unlock()
|
||||||
|
|||||||
@@ -27,7 +27,6 @@ import (
|
|||||||
"github.com/xtls/xray-core/features/stats"
|
"github.com/xtls/xray-core/features/stats"
|
||||||
"github.com/xtls/xray-core/transport"
|
"github.com/xtls/xray-core/transport"
|
||||||
"github.com/xtls/xray-core/transport/internet"
|
"github.com/xtls/xray-core/transport/internet"
|
||||||
"github.com/xtls/xray-core/transport/internet/finalmask"
|
|
||||||
"golang.zx2c4.com/wireguard/device"
|
"golang.zx2c4.com/wireguard/device"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -200,7 +199,7 @@ func (h *Handler) Process(ctx context.Context, link *transport.Link, dialer inte
|
|||||||
}
|
}
|
||||||
defer conn.Close()
|
defer conn.Close()
|
||||||
c := &UDPConnClient{
|
c := &UDPConnClient{
|
||||||
PacketConn: conn.(*internet.PacketConnWrapper).PacketConn,
|
PacketConn: conn.(*net.PacketConnWrapper).PacketConn,
|
||||||
Dest: conn.RemoteAddr().(*net.UDPAddr),
|
Dest: conn.RemoteAddr().(*net.UDPAddr),
|
||||||
}
|
}
|
||||||
reader = c
|
reader = c
|
||||||
@@ -264,14 +263,14 @@ func (h *Handler) init(ctx context.Context) error {
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, errors.New("failed to dial to dest").Base(err)
|
return nil, errors.New("failed to dial to dest").Base(err)
|
||||||
}
|
}
|
||||||
pktConn = conn.(*finalmask.PacketConnWrapper).PacketConn
|
pktConn = conn.(*net.PacketConnWrapper).PacketConn
|
||||||
} else {
|
} else {
|
||||||
conn, err := internet.DialSystem(ctx, dest, h.streamSettings.SocketSettings)
|
conn, err := internet.DialSystem(ctx, dest, h.streamSettings.SocketSettings)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, errors.New("failed to dial to dest").Base(err)
|
return nil, errors.New("failed to dial to dest").Base(err)
|
||||||
}
|
}
|
||||||
switch c := conn.(type) {
|
switch c := conn.(type) {
|
||||||
case *internet.PacketConnWrapper:
|
case *net.PacketConnWrapper:
|
||||||
pktConn = c.PacketConn
|
pktConn = c.PacketConn
|
||||||
case *cnc.Connection:
|
case *cnc.Connection:
|
||||||
pktConn = &internet.FakePacketConn{Conn: c}
|
pktConn = &internet.FakePacketConn{Conn: c}
|
||||||
@@ -288,7 +287,13 @@ func (h *Handler) init(ctx context.Context) error {
|
|||||||
}
|
}
|
||||||
return pktConn, nil
|
return pktConn, nil
|
||||||
}
|
}
|
||||||
bind := &bind{}
|
// device.NewDevice may use the bind right away (Up -> BindUpdate -> Open),
|
||||||
|
// so everything it reads must be set before creating the device.
|
||||||
|
bind := &bind{
|
||||||
|
resolveFunc: resolveFunc,
|
||||||
|
listenFunc: listenFunc,
|
||||||
|
reserved: h.conf.Reserved,
|
||||||
|
}
|
||||||
logger := &device.Logger{
|
logger := &device.Logger{
|
||||||
Verbosef: func(format string, args ...any) {
|
Verbosef: func(format string, args ...any) {
|
||||||
log.Record(&log.GeneralMessage{
|
log.Record(&log.GeneralMessage{
|
||||||
@@ -304,10 +309,7 @@ func (h *Handler) init(ctx context.Context) error {
|
|||||||
},
|
},
|
||||||
}
|
}
|
||||||
dev := device.NewDevice(h.tun, bind, logger)
|
dev := device.NewDevice(h.tun, bind, logger)
|
||||||
bind.resolveFunc = resolveFunc
|
bind.setDownFunc(dev.Down)
|
||||||
bind.listenFunc = listenFunc
|
|
||||||
bind.downFunc = dev.Down
|
|
||||||
bind.reserved = h.conf.Reserved
|
|
||||||
var cfg strings.Builder
|
var cfg strings.Builder
|
||||||
cfg.WriteString("private_key=" + h.conf.SecretKey + "\n")
|
cfg.WriteString("private_key=" + h.conf.SecretKey + "\n")
|
||||||
for _, peer := range h.conf.Peers {
|
for _, peer := range h.conf.Peers {
|
||||||
|
|||||||
@@ -21,7 +21,7 @@ import (
|
|||||||
"syscall"
|
"syscall"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/xtls/xray-core/transport/internet"
|
xnet "github.com/xtls/xray-core/common/net"
|
||||||
"golang.zx2c4.com/wireguard/tun"
|
"golang.zx2c4.com/wireguard/tun"
|
||||||
|
|
||||||
"golang.org/x/net/dns/dnsmessage"
|
"golang.org/x/net/dns/dnsmessage"
|
||||||
@@ -220,7 +220,7 @@ func (tun *netTun) DialUDPAddrPort(laddr, raddr netip.AddrPort) (net.Conn, error
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
return &internet.PacketConnWrapper{
|
return &xnet.PacketConnWrapper{
|
||||||
PacketConn: conn,
|
PacketConn: conn,
|
||||||
Dest: net.UDPAddrFromAddrPort(raddr),
|
Dest: net.UDPAddrFromAddrPort(raddr),
|
||||||
}, nil
|
}, nil
|
||||||
|
|||||||
@@ -113,7 +113,7 @@ func NewServer(ctx context.Context, conf *DeviceConfig) (*Server, error) {
|
|||||||
users.Store(user.Account.(*MemoryAccount).Pub, user)
|
users.Store(user.Account.(*MemoryAccount).Pub, user)
|
||||||
}
|
}
|
||||||
|
|
||||||
return &Server{
|
s := &Server{
|
||||||
conf: conf,
|
conf: conf,
|
||||||
ctx: core.ToBackgroundDetachedContext(ctx),
|
ctx: core.ToBackgroundDetachedContext(ctx),
|
||||||
policyManager: p,
|
policyManager: p,
|
||||||
@@ -131,7 +131,10 @@ func NewServer(ctx context.Context, conf *DeviceConfig) (*Server, error) {
|
|||||||
|
|
||||||
pub: pub,
|
pub: pub,
|
||||||
users: users,
|
users: users,
|
||||||
}, nil
|
}
|
||||||
|
// Install the stack's protocol handlers before the device can deliver packets to it (Start -> dev.Up).
|
||||||
|
CreateForwarder(stack, s.HandleConnection)
|
||||||
|
return s, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (s *Server) AddUser(ctx context.Context, user *protocol.MemoryUser) error {
|
func (s *Server) AddUser(ctx context.Context, user *protocol.MemoryUser) error {
|
||||||
@@ -320,7 +323,6 @@ func (s *Server) Start() error {
|
|||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
s.dev = dev
|
s.dev = dev
|
||||||
CreateForwarder(s.stack, s.HandleConnection)
|
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -16,7 +16,7 @@ import (
|
|||||||
|
|
||||||
"github.com/vishvananda/netlink"
|
"github.com/vishvananda/netlink"
|
||||||
"github.com/xtls/xray-core/common/errors"
|
"github.com/xtls/xray-core/common/errors"
|
||||||
"github.com/xtls/xray-core/transport/internet"
|
xnet "github.com/xtls/xray-core/common/net"
|
||||||
"golang.zx2c4.com/wireguard/tun"
|
"golang.zx2c4.com/wireguard/tun"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -263,7 +263,7 @@ func (tun *kernelTun) DialUDPAddrPort(laddr, raddr netip.AddrPort) (net.Conn, er
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
return &internet.PacketConnWrapper{
|
return &xnet.PacketConnWrapper{
|
||||||
PacketConn: conn,
|
PacketConn: conn,
|
||||||
Dest: net.UDPAddrFromAddrPort(raddr),
|
Dest: net.UDPAddrFromAddrPort(raddr),
|
||||||
}, nil
|
}, nil
|
||||||
|
|||||||
@@ -0,0 +1,32 @@
|
|||||||
|
package internet
|
||||||
|
|
||||||
|
import (
|
||||||
|
"net/netip"
|
||||||
|
"slices"
|
||||||
|
"sync/atomic"
|
||||||
|
)
|
||||||
|
|
||||||
|
var skippedDNSServers atomic.Pointer[[]netip.Addr]
|
||||||
|
|
||||||
|
// SkipDNSServers has the queries Xray sends to the system's DNS servers on its
|
||||||
|
// own, like those of localdns, skip servers until it is called again. The DNS
|
||||||
|
// servers of a TUN are only meant for what goes through it: queried by Xray
|
||||||
|
// itself they lead back into it, or nowhere.
|
||||||
|
func SkipDNSServers(servers []netip.Addr) {
|
||||||
|
skipped := make([]netip.Addr, len(servers))
|
||||||
|
for i, server := range servers {
|
||||||
|
skipped[i] = server.Unmap()
|
||||||
|
}
|
||||||
|
skippedDNSServers.Store(&skipped)
|
||||||
|
}
|
||||||
|
|
||||||
|
// IsSkippedDNSServer reports whether address, a DNS server as host:port, is to
|
||||||
|
// be skipped, see SkipDNSServers.
|
||||||
|
func IsSkippedDNSServer(address string) bool {
|
||||||
|
skipped := skippedDNSServers.Load()
|
||||||
|
if skipped == nil {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
server, err := netip.ParseAddrPort(address)
|
||||||
|
return err == nil && slices.Contains(*skipped, server.Addr().Unmap())
|
||||||
|
}
|
||||||
@@ -0,0 +1,27 @@
|
|||||||
|
package internet_test
|
||||||
|
|
||||||
|
import (
|
||||||
|
"net/netip"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/xtls/xray-core/transport/internet"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestSkipDNSServers(t *testing.T) {
|
||||||
|
internet.SkipDNSServers([]netip.Addr{netip.MustParseAddr("::ffff:203.0.113.53"), netip.MustParseAddr("2001:db8::53")})
|
||||||
|
t.Cleanup(func() { internet.SkipDNSServers(nil) })
|
||||||
|
for address, want := range map[string]bool{
|
||||||
|
"203.0.113.53:53": true,
|
||||||
|
"[2001:db8::53]:53": true,
|
||||||
|
"198.51.100.53:53": false,
|
||||||
|
"localhost:53": false,
|
||||||
|
} {
|
||||||
|
if got := internet.IsSkippedDNSServer(address); got != want {
|
||||||
|
t.Errorf("IsSkippedDNSServer(%q) = %v, want %v", address, got, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
internet.SkipDNSServers(nil)
|
||||||
|
if internet.IsSkippedDNSServer("203.0.113.53:53") {
|
||||||
|
t.Error("still skipped after SkipDNSServers(nil)")
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -82,7 +82,7 @@ func (fm *FinalMask) DialTCP(ctx context.Context, dest net.Destination) (net.Con
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
return &PacketConnWrapper{PacketConn: conn, udpAddr: addr}, err
|
return &net.PacketConnWrapper{PacketConn: conn, Dest: addr}, err
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
for i := range fm.tcpMasks {
|
for i := range fm.tcpMasks {
|
||||||
@@ -144,7 +144,7 @@ func (fm *FinalMask) DialUDP(ctx context.Context, dest net.Destination) (net.Con
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
return &PacketConnWrapper{PacketConn: conn, udpAddr: addr}, nil
|
return &net.PacketConnWrapper{PacketConn: conn, Dest: addr}, nil
|
||||||
}
|
}
|
||||||
for i := range fm.udpMasks {
|
for i := range fm.udpMasks {
|
||||||
if i > 0 {
|
if i > 0 {
|
||||||
@@ -171,7 +171,7 @@ func (fm *FinalMask) DialUDP(ctx context.Context, dest net.Destination) (net.Con
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
return &PacketConnWrapper{PacketConn: conn, udpAddr: addr}, err
|
return &net.PacketConnWrapper{PacketConn: conn, Dest: addr}, err
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
var sizes []int
|
var sizes []int
|
||||||
@@ -208,7 +208,7 @@ func (fm *FinalMask) DialUDP(ctx context.Context, dest net.Destination) (net.Con
|
|||||||
if addr == nil {
|
if addr == nil {
|
||||||
addr = &net.UDPAddr{IP: []byte{0, 0, 0, 0}}
|
addr = &net.UDPAddr{IP: []byte{0, 0, 0, 0}}
|
||||||
}
|
}
|
||||||
return &PacketConnWrapper{PacketConn: conn, udpAddr: addr}, nil
|
return &net.PacketConnWrapper{PacketConn: conn, Dest: addr}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (fm *FinalMask) ListenPacket(ctx context.Context, addr net.Addr) (net.PacketConn, error) {
|
func (fm *FinalMask) ListenPacket(ctx context.Context, addr net.Addr) (net.PacketConn, error) {
|
||||||
@@ -272,24 +272,6 @@ const (
|
|||||||
UDPSize = 4096
|
UDPSize = 4096
|
||||||
)
|
)
|
||||||
|
|
||||||
type PacketConnWrapper struct {
|
|
||||||
net.PacketConn
|
|
||||||
udpAddr net.Addr
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *PacketConnWrapper) RemoteAddr() net.Addr {
|
|
||||||
return c.udpAddr
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *PacketConnWrapper) Read(b []byte) (n int, err error) {
|
|
||||||
n, _, err = c.PacketConn.ReadFrom(b)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *PacketConnWrapper) Write(b []byte) (n int, err error) {
|
|
||||||
return c.PacketConn.WriteTo(b, c.udpAddr)
|
|
||||||
}
|
|
||||||
|
|
||||||
type headerManagerConn struct {
|
type headerManagerConn struct {
|
||||||
net.PacketConn
|
net.PacketConn
|
||||||
|
|
||||||
|
|||||||
@@ -21,6 +21,135 @@ const (
|
|||||||
_ = protoimpl.EnforceVersion(protoimpl.MaxVersion - 20)
|
_ = protoimpl.EnforceVersion(protoimpl.MaxVersion - 20)
|
||||||
)
|
)
|
||||||
|
|
||||||
|
type Segment_Kind int32
|
||||||
|
|
||||||
|
const (
|
||||||
|
Segment_BYTES Segment_Kind = 0
|
||||||
|
Segment_RANDOM Segment_Kind = 1
|
||||||
|
Segment_RANDOM_ASCII Segment_Kind = 2
|
||||||
|
Segment_RANDOM_DIGIT Segment_Kind = 3
|
||||||
|
Segment_TIMESTAMP Segment_Kind = 4
|
||||||
|
Segment_COUNTER Segment_Kind = 5
|
||||||
|
Segment_NONCE Segment_Kind = 6
|
||||||
|
)
|
||||||
|
|
||||||
|
// Enum value maps for Segment_Kind.
|
||||||
|
var (
|
||||||
|
Segment_Kind_name = map[int32]string{
|
||||||
|
0: "BYTES",
|
||||||
|
1: "RANDOM",
|
||||||
|
2: "RANDOM_ASCII",
|
||||||
|
3: "RANDOM_DIGIT",
|
||||||
|
4: "TIMESTAMP",
|
||||||
|
5: "COUNTER",
|
||||||
|
6: "NONCE",
|
||||||
|
}
|
||||||
|
Segment_Kind_value = map[string]int32{
|
||||||
|
"BYTES": 0,
|
||||||
|
"RANDOM": 1,
|
||||||
|
"RANDOM_ASCII": 2,
|
||||||
|
"RANDOM_DIGIT": 3,
|
||||||
|
"TIMESTAMP": 4,
|
||||||
|
"COUNTER": 5,
|
||||||
|
"NONCE": 6,
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
func (x Segment_Kind) Enum() *Segment_Kind {
|
||||||
|
p := new(Segment_Kind)
|
||||||
|
*p = x
|
||||||
|
return p
|
||||||
|
}
|
||||||
|
|
||||||
|
func (x Segment_Kind) String() string {
|
||||||
|
return protoimpl.X.EnumStringOf(x.Descriptor(), protoreflect.EnumNumber(x))
|
||||||
|
}
|
||||||
|
|
||||||
|
func (Segment_Kind) Descriptor() protoreflect.EnumDescriptor {
|
||||||
|
return file_transport_internet_finalmask_noise_config_proto_enumTypes[0].Descriptor()
|
||||||
|
}
|
||||||
|
|
||||||
|
func (Segment_Kind) Type() protoreflect.EnumType {
|
||||||
|
return &file_transport_internet_finalmask_noise_config_proto_enumTypes[0]
|
||||||
|
}
|
||||||
|
|
||||||
|
func (x Segment_Kind) Number() protoreflect.EnumNumber {
|
||||||
|
return protoreflect.EnumNumber(x)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Deprecated: Use Segment_Kind.Descriptor instead.
|
||||||
|
func (Segment_Kind) EnumDescriptor() ([]byte, []int) {
|
||||||
|
return file_transport_internet_finalmask_noise_config_proto_rawDescGZIP(), []int{0, 0}
|
||||||
|
}
|
||||||
|
|
||||||
|
type Segment struct {
|
||||||
|
state protoimpl.MessageState `protogen:"open.v1"`
|
||||||
|
Kind Segment_Kind `protobuf:"varint,1,opt,name=kind,proto3,enum=xray.transport.internet.finalmask.noise.Segment_Kind" json:"kind,omitempty"`
|
||||||
|
Bytes []byte `protobuf:"bytes,2,opt,name=bytes,proto3" json:"bytes,omitempty"`
|
||||||
|
MinSize int64 `protobuf:"varint,3,opt,name=min_size,json=minSize,proto3" json:"min_size,omitempty"`
|
||||||
|
MaxSize int64 `protobuf:"varint,4,opt,name=max_size,json=maxSize,proto3" json:"max_size,omitempty"`
|
||||||
|
unknownFields protoimpl.UnknownFields
|
||||||
|
sizeCache protoimpl.SizeCache
|
||||||
|
}
|
||||||
|
|
||||||
|
func (x *Segment) Reset() {
|
||||||
|
*x = Segment{}
|
||||||
|
mi := &file_transport_internet_finalmask_noise_config_proto_msgTypes[0]
|
||||||
|
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
|
||||||
|
ms.StoreMessageInfo(mi)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (x *Segment) String() string {
|
||||||
|
return protoimpl.X.MessageStringOf(x)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (*Segment) ProtoMessage() {}
|
||||||
|
|
||||||
|
func (x *Segment) ProtoReflect() protoreflect.Message {
|
||||||
|
mi := &file_transport_internet_finalmask_noise_config_proto_msgTypes[0]
|
||||||
|
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 Segment.ProtoReflect.Descriptor instead.
|
||||||
|
func (*Segment) Descriptor() ([]byte, []int) {
|
||||||
|
return file_transport_internet_finalmask_noise_config_proto_rawDescGZIP(), []int{0}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (x *Segment) GetKind() Segment_Kind {
|
||||||
|
if x != nil {
|
||||||
|
return x.Kind
|
||||||
|
}
|
||||||
|
return Segment_BYTES
|
||||||
|
}
|
||||||
|
|
||||||
|
func (x *Segment) GetBytes() []byte {
|
||||||
|
if x != nil {
|
||||||
|
return x.Bytes
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (x *Segment) GetMinSize() int64 {
|
||||||
|
if x != nil {
|
||||||
|
return x.MinSize
|
||||||
|
}
|
||||||
|
return 0
|
||||||
|
}
|
||||||
|
|
||||||
|
func (x *Segment) GetMaxSize() int64 {
|
||||||
|
if x != nil {
|
||||||
|
return x.MaxSize
|
||||||
|
}
|
||||||
|
return 0
|
||||||
|
}
|
||||||
|
|
||||||
type Item struct {
|
type Item struct {
|
||||||
state protoimpl.MessageState `protogen:"open.v1"`
|
state protoimpl.MessageState `protogen:"open.v1"`
|
||||||
RandMin int64 `protobuf:"varint,1,opt,name=rand_min,json=randMin,proto3" json:"rand_min,omitempty"`
|
RandMin int64 `protobuf:"varint,1,opt,name=rand_min,json=randMin,proto3" json:"rand_min,omitempty"`
|
||||||
@@ -30,13 +159,14 @@ type Item struct {
|
|||||||
Packet []byte `protobuf:"bytes,5,opt,name=packet,proto3" json:"packet,omitempty"`
|
Packet []byte `protobuf:"bytes,5,opt,name=packet,proto3" json:"packet,omitempty"`
|
||||||
DelayMin int64 `protobuf:"varint,6,opt,name=delay_min,json=delayMin,proto3" json:"delay_min,omitempty"`
|
DelayMin int64 `protobuf:"varint,6,opt,name=delay_min,json=delayMin,proto3" json:"delay_min,omitempty"`
|
||||||
DelayMax int64 `protobuf:"varint,7,opt,name=delay_max,json=delayMax,proto3" json:"delay_max,omitempty"`
|
DelayMax int64 `protobuf:"varint,7,opt,name=delay_max,json=delayMax,proto3" json:"delay_max,omitempty"`
|
||||||
|
Segments []*Segment `protobuf:"bytes,8,rep,name=segments,proto3" json:"segments,omitempty"`
|
||||||
unknownFields protoimpl.UnknownFields
|
unknownFields protoimpl.UnknownFields
|
||||||
sizeCache protoimpl.SizeCache
|
sizeCache protoimpl.SizeCache
|
||||||
}
|
}
|
||||||
|
|
||||||
func (x *Item) Reset() {
|
func (x *Item) Reset() {
|
||||||
*x = Item{}
|
*x = Item{}
|
||||||
mi := &file_transport_internet_finalmask_noise_config_proto_msgTypes[0]
|
mi := &file_transport_internet_finalmask_noise_config_proto_msgTypes[1]
|
||||||
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
|
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
|
||||||
ms.StoreMessageInfo(mi)
|
ms.StoreMessageInfo(mi)
|
||||||
}
|
}
|
||||||
@@ -48,7 +178,7 @@ func (x *Item) String() string {
|
|||||||
func (*Item) ProtoMessage() {}
|
func (*Item) ProtoMessage() {}
|
||||||
|
|
||||||
func (x *Item) ProtoReflect() protoreflect.Message {
|
func (x *Item) ProtoReflect() protoreflect.Message {
|
||||||
mi := &file_transport_internet_finalmask_noise_config_proto_msgTypes[0]
|
mi := &file_transport_internet_finalmask_noise_config_proto_msgTypes[1]
|
||||||
if x != nil {
|
if x != nil {
|
||||||
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
|
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
|
||||||
if ms.LoadMessageInfo() == nil {
|
if ms.LoadMessageInfo() == nil {
|
||||||
@@ -61,7 +191,7 @@ func (x *Item) ProtoReflect() protoreflect.Message {
|
|||||||
|
|
||||||
// Deprecated: Use Item.ProtoReflect.Descriptor instead.
|
// Deprecated: Use Item.ProtoReflect.Descriptor instead.
|
||||||
func (*Item) Descriptor() ([]byte, []int) {
|
func (*Item) Descriptor() ([]byte, []int) {
|
||||||
return file_transport_internet_finalmask_noise_config_proto_rawDescGZIP(), []int{0}
|
return file_transport_internet_finalmask_noise_config_proto_rawDescGZIP(), []int{1}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (x *Item) GetRandMin() int64 {
|
func (x *Item) GetRandMin() int64 {
|
||||||
@@ -113,6 +243,13 @@ func (x *Item) GetDelayMax() int64 {
|
|||||||
return 0
|
return 0
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (x *Item) GetSegments() []*Segment {
|
||||||
|
if x != nil {
|
||||||
|
return x.Segments
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
type Config struct {
|
type Config struct {
|
||||||
state protoimpl.MessageState `protogen:"open.v1"`
|
state protoimpl.MessageState `protogen:"open.v1"`
|
||||||
ResetMin int64 `protobuf:"varint,1,opt,name=reset_min,json=resetMin,proto3" json:"reset_min,omitempty"`
|
ResetMin int64 `protobuf:"varint,1,opt,name=reset_min,json=resetMin,proto3" json:"reset_min,omitempty"`
|
||||||
@@ -124,7 +261,7 @@ type Config struct {
|
|||||||
|
|
||||||
func (x *Config) Reset() {
|
func (x *Config) Reset() {
|
||||||
*x = Config{}
|
*x = Config{}
|
||||||
mi := &file_transport_internet_finalmask_noise_config_proto_msgTypes[1]
|
mi := &file_transport_internet_finalmask_noise_config_proto_msgTypes[2]
|
||||||
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
|
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
|
||||||
ms.StoreMessageInfo(mi)
|
ms.StoreMessageInfo(mi)
|
||||||
}
|
}
|
||||||
@@ -136,7 +273,7 @@ func (x *Config) String() string {
|
|||||||
func (*Config) ProtoMessage() {}
|
func (*Config) ProtoMessage() {}
|
||||||
|
|
||||||
func (x *Config) ProtoReflect() protoreflect.Message {
|
func (x *Config) ProtoReflect() protoreflect.Message {
|
||||||
mi := &file_transport_internet_finalmask_noise_config_proto_msgTypes[1]
|
mi := &file_transport_internet_finalmask_noise_config_proto_msgTypes[2]
|
||||||
if x != nil {
|
if x != nil {
|
||||||
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
|
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
|
||||||
if ms.LoadMessageInfo() == nil {
|
if ms.LoadMessageInfo() == nil {
|
||||||
@@ -149,7 +286,7 @@ func (x *Config) ProtoReflect() protoreflect.Message {
|
|||||||
|
|
||||||
// Deprecated: Use Config.ProtoReflect.Descriptor instead.
|
// Deprecated: Use Config.ProtoReflect.Descriptor instead.
|
||||||
func (*Config) Descriptor() ([]byte, []int) {
|
func (*Config) Descriptor() ([]byte, []int) {
|
||||||
return file_transport_internet_finalmask_noise_config_proto_rawDescGZIP(), []int{1}
|
return file_transport_internet_finalmask_noise_config_proto_rawDescGZIP(), []int{2}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (x *Config) GetResetMin() int64 {
|
func (x *Config) GetResetMin() int64 {
|
||||||
@@ -177,7 +314,21 @@ var File_transport_internet_finalmask_noise_config_proto protoreflect.FileDescri
|
|||||||
|
|
||||||
const file_transport_internet_finalmask_noise_config_proto_rawDesc = "" +
|
const file_transport_internet_finalmask_noise_config_proto_rawDesc = "" +
|
||||||
"\n" +
|
"\n" +
|
||||||
"/transport/internet/finalmask/noise/config.proto\x12'xray.transport.internet.finalmask.noise\"\xda\x01\n" +
|
"/transport/internet/finalmask/noise/config.proto\x12'xray.transport.internet.finalmask.noise\"\x8a\x02\n" +
|
||||||
|
"\aSegment\x12I\n" +
|
||||||
|
"\x04kind\x18\x01 \x01(\x0e25.xray.transport.internet.finalmask.noise.Segment.KindR\x04kind\x12\x14\n" +
|
||||||
|
"\x05bytes\x18\x02 \x01(\fR\x05bytes\x12\x19\n" +
|
||||||
|
"\bmin_size\x18\x03 \x01(\x03R\aminSize\x12\x19\n" +
|
||||||
|
"\bmax_size\x18\x04 \x01(\x03R\amaxSize\"h\n" +
|
||||||
|
"\x04Kind\x12\t\n" +
|
||||||
|
"\x05BYTES\x10\x00\x12\n" +
|
||||||
|
"\n" +
|
||||||
|
"\x06RANDOM\x10\x01\x12\x10\n" +
|
||||||
|
"\fRANDOM_ASCII\x10\x02\x12\x10\n" +
|
||||||
|
"\fRANDOM_DIGIT\x10\x03\x12\r\n" +
|
||||||
|
"\tTIMESTAMP\x10\x04\x12\v\n" +
|
||||||
|
"\aCOUNTER\x10\x05\x12\t\n" +
|
||||||
|
"\x05NONCE\x10\x06\"\xa8\x02\n" +
|
||||||
"\x04Item\x12\x19\n" +
|
"\x04Item\x12\x19\n" +
|
||||||
"\brand_min\x18\x01 \x01(\x03R\arandMin\x12\x19\n" +
|
"\brand_min\x18\x01 \x01(\x03R\arandMin\x12\x19\n" +
|
||||||
"\brand_max\x18\x02 \x01(\x03R\arandMax\x12$\n" +
|
"\brand_max\x18\x02 \x01(\x03R\arandMax\x12$\n" +
|
||||||
@@ -185,7 +336,8 @@ const file_transport_internet_finalmask_noise_config_proto_rawDesc = "" +
|
|||||||
"\x0erand_range_max\x18\x04 \x01(\x05R\frandRangeMax\x12\x16\n" +
|
"\x0erand_range_max\x18\x04 \x01(\x05R\frandRangeMax\x12\x16\n" +
|
||||||
"\x06packet\x18\x05 \x01(\fR\x06packet\x12\x1b\n" +
|
"\x06packet\x18\x05 \x01(\fR\x06packet\x12\x1b\n" +
|
||||||
"\tdelay_min\x18\x06 \x01(\x03R\bdelayMin\x12\x1b\n" +
|
"\tdelay_min\x18\x06 \x01(\x03R\bdelayMin\x12\x1b\n" +
|
||||||
"\tdelay_max\x18\a \x01(\x03R\bdelayMax\"\x87\x01\n" +
|
"\tdelay_max\x18\a \x01(\x03R\bdelayMax\x12L\n" +
|
||||||
|
"\bsegments\x18\b \x03(\v20.xray.transport.internet.finalmask.noise.SegmentR\bsegments\"\x87\x01\n" +
|
||||||
"\x06Config\x12\x1b\n" +
|
"\x06Config\x12\x1b\n" +
|
||||||
"\treset_min\x18\x01 \x01(\x03R\bresetMin\x12\x1b\n" +
|
"\treset_min\x18\x01 \x01(\x03R\bresetMin\x12\x1b\n" +
|
||||||
"\treset_max\x18\x02 \x01(\x03R\bresetMax\x12C\n" +
|
"\treset_max\x18\x02 \x01(\x03R\bresetMax\x12C\n" +
|
||||||
@@ -204,18 +356,23 @@ func file_transport_internet_finalmask_noise_config_proto_rawDescGZIP() []byte {
|
|||||||
return file_transport_internet_finalmask_noise_config_proto_rawDescData
|
return file_transport_internet_finalmask_noise_config_proto_rawDescData
|
||||||
}
|
}
|
||||||
|
|
||||||
var file_transport_internet_finalmask_noise_config_proto_msgTypes = make([]protoimpl.MessageInfo, 2)
|
var file_transport_internet_finalmask_noise_config_proto_enumTypes = make([]protoimpl.EnumInfo, 1)
|
||||||
|
var file_transport_internet_finalmask_noise_config_proto_msgTypes = make([]protoimpl.MessageInfo, 3)
|
||||||
var file_transport_internet_finalmask_noise_config_proto_goTypes = []any{
|
var file_transport_internet_finalmask_noise_config_proto_goTypes = []any{
|
||||||
(*Item)(nil), // 0: xray.transport.internet.finalmask.noise.Item
|
(Segment_Kind)(0), // 0: xray.transport.internet.finalmask.noise.Segment.Kind
|
||||||
(*Config)(nil), // 1: xray.transport.internet.finalmask.noise.Config
|
(*Segment)(nil), // 1: xray.transport.internet.finalmask.noise.Segment
|
||||||
|
(*Item)(nil), // 2: xray.transport.internet.finalmask.noise.Item
|
||||||
|
(*Config)(nil), // 3: xray.transport.internet.finalmask.noise.Config
|
||||||
}
|
}
|
||||||
var file_transport_internet_finalmask_noise_config_proto_depIdxs = []int32{
|
var file_transport_internet_finalmask_noise_config_proto_depIdxs = []int32{
|
||||||
0, // 0: xray.transport.internet.finalmask.noise.Config.items:type_name -> xray.transport.internet.finalmask.noise.Item
|
0, // 0: xray.transport.internet.finalmask.noise.Segment.kind:type_name -> xray.transport.internet.finalmask.noise.Segment.Kind
|
||||||
1, // [1:1] is the sub-list for method output_type
|
1, // 1: xray.transport.internet.finalmask.noise.Item.segments:type_name -> xray.transport.internet.finalmask.noise.Segment
|
||||||
1, // [1:1] is the sub-list for method input_type
|
2, // 2: xray.transport.internet.finalmask.noise.Config.items:type_name -> xray.transport.internet.finalmask.noise.Item
|
||||||
1, // [1:1] is the sub-list for extension type_name
|
3, // [3:3] is the sub-list for method output_type
|
||||||
1, // [1:1] is the sub-list for extension extendee
|
3, // [3:3] is the sub-list for method input_type
|
||||||
0, // [0:1] is the sub-list for field type_name
|
3, // [3:3] is the sub-list for extension type_name
|
||||||
|
3, // [3:3] is the sub-list for extension extendee
|
||||||
|
0, // [0:3] is the sub-list for field type_name
|
||||||
}
|
}
|
||||||
|
|
||||||
func init() { file_transport_internet_finalmask_noise_config_proto_init() }
|
func init() { file_transport_internet_finalmask_noise_config_proto_init() }
|
||||||
@@ -228,13 +385,14 @@ func file_transport_internet_finalmask_noise_config_proto_init() {
|
|||||||
File: protoimpl.DescBuilder{
|
File: protoimpl.DescBuilder{
|
||||||
GoPackagePath: reflect.TypeOf(x{}).PkgPath(),
|
GoPackagePath: reflect.TypeOf(x{}).PkgPath(),
|
||||||
RawDescriptor: unsafe.Slice(unsafe.StringData(file_transport_internet_finalmask_noise_config_proto_rawDesc), len(file_transport_internet_finalmask_noise_config_proto_rawDesc)),
|
RawDescriptor: unsafe.Slice(unsafe.StringData(file_transport_internet_finalmask_noise_config_proto_rawDesc), len(file_transport_internet_finalmask_noise_config_proto_rawDesc)),
|
||||||
NumEnums: 0,
|
NumEnums: 1,
|
||||||
NumMessages: 2,
|
NumMessages: 3,
|
||||||
NumExtensions: 0,
|
NumExtensions: 0,
|
||||||
NumServices: 0,
|
NumServices: 0,
|
||||||
},
|
},
|
||||||
GoTypes: file_transport_internet_finalmask_noise_config_proto_goTypes,
|
GoTypes: file_transport_internet_finalmask_noise_config_proto_goTypes,
|
||||||
DependencyIndexes: file_transport_internet_finalmask_noise_config_proto_depIdxs,
|
DependencyIndexes: file_transport_internet_finalmask_noise_config_proto_depIdxs,
|
||||||
|
EnumInfos: file_transport_internet_finalmask_noise_config_proto_enumTypes,
|
||||||
MessageInfos: file_transport_internet_finalmask_noise_config_proto_msgTypes,
|
MessageInfos: file_transport_internet_finalmask_noise_config_proto_msgTypes,
|
||||||
}.Build()
|
}.Build()
|
||||||
File_transport_internet_finalmask_noise_config_proto = out.File
|
File_transport_internet_finalmask_noise_config_proto = out.File
|
||||||
|
|||||||
@@ -6,6 +6,22 @@ option go_package = "github.com/xtls/xray-core/transport/internet/finalmask/nois
|
|||||||
option java_package = "com.xray.transport.internet.finalmask.noise";
|
option java_package = "com.xray.transport.internet.finalmask.noise";
|
||||||
option java_multiple_files = true;
|
option java_multiple_files = true;
|
||||||
|
|
||||||
|
message Segment {
|
||||||
|
enum Kind {
|
||||||
|
BYTES = 0;
|
||||||
|
RANDOM = 1;
|
||||||
|
RANDOM_ASCII = 2;
|
||||||
|
RANDOM_DIGIT = 3;
|
||||||
|
TIMESTAMP = 4;
|
||||||
|
COUNTER = 5;
|
||||||
|
NONCE = 6;
|
||||||
|
}
|
||||||
|
Kind kind = 1;
|
||||||
|
bytes bytes = 2;
|
||||||
|
int64 min_size = 3;
|
||||||
|
int64 max_size = 4;
|
||||||
|
}
|
||||||
|
|
||||||
message Item {
|
message Item {
|
||||||
int64 rand_min = 1;
|
int64 rand_min = 1;
|
||||||
int64 rand_max = 2;
|
int64 rand_max = 2;
|
||||||
@@ -14,6 +30,7 @@ message Item {
|
|||||||
bytes packet = 5;
|
bytes packet = 5;
|
||||||
int64 delay_min = 6;
|
int64 delay_min = 6;
|
||||||
int64 delay_max = 7;
|
int64 delay_max = 7;
|
||||||
|
repeated Segment segments = 8;
|
||||||
}
|
}
|
||||||
|
|
||||||
message Config {
|
message Config {
|
||||||
|
|||||||
@@ -1,18 +1,25 @@
|
|||||||
package noise
|
package noise
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"crypto/rand"
|
||||||
|
"encoding/binary"
|
||||||
"net"
|
"net"
|
||||||
"sync"
|
"sync"
|
||||||
|
"sync/atomic"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
|
"github.com/xtls/xray-core/common"
|
||||||
"github.com/xtls/xray-core/common/crypto"
|
"github.com/xtls/xray-core/common/crypto"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
const asciiLetters = "abcdefghijklmnopqrstuvwxyzABCDEFGHIJKLMNOPQRSTUVWXYZ"
|
||||||
|
|
||||||
type noiseConn struct {
|
type noiseConn struct {
|
||||||
net.PacketConn
|
net.PacketConn
|
||||||
config *Config
|
config *Config
|
||||||
m map[string]time.Time
|
m map[string]time.Time
|
||||||
mu sync.Mutex
|
mu sync.Mutex
|
||||||
|
counter atomic.Uint32
|
||||||
}
|
}
|
||||||
|
|
||||||
func NewConnClient(c *Config, raw net.PacketConn) (net.PacketConn, error) {
|
func NewConnClient(c *Config, raw net.PacketConn) (net.PacketConn, error) {
|
||||||
@@ -27,6 +34,62 @@ func NewConnServer(c *Config, raw net.PacketConn) (net.PacketConn, error) {
|
|||||||
return NewConnClient(c, raw)
|
return NewConnClient(c, raw)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (c *noiseConn) buildPacket(item *Item) []byte {
|
||||||
|
if len(item.Segments) == 0 {
|
||||||
|
if item.RandMax > 0 {
|
||||||
|
buf := make([]byte, crypto.RandBetween(item.RandMin, item.RandMax))
|
||||||
|
crypto.RandBytesBetween(buf, byte(item.RandRangeMin), byte(item.RandRangeMax))
|
||||||
|
return buf
|
||||||
|
}
|
||||||
|
return item.Packet
|
||||||
|
}
|
||||||
|
var out []byte
|
||||||
|
for _, seg := range item.Segments {
|
||||||
|
out = append(out, c.buildSegment(seg)...)
|
||||||
|
}
|
||||||
|
return out
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *noiseConn) buildSegment(seg *Segment) []byte {
|
||||||
|
switch seg.Kind {
|
||||||
|
case Segment_BYTES:
|
||||||
|
return seg.Bytes
|
||||||
|
case Segment_TIMESTAMP:
|
||||||
|
b := make([]byte, 4)
|
||||||
|
binary.BigEndian.PutUint32(b, uint32(time.Now().Unix()))
|
||||||
|
return b
|
||||||
|
case Segment_COUNTER:
|
||||||
|
b := make([]byte, 4)
|
||||||
|
binary.BigEndian.PutUint32(b, c.counter.Add(1))
|
||||||
|
return b
|
||||||
|
case Segment_NONCE:
|
||||||
|
b := make([]byte, 8)
|
||||||
|
common.Must2(rand.Read(b))
|
||||||
|
return b
|
||||||
|
default:
|
||||||
|
size := crypto.RandBetween(seg.MinSize, seg.MaxSize+1)
|
||||||
|
if size <= 0 {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
buf := make([]byte, size)
|
||||||
|
switch seg.Kind {
|
||||||
|
case Segment_RANDOM_ASCII:
|
||||||
|
common.Must2(rand.Read(buf))
|
||||||
|
for i := range buf {
|
||||||
|
buf[i] = asciiLetters[int(buf[i])%len(asciiLetters)]
|
||||||
|
}
|
||||||
|
case Segment_RANDOM_DIGIT:
|
||||||
|
common.Must2(rand.Read(buf))
|
||||||
|
for i := range buf {
|
||||||
|
buf[i] = '0' + buf[i]%10
|
||||||
|
}
|
||||||
|
default:
|
||||||
|
common.Must2(rand.Read(buf))
|
||||||
|
}
|
||||||
|
return buf
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func (c *noiseConn) WriteTo(p []byte, addr net.Addr) (n int, err error) {
|
func (c *noiseConn) WriteTo(p []byte, addr net.Addr) (n int, err error) {
|
||||||
c.mu.Lock()
|
c.mu.Lock()
|
||||||
defer c.mu.Unlock()
|
defer c.mu.Unlock()
|
||||||
@@ -35,13 +98,7 @@ func (c *noiseConn) WriteTo(p []byte, addr net.Addr) (n int, err error) {
|
|||||||
|
|
||||||
if t.IsZero() || (c.config.ResetMax > 0 && time.Now().After(t)) {
|
if t.IsZero() || (c.config.ResetMax > 0 && time.Now().After(t)) {
|
||||||
for _, item := range c.config.Items {
|
for _, item := range c.config.Items {
|
||||||
if item.RandMax > 0 {
|
c.PacketConn.WriteTo(c.buildPacket(item), addr)
|
||||||
buf := make([]byte, crypto.RandBetween(item.RandMin, item.RandMax))
|
|
||||||
crypto.RandBytesBetween(buf, byte(item.RandRangeMin), byte(item.RandRangeMax))
|
|
||||||
c.PacketConn.WriteTo(buf, addr)
|
|
||||||
} else {
|
|
||||||
c.PacketConn.WriteTo(item.Packet, addr)
|
|
||||||
}
|
|
||||||
time.Sleep(time.Duration(crypto.RandBetween(item.DelayMin, item.DelayMax)) * time.Millisecond)
|
time.Sleep(time.Duration(crypto.RandBetween(item.DelayMin, item.DelayMax)) * time.Millisecond)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -0,0 +1,137 @@
|
|||||||
|
package noise
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"encoding/binary"
|
||||||
|
"net"
|
||||||
|
"sync"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
)
|
||||||
|
|
||||||
|
type fakePacketConn struct {
|
||||||
|
mu sync.Mutex
|
||||||
|
written [][]byte
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *fakePacketConn) WriteTo(p []byte, _ net.Addr) (int, error) {
|
||||||
|
c.mu.Lock()
|
||||||
|
defer c.mu.Unlock()
|
||||||
|
c.written = append(c.written, bytes.Clone(p))
|
||||||
|
return len(p), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *fakePacketConn) packets() [][]byte {
|
||||||
|
c.mu.Lock()
|
||||||
|
defer c.mu.Unlock()
|
||||||
|
return c.written
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *fakePacketConn) ReadFrom(_ []byte) (int, net.Addr, error) { return 0, nil, nil }
|
||||||
|
func (c *fakePacketConn) Close() error { return nil }
|
||||||
|
func (c *fakePacketConn) LocalAddr() net.Addr { return &net.UDPAddr{} }
|
||||||
|
func (c *fakePacketConn) SetDeadline(time.Time) error { return nil }
|
||||||
|
func (c *fakePacketConn) SetReadDeadline(time.Time) error { return nil }
|
||||||
|
func (c *fakePacketConn) SetWriteDeadline(time.Time) error { return nil }
|
||||||
|
|
||||||
|
func newConn() *noiseConn {
|
||||||
|
return &noiseConn{PacketConn: &fakePacketConn{}, config: &Config{}, m: make(map[string]time.Time)}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestBuildSegmentBytes(t *testing.T) {
|
||||||
|
c := newConn()
|
||||||
|
got := c.buildSegment(&Segment{Kind: Segment_BYTES, Bytes: []byte{0x0d, 0x0a, 0x0d, 0x0a}})
|
||||||
|
require.Equal(t, []byte{0x0d, 0x0a, 0x0d, 0x0a}, got)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestBuildSegmentTimestamp(t *testing.T) {
|
||||||
|
c := newConn()
|
||||||
|
before := time.Now().Unix()
|
||||||
|
got := c.buildSegment(&Segment{Kind: Segment_TIMESTAMP})
|
||||||
|
require.Len(t, got, 4)
|
||||||
|
ts := int64(binary.BigEndian.Uint32(got))
|
||||||
|
require.GreaterOrEqual(t, ts, before)
|
||||||
|
require.LessOrEqual(t, ts, time.Now().Unix())
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestBuildSegmentCounter(t *testing.T) {
|
||||||
|
c := newConn()
|
||||||
|
first := binary.BigEndian.Uint32(c.buildSegment(&Segment{Kind: Segment_COUNTER}))
|
||||||
|
second := binary.BigEndian.Uint32(c.buildSegment(&Segment{Kind: Segment_COUNTER}))
|
||||||
|
require.Equal(t, uint32(1), first)
|
||||||
|
require.Equal(t, uint32(2), second)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestBuildSegmentNonce(t *testing.T) {
|
||||||
|
c := newConn()
|
||||||
|
a := c.buildSegment(&Segment{Kind: Segment_NONCE})
|
||||||
|
b := c.buildSegment(&Segment{Kind: Segment_NONCE})
|
||||||
|
require.Len(t, a, 8)
|
||||||
|
require.Len(t, b, 8)
|
||||||
|
require.NotEqual(t, a, b)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestBuildSegmentRandomSizes(t *testing.T) {
|
||||||
|
c := newConn()
|
||||||
|
for range 200 {
|
||||||
|
require.Len(t, c.buildSegment(&Segment{Kind: Segment_RANDOM, MinSize: 24, MaxSize: 24}), 24)
|
||||||
|
|
||||||
|
n := len(c.buildSegment(&Segment{Kind: Segment_RANDOM, MinSize: 20, MaxSize: 32}))
|
||||||
|
require.GreaterOrEqual(t, n, 20)
|
||||||
|
require.LessOrEqual(t, n, 32)
|
||||||
|
|
||||||
|
for _, b := range c.buildSegment(&Segment{Kind: Segment_RANDOM_ASCII, MinSize: 40, MaxSize: 40}) {
|
||||||
|
require.True(t, (b >= 'a' && b <= 'z') || (b >= 'A' && b <= 'Z'), "not a letter: %q", b)
|
||||||
|
}
|
||||||
|
for _, b := range c.buildSegment(&Segment{Kind: Segment_RANDOM_DIGIT, MinSize: 40, MaxSize: 40}) {
|
||||||
|
require.True(t, b >= '0' && b <= '9', "not a digit: %q", b)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestBuildPacketComposite(t *testing.T) {
|
||||||
|
c := newConn()
|
||||||
|
item := &Item{Segments: []*Segment{
|
||||||
|
{Kind: Segment_BYTES, Bytes: []byte{0x0d, 0x0a, 0x0d, 0x0a}},
|
||||||
|
{Kind: Segment_TIMESTAMP},
|
||||||
|
{Kind: Segment_RANDOM, MinSize: 24, MaxSize: 24},
|
||||||
|
}}
|
||||||
|
got := c.buildPacket(item)
|
||||||
|
require.Len(t, got, 4+4+24)
|
||||||
|
require.Equal(t, []byte{0x0d, 0x0a, 0x0d, 0x0a}, got[:4])
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestBuildPacketLegacy(t *testing.T) {
|
||||||
|
c := newConn()
|
||||||
|
require.Equal(t, []byte{1, 2, 3}, c.buildPacket(&Item{Packet: []byte{1, 2, 3}}))
|
||||||
|
require.Len(t, c.buildPacket(&Item{RandMin: 16, RandMax: 17}), 16)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestWriteToSendsNoiseThenPayload(t *testing.T) {
|
||||||
|
raw := &fakePacketConn{}
|
||||||
|
c := &noiseConn{
|
||||||
|
PacketConn: raw,
|
||||||
|
m: make(map[string]time.Time),
|
||||||
|
config: &Config{Items: []*Item{
|
||||||
|
{Segments: []*Segment{{Kind: Segment_BYTES, Bytes: []byte{0x0d, 0x0a, 0x0d, 0x0a}}, {Kind: Segment_RANDOM, MinSize: 8, MaxSize: 8}}},
|
||||||
|
{RandMin: 40, RandMax: 41},
|
||||||
|
}},
|
||||||
|
}
|
||||||
|
addr := &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1), Port: 51820}
|
||||||
|
payload := []byte("real-handshake")
|
||||||
|
_, err := c.WriteTo(payload, addr)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
sent := raw.packets()
|
||||||
|
require.Len(t, sent, 3)
|
||||||
|
require.Len(t, sent[0], 12)
|
||||||
|
require.Equal(t, []byte{0x0d, 0x0a, 0x0d, 0x0a}, sent[0][:4])
|
||||||
|
require.Len(t, sent[1], 40)
|
||||||
|
require.Equal(t, payload, sent[2])
|
||||||
|
|
||||||
|
_, err = c.WriteTo(payload, addr)
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.Len(t, raw.packets(), 4)
|
||||||
|
}
|
||||||
@@ -380,7 +380,7 @@ func TestPacketConnReadWrite(t *testing.T) {
|
|||||||
t.Fatal(err)
|
t.Fatal(err)
|
||||||
}
|
}
|
||||||
t.Cleanup(func() { clientConn.Close() })
|
t.Cleanup(func() { clientConn.Close() })
|
||||||
client := clientConn.(*finalmask.PacketConnWrapper).PacketConn
|
client := clientConn.(*net.PacketConnWrapper).PacketConn
|
||||||
|
|
||||||
_ = client.SetDeadline(time.Now().Add(time.Second))
|
_ = client.SetDeadline(time.Now().Add(time.Second))
|
||||||
_ = server.SetDeadline(time.Now().Add(time.Second))
|
_ = server.SetDeadline(time.Now().Add(time.Second))
|
||||||
|
|||||||
@@ -73,7 +73,7 @@ func NewUDPHopConn(c *Config, dest *net.Destination, dialer *finalmask.Dialer) (
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
cur := conn.(*finalmask.PacketConnWrapper).PacketConn
|
cur := conn.(*net.PacketConnWrapper).PacketConn
|
||||||
addr := conn.RemoteAddr().(*net.UDPAddr)
|
addr := conn.RemoteAddr().(*net.UDPAddr)
|
||||||
client := &udpHopConn{
|
client := &udpHopConn{
|
||||||
dialer: dialer,
|
dialer: dialer,
|
||||||
@@ -150,7 +150,7 @@ func (c *udpHopConn) hop() {
|
|||||||
_ = c.pre.Close()
|
_ = c.pre.Close()
|
||||||
}
|
}
|
||||||
c.pre = c.cur
|
c.pre = c.cur
|
||||||
c.cur = conn.(*finalmask.PacketConnWrapper).PacketConn
|
c.cur = conn.(*net.PacketConnWrapper).PacketConn
|
||||||
c.wg.Add(1)
|
c.wg.Add(1)
|
||||||
go c.recv(c.cur)
|
go c.recv(c.cur)
|
||||||
}
|
}
|
||||||
@@ -223,13 +223,6 @@ func (c *udpHopConn) Close() error {
|
|||||||
}
|
}
|
||||||
_ = c.cur.Close()
|
_ = c.cur.Close()
|
||||||
c.wg.Wait()
|
c.wg.Wait()
|
||||||
select {
|
|
||||||
case packet := <-c.readCh:
|
|
||||||
if packet.p != nil {
|
|
||||||
pool.Put(packet.p[:cap(packet.p)])
|
|
||||||
}
|
|
||||||
default:
|
|
||||||
}
|
|
||||||
close(c.readCh)
|
close(c.readCh)
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1,417 +1,441 @@
|
|||||||
package xdns
|
package xdns
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"bytes"
|
|
||||||
"context"
|
"context"
|
||||||
"crypto/rand"
|
"crypto/rand"
|
||||||
"encoding/base32"
|
|
||||||
"encoding/binary"
|
|
||||||
go_errors "errors"
|
|
||||||
"io"
|
"io"
|
||||||
"net"
|
mrand "math/rand"
|
||||||
"strconv"
|
|
||||||
"sync"
|
"sync"
|
||||||
"sync/atomic"
|
"sync/atomic"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/xtls/xray-core/common"
|
"github.com/xtls/xray-core/common"
|
||||||
"github.com/xtls/xray-core/common/errors"
|
"github.com/xtls/xray-core/common/errors"
|
||||||
|
"github.com/xtls/xray-core/common/net"
|
||||||
"github.com/xtls/xray-core/transport/internet/finalmask"
|
"github.com/xtls/xray-core/transport/internet/finalmask"
|
||||||
|
"golang.org/x/net/dns/dnsmessage"
|
||||||
)
|
)
|
||||||
|
|
||||||
const (
|
const (
|
||||||
numPadding = 3
|
|
||||||
numPaddingForPoll = 8
|
|
||||||
initPollDelay = 500 * time.Millisecond
|
initPollDelay = 500 * time.Millisecond
|
||||||
maxPollDelay = 10 * time.Second
|
maxPollDelay = 10 * time.Second
|
||||||
pollDelayMultiplier = 2.0
|
pollDelayMultiplier = 2.0
|
||||||
pollLimit = 16
|
pollLimit = 16
|
||||||
)
|
)
|
||||||
|
|
||||||
var base32Encoding = base32.StdEncoding.WithPadding(base32.NoPadding)
|
var pool4K = sync.Pool{
|
||||||
|
New: func() any {
|
||||||
|
return make([]byte, 4096)
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
type packet struct {
|
type packet struct {
|
||||||
p []byte
|
p []byte
|
||||||
addr net.Addr
|
addr net.Addr
|
||||||
}
|
}
|
||||||
|
|
||||||
type xdnsConnClient struct {
|
type xdnsClient struct {
|
||||||
net.PacketConn
|
dialer *finalmask.Dialer
|
||||||
|
|
||||||
resolverAddrs []*net.UDPAddr
|
clientID ClientID
|
||||||
resolverTypes []uint16
|
fragID atomic.Uint32
|
||||||
resolverIdx uint32
|
domains []*Domain
|
||||||
resolverSend map[string]*atomic.Uint32
|
extraPoll int32
|
||||||
|
|
||||||
clientID []byte
|
resolvers []Resolver
|
||||||
domains []Name
|
resolverSends []atomic.Uint32
|
||||||
|
resolverIndex atomic.Uint32
|
||||||
|
|
||||||
pollChan chan struct{}
|
readCh chan packet
|
||||||
readQueue chan *packet
|
sendCh chan []byte
|
||||||
writeQueue chan *packet
|
poolCh chan struct{}
|
||||||
|
closeCh chan struct{}
|
||||||
closed bool
|
wg sync.WaitGroup
|
||||||
mutex sync.Mutex
|
mu sync.Mutex
|
||||||
}
|
}
|
||||||
|
|
||||||
func NewConnClient(c *Config, raw net.PacketConn) (net.PacketConn, error) {
|
func NewClient(c *Config, dialer *finalmask.Dialer) (net.PacketConn, error) {
|
||||||
|
if len(c.Domains) == 0 {
|
||||||
|
return nil, errors.New("empty domains")
|
||||||
|
}
|
||||||
if len(c.Resolvers) == 0 {
|
if len(c.Resolvers) == 0 {
|
||||||
return nil, errors.New("empty resolvers")
|
return nil, errors.New("empty resolvers")
|
||||||
}
|
}
|
||||||
|
if c.ExtraPoll < 0 || c.ExtraPoll > 3 {
|
||||||
var domains []Name
|
return nil, errors.New("c.ExtraPoll < 0 || c.ExtraPoll > 3")
|
||||||
var servers []string
|
|
||||||
var resolverTypes []uint16
|
|
||||||
for _, rs := range c.Resolvers {
|
|
||||||
domain, server, resolverType, err := parseResolver(rs)
|
|
||||||
if err != nil {
|
|
||||||
return nil, errors.New("invalid resolvers").Base(err)
|
|
||||||
}
|
|
||||||
domains = append(domains, domain)
|
|
||||||
servers = append(servers, server)
|
|
||||||
resolverTypes = append(resolverTypes, resolverType)
|
|
||||||
}
|
}
|
||||||
|
domains := make([]*Domain, 0, len(c.Domains))
|
||||||
var resolverAddrs []*net.UDPAddr
|
for i := range c.Domains {
|
||||||
resolverSend := make(map[string]*atomic.Uint32)
|
types := make([]uint16, 0, len(c.Domains[i].Types))
|
||||||
for _, rs := range servers {
|
for j := range c.Domains[i].Types {
|
||||||
h, p, err := net.SplitHostPort(rs)
|
types = append(types, uint16(c.Domains[i].Types[j]))
|
||||||
|
}
|
||||||
|
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 {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
ip := net.ParseIP(h)
|
domains = append(domains, domain)
|
||||||
if ip == nil {
|
}
|
||||||
return nil, errors.New("invalid ip address")
|
resolvers := make([]Resolver, 0, len(c.Resolvers))
|
||||||
}
|
for i := range c.Resolvers {
|
||||||
port, err := strconv.Atoi(p)
|
resolver, err := NewResolver(c.Resolvers[i], dialer)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, errors.New("invalid port").Base(err)
|
return nil, err
|
||||||
}
|
}
|
||||||
addr := &net.UDPAddr{IP: ip, Port: port}
|
resolvers = append(resolvers, resolver)
|
||||||
resolverAddrs = append(resolverAddrs, addr)
|
|
||||||
resolverSend[addr.String()] = &atomic.Uint32{}
|
|
||||||
}
|
}
|
||||||
|
client := &xdnsClient{
|
||||||
|
dialer: dialer,
|
||||||
|
|
||||||
conn := &xdnsConnClient{
|
clientID: NewClientID(),
|
||||||
PacketConn: raw,
|
domains: domains,
|
||||||
|
extraPoll: c.ExtraPoll,
|
||||||
|
|
||||||
resolverAddrs: resolverAddrs,
|
resolvers: resolvers,
|
||||||
resolverTypes: resolverTypes,
|
resolverSends: make([]atomic.Uint32, len(c.Resolvers)),
|
||||||
resolverIdx: 0,
|
|
||||||
resolverSend: resolverSend,
|
|
||||||
|
|
||||||
clientID: make([]byte, 8),
|
readCh: make(chan packet),
|
||||||
domains: domains,
|
sendCh: make(chan []byte, 16),
|
||||||
|
poolCh: make(chan struct{}, pollLimit),
|
||||||
pollChan: make(chan struct{}, pollLimit),
|
closeCh: make(chan struct{}),
|
||||||
readQueue: make(chan *packet, 256),
|
|
||||||
writeQueue: make(chan *packet, 256),
|
|
||||||
}
|
}
|
||||||
|
go client.run()
|
||||||
common.Must2(rand.Read(conn.clientID))
|
return client, nil
|
||||||
|
|
||||||
go conn.recvLoop()
|
|
||||||
go conn.sendLoop()
|
|
||||||
|
|
||||||
return conn, nil
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *xdnsConnClient) recvLoop() {
|
func (c *xdnsClient) closed() bool {
|
||||||
var buf [finalmask.UDPSize]byte
|
select {
|
||||||
|
case <-c.closeCh:
|
||||||
|
return true
|
||||||
|
default:
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
for {
|
func (c *xdnsClient) read(buf []byte, addr net.Addr) bool {
|
||||||
if c.closed {
|
msg := dnsmessage.Message{}
|
||||||
|
if err := msg.Unpack(buf); err != nil {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if !msg.Header.Response || msg.Header.Truncated || msg.Header.RCode != dnsmessage.RCodeSuccess || len(msg.Questions) != 1 {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
var domain *Domain
|
||||||
|
for i := range c.domains {
|
||||||
|
if c.domains[i].IsDomain(msg.Questions[0].Name) {
|
||||||
|
domain = c.domains[i]
|
||||||
break
|
break
|
||||||
}
|
}
|
||||||
|
}
|
||||||
|
if domain == nil || !domain.HasType(uint16(msg.Questions[0].Type)) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
n, addr, err := c.PacketConn.ReadFrom(buf[:])
|
edns0 := uint16(0)
|
||||||
|
for i := range msg.Additionals {
|
||||||
|
if msg.Additionals[i].Header.Type == dnsmessage.TypeOPT {
|
||||||
|
edns0 = uint16(msg.Additionals[i].Header.Class)
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
errors.LogDebug(context.Background(), addr, " edns0 ", edns0, " buf ", len(buf), " ", msg.Questions[0].Type)
|
||||||
|
|
||||||
|
resp := NewResp(msg, domain, 0)
|
||||||
|
|
||||||
|
p := pool4K.Get().([]byte)
|
||||||
|
n := resp.Decode(p)
|
||||||
|
p = p[:n]
|
||||||
|
|
||||||
|
b := p
|
||||||
|
var bs [][]byte
|
||||||
|
for len(b) > 1 {
|
||||||
|
last := b[0]&0xC0 == 0xC0
|
||||||
|
length := int(b[0]&0x3F)<<8 | int(b[1])
|
||||||
|
b = b[2:]
|
||||||
|
if length > len(b) {
|
||||||
|
bs = nil
|
||||||
|
break
|
||||||
|
}
|
||||||
|
packet := make([]byte, length)
|
||||||
|
copy(packet, b)
|
||||||
|
bs = append(bs, packet)
|
||||||
|
if last {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
b = b[length:]
|
||||||
|
if len(b) < 2 {
|
||||||
|
bs = nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
pool4K.Put(p[:cap(p)])
|
||||||
|
|
||||||
|
for i := range bs {
|
||||||
|
select {
|
||||||
|
case <-c.closeCh:
|
||||||
|
return true
|
||||||
|
case c.readCh <- packet{p: bs[i], addr: addr}:
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return len(bs) > 0
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *xdnsClient) run() {
|
||||||
|
for i := range len(c.resolvers) {
|
||||||
|
c.wg.Add(1)
|
||||||
|
go c.recv(i)
|
||||||
|
}
|
||||||
|
|
||||||
|
c.wg.Add(1)
|
||||||
|
go c.send()
|
||||||
|
|
||||||
|
c.wg.Wait()
|
||||||
|
close(c.readCh)
|
||||||
|
close(c.sendCh)
|
||||||
|
close(c.poolCh)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *xdnsClient) recv(i int) {
|
||||||
|
defer c.wg.Done()
|
||||||
|
|
||||||
|
var buf [4096]byte
|
||||||
|
for {
|
||||||
|
n, err := c.resolvers[i].Read(buf[:])
|
||||||
if err != nil {
|
if err != nil {
|
||||||
if go_errors.Is(err, net.ErrClosed) {
|
if c.closed() {
|
||||||
break
|
return
|
||||||
}
|
}
|
||||||
continue
|
errors.LogErrorInner(context.Background(), err, "recv err ", i)
|
||||||
|
return
|
||||||
}
|
}
|
||||||
|
if c.read(buf[:n], c.resolvers[i].Addr()) {
|
||||||
if addr == nil {
|
c.resolverSends[i].Store(0)
|
||||||
continue
|
|
||||||
}
|
|
||||||
|
|
||||||
send := c.resolverSend[addr.String()]
|
|
||||||
if send == nil {
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
|
|
||||||
resp, err := MessageFromWireFormat(buf[:n])
|
|
||||||
if err != nil {
|
|
||||||
errors.LogDebug(context.Background(), addr, " xdns from wireformat err ", err)
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
|
|
||||||
payload := dnsResponsePayload(&resp, c.domains)
|
|
||||||
|
|
||||||
r := bytes.NewReader(payload)
|
|
||||||
anyPacket := false
|
|
||||||
for {
|
|
||||||
p, err := nextPacket(r)
|
|
||||||
if err != nil {
|
|
||||||
break
|
|
||||||
}
|
|
||||||
anyPacket = true
|
|
||||||
|
|
||||||
buf := make([]byte, len(p))
|
|
||||||
copy(buf, p)
|
|
||||||
select {
|
select {
|
||||||
case c.readQueue <- &packet{
|
case c.poolCh <- struct{}{}:
|
||||||
p: buf,
|
|
||||||
addr: addr,
|
|
||||||
}:
|
|
||||||
default:
|
|
||||||
errors.LogDebug(context.Background(), addr, " mask read err queue full")
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
if anyPacket {
|
|
||||||
send.Store(0)
|
|
||||||
select {
|
|
||||||
case c.pollChan <- struct{}{}:
|
|
||||||
default:
|
default:
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
errors.LogDebug(context.Background(), "xdns closed")
|
|
||||||
|
|
||||||
close(c.pollChan)
|
|
||||||
close(c.readQueue)
|
|
||||||
|
|
||||||
c.mutex.Lock()
|
|
||||||
defer c.mutex.Unlock()
|
|
||||||
|
|
||||||
c.closed = true
|
|
||||||
close(c.writeQueue)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *xdnsConnClient) sendLoop() {
|
func (c *xdnsClient) send() {
|
||||||
pollDelay := initPollDelay
|
defer c.wg.Done()
|
||||||
pollTimer := time.NewTimer(pollDelay)
|
|
||||||
for {
|
|
||||||
var p *packet
|
|
||||||
pollTimerExpired := false
|
|
||||||
|
|
||||||
select {
|
var buf [512]byte
|
||||||
case p = <-c.writeQueue:
|
var data [255]byte
|
||||||
default:
|
|
||||||
select {
|
sendMsg := func(p []byte, domain *Domain, qtype uint16) {
|
||||||
case p = <-c.writeQueue:
|
msg := dnsmessage.Message{
|
||||||
case <-c.pollChan:
|
Header: dnsmessage.Header{
|
||||||
case <-pollTimer.C:
|
RecursionDesired: true,
|
||||||
pollTimerExpired = true
|
},
|
||||||
|
Questions: []dnsmessage.Question{
|
||||||
|
{
|
||||||
|
Name: domain.Encode(p),
|
||||||
|
Type: dnsmessage.Type(qtype),
|
||||||
|
Class: dnsmessage.ClassINET,
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
if domain.edns0 > 0 {
|
||||||
|
msg.Additionals = []dnsmessage.Resource{
|
||||||
|
{
|
||||||
|
Header: dnsmessage.ResourceHeader{
|
||||||
|
Name: dnsmessage.MustNewName("."),
|
||||||
|
Type: dnsmessage.TypeOPT,
|
||||||
|
Class: dnsmessage.Class(domain.edns0),
|
||||||
|
TTL: 0,
|
||||||
|
},
|
||||||
|
Body: &dnsmessage.OPTResource{},
|
||||||
|
},
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
pack := common.Must2(msg.AppendPack(buf[:0]))
|
||||||
|
common.Must2(rand.Read(pack[:2]))
|
||||||
|
|
||||||
if p != nil {
|
index := c.resolverIndex.Load()
|
||||||
select {
|
cur := c.resolverSends[index].Add(1)
|
||||||
case <-c.pollChan:
|
i := index
|
||||||
default:
|
for {
|
||||||
|
i++
|
||||||
|
if i == uint32(len(c.resolvers)) {
|
||||||
|
i = 0
|
||||||
}
|
}
|
||||||
} else {
|
if i == index {
|
||||||
encoded, _ := encode(nil, c.clientID, c.domains[c.resolverIdx], c.resolverTypes[c.resolverIdx])
|
break
|
||||||
p = &packet{
|
}
|
||||||
p: encoded,
|
if cur > c.resolverSends[i].Load() {
|
||||||
|
break
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
c.resolverIndex.Store(i)
|
||||||
|
c.resolvers[index].Send(pack)
|
||||||
|
}
|
||||||
|
|
||||||
if pollTimerExpired {
|
send := func(p []byte) {
|
||||||
pollDelay = time.Duration(float64(pollDelay) * pollDelayMultiplier)
|
domain := c.domains[mrand.Intn(len(c.domains))]
|
||||||
if pollDelay > maxPollDelay {
|
qtype := domain.types[mrand.Intn(len(domain.types))]
|
||||||
pollDelay = maxPollDelay
|
|
||||||
}
|
|
||||||
} else {
|
|
||||||
if !pollTimer.Stop() {
|
|
||||||
<-pollTimer.C
|
|
||||||
}
|
|
||||||
pollDelay = initPollDelay
|
|
||||||
}
|
|
||||||
pollTimer.Reset(pollDelay)
|
|
||||||
|
|
||||||
if c.closed {
|
if len(p) == 0 {
|
||||||
|
copy(data[:], c.clientID[:])
|
||||||
|
data[0] |= TypeMap[qtype]
|
||||||
|
data[8] = 8
|
||||||
|
common.Must2(rand.Read(data[9:17]))
|
||||||
|
sendMsg(data[:17], domain, qtype)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
cur := c.resolverIdx
|
if len(p) <= domain.cap-12 {
|
||||||
curSend := c.resolverSend[c.resolverAddrs[cur].String()].Add(1)
|
copy(data[:], c.clientID[:])
|
||||||
_, _ = c.PacketConn.WriteTo(p.p, c.resolverAddrs[cur])
|
data[0] |= TypeMap[qtype]
|
||||||
for {
|
data[8] = 3
|
||||||
c.resolverIdx += 1
|
common.Must2(rand.Read(data[9:12]))
|
||||||
c.resolverIdx %= uint32(len(c.resolverAddrs))
|
copy(data[12:], p)
|
||||||
if c.resolverIdx == cur {
|
sendMsg(data[:12+len(p)], domain, qtype)
|
||||||
break
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(p) <= 255*(domain.cap-15) {
|
||||||
|
copy(data[:], c.clientID[:])
|
||||||
|
data[0] |= TypeMap[qtype]
|
||||||
|
data[8] = 3 | 0xC0
|
||||||
|
common.Must2(rand.Read(data[9:12]))
|
||||||
|
|
||||||
|
fragID := byte(c.fragID.Add(1))
|
||||||
|
fragN := len(p) / (domain.cap - 15)
|
||||||
|
if len(p)%(domain.cap-15) > 0 {
|
||||||
|
fragN++
|
||||||
}
|
}
|
||||||
if c.resolverSend[c.resolverAddrs[c.resolverIdx].String()].Load() < curSend {
|
|
||||||
break
|
for i := range fragN {
|
||||||
|
data[12] = fragID
|
||||||
|
data[13] = byte(i)
|
||||||
|
data[14] = byte(fragN)
|
||||||
|
size := min(len(p), domain.cap-15)
|
||||||
|
copy(data[15:], p[:size])
|
||||||
|
sendMsg(data[:15+size], domain, qtype)
|
||||||
|
p = p[size:]
|
||||||
|
}
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
errors.LogError(context.Background(), "err size ", len(p))
|
||||||
|
}
|
||||||
|
|
||||||
|
ticker := time.NewTicker(initPollDelay)
|
||||||
|
defer ticker.Stop()
|
||||||
|
delay := initPollDelay
|
||||||
|
p := []byte(nil)
|
||||||
|
timeout := false
|
||||||
|
for {
|
||||||
|
select {
|
||||||
|
case <-c.closeCh:
|
||||||
|
return
|
||||||
|
default:
|
||||||
|
select {
|
||||||
|
case <-c.closeCh:
|
||||||
|
return
|
||||||
|
case p = <-c.sendCh:
|
||||||
|
case <-c.poolCh:
|
||||||
|
case <-ticker.C:
|
||||||
|
timeout = true
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
if len(p) > 0 {
|
||||||
|
select {
|
||||||
|
case <-c.poolCh:
|
||||||
|
default:
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
send(p)
|
||||||
|
for range c.extraPoll {
|
||||||
|
send(nil)
|
||||||
|
}
|
||||||
|
|
||||||
|
if timeout {
|
||||||
|
delay *= pollDelayMultiplier
|
||||||
|
if delay > maxPollDelay {
|
||||||
|
delay = maxPollDelay
|
||||||
|
}
|
||||||
|
timeout = false
|
||||||
|
} else {
|
||||||
|
delay = initPollDelay
|
||||||
|
}
|
||||||
|
ticker.Reset(delay)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *xdnsConnClient) ReadFrom(p []byte) (n int, addr net.Addr, err error) {
|
func (c *xdnsClient) ReadFrom(p []byte) (n int, addr net.Addr, err error) {
|
||||||
packet, ok := <-c.readQueue
|
packet, ok := <-c.readCh
|
||||||
if !ok {
|
if ok {
|
||||||
return 0, nil, net.ErrClosed
|
return copy(p, packet.p), packet.addr, nil
|
||||||
}
|
}
|
||||||
if len(p) < len(packet.p) {
|
return 0, nil, io.ErrClosedPipe
|
||||||
errors.LogDebug(context.Background(), packet.addr, " mask read err short buffer ", len(p), " ", len(packet.p))
|
|
||||||
return 0, packet.addr, nil
|
|
||||||
}
|
|
||||||
copy(p, packet.p)
|
|
||||||
return len(packet.p), packet.addr, nil
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *xdnsConnClient) WriteTo(p []byte, addr net.Addr) (n int, err error) {
|
func (c *xdnsClient) WriteTo(p []byte, addr net.Addr) (n int, err error) {
|
||||||
c.mutex.Lock()
|
c.mu.Lock()
|
||||||
defer c.mutex.Unlock()
|
defer c.mu.Unlock()
|
||||||
|
if c.closed() {
|
||||||
if c.closed {
|
|
||||||
return 0, io.ErrClosedPipe
|
return 0, io.ErrClosedPipe
|
||||||
}
|
}
|
||||||
|
if len(p) == 0 || len(p) > 4096 {
|
||||||
idx := c.resolverIdx % uint32(len(c.resolverAddrs))
|
errors.LogError(context.Background(), "err size ", len(p))
|
||||||
encoded, err := encode(p, c.clientID, c.domains[idx], c.resolverTypes[idx])
|
return 0, errors.New("err size")
|
||||||
if err != nil {
|
|
||||||
errors.LogDebug(context.Background(), addr, " xdns wireformat err ", err, " ", len(p))
|
|
||||||
return 0, nil
|
|
||||||
}
|
}
|
||||||
|
b := make([]byte, len(p))
|
||||||
|
copy(b, p)
|
||||||
select {
|
select {
|
||||||
case c.writeQueue <- &packet{
|
case c.sendCh <- b:
|
||||||
p: encoded,
|
|
||||||
addr: addr,
|
|
||||||
}:
|
|
||||||
return len(p), nil
|
|
||||||
default:
|
default:
|
||||||
errors.LogDebug(context.Background(), addr, " mask write err queue full")
|
|
||||||
return 0, nil
|
|
||||||
}
|
}
|
||||||
|
return len(p), nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *xdnsConnClient) Close() error {
|
func (c *xdnsClient) Close() error {
|
||||||
c.closed = true
|
c.mu.Lock()
|
||||||
return c.PacketConn.Close()
|
defer c.mu.Unlock()
|
||||||
}
|
if c.closed() {
|
||||||
|
|
||||||
func encode(p []byte, clientID []byte, domain Name, qtype uint16) ([]byte, error) {
|
|
||||||
var decoded []byte
|
|
||||||
{
|
|
||||||
if len(p) >= 224 {
|
|
||||||
return nil, errors.New("too long")
|
|
||||||
}
|
|
||||||
var buf bytes.Buffer
|
|
||||||
buf.Write(clientID[:])
|
|
||||||
n := numPadding
|
|
||||||
if len(p) == 0 {
|
|
||||||
n = numPaddingForPoll
|
|
||||||
}
|
|
||||||
buf.WriteByte(byte(224 + n))
|
|
||||||
_, _ = io.CopyN(&buf, rand.Reader, int64(n))
|
|
||||||
if len(p) > 0 {
|
|
||||||
buf.WriteByte(byte(len(p)))
|
|
||||||
buf.Write(p)
|
|
||||||
}
|
|
||||||
decoded = buf.Bytes()
|
|
||||||
}
|
|
||||||
|
|
||||||
encoded := make([]byte, base32Encoding.EncodedLen(len(decoded)))
|
|
||||||
base32Encoding.Encode(encoded, decoded)
|
|
||||||
encoded = bytes.ToLower(encoded)
|
|
||||||
labels := chunks(encoded, 63)
|
|
||||||
labels = append(labels, domain...)
|
|
||||||
name, err := NewName(labels)
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
var id uint16
|
|
||||||
_ = binary.Read(rand.Reader, binary.BigEndian, &id)
|
|
||||||
query := &Message{
|
|
||||||
ID: id,
|
|
||||||
Flags: 0x0100,
|
|
||||||
Question: []Question{
|
|
||||||
{
|
|
||||||
Name: name,
|
|
||||||
Type: qtype,
|
|
||||||
Class: ClassIN,
|
|
||||||
},
|
|
||||||
},
|
|
||||||
Additional: []RR{
|
|
||||||
{
|
|
||||||
Name: Name{},
|
|
||||||
Type: RRTypeOPT,
|
|
||||||
Class: 4096,
|
|
||||||
TTL: 0,
|
|
||||||
Data: []byte{},
|
|
||||||
},
|
|
||||||
},
|
|
||||||
}
|
|
||||||
|
|
||||||
buf, err := query.WireFormat()
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
return buf, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func chunks(p []byte, n int) [][]byte {
|
|
||||||
var result [][]byte
|
|
||||||
for len(p) > 0 {
|
|
||||||
sz := len(p)
|
|
||||||
if sz > n {
|
|
||||||
sz = n
|
|
||||||
}
|
|
||||||
result = append(result, p[:sz])
|
|
||||||
p = p[sz:]
|
|
||||||
}
|
|
||||||
return result
|
|
||||||
}
|
|
||||||
|
|
||||||
func nextPacket(r *bytes.Reader) ([]byte, error) {
|
|
||||||
var n uint16
|
|
||||||
err := binary.Read(r, binary.BigEndian, &n)
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
p := make([]byte, n)
|
|
||||||
_, err = io.ReadFull(r, p)
|
|
||||||
if err == io.EOF {
|
|
||||||
err = io.ErrUnexpectedEOF
|
|
||||||
}
|
|
||||||
return p, err
|
|
||||||
}
|
|
||||||
|
|
||||||
func dnsResponsePayload(resp *Message, domains []Name) []byte {
|
|
||||||
if resp.Flags&0x8000 != 0x8000 {
|
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
if resp.Flags&0x000f != RcodeNoError {
|
close(c.closeCh)
|
||||||
return nil
|
for i := range c.resolvers {
|
||||||
|
c.resolvers[i].Close()
|
||||||
}
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
if len(resp.Answer) == 0 {
|
func (c *xdnsClient) LocalAddr() net.Addr { return &net.UDPAddr{IP: []byte{0, 0, 0, 0}} }
|
||||||
return nil
|
|
||||||
}
|
func (c *xdnsClient) SetDeadline(t time.Time) error { return errors.New("not support") }
|
||||||
|
|
||||||
for _, answer := range resp.Answer {
|
func (c *xdnsClient) SetReadDeadline(t time.Time) error { return errors.New("not support") }
|
||||||
var ok bool
|
|
||||||
for _, domain := range domains {
|
func (c *xdnsClient) SetWriteDeadline(t time.Time) error { return errors.New("not support") }
|
||||||
_, ok = answer.Name.TrimSuffix(domain)
|
|
||||||
if ok {
|
type ClientID [8]byte
|
||||||
break
|
|
||||||
}
|
func NewClientID() ClientID {
|
||||||
}
|
var id ClientID
|
||||||
if !ok {
|
common.Must2(rand.Read(id[:]))
|
||||||
return nil
|
id[0] &= 0xFC
|
||||||
}
|
return id
|
||||||
}
|
}
|
||||||
|
|
||||||
return decodeResponsePayload(resp.Answer)
|
func ClientIDFromRaw(id [8]byte) ClientID {
|
||||||
|
id[0] &= 0xFC
|
||||||
|
return id
|
||||||
|
}
|
||||||
|
|
||||||
|
func ClientIDFromAddr(addr *net.UDPAddr) ClientID {
|
||||||
|
return ClientID(addr.IP[8:])
|
||||||
|
}
|
||||||
|
|
||||||
|
func (id ClientID) Addr() *net.UDPAddr {
|
||||||
|
var ip [16]byte
|
||||||
|
ip[0] = 0xFD
|
||||||
|
copy(ip[8:], id[:])
|
||||||
|
return &net.UDPAddr{IP: ip[:]}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -6,9 +6,9 @@ import (
|
|||||||
)
|
)
|
||||||
|
|
||||||
func (c *Config) WrapPacketConnClient(conn net.PacketConn, dest *net.Destination, dialer *finalmask.Dialer) (net.PacketConn, error) {
|
func (c *Config) WrapPacketConnClient(conn net.PacketConn, dest *net.Destination, dialer *finalmask.Dialer) (net.PacketConn, error) {
|
||||||
return NewConnClient(c, conn)
|
return NewClient(c, dialer)
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *Config) WrapPacketConnServer(conn net.PacketConn, addr net.Addr, lc *finalmask.ListenConfig) (net.PacketConn, error) {
|
func (c *Config) WrapPacketConnServer(conn net.PacketConn, addr net.Addr, lc *finalmask.ListenConfig) (net.PacketConn, error) {
|
||||||
return NewConnServer(c, conn)
|
return NewServer(c, conn)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -7,6 +7,7 @@
|
|||||||
package xdns
|
package xdns
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
serial "github.com/xtls/xray-core/common/serial"
|
||||||
protoreflect "google.golang.org/protobuf/reflect/protoreflect"
|
protoreflect "google.golang.org/protobuf/reflect/protoreflect"
|
||||||
protoimpl "google.golang.org/protobuf/runtime/protoimpl"
|
protoimpl "google.golang.org/protobuf/runtime/protoimpl"
|
||||||
reflect "reflect"
|
reflect "reflect"
|
||||||
@@ -21,17 +22,94 @@ const (
|
|||||||
_ = protoimpl.EnforceVersion(protoimpl.MaxVersion - 20)
|
_ = protoimpl.EnforceVersion(protoimpl.MaxVersion - 20)
|
||||||
)
|
)
|
||||||
|
|
||||||
|
type DomainProto struct {
|
||||||
|
state protoimpl.MessageState `protogen:"open.v1"`
|
||||||
|
Name string `protobuf:"bytes,1,opt,name=name,proto3" json:"name,omitempty"`
|
||||||
|
LenLimit int32 `protobuf:"varint,2,opt,name=len_limit,json=lenLimit,proto3" json:"len_limit,omitempty"`
|
||||||
|
LabelLimit int32 `protobuf:"varint,3,opt,name=label_limit,json=labelLimit,proto3" json:"label_limit,omitempty"`
|
||||||
|
Types []int32 `protobuf:"varint,4,rep,packed,name=types,proto3" json:"types,omitempty"`
|
||||||
|
Edns0 int32 `protobuf:"varint,5,opt,name=edns0,proto3" json:"edns0,omitempty"`
|
||||||
|
unknownFields protoimpl.UnknownFields
|
||||||
|
sizeCache protoimpl.SizeCache
|
||||||
|
}
|
||||||
|
|
||||||
|
func (x *DomainProto) Reset() {
|
||||||
|
*x = DomainProto{}
|
||||||
|
mi := &file_transport_internet_finalmask_xdns_config_proto_msgTypes[0]
|
||||||
|
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
|
||||||
|
ms.StoreMessageInfo(mi)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (x *DomainProto) String() string {
|
||||||
|
return protoimpl.X.MessageStringOf(x)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (*DomainProto) ProtoMessage() {}
|
||||||
|
|
||||||
|
func (x *DomainProto) ProtoReflect() protoreflect.Message {
|
||||||
|
mi := &file_transport_internet_finalmask_xdns_config_proto_msgTypes[0]
|
||||||
|
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 DomainProto.ProtoReflect.Descriptor instead.
|
||||||
|
func (*DomainProto) Descriptor() ([]byte, []int) {
|
||||||
|
return file_transport_internet_finalmask_xdns_config_proto_rawDescGZIP(), []int{0}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (x *DomainProto) GetName() string {
|
||||||
|
if x != nil {
|
||||||
|
return x.Name
|
||||||
|
}
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
|
||||||
|
func (x *DomainProto) GetLenLimit() int32 {
|
||||||
|
if x != nil {
|
||||||
|
return x.LenLimit
|
||||||
|
}
|
||||||
|
return 0
|
||||||
|
}
|
||||||
|
|
||||||
|
func (x *DomainProto) GetLabelLimit() int32 {
|
||||||
|
if x != nil {
|
||||||
|
return x.LabelLimit
|
||||||
|
}
|
||||||
|
return 0
|
||||||
|
}
|
||||||
|
|
||||||
|
func (x *DomainProto) GetTypes() []int32 {
|
||||||
|
if x != nil {
|
||||||
|
return x.Types
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (x *DomainProto) GetEdns0() int32 {
|
||||||
|
if x != nil {
|
||||||
|
return x.Edns0
|
||||||
|
}
|
||||||
|
return 0
|
||||||
|
}
|
||||||
|
|
||||||
type Config struct {
|
type Config struct {
|
||||||
state protoimpl.MessageState `protogen:"open.v1"`
|
state protoimpl.MessageState `protogen:"open.v1"`
|
||||||
Domains []string `protobuf:"bytes,1,rep,name=domains,proto3" json:"domains,omitempty"`
|
Domains []*DomainProto `protobuf:"bytes,1,rep,name=domains,proto3" json:"domains,omitempty"`
|
||||||
Resolvers []string `protobuf:"bytes,2,rep,name=resolvers,proto3" json:"resolvers,omitempty"`
|
Resolvers []*serial.TypedMessage `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
|
unknownFields protoimpl.UnknownFields
|
||||||
sizeCache protoimpl.SizeCache
|
sizeCache protoimpl.SizeCache
|
||||||
}
|
}
|
||||||
|
|
||||||
func (x *Config) Reset() {
|
func (x *Config) Reset() {
|
||||||
*x = Config{}
|
*x = Config{}
|
||||||
mi := &file_transport_internet_finalmask_xdns_config_proto_msgTypes[0]
|
mi := &file_transport_internet_finalmask_xdns_config_proto_msgTypes[1]
|
||||||
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
|
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
|
||||||
ms.StoreMessageInfo(mi)
|
ms.StoreMessageInfo(mi)
|
||||||
}
|
}
|
||||||
@@ -43,7 +121,7 @@ func (x *Config) String() string {
|
|||||||
func (*Config) ProtoMessage() {}
|
func (*Config) ProtoMessage() {}
|
||||||
|
|
||||||
func (x *Config) ProtoReflect() protoreflect.Message {
|
func (x *Config) ProtoReflect() protoreflect.Message {
|
||||||
mi := &file_transport_internet_finalmask_xdns_config_proto_msgTypes[0]
|
mi := &file_transport_internet_finalmask_xdns_config_proto_msgTypes[1]
|
||||||
if x != nil {
|
if x != nil {
|
||||||
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
|
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
|
||||||
if ms.LoadMessageInfo() == nil {
|
if ms.LoadMessageInfo() == nil {
|
||||||
@@ -56,31 +134,139 @@ func (x *Config) ProtoReflect() protoreflect.Message {
|
|||||||
|
|
||||||
// Deprecated: Use Config.ProtoReflect.Descriptor instead.
|
// Deprecated: Use Config.ProtoReflect.Descriptor instead.
|
||||||
func (*Config) Descriptor() ([]byte, []int) {
|
func (*Config) Descriptor() ([]byte, []int) {
|
||||||
return file_transport_internet_finalmask_xdns_config_proto_rawDescGZIP(), []int{0}
|
return file_transport_internet_finalmask_xdns_config_proto_rawDescGZIP(), []int{1}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (x *Config) GetDomains() []string {
|
func (x *Config) GetDomains() []*DomainProto {
|
||||||
if x != nil {
|
if x != nil {
|
||||||
return x.Domains
|
return x.Domains
|
||||||
}
|
}
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (x *Config) GetResolvers() []string {
|
func (x *Config) GetResolvers() []*serial.TypedMessage {
|
||||||
if x != nil {
|
if x != nil {
|
||||||
return x.Resolvers
|
return x.Resolvers
|
||||||
}
|
}
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (x *Config) GetExtraPoll() int32 {
|
||||||
|
if x != nil {
|
||||||
|
return x.ExtraPoll
|
||||||
|
}
|
||||||
|
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
|
var File_transport_internet_finalmask_xdns_config_proto protoreflect.FileDescriptor
|
||||||
|
|
||||||
const file_transport_internet_finalmask_xdns_config_proto_rawDesc = "" +
|
const file_transport_internet_finalmask_xdns_config_proto_rawDesc = "" +
|
||||||
"\n" +
|
"\n" +
|
||||||
".transport/internet/finalmask/xdns/config.proto\x12&xray.transport.internet.finalmask.xdns\"@\n" +
|
".transport/internet/finalmask/xdns/config.proto\x12&xray.transport.internet.finalmask.xdns\x1a!common/serial/typed_message.proto\"\x8b\x01\n" +
|
||||||
"\x06Config\x12\x18\n" +
|
"\vDomainProto\x12\x12\n" +
|
||||||
"\adomains\x18\x01 \x03(\tR\adomains\x12\x1c\n" +
|
"\x04name\x18\x01 \x01(\tR\x04name\x12\x1b\n" +
|
||||||
"\tresolvers\x18\x02 \x03(\tR\tresolversB\x94\x01\n" +
|
"\tlen_limit\x18\x02 \x01(\x05R\blenLimit\x12\x1f\n" +
|
||||||
|
"\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" +
|
||||||
|
"\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" +
|
||||||
|
"\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" +
|
||||||
"*com.xray.transport.internet.finalmask.xdnsP\x01Z;github.com/xtls/xray-core/transport/internet/finalmask/xdns\xaa\x02&Xray.Transport.Internet.Finalmask.Xdnsb\x06proto3"
|
"*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 (
|
var (
|
||||||
@@ -95,16 +281,22 @@ func file_transport_internet_finalmask_xdns_config_proto_rawDescGZIP() []byte {
|
|||||||
return file_transport_internet_finalmask_xdns_config_proto_rawDescData
|
return file_transport_internet_finalmask_xdns_config_proto_rawDescData
|
||||||
}
|
}
|
||||||
|
|
||||||
var file_transport_internet_finalmask_xdns_config_proto_msgTypes = make([]protoimpl.MessageInfo, 1)
|
var file_transport_internet_finalmask_xdns_config_proto_msgTypes = make([]protoimpl.MessageInfo, 4)
|
||||||
var file_transport_internet_finalmask_xdns_config_proto_goTypes = []any{
|
var file_transport_internet_finalmask_xdns_config_proto_goTypes = []any{
|
||||||
(*Config)(nil), // 0: xray.transport.internet.finalmask.xdns.Config
|
(*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
|
||||||
}
|
}
|
||||||
var file_transport_internet_finalmask_xdns_config_proto_depIdxs = []int32{
|
var file_transport_internet_finalmask_xdns_config_proto_depIdxs = []int32{
|
||||||
0, // [0:0] is the sub-list for method output_type
|
0, // 0: xray.transport.internet.finalmask.xdns.Config.domains:type_name -> xray.transport.internet.finalmask.xdns.DomainProto
|
||||||
0, // [0:0] is the sub-list for method input_type
|
4, // 1: xray.transport.internet.finalmask.xdns.Config.resolvers:type_name -> xray.common.serial.TypedMessage
|
||||||
0, // [0:0] is the sub-list for extension type_name
|
2, // [2:2] is the sub-list for method output_type
|
||||||
0, // [0:0] is the sub-list for extension extendee
|
2, // [2:2] is the sub-list for method input_type
|
||||||
0, // [0:0] is the sub-list for field type_name
|
2, // [2:2] is the sub-list for extension type_name
|
||||||
|
2, // [2:2] is the sub-list for extension extendee
|
||||||
|
0, // [0:2] is the sub-list for field type_name
|
||||||
}
|
}
|
||||||
|
|
||||||
func init() { file_transport_internet_finalmask_xdns_config_proto_init() }
|
func init() { file_transport_internet_finalmask_xdns_config_proto_init() }
|
||||||
@@ -118,7 +310,7 @@ func file_transport_internet_finalmask_xdns_config_proto_init() {
|
|||||||
GoPackagePath: reflect.TypeOf(x{}).PkgPath(),
|
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)),
|
RawDescriptor: unsafe.Slice(unsafe.StringData(file_transport_internet_finalmask_xdns_config_proto_rawDesc), len(file_transport_internet_finalmask_xdns_config_proto_rawDesc)),
|
||||||
NumEnums: 0,
|
NumEnums: 0,
|
||||||
NumMessages: 1,
|
NumMessages: 4,
|
||||||
NumExtensions: 0,
|
NumExtensions: 0,
|
||||||
NumServices: 0,
|
NumServices: 0,
|
||||||
},
|
},
|
||||||
|
|||||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user