mirror of
https://github.com/XTLS/Xray-core.git
synced 2026-10-04 21:08:11 +03:00
Compare commits
79
Commits
http-sniffer
..
lua
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
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 | ||
|
|
7780db9bbe | ||
|
|
2953d44734 | ||
|
|
7b8ade3ec5 | ||
|
|
5dda894e29 | ||
|
|
47a2c2ffdc | ||
|
|
7a018833ec | ||
|
|
65e853ed84 | ||
|
|
459301d42e | ||
|
|
3982028a9c | ||
|
|
70b8e9a61d | ||
|
|
219f758060 | ||
|
|
3519dfecbd | ||
|
|
df261e4479 | ||
|
|
9628003594 | ||
|
|
72d9ab50b9 | ||
|
|
235843c5d2 | ||
|
|
61cad5ec8b | ||
|
|
a642a190ed | ||
|
|
60e2a0c502 | ||
|
|
7d3e44fee2 | ||
|
|
9927942aaa | ||
|
|
a308ded2e6 | ||
|
|
7741e9e77e | ||
|
|
d562d8947d | ||
|
|
dbb1ea30ba | ||
|
|
efc9e6da62 | ||
|
|
8267cf953a | ||
|
|
24e6f6d551 | ||
|
|
dcdfc57ccd | ||
|
|
3461c511aa | ||
|
|
c412e77a9b | ||
|
|
ccb69ea5e2 | ||
|
|
52a412d9e2 | ||
|
|
18a1b5042a | ||
|
|
c26d2eda24 | ||
|
|
a1bf968be9 | ||
|
|
c037ccd98d | ||
|
|
37ceb8b4b6 | ||
|
|
fd2ca74822 | ||
|
|
47cfe9994a | ||
|
|
3e2f040cd8 | ||
|
|
c7245c0336 | ||
|
|
eef6e63bc1 | ||
|
|
6ce8dc53e7 | ||
|
|
01a034be53 | ||
|
|
de2caf3cef | ||
|
|
cecc88f43c | ||
|
|
cd4ce973e9 | ||
|
|
fc7b980636 | ||
|
|
8ee131cbbb | ||
|
|
2776ea6d74 |
@@ -67,9 +67,7 @@ jobs:
|
|||||||
check-latest: true
|
check-latest: true
|
||||||
cache: false
|
cache: false
|
||||||
- name: Check Format
|
- name: Check Format
|
||||||
run: |
|
run: go run ./infra/vformat/main.go -mode check -pwd ./
|
||||||
go install -v mvdan.cc/gofumpt@latest
|
|
||||||
go run ./infra/vformat/main.go -mode check -pwd ./
|
|
||||||
|
|
||||||
test:
|
test:
|
||||||
needs: check-assets
|
needs: check-assets
|
||||||
|
|||||||
@@ -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)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -23,7 +23,7 @@ func newFakeDNSSniffer(ctx context.Context) (protocolSnifferWithMetadata, error)
|
|||||||
}
|
}
|
||||||
|
|
||||||
if fakeDNSEngine == nil {
|
if fakeDNSEngine == nil {
|
||||||
errNotInit := errors.New("FakeDNSEngine is not initialized, but such a sniffer is used").AtError()
|
errNotInit := errors.New("FakeDNSEngine is not initialized, but such a sniffer is used")
|
||||||
return protocolSnifferWithMetadata{}, errNotInit
|
return protocolSnifferWithMetadata{}, errNotInit
|
||||||
}
|
}
|
||||||
return protocolSnifferWithMetadata{protocolSniffer: func(ctx context.Context, bytes []byte) (SniffResult, error) {
|
return protocolSnifferWithMetadata{protocolSniffer: func(ctx context.Context, bytes []byte) (SniffResult, error) {
|
||||||
|
|||||||
+1
-1
@@ -28,7 +28,7 @@ func toNetIP(addrs []net.Address) ([]net.IP, error) {
|
|||||||
if addr.Family().IsIP() {
|
if addr.Family().IsIP() {
|
||||||
ips = append(ips, addr.IP())
|
ips = append(ips, addr.IP())
|
||||||
} else {
|
} else {
|
||||||
return nil, errors.New("Failed to convert address", addr, "to Net IP.").AtWarning()
|
return nil, errors.New("Failed to convert address", addr, "to Net IP.")
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
return ips, nil
|
return ips, nil
|
||||||
|
|||||||
+23
-4
@@ -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,6 +266,8 @@ type Config struct {
|
|||||||
DisableFallback bool `protobuf:"varint,10,opt,name=disableFallback,proto3" json:"disableFallback,omitempty"`
|
DisableFallback bool `protobuf:"varint,10,opt,name=disableFallback,proto3" json:"disableFallback,omitempty"`
|
||||||
DisableFallbackIfMatch bool `protobuf:"varint,11,opt,name=disableFallbackIfMatch,proto3" json:"disableFallbackIfMatch,omitempty"`
|
DisableFallbackIfMatch bool `protobuf:"varint,11,opt,name=disableFallbackIfMatch,proto3" json:"disableFallbackIfMatch,omitempty"`
|
||||||
EnableParallelQuery bool `protobuf:"varint,14,opt,name=enableParallelQuery,proto3" json:"enableParallelQuery,omitempty"`
|
EnableParallelQuery bool `protobuf:"varint,14,opt,name=enableParallelQuery,proto3" json:"enableParallelQuery,omitempty"`
|
||||||
|
// Absolute path to the Lua DNS query script.
|
||||||
|
Script string `protobuf:"bytes,15,opt,name=script,proto3" json:"script,omitempty"`
|
||||||
unknownFields protoimpl.UnknownFields
|
unknownFields protoimpl.UnknownFields
|
||||||
sizeCache protoimpl.SizeCache
|
sizeCache protoimpl.SizeCache
|
||||||
}
|
}
|
||||||
@@ -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
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -212,6 +225,28 @@ func (s *DNS) IsOwnLink(ctx context.Context) bool {
|
|||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// MayUseSystemResolver reports whether any name server configured here could
|
||||||
|
// still resolve through the system resolver. That is what happens when no name
|
||||||
|
// server is configured at all, and it is also what a name server pointed at
|
||||||
|
// "localhost" does. Callers that are about to redirect the system resolver need
|
||||||
|
// to know, because a resolution path that reaches it would then loop back to
|
||||||
|
// them.
|
||||||
|
//
|
||||||
|
// Any such server is enough: name servers can be selected per domain, so a
|
||||||
|
// single local one makes some query reach the system resolver even when
|
||||||
|
// independent upstreams are configured alongside it.
|
||||||
|
func (s *DNS) MayUseSystemResolver() bool {
|
||||||
|
if len(s.clients) == 0 {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
for _, client := range s.clients {
|
||||||
|
if _, isLocal := client.server.(*LocalNameServer); isLocal {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
// LookupIP implements dns.Client.
|
// LookupIP implements dns.Client.
|
||||||
func (s *DNS) LookupIP(domain string, option dns.IPOption) ([]net.IP, uint32, error) {
|
func (s *DNS) LookupIP(domain string, option dns.IPOption) ([]net.IP, uint32, error) {
|
||||||
// Normalize the FQDN form query
|
// Normalize the FQDN form query
|
||||||
@@ -257,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 {
|
||||||
|
|||||||
@@ -0,0 +1,59 @@
|
|||||||
|
package dns
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/xtls/xray-core/common/net"
|
||||||
|
feature_dns "github.com/xtls/xray-core/features/dns"
|
||||||
|
)
|
||||||
|
|
||||||
|
// fakeServer stands in for any name server that is not the system resolver.
|
||||||
|
type fakeServer struct{}
|
||||||
|
|
||||||
|
func (fakeServer) Name() string { return "fake" }
|
||||||
|
func (fakeServer) IsDisableCache() bool { return false }
|
||||||
|
func (fakeServer) QueryIP(context.Context, string, feature_dns.IPOption) ([]net.IP, uint32, error) {
|
||||||
|
return nil, 0, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Callers that are about to redirect the system resolver rely on this to tell
|
||||||
|
// whether any resolution path could still reach the system resolver, so the
|
||||||
|
// mixed shape has to be reported as reachable: a domain-specific rule can
|
||||||
|
// select the system resolver even when an independent upstream also exists.
|
||||||
|
func TestMayUseSystemResolver(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
clients []*Client
|
||||||
|
want bool
|
||||||
|
}{
|
||||||
|
{
|
||||||
|
name: "no clients at all",
|
||||||
|
want: true,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "only the system resolver",
|
||||||
|
clients: []*Client{{server: NewLocalNameServer()}},
|
||||||
|
want: true,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "the system resolver alongside an independent name server",
|
||||||
|
clients: []*Client{{server: fakeServer{}}, {server: NewLocalNameServer()}},
|
||||||
|
want: true,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "only independent name servers",
|
||||||
|
clients: []*Client{{server: fakeServer{}}},
|
||||||
|
want: false,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
server := &DNS{clients: tt.clients}
|
||||||
|
if got := server.MayUseSystemResolver(); got != tt.want {
|
||||||
|
t.Errorf("MayUseSystemResolver() = %v, want %v", got, tt.want)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -188,10 +188,10 @@ func parseResponse(payload []byte) (*IPRecord, error) {
|
|||||||
var parser dnsmessage.Parser
|
var parser dnsmessage.Parser
|
||||||
h, err := parser.Start(payload)
|
h, err := parser.Start(payload)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, errors.New("failed to parse DNS response").Base(err).AtWarning()
|
return nil, errors.New("failed to parse DNS response").Base(err)
|
||||||
}
|
}
|
||||||
if err := parser.SkipAllQuestions(); err != nil {
|
if err := parser.SkipAllQuestions(); err != nil {
|
||||||
return nil, errors.New("failed to skip questions in DNS response").Base(err).AtWarning()
|
return nil, errors.New("failed to skip questions in DNS response").Base(err)
|
||||||
}
|
}
|
||||||
|
|
||||||
now := time.Now()
|
now := time.Now()
|
||||||
|
|||||||
@@ -58,7 +58,7 @@ func NewFakeDNSHolder() (*Holder, error) {
|
|||||||
var err error
|
var err error
|
||||||
|
|
||||||
if fkdns, err = NewFakeDNSHolderConfigOnly(nil); err != nil {
|
if fkdns, err = NewFakeDNSHolderConfigOnly(nil); err != nil {
|
||||||
return nil, errors.New("Unable to create Fake Dns Engine").Base(err).AtError()
|
return nil, errors.New("Unable to create Fake Dns Engine").Base(err)
|
||||||
}
|
}
|
||||||
err = fkdns.initialize(dns.FakeIPv4Pool, 65535)
|
err = fkdns.initialize(dns.FakeIPv4Pool, 65535)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -80,13 +80,13 @@ func (fkdns *Holder) initialize(ipPoolCidr string, lruSize int) error {
|
|||||||
var err error
|
var err error
|
||||||
|
|
||||||
if _, ipRange, err = net.ParseCIDR(ipPoolCidr); err != nil {
|
if _, ipRange, err = net.ParseCIDR(ipPoolCidr); err != nil {
|
||||||
return errors.New("Unable to parse CIDR for Fake DNS IP assignment").Base(err).AtError()
|
return errors.New("Unable to parse CIDR for Fake DNS IP assignment").Base(err)
|
||||||
}
|
}
|
||||||
|
|
||||||
ones, bits := ipRange.Mask.Size()
|
ones, bits := ipRange.Mask.Size()
|
||||||
rooms := bits - ones
|
rooms := bits - ones
|
||||||
if math.Log2(float64(lruSize)) >= float64(rooms) {
|
if math.Log2(float64(lruSize)) >= float64(rooms) {
|
||||||
return errors.New("LRU size is bigger than subnet size").AtError()
|
return errors.New("LRU size is bigger than subnet size")
|
||||||
}
|
}
|
||||||
fkdns.domainToIP = cache.NewLru(lruSize)
|
fkdns.domainToIP = cache.NewLru(lruSize)
|
||||||
fkdns.ipRange = ipRange
|
fkdns.ipRange = ipRange
|
||||||
|
|||||||
+165
@@ -0,0 +1,165 @@
|
|||||||
|
package dns
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"strings"
|
||||||
|
|
||||||
|
"github.com/xtls/xray-core/common/errors"
|
||||||
|
xlua "github.com/xtls/xray-core/common/lua"
|
||||||
|
"github.com/xtls/xray-core/common/net"
|
||||||
|
featureDNS "github.com/xtls/xray-core/features/dns"
|
||||||
|
"github.com/xtls/xray-core/features/dns/localdns"
|
||||||
|
lua "github.com/yuin/gopher-lua"
|
||||||
|
)
|
||||||
|
|
||||||
|
// luaDNSServer adapts configured and local DNS to the same Lua API.
|
||||||
|
type luaDNSServer struct {
|
||||||
|
id string
|
||||||
|
name string
|
||||||
|
query func(context.Context, string, featureDNS.IPOption) ([]net.IP, uint32, error)
|
||||||
|
}
|
||||||
|
|
||||||
|
// RegisterLua makes xray.dns available to scripts backed by client.
|
||||||
|
func RegisterLua(L *lua.LState, client featureDNS.Client) {
|
||||||
|
var servers []luaDNSServer
|
||||||
|
switch client := client.(type) {
|
||||||
|
case *DNS:
|
||||||
|
servers = luaServers(client)
|
||||||
|
case *localdns.Client:
|
||||||
|
servers = []luaDNSServer{{
|
||||||
|
id: "localhost",
|
||||||
|
name: "localhost",
|
||||||
|
query: func(_ context.Context, domain string, option featureDNS.IPOption) ([]net.IP, uint32, error) {
|
||||||
|
return client.LookupIP(domain, option)
|
||||||
|
},
|
||||||
|
}}
|
||||||
|
}
|
||||||
|
registerLua(L, servers, client)
|
||||||
|
}
|
||||||
|
|
||||||
|
// registerLua makes xray.dns available to DNS scripts.
|
||||||
|
func (s *DNS) registerLua(L *lua.LState) {
|
||||||
|
registerLua(L, luaServers(s), nil)
|
||||||
|
}
|
||||||
|
|
||||||
|
func luaServers(s *DNS) []luaDNSServer {
|
||||||
|
servers := make([]luaDNSServer, len(s.clients))
|
||||||
|
for i, client := range s.clients {
|
||||||
|
servers[i] = luaDNSServer{id: client.id, name: client.Name(), query: client.QueryIP}
|
||||||
|
}
|
||||||
|
return servers
|
||||||
|
}
|
||||||
|
|
||||||
|
func registerLua(L *lua.LState, servers []luaDNSServer, client featureDNS.Client) {
|
||||||
|
L.PreloadModule("xray.dns", func(L *lua.LState) int {
|
||||||
|
serverList := L.NewTable()
|
||||||
|
for i, client := range servers {
|
||||||
|
server := L.NewTable()
|
||||||
|
|
||||||
|
server.RawSetString("ID", lua.LString(client.id))
|
||||||
|
|
||||||
|
server.RawSetString("Query", L.NewFunction(func(L *lua.LState) int {
|
||||||
|
domain, ok := L.Get(2).(lua.LString)
|
||||||
|
if !ok {
|
||||||
|
L.RaiseError("server:Query requires a domain")
|
||||||
|
return 0
|
||||||
|
}
|
||||||
|
option := featureDNS.IPOption{
|
||||||
|
IPv4Enable: L.CheckBool(3),
|
||||||
|
IPv6Enable: L.CheckBool(4),
|
||||||
|
FakeEnable: L.CheckBool(5),
|
||||||
|
}
|
||||||
|
ctx := L.Context()
|
||||||
|
if ctx == nil {
|
||||||
|
L.RaiseError("server:Query requires an active DNS query")
|
||||||
|
return 0
|
||||||
|
}
|
||||||
|
var ips []net.IP
|
||||||
|
var ttl uint32
|
||||||
|
var err error
|
||||||
|
if !option.FakeEnable && strings.EqualFold(client.name, "FakeDNS") {
|
||||||
|
err = featureDNS.ErrEmptyResponse
|
||||||
|
} else {
|
||||||
|
ips, ttl, err = client.query(ctx, string(domain), option)
|
||||||
|
}
|
||||||
|
xlua.PushUserData(L, ips)
|
||||||
|
xlua.PushNumber(L, ttl)
|
||||||
|
xlua.PushError(L, err)
|
||||||
|
return 3
|
||||||
|
}))
|
||||||
|
serverList.RawSetInt(i+1, server)
|
||||||
|
}
|
||||||
|
|
||||||
|
module := L.NewTable()
|
||||||
|
if servers != nil {
|
||||||
|
module.RawSetString("Servers", serverList)
|
||||||
|
}
|
||||||
|
if client != nil {
|
||||||
|
module.RawSetString("Query", newLuaClientQuery(L, client))
|
||||||
|
}
|
||||||
|
L.Push(module)
|
||||||
|
return 1
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func newLuaClientQuery(L *lua.LState, client featureDNS.Client) *lua.LFunction {
|
||||||
|
return L.NewFunction(func(L *lua.LState) int {
|
||||||
|
domain, ok := L.Get(1).(lua.LString)
|
||||||
|
if !ok {
|
||||||
|
L.RaiseError("dns.Query requires a domain")
|
||||||
|
return 0
|
||||||
|
}
|
||||||
|
option := featureDNS.IPOption{
|
||||||
|
IPv4Enable: L.CheckBool(2),
|
||||||
|
IPv6Enable: L.CheckBool(3),
|
||||||
|
FakeEnable: L.CheckBool(4),
|
||||||
|
}
|
||||||
|
if L.Context() == nil {
|
||||||
|
L.RaiseError("dns.Query requires an active DNS query")
|
||||||
|
return 0
|
||||||
|
}
|
||||||
|
ips, ttl, err := client.LookupIP(string(domain), option)
|
||||||
|
xlua.PushUserData(L, ips)
|
||||||
|
xlua.PushNumber(L, ttl)
|
||||||
|
xlua.PushError(L, err)
|
||||||
|
return 3
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
// callLuaHook invokes HandleDNSQuery in the supplied state.
|
||||||
|
// Returned slices and IP bytes may share storage with DNS caches or matcher inputs.
|
||||||
|
func (s *DNS) callLuaHook(L *lua.LState, domain string, option featureDNS.IPOption) ([]net.IP, uint32, error) {
|
||||||
|
top := L.GetTop()
|
||||||
|
defer L.SetTop(top)
|
||||||
|
fn := L.GetGlobal("HandleDNSQuery")
|
||||||
|
if fn.Type() != lua.LTFunction {
|
||||||
|
return nil, 0, errors.New("DNS script must define HandleDNSQuery(...)")
|
||||||
|
}
|
||||||
|
if err := L.CallByParam(lua.P{Fn: fn, NRet: 3, Protect: true},
|
||||||
|
lua.LString(strings.ToLower(domain)), lua.LBool(option.IPv4Enable),
|
||||||
|
lua.LBool(option.IPv6Enable), lua.LBool(option.FakeEnable)); err != nil {
|
||||||
|
return nil, 0, err
|
||||||
|
}
|
||||||
|
return readLuaDNSResult(L.Get(-3), L.Get(-2), L.Get(-1))
|
||||||
|
}
|
||||||
|
|
||||||
|
func readLuaDNSResult(addresses, ttlValue, errorValue lua.LValue) ([]net.IP, uint32, error) {
|
||||||
|
if err := xlua.ReadError(errorValue, "DNS script error must be an error or string"); err != nil {
|
||||||
|
return nil, 0, err
|
||||||
|
}
|
||||||
|
ttl, err := xlua.ReadUint32(ttlValue, "DNS script returned invalid TTL")
|
||||||
|
if err != nil {
|
||||||
|
return nil, 0, err
|
||||||
|
}
|
||||||
|
if addresses == lua.LNil {
|
||||||
|
return nil, 0, featureDNS.ErrEmptyResponse
|
||||||
|
}
|
||||||
|
ips, err := xlua.ReadUserData[[]net.IP](addresses, "DNS script IPs must be native IP slice userdata")
|
||||||
|
if err != nil {
|
||||||
|
return nil, 0, err
|
||||||
|
}
|
||||||
|
if len(ips) == 0 {
|
||||||
|
return nil, 0, featureDNS.ErrEmptyResponse
|
||||||
|
}
|
||||||
|
return ips, ttl, nil
|
||||||
|
}
|
||||||
@@ -0,0 +1,294 @@
|
|||||||
|
package dns
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
go_errors "errors"
|
||||||
|
"math"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/xtls/xray-core/common/geodata"
|
||||||
|
"github.com/xtls/xray-core/common/net"
|
||||||
|
featureDNS "github.com/xtls/xray-core/features/dns"
|
||||||
|
"github.com/xtls/xray-core/features/dns/localdns"
|
||||||
|
lua "github.com/yuin/gopher-lua"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestReadLuaDNSResult(t *testing.T) {
|
||||||
|
L := lua.NewState()
|
||||||
|
defer L.Close()
|
||||||
|
want := []net.IP{net.ParseIP("8.8.8.8"), {127, 0, 0, 1}, net.ParseIP("::1")}
|
||||||
|
addresses := L.NewUserData()
|
||||||
|
addresses.Value = want
|
||||||
|
ips, ttl, err := readLuaDNSResult(addresses, lua.LNumber(45), lua.LNil)
|
||||||
|
if err != nil || ttl != 45 || len(ips) != len(want) {
|
||||||
|
t.Fatalf("readLuaDNSResult() = %v, TTL %d, %v", ips, ttl, err)
|
||||||
|
}
|
||||||
|
for i := range want {
|
||||||
|
if !ips[i].Equal(want[i]) {
|
||||||
|
t.Fatalf("IP %d = %v, want %v", i, ips[i], want[i])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestReadLuaDNSResultValidation(t *testing.T) {
|
||||||
|
L := lua.NewState()
|
||||||
|
defer L.Close()
|
||||||
|
|
||||||
|
for _, tc := range []struct {
|
||||||
|
name string
|
||||||
|
change func(*[3]lua.LValue)
|
||||||
|
want string
|
||||||
|
}{
|
||||||
|
{"fractional TTL", func(v *[3]lua.LValue) { v[1] = lua.LNumber(1.5) }, "invalid TTL"},
|
||||||
|
{"oversized TTL", func(v *[3]lua.LValue) { v[1] = lua.LNumber(4294967296) }, "invalid TTL"},
|
||||||
|
{"negative TTL", func(v *[3]lua.LValue) { v[1] = lua.LNumber(-1) }, "invalid TTL"},
|
||||||
|
{"NaN TTL", func(v *[3]lua.LValue) { v[1] = lua.LNumber(math.NaN()) }, "invalid TTL"},
|
||||||
|
{"missing TTL", func(v *[3]lua.LValue) { v[1] = lua.LNil }, "invalid TTL"},
|
||||||
|
{"string IPs", func(v *[3]lua.LValue) { v[0] = lua.LString("127.0.0.1") }, "native IP slice"},
|
||||||
|
{"wrong userdata", func(v *[3]lua.LValue) { v[0].(*lua.LUserData).Value = net.ParseIP("127.0.0.1") }, "native IP slice"},
|
||||||
|
{"script error", func(v *[3]lua.LValue) { v[2] = lua.LString("blocked by script") }, "blocked by script"},
|
||||||
|
{"invalid error", func(v *[3]lua.LValue) { v[2] = lua.LTrue }, "error or string"},
|
||||||
|
} {
|
||||||
|
t.Run(tc.name, func(t *testing.T) {
|
||||||
|
addresses := L.NewUserData()
|
||||||
|
addresses.Value = []net.IP{net.ParseIP("127.0.0.1")}
|
||||||
|
values := [3]lua.LValue{addresses, lua.LNumber(60), lua.LNil}
|
||||||
|
tc.change(&values)
|
||||||
|
_, _, err := readLuaDNSResult(values[0], values[1], values[2])
|
||||||
|
if err == nil || !strings.Contains(err.Error(), tc.want) {
|
||||||
|
t.Fatalf("readLuaDNSResult error = %v, want %q", err, tc.want)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
addresses := L.NewUserData()
|
||||||
|
addresses.Value = []net.IP(nil)
|
||||||
|
for _, empty := range []lua.LValue{addresses, lua.LNil} {
|
||||||
|
if _, _, err := readLuaDNSResult(empty, lua.LNumber(0), lua.LNil); !go_errors.Is(err, featureDNS.ErrEmptyResponse) {
|
||||||
|
t.Fatalf("empty result error = %v, want ErrEmptyResponse", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
wantErr := go_errors.New("upstream failed")
|
||||||
|
errorValue := L.NewUserData()
|
||||||
|
errorValue.Value = wantErr
|
||||||
|
if _, _, err := readLuaDNSResult(lua.LNil, lua.LNil, errorValue); err != wantErr {
|
||||||
|
t.Fatalf("upstream error = %v, want original error %v", err, wantErr)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCallLuaHookCancellation(t *testing.T) {
|
||||||
|
L := lua.NewState()
|
||||||
|
defer L.Close()
|
||||||
|
if err := L.DoString(`function HandleDNSQuery(domain, ipv4, ipv6, fake) while true do end end`); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
ctx, cancel := context.WithCancel(context.Background())
|
||||||
|
cancel()
|
||||||
|
L.SetContext(ctx)
|
||||||
|
_, _, err := (&DNS{}).callLuaHook(L, "example.com", featureDNS.IPOption{IPv4Enable: true})
|
||||||
|
if err == nil {
|
||||||
|
t.Fatal("CallLuaHook did not stop after context cancellation")
|
||||||
|
}
|
||||||
|
if L.Context() != ctx {
|
||||||
|
t.Fatal("CallLuaHook changed the Lua state's context")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCallLuaHookNormalizesDomain(t *testing.T) {
|
||||||
|
L := lua.NewState()
|
||||||
|
defer L.Close()
|
||||||
|
addresses := L.NewUserData()
|
||||||
|
addresses.Value = []net.IP{net.ParseIP("127.0.0.1")}
|
||||||
|
L.SetGlobal("ips", addresses)
|
||||||
|
if err := L.DoString(`
|
||||||
|
function HandleDNSQuery(domain, ipv4, ipv6, fake)
|
||||||
|
assert(domain == "example.com")
|
||||||
|
assert(ipv4 and not ipv6 and not fake)
|
||||||
|
return ips, 60, nil
|
||||||
|
end
|
||||||
|
`); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
s := &DNS{}
|
||||||
|
if _, _, err := s.callLuaHook(L, "ExAmPlE.CoM", featureDNS.IPOption{IPv4Enable: true}); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCallLuaHookRestoresStack(t *testing.T) {
|
||||||
|
for _, tc := range []struct {
|
||||||
|
name string
|
||||||
|
body string
|
||||||
|
wantErr bool
|
||||||
|
}{
|
||||||
|
{"success", `return ips, 60`, false},
|
||||||
|
{"error", `error("failed")`, true},
|
||||||
|
} {
|
||||||
|
t.Run(tc.name, func(t *testing.T) {
|
||||||
|
L := lua.NewState()
|
||||||
|
defer L.Close()
|
||||||
|
addresses := L.NewUserData()
|
||||||
|
addresses.Value = []net.IP{net.ParseIP("127.0.0.1")}
|
||||||
|
L.SetGlobal("ips", addresses)
|
||||||
|
if err := L.DoString("function HandleDNSQuery() " + tc.body + " end"); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
L.Push(lua.LTrue)
|
||||||
|
_, _, err := (&DNS{}).callLuaHook(L, "example.com", featureDNS.IPOption{IPv4Enable: true})
|
||||||
|
if (err != nil) != tc.wantErr {
|
||||||
|
t.Fatalf("hook error = %v, want error %t", err, tc.wantErr)
|
||||||
|
}
|
||||||
|
if L.GetTop() != 1 || L.Get(1) != lua.LTrue {
|
||||||
|
t.Fatal("hook did not restore the stack")
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestLuaDNSServerQuery(t *testing.T) {
|
||||||
|
L := lua.NewState()
|
||||||
|
defer L.Close()
|
||||||
|
geodata.RegisterLua(L)
|
||||||
|
option := featureDNS.IPOption{IPv4Enable: true}
|
||||||
|
ips := []net.IP{net.ParseIP("127.0.0.1"), net.ParseIP("8.8.8.8")}
|
||||||
|
server := &DNS{clients: []*Client{{server: &benchmarkLuaNameServer{ips: ips}, ipOption: &option, timeoutMs: time.Second}}}
|
||||||
|
server.registerLua(L)
|
||||||
|
if err := L.DoString(`
|
||||||
|
local server = require("xray.dns").Servers[1]
|
||||||
|
local matcher = require("xray.geodata").BuildIPMatcher("127.0.0.0/8")
|
||||||
|
function HandleDNSQuery(domain, ipv4, ipv6, fake)
|
||||||
|
local ips, ttl, err = server:Query(domain, ipv4, ipv6, fake)
|
||||||
|
assert(type(ips) == "userdata" and not err)
|
||||||
|
assert(matcher:AnyMatch(ips))
|
||||||
|
local matched = matcher:FilterIPs(ips)
|
||||||
|
return matched, ttl, err
|
||||||
|
end
|
||||||
|
`); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
L.SetContext(context.Background())
|
||||||
|
got, ttl, err := server.callLuaHook(L, "example.com", option)
|
||||||
|
if err != nil || ttl != 60 || len(got) != 1 || !got[0].Equal(ips[0]) {
|
||||||
|
t.Fatalf("server query = %v, TTL %d, %v", got, ttl, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
type luaDNSClient struct {
|
||||||
|
featureDNS.Client
|
||||||
|
lookup func(string, featureDNS.IPOption) ([]net.IP, uint32, error)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *luaDNSClient) LookupIP(domain string, option featureDNS.IPOption) ([]net.IP, uint32, error) {
|
||||||
|
return c.lookup(domain, option)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestLuaDNSClientQuery(t *testing.T) {
|
||||||
|
L := lua.NewState()
|
||||||
|
defer L.Close()
|
||||||
|
L.SetContext(context.Background())
|
||||||
|
geodata.RegisterLua(L)
|
||||||
|
want := []net.IP{{127, 0, 0, 1}}
|
||||||
|
client := &luaDNSClient{lookup: func(domain string, option featureDNS.IPOption) ([]net.IP, uint32, error) {
|
||||||
|
if domain != "MiXeD.Example." || !option.IPv4Enable || option.IPv6Enable || !option.FakeEnable {
|
||||||
|
t.Fatalf("dns.Query arguments = %q, %+v", domain, option)
|
||||||
|
}
|
||||||
|
return want, 42, nil
|
||||||
|
}}
|
||||||
|
RegisterLua(L, client)
|
||||||
|
if err := L.DoString(`
|
||||||
|
local dns = require("xray.dns")
|
||||||
|
local matcher = require("xray.geodata").BuildIPMatcher("127.0.0.1")
|
||||||
|
assert(dns.Servers == nil)
|
||||||
|
ips, ttl, err = dns.Query("MiXeD.Example.", true, false, true)
|
||||||
|
assert(not err and ttl == 42 and matcher:AnyMatch(ips))
|
||||||
|
`); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
got := L.GetGlobal("ips").(*lua.LUserData).Value.([]net.IP)
|
||||||
|
if &got[0] != &want[0] {
|
||||||
|
t.Fatal("dns.Query copied the IP slice")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestLuaDNSLocalClient(t *testing.T) {
|
||||||
|
L := lua.NewState()
|
||||||
|
defer L.Close()
|
||||||
|
L.SetContext(context.Background())
|
||||||
|
RegisterLua(L, localdns.New())
|
||||||
|
if err := L.DoString(`
|
||||||
|
local dns = require("xray.dns")
|
||||||
|
assert(dns.Servers[1].ID == "localhost")
|
||||||
|
serverIPs, _, serverErr = dns.Servers[1]:Query("127.0.0.1", true, false, false)
|
||||||
|
clientIPs, _, clientErr = dns.Query("127.0.0.1", true, false, false)
|
||||||
|
assert(not serverErr and not clientErr)
|
||||||
|
`); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
for _, name := range []string{"serverIPs", "clientIPs"} {
|
||||||
|
ips := L.GetGlobal(name).(*lua.LUserData).Value.([]net.IP)
|
||||||
|
if len(ips) != 1 || !ips[0].Equal(net.ParseIP("127.0.0.1")) {
|
||||||
|
t.Fatalf("%s = %v", name, ips)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
type benchmarkLuaNameServer struct {
|
||||||
|
ips []net.IP
|
||||||
|
}
|
||||||
|
|
||||||
|
func (*benchmarkLuaNameServer) Name() string { return "benchmark" }
|
||||||
|
func (*benchmarkLuaNameServer) IsDisableCache() bool { return true }
|
||||||
|
func (s *benchmarkLuaNameServer) QueryIP(context.Context, string, featureDNS.IPOption) ([]net.IP, uint32, error) {
|
||||||
|
return s.ips, 60, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// BenchmarkLuaDNSHookCall isolates a preloaded Lua hook and its server:Query bridge.
|
||||||
|
// The direct case measures the same DNS client without Lua.
|
||||||
|
func BenchmarkLuaDNSHookCall(b *testing.B) {
|
||||||
|
option := featureDNS.IPOption{IPv4Enable: true}
|
||||||
|
ip := net.ParseIP("127.0.0.1")
|
||||||
|
upstream := &benchmarkLuaNameServer{ips: []net.IP{ip}}
|
||||||
|
client := &Client{server: upstream, ipOption: &option, timeoutMs: time.Second}
|
||||||
|
server := &DNS{clients: []*Client{client}}
|
||||||
|
L := lua.NewState()
|
||||||
|
defer L.Close()
|
||||||
|
server.registerLua(L)
|
||||||
|
if err := L.DoString(`
|
||||||
|
local server = require("xray.dns").Servers[1]
|
||||||
|
function HandleDNSQuery(domain, ipv4, ipv6, fake)
|
||||||
|
return server:Query(domain, ipv4, ipv6, fake)
|
||||||
|
end
|
||||||
|
`); err != nil {
|
||||||
|
b.Fatal(err)
|
||||||
|
}
|
||||||
|
|
||||||
|
ctx := context.Background()
|
||||||
|
L.SetContext(ctx)
|
||||||
|
for _, bench := range []struct {
|
||||||
|
name string
|
||||||
|
query func() ([]net.IP, uint32, error)
|
||||||
|
}{
|
||||||
|
{"direct", func() ([]net.IP, uint32, error) { return client.QueryIP(ctx, "example.com", option) }},
|
||||||
|
{"lua_hook", func() ([]net.IP, uint32, error) { return server.callLuaHook(L, "example.com", option) }},
|
||||||
|
} {
|
||||||
|
b.Run(bench.name, func(b *testing.B) {
|
||||||
|
b.ReportAllocs()
|
||||||
|
b.ResetTimer()
|
||||||
|
var ips []net.IP
|
||||||
|
var ttl uint32
|
||||||
|
var err error
|
||||||
|
for i := 0; i < b.N; i++ {
|
||||||
|
ips, ttl, err = bench.query()
|
||||||
|
if err != nil {
|
||||||
|
b.Fatal(err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
b.StopTimer()
|
||||||
|
if ttl != 60 || len(ips) != 1 || !ips[0].Equal(ip) {
|
||||||
|
b.Fatalf("query() = %v, TTL %d; want %v, TTL 60", ips, ttl, ip)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -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
|
||||||
@@ -84,7 +85,7 @@ func NewServer(ctx context.Context, dest net.Destination, dispatcher routing.Dis
|
|||||||
if dest.Network == net.Network_UDP { // UDP classic DNS mode
|
if dest.Network == net.Network_UDP { // UDP classic DNS mode
|
||||||
return NewClassicNameServer(dest, dispatcher, disableCache, serveStale, serveExpiredTTL, clientIP), nil
|
return NewClassicNameServer(dest, dispatcher, disableCache, serveStale, serveExpiredTTL, clientIP), nil
|
||||||
}
|
}
|
||||||
return nil, errors.New("No available name server could be created from ", dest).AtWarning()
|
return nil, errors.New("No available name server could be created from ", dest)
|
||||||
}
|
}
|
||||||
|
|
||||||
// NewClient creates a DNS client managing a name server with client IP, domain rules and expected IPs.
|
// NewClient creates a DNS client managing a name server with client IP, domain rules and expected IPs.
|
||||||
@@ -97,12 +98,12 @@ 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)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return errors.New("failed to create nameserver").Base(err).AtWarning()
|
return errors.New("failed to create nameserver").Base(err)
|
||||||
}
|
}
|
||||||
|
|
||||||
_, isLocalDNS := server.(*LocalNameServer)
|
_, isLocalDNS := server.(*LocalNameServer)
|
||||||
@@ -113,7 +114,7 @@ func NewClient(
|
|||||||
if len(ns.ExpectedIp) > 0 {
|
if len(ns.ExpectedIp) > 0 {
|
||||||
expectedMatcher, err = geodata.IPReg.BuildIPMatcher(ns.ExpectedIp)
|
expectedMatcher, err = geodata.IPReg.BuildIPMatcher(ns.ExpectedIp)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return errors.New("failed to create expected ip matcher").Base(err).AtWarning()
|
return errors.New("failed to create expected ip matcher").Base(err)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -122,7 +123,7 @@ func NewClient(
|
|||||||
if len(ns.UnexpectedIp) > 0 {
|
if len(ns.UnexpectedIp) > 0 {
|
||||||
unexpectedMatcher, err = geodata.IPReg.BuildIPMatcher(ns.UnexpectedIp)
|
unexpectedMatcher, err = geodata.IPReg.BuildIPMatcher(ns.UnexpectedIp)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return errors.New("failed to create unexpected ip matcher").Base(err).AtWarning()
|
return errors.New("failed to create unexpected ip matcher").Base(err)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -27,7 +27,7 @@ func (s *FakeDNSServer) IsDisableCache() bool {
|
|||||||
|
|
||||||
func (f *FakeDNSServer) QueryIP(ctx context.Context, domain string, opt dns.IPOption) ([]net.IP, uint32, error) {
|
func (f *FakeDNSServer) QueryIP(ctx context.Context, domain string, opt dns.IPOption) ([]net.IP, uint32, error) {
|
||||||
if f.fakeDNSEngine == nil {
|
if f.fakeDNSEngine == nil {
|
||||||
return nil, 0, errors.New("Unable to locate a fake DNS Engine").AtError()
|
return nil, 0, errors.New("Unable to locate a fake DNS Engine")
|
||||||
}
|
}
|
||||||
|
|
||||||
var ips []net.Address
|
var ips []net.Address
|
||||||
@@ -39,7 +39,7 @@ func (f *FakeDNSServer) QueryIP(ctx context.Context, domain string, opt dns.IPOp
|
|||||||
|
|
||||||
netIP, err := toNetIP(ips)
|
netIP, err := toNetIP(ips)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, 0, errors.New("Unable to convert IP to net ip").Base(err).AtError()
|
return nil, 0, errors.New("Unable to convert IP to net ip").Base(err)
|
||||||
}
|
}
|
||||||
|
|
||||||
errors.LogInfo(ctx, f.Name(), " got answer: ", domain, " -> ", ips)
|
errors.LogInfo(ctx, f.Name(), " got answer: ", domain, " -> ", ips)
|
||||||
|
|||||||
@@ -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,59 @@
|
|||||||
|
package dns
|
||||||
|
|
||||||
|
import (
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/xtls/xray-core/common/errors"
|
||||||
|
"github.com/xtls/xray-core/common/geodata"
|
||||||
|
"github.com/xtls/xray-core/common/log"
|
||||||
|
xlua "github.com/xtls/xray-core/common/lua"
|
||||||
|
"github.com/xtls/xray-core/common/net"
|
||||||
|
"github.com/xtls/xray-core/features/dns"
|
||||||
|
lua "github.com/yuin/gopher-lua"
|
||||||
|
)
|
||||||
|
|
||||||
|
const scriptExecutionTimeout = 6 * time.Second
|
||||||
|
|
||||||
|
type scriptEngine struct {
|
||||||
|
dns *DNS
|
||||||
|
pool *xlua.Pool
|
||||||
|
}
|
||||||
|
|
||||||
|
func newScriptEngine(path string, server *DNS) (*scriptEngine, error) {
|
||||||
|
program, err := xlua.CompileFile(path)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
e := &scriptEngine{dns: server}
|
||||||
|
e.pool, err = xlua.NewPool(server.ctx, scriptExecutionTimeout, program.NewStateFactory(
|
||||||
|
scriptExecutionTimeout*20,
|
||||||
|
func(L *lua.LState) {
|
||||||
|
geodata.RegisterLua(L)
|
||||||
|
log.RegisterLua(L)
|
||||||
|
server.registerLua(L)
|
||||||
|
},
|
||||||
|
func(L *lua.LState) error {
|
||||||
|
if L.GetGlobal("HandleDNSQuery").Type() != lua.LTFunction {
|
||||||
|
return errors.New("DNS script must define HandleDNSQuery(...)")
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}))
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
errors.LogInfo(server.ctx, "DNS script initialized from ", path)
|
||||||
|
return e, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (e *scriptEngine) close() {
|
||||||
|
e.pool.Close()
|
||||||
|
}
|
||||||
|
|
||||||
|
func (e *scriptEngine) query(domain string, option dns.IPOption) (ips []net.IP, ttl uint32, err error) {
|
||||||
|
err = e.pool.WithState(nil, 0, func(L *lua.LState) error {
|
||||||
|
var hookErr error
|
||||||
|
ips, ttl, hookErr = e.dns.callLuaHook(L, domain, option)
|
||||||
|
return hookErr
|
||||||
|
})
|
||||||
|
return
|
||||||
|
}
|
||||||
@@ -0,0 +1,197 @@
|
|||||||
|
package dns
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/xtls/xray-core/common/net"
|
||||||
|
featureDNS "github.com/xtls/xray-core/features/dns"
|
||||||
|
)
|
||||||
|
|
||||||
|
type geoIPScriptNameServer struct {
|
||||||
|
name string
|
||||||
|
answers map[string]net.IP
|
||||||
|
ttl uint32
|
||||||
|
calls int
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *geoIPScriptNameServer) Name() string { return s.name }
|
||||||
|
func (s *geoIPScriptNameServer) IsDisableCache() bool { return true }
|
||||||
|
|
||||||
|
func (s *geoIPScriptNameServer) QueryIP(ctx context.Context, domain string, _ featureDNS.IPOption) ([]net.IP, uint32, error) {
|
||||||
|
if err := ctx.Err(); err != nil {
|
||||||
|
return nil, 0, err
|
||||||
|
}
|
||||||
|
s.calls++
|
||||||
|
ip, ok := s.answers[domain]
|
||||||
|
if !ok {
|
||||||
|
return nil, 0, featureDNS.ErrEmptyResponse
|
||||||
|
}
|
||||||
|
return []net.IP{ip}, s.ttl, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestDNSScriptGeoIPFallback(t *testing.T) {
|
||||||
|
t.Setenv("xray.location.asset", filepath.Join("..", "..", "resources"))
|
||||||
|
script := `
|
||||||
|
local servers = require("xray.dns").Servers
|
||||||
|
local us_ips = require("xray.geodata").BuildIPMatcher("geoip:us")
|
||||||
|
|
||||||
|
local by_id = {}
|
||||||
|
for _, server in ipairs(servers) do
|
||||||
|
by_id[server.ID] = server
|
||||||
|
end
|
||||||
|
assert(by_id.primary and by_id.fallback, "primary and fallback DNS servers are required")
|
||||||
|
|
||||||
|
function HandleDNSQuery(domain, ipv4, ipv6, fake)
|
||||||
|
local ips, ttl, err = by_id.primary:Query(domain, ipv4, ipv6, fake)
|
||||||
|
if not err and us_ips:AnyMatch(ips) then
|
||||||
|
return ips, ttl, nil
|
||||||
|
end
|
||||||
|
return by_id.fallback:Query(domain, ipv4, ipv6, fake)
|
||||||
|
end
|
||||||
|
`
|
||||||
|
scriptPath := filepath.Join(t.TempDir(), "geoip_fallback.lua")
|
||||||
|
if err := os.WriteFile(scriptPath, []byte(script), 0o600); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
|
||||||
|
primary := &geoIPScriptNameServer{
|
||||||
|
name: "primary",
|
||||||
|
answers: map[string]net.IP{
|
||||||
|
"us.example": net.ParseIP("2001:4860:4860::8888"),
|
||||||
|
"other.example": net.ParseIP("127.0.0.1"),
|
||||||
|
},
|
||||||
|
ttl: 30,
|
||||||
|
}
|
||||||
|
fallback := &geoIPScriptNameServer{
|
||||||
|
name: "fallback",
|
||||||
|
answers: map[string]net.IP{"other.example": net.ParseIP("9.9.9.9")},
|
||||||
|
ttl: 60,
|
||||||
|
}
|
||||||
|
option := featureDNS.IPOption{IPv4Enable: true, IPv6Enable: true}
|
||||||
|
hosts, err := NewStaticHosts(nil)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
server := &DNS{
|
||||||
|
ctx: context.Background(),
|
||||||
|
hosts: hosts,
|
||||||
|
ipOption: &option,
|
||||||
|
scriptPath: scriptPath,
|
||||||
|
clients: []*Client{
|
||||||
|
{id: "primary", server: primary, ipOption: &option, timeoutMs: 2 * time.Second},
|
||||||
|
{id: "fallback", server: fallback, ipOption: &option, timeoutMs: 2 * time.Second},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
if err := server.Start(); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
defer server.Close()
|
||||||
|
|
||||||
|
for _, tc := range []struct {
|
||||||
|
domain string
|
||||||
|
ip net.IP
|
||||||
|
ttl uint32
|
||||||
|
}{
|
||||||
|
{"Us.Example.", net.ParseIP("2001:4860:4860::8888"), 30},
|
||||||
|
{"other.example", net.ParseIP("9.9.9.9"), 60},
|
||||||
|
} {
|
||||||
|
ips, ttl, err := server.LookupIP(tc.domain, option)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("LookupIP(%q): %v", tc.domain, err)
|
||||||
|
}
|
||||||
|
if ttl != tc.ttl || len(ips) != 1 || !ips[0].Equal(tc.ip) {
|
||||||
|
t.Fatalf("LookupIP(%q) = %v, TTL %d; want %v, TTL %d", tc.domain, ips, ttl, tc.ip, tc.ttl)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if primary.calls != 2 || fallback.calls != 1 {
|
||||||
|
t.Fatalf("upstream calls: primary %d, fallback %d; want 2 and 1", primary.calls, fallback.calls)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestDNSScriptRejectsInvalidStartup(t *testing.T) {
|
||||||
|
for _, tc := range []struct {
|
||||||
|
name string
|
||||||
|
script string
|
||||||
|
}{
|
||||||
|
{"syntax", "function HandleDNSQuery("},
|
||||||
|
{"missing hook", "value = 1"},
|
||||||
|
{"top-level error", `error("setup failed")`},
|
||||||
|
} {
|
||||||
|
t.Run(tc.name, func(t *testing.T) {
|
||||||
|
path := filepath.Join(t.TempDir(), "script.lua")
|
||||||
|
if err := os.WriteFile(path, []byte(tc.script), 0o600); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
server := &DNS{ctx: context.Background(), scriptPath: path}
|
||||||
|
if err := server.Start(); err == nil {
|
||||||
|
t.Fatal("Start accepted an invalid DNS script")
|
||||||
|
}
|
||||||
|
if server.script != nil {
|
||||||
|
t.Fatal("Start retained a script engine after failure")
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestDNSScriptHookErrorAndFakeDNSOption(t *testing.T) {
|
||||||
|
path := filepath.Join(t.TempDir(), "script.lua")
|
||||||
|
script := `
|
||||||
|
local server = require("xray.dns").Servers[1]
|
||||||
|
local log = require("xray.log")
|
||||||
|
log.Info("DNS script loaded")
|
||||||
|
function HandleDNSQuery(domain, ipv4, ipv6, fake)
|
||||||
|
log.Debug("DNS query: ", domain)
|
||||||
|
if domain == "bad.example" then error("script failure") end
|
||||||
|
local ips, ttl, err = server:Query(domain, ipv4, ipv6, fake)
|
||||||
|
if err then log.Error("DNS failed: ", err) end
|
||||||
|
return ips, ttl, err
|
||||||
|
end
|
||||||
|
`
|
||||||
|
if err := os.WriteFile(path, []byte(script), 0o600); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
option := featureDNS.IPOption{IPv4Enable: true}
|
||||||
|
hosts, err := NewStaticHosts(nil)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
upstream := &geoIPScriptNameServer{
|
||||||
|
name: "FakeDNS",
|
||||||
|
answers: map[string]net.IP{"good.example": net.ParseIP("198.18.0.1")},
|
||||||
|
ttl: 30,
|
||||||
|
}
|
||||||
|
server := &DNS{
|
||||||
|
ctx: context.Background(),
|
||||||
|
hosts: hosts,
|
||||||
|
ipOption: &option,
|
||||||
|
scriptPath: path,
|
||||||
|
clients: []*Client{{id: "fake", server: upstream, ipOption: &option, timeoutMs: time.Second}},
|
||||||
|
}
|
||||||
|
if err := server.Start(); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
defer server.Close()
|
||||||
|
|
||||||
|
if _, _, err := server.LookupIP("bad.example", option); err == nil || !strings.Contains(err.Error(), "script failure") {
|
||||||
|
t.Fatalf("hook failure = %v, want script failure", err)
|
||||||
|
}
|
||||||
|
if _, _, err := server.LookupIP("good.example", option); err != featureDNS.ErrEmptyResponse {
|
||||||
|
t.Fatalf("FakeDNS without FakeEnable = %v, want ErrEmptyResponse", err)
|
||||||
|
}
|
||||||
|
if upstream.calls != 0 {
|
||||||
|
t.Fatalf("FakeDNS was queried without FakeEnable: %d calls", upstream.calls)
|
||||||
|
}
|
||||||
|
withFake := featureDNS.IPOption{IPv4Enable: true, FakeEnable: true}
|
||||||
|
ips, ttl, err := server.LookupIP("good.example", withFake)
|
||||||
|
if err != nil || ttl != 30 || len(ips) != 1 || !ips[0].Equal(net.ParseIP("198.18.0.1")) {
|
||||||
|
t.Fatalf("FakeDNS with FakeEnable = %v, TTL %d, %v", ips, ttl, err)
|
||||||
|
}
|
||||||
|
if upstream.calls != 1 {
|
||||||
|
t.Fatalf("FakeDNS query count = %d, want 1", upstream.calls)
|
||||||
|
}
|
||||||
|
}
|
||||||
+6
-2
@@ -89,10 +89,10 @@ func (g *Instance) startInternal() error {
|
|||||||
g.active = true
|
g.active = true
|
||||||
|
|
||||||
if err := g.initAccessLogger(); err != nil {
|
if err := g.initAccessLogger(); err != nil {
|
||||||
return errors.New("failed to initialize access logger").Base(err).AtWarning()
|
return errors.New("failed to initialize access logger").Base(err)
|
||||||
}
|
}
|
||||||
if err := g.initErrorLogger(); err != nil {
|
if err := g.initErrorLogger(); err != nil {
|
||||||
return errors.New("failed to initialize error logger").Base(err).AtWarning()
|
return errors.New("failed to initialize error logger").Base(err)
|
||||||
}
|
}
|
||||||
|
|
||||||
return nil
|
return nil
|
||||||
@@ -141,6 +141,10 @@ func (g *Instance) Handle(msg log.Message) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (g *Instance) Severity() log.Severity {
|
||||||
|
return g.config.ErrorLogLevel
|
||||||
|
}
|
||||||
|
|
||||||
// Close implements common.Closable.Close().
|
// Close implements common.Closable.Close().
|
||||||
func (g *Instance) Close() error {
|
func (g *Instance) Close() error {
|
||||||
errors.LogDebug(context.Background(), "Logger closing")
|
errors.LogDebug(context.Background(), "Logger closing")
|
||||||
|
|||||||
+11
-22
@@ -330,7 +330,6 @@ type SenderConfig struct {
|
|||||||
// Send traffic through the given IP. Only IP is allowed.
|
// Send traffic through the given IP. Only IP is allowed.
|
||||||
Via *net.IPOrDomain `protobuf:"bytes,1,opt,name=via,proto3" json:"via,omitempty"`
|
Via *net.IPOrDomain `protobuf:"bytes,1,opt,name=via,proto3" json:"via,omitempty"`
|
||||||
StreamSettings *internet.StreamConfig `protobuf:"bytes,2,opt,name=stream_settings,json=streamSettings,proto3" json:"stream_settings,omitempty"`
|
StreamSettings *internet.StreamConfig `protobuf:"bytes,2,opt,name=stream_settings,json=streamSettings,proto3" json:"stream_settings,omitempty"`
|
||||||
ProxySettings *internet.ProxyConfig `protobuf:"bytes,3,opt,name=proxy_settings,json=proxySettings,proto3" json:"proxy_settings,omitempty"`
|
|
||||||
MultiplexSettings *MultiplexingConfig `protobuf:"bytes,4,opt,name=multiplex_settings,json=multiplexSettings,proto3" json:"multiplex_settings,omitempty"`
|
MultiplexSettings *MultiplexingConfig `protobuf:"bytes,4,opt,name=multiplex_settings,json=multiplexSettings,proto3" json:"multiplex_settings,omitempty"`
|
||||||
ViaCidr string `protobuf:"bytes,5,opt,name=via_cidr,json=viaCidr,proto3" json:"via_cidr,omitempty"`
|
ViaCidr string `protobuf:"bytes,5,opt,name=via_cidr,json=viaCidr,proto3" json:"via_cidr,omitempty"`
|
||||||
TargetStrategy internet.DomainStrategy `protobuf:"varint,6,opt,name=target_strategy,json=targetStrategy,proto3,enum=xray.transport.internet.DomainStrategy" json:"target_strategy,omitempty"`
|
TargetStrategy internet.DomainStrategy `protobuf:"varint,6,opt,name=target_strategy,json=targetStrategy,proto3,enum=xray.transport.internet.DomainStrategy" json:"target_strategy,omitempty"`
|
||||||
@@ -382,13 +381,6 @@ func (x *SenderConfig) GetStreamSettings() *internet.StreamConfig {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (x *SenderConfig) GetProxySettings() *internet.ProxyConfig {
|
|
||||||
if x != nil {
|
|
||||||
return x.ProxySettings
|
|
||||||
}
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (x *SenderConfig) GetMultiplexSettings() *MultiplexingConfig {
|
func (x *SenderConfig) GetMultiplexSettings() *MultiplexingConfig {
|
||||||
if x != nil {
|
if x != nil {
|
||||||
return x.MultiplexSettings
|
return x.MultiplexSettings
|
||||||
@@ -506,14 +498,13 @@ const file_app_proxyman_config_proto_rawDesc = "" +
|
|||||||
"\x03tag\x18\x01 \x01(\tR\x03tag\x12M\n" +
|
"\x03tag\x18\x01 \x01(\tR\x03tag\x12M\n" +
|
||||||
"\x11receiver_settings\x18\x02 \x01(\v2 .xray.common.serial.TypedMessageR\x10receiverSettings\x12G\n" +
|
"\x11receiver_settings\x18\x02 \x01(\v2 .xray.common.serial.TypedMessageR\x10receiverSettings\x12G\n" +
|
||||||
"\x0eproxy_settings\x18\x03 \x01(\v2 .xray.common.serial.TypedMessageR\rproxySettings\"\x10\n" +
|
"\x0eproxy_settings\x18\x03 \x01(\v2 .xray.common.serial.TypedMessageR\rproxySettings\"\x10\n" +
|
||||||
"\x0eOutboundConfig\"\x9d\x03\n" +
|
"\x0eOutboundConfig\"\xd6\x02\n" +
|
||||||
"\fSenderConfig\x12-\n" +
|
"\fSenderConfig\x12-\n" +
|
||||||
"\x03via\x18\x01 \x01(\v2\x1b.xray.common.net.IPOrDomainR\x03via\x12N\n" +
|
"\x03via\x18\x01 \x01(\v2\x1b.xray.common.net.IPOrDomainR\x03via\x12N\n" +
|
||||||
"\x0fstream_settings\x18\x02 \x01(\v2%.xray.transport.internet.StreamConfigR\x0estreamSettings\x12K\n" +
|
"\x0fstream_settings\x18\x02 \x01(\v2%.xray.transport.internet.StreamConfigR\x0estreamSettings\x12T\n" +
|
||||||
"\x0eproxy_settings\x18\x03 \x01(\v2$.xray.transport.internet.ProxyConfigR\rproxySettings\x12T\n" +
|
|
||||||
"\x12multiplex_settings\x18\x04 \x01(\v2%.xray.app.proxyman.MultiplexingConfigR\x11multiplexSettings\x12\x19\n" +
|
"\x12multiplex_settings\x18\x04 \x01(\v2%.xray.app.proxyman.MultiplexingConfigR\x11multiplexSettings\x12\x19\n" +
|
||||||
"\bvia_cidr\x18\x05 \x01(\tR\aviaCidr\x12P\n" +
|
"\bvia_cidr\x18\x05 \x01(\tR\aviaCidr\x12P\n" +
|
||||||
"\x0ftarget_strategy\x18\x06 \x01(\x0e2'.xray.transport.internet.DomainStrategyR\x0etargetStrategy\"\xa4\x01\n" +
|
"\x0ftarget_strategy\x18\x06 \x01(\x0e2'.xray.transport.internet.DomainStrategyR\x0etargetStrategyJ\x04\b\x03\x10\x04\"\xa4\x01\n" +
|
||||||
"\x12MultiplexingConfig\x12\x18\n" +
|
"\x12MultiplexingConfig\x12\x18\n" +
|
||||||
"\aenabled\x18\x01 \x01(\bR\aenabled\x12 \n" +
|
"\aenabled\x18\x01 \x01(\bR\aenabled\x12 \n" +
|
||||||
"\vconcurrency\x18\x02 \x01(\x05R\vconcurrency\x12(\n" +
|
"\vconcurrency\x18\x02 \x01(\x05R\vconcurrency\x12(\n" +
|
||||||
@@ -548,8 +539,7 @@ var file_app_proxyman_config_proto_goTypes = []any{
|
|||||||
(*net.IPOrDomain)(nil), // 10: xray.common.net.IPOrDomain
|
(*net.IPOrDomain)(nil), // 10: xray.common.net.IPOrDomain
|
||||||
(*internet.StreamConfig)(nil), // 11: xray.transport.internet.StreamConfig
|
(*internet.StreamConfig)(nil), // 11: xray.transport.internet.StreamConfig
|
||||||
(*serial.TypedMessage)(nil), // 12: xray.common.serial.TypedMessage
|
(*serial.TypedMessage)(nil), // 12: xray.common.serial.TypedMessage
|
||||||
(*internet.ProxyConfig)(nil), // 13: xray.transport.internet.ProxyConfig
|
(internet.DomainStrategy)(0), // 13: xray.transport.internet.DomainStrategy
|
||||||
(internet.DomainStrategy)(0), // 14: xray.transport.internet.DomainStrategy
|
|
||||||
}
|
}
|
||||||
var file_app_proxyman_config_proto_depIdxs = []int32{
|
var file_app_proxyman_config_proto_depIdxs = []int32{
|
||||||
7, // 0: xray.app.proxyman.SniffingConfig.domains_excluded:type_name -> xray.common.geodata.DomainRule
|
7, // 0: xray.app.proxyman.SniffingConfig.domains_excluded:type_name -> xray.common.geodata.DomainRule
|
||||||
@@ -562,14 +552,13 @@ var file_app_proxyman_config_proto_depIdxs = []int32{
|
|||||||
12, // 7: xray.app.proxyman.InboundHandlerConfig.proxy_settings:type_name -> xray.common.serial.TypedMessage
|
12, // 7: xray.app.proxyman.InboundHandlerConfig.proxy_settings:type_name -> xray.common.serial.TypedMessage
|
||||||
10, // 8: xray.app.proxyman.SenderConfig.via:type_name -> xray.common.net.IPOrDomain
|
10, // 8: xray.app.proxyman.SenderConfig.via:type_name -> xray.common.net.IPOrDomain
|
||||||
11, // 9: xray.app.proxyman.SenderConfig.stream_settings:type_name -> xray.transport.internet.StreamConfig
|
11, // 9: xray.app.proxyman.SenderConfig.stream_settings:type_name -> xray.transport.internet.StreamConfig
|
||||||
13, // 10: xray.app.proxyman.SenderConfig.proxy_settings:type_name -> xray.transport.internet.ProxyConfig
|
6, // 10: xray.app.proxyman.SenderConfig.multiplex_settings:type_name -> xray.app.proxyman.MultiplexingConfig
|
||||||
6, // 11: xray.app.proxyman.SenderConfig.multiplex_settings:type_name -> xray.app.proxyman.MultiplexingConfig
|
13, // 11: xray.app.proxyman.SenderConfig.target_strategy:type_name -> xray.transport.internet.DomainStrategy
|
||||||
14, // 12: xray.app.proxyman.SenderConfig.target_strategy:type_name -> xray.transport.internet.DomainStrategy
|
12, // [12:12] is the sub-list for method output_type
|
||||||
13, // [13:13] is the sub-list for method output_type
|
12, // [12:12] is the sub-list for method input_type
|
||||||
13, // [13:13] is the sub-list for method input_type
|
12, // [12:12] is the sub-list for extension type_name
|
||||||
13, // [13:13] is the sub-list for extension type_name
|
12, // [12:12] is the sub-list for extension extendee
|
||||||
13, // [13:13] is the sub-list for extension extendee
|
0, // [0:12] is the sub-list for field type_name
|
||||||
0, // [0:13] is the sub-list for field type_name
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func init() { file_app_proxyman_config_proto_init() }
|
func init() { file_app_proxyman_config_proto_init() }
|
||||||
|
|||||||
@@ -57,7 +57,7 @@ message SenderConfig {
|
|||||||
// Send traffic through the given IP. Only IP is allowed.
|
// Send traffic through the given IP. Only IP is allowed.
|
||||||
xray.common.net.IPOrDomain via = 1;
|
xray.common.net.IPOrDomain via = 1;
|
||||||
xray.transport.internet.StreamConfig stream_settings = 2;
|
xray.transport.internet.StreamConfig stream_settings = 2;
|
||||||
xray.transport.internet.ProxyConfig proxy_settings = 3;
|
reserved 3;
|
||||||
MultiplexingConfig multiplex_settings = 4;
|
MultiplexingConfig multiplex_settings = 4;
|
||||||
string via_cidr = 5;
|
string via_cidr = 5;
|
||||||
xray.transport.internet.DomainStrategy target_strategy = 6;
|
xray.transport.internet.DomainStrategy target_strategy = 6;
|
||||||
|
|||||||
@@ -66,7 +66,7 @@ func NewAlwaysOnInboundHandler(ctx context.Context, tag string, receiverConfig *
|
|||||||
}
|
}
|
||||||
mss, err := internet.ToMemoryStreamConfig(receiverConfig.StreamSettings)
|
mss, err := internet.ToMemoryStreamConfig(receiverConfig.StreamSettings)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, errors.New("failed to parse stream config").Base(err).AtWarning()
|
return nil, errors.New("failed to parse stream config").Base(err)
|
||||||
}
|
}
|
||||||
|
|
||||||
newCtx := session.ContextWithInbound(ctx, &session.Inbound{Tag: tag, Source: src})
|
newCtx := session.ContextWithInbound(ctx, &session.Inbound{Tag: tag, Source: src})
|
||||||
|
|||||||
@@ -165,7 +165,7 @@ func NewHandler(ctx context.Context, config *core.InboundHandlerConfig) (inbound
|
|||||||
|
|
||||||
receiverSettings, ok := rawReceiverSettings.(*proxyman.ReceiverConfig)
|
receiverSettings, ok := rawReceiverSettings.(*proxyman.ReceiverConfig)
|
||||||
if !ok {
|
if !ok {
|
||||||
return nil, errors.New("not a ReceiverConfig").AtError()
|
return nil, errors.New("not a ReceiverConfig")
|
||||||
}
|
}
|
||||||
|
|
||||||
streamSettings := receiverSettings.StreamSettings
|
streamSettings := receiverSettings.StreamSettings
|
||||||
|
|||||||
@@ -142,7 +142,7 @@ func (w *tcpWorker) Start() error {
|
|||||||
go w.callback(conn)
|
go w.callback(conn)
|
||||||
})
|
})
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return errors.New("failed to listen TCP on ", w.port).AtWarning().Base(err)
|
return errors.New("failed to listen TCP on ", w.port).Base(err)
|
||||||
}
|
}
|
||||||
w.hub = hub
|
w.hub = hub
|
||||||
return nil
|
return nil
|
||||||
@@ -528,7 +528,7 @@ func (w *dsWorker) Start() error {
|
|||||||
go w.callback(conn)
|
go w.callback(conn)
|
||||||
})
|
})
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return errors.New("failed to listen Unix Domain Socket on ", w.address).AtWarning().Base(err)
|
return errors.New("failed to listen Unix Domain Socket on ", w.address).Base(err)
|
||||||
}
|
}
|
||||||
w.hub = hub
|
w.hub = hub
|
||||||
return nil
|
return nil
|
||||||
|
|||||||
@@ -15,7 +15,6 @@ import (
|
|||||||
"github.com/xtls/xray-core/common/errors"
|
"github.com/xtls/xray-core/common/errors"
|
||||||
"github.com/xtls/xray-core/common/mux"
|
"github.com/xtls/xray-core/common/mux"
|
||||||
"github.com/xtls/xray-core/common/net"
|
"github.com/xtls/xray-core/common/net"
|
||||||
"github.com/xtls/xray-core/common/net/cnc"
|
|
||||||
"github.com/xtls/xray-core/common/serial"
|
"github.com/xtls/xray-core/common/serial"
|
||||||
"github.com/xtls/xray-core/common/session"
|
"github.com/xtls/xray-core/common/session"
|
||||||
"github.com/xtls/xray-core/core"
|
"github.com/xtls/xray-core/core"
|
||||||
@@ -26,8 +25,6 @@ import (
|
|||||||
"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/stat"
|
"github.com/xtls/xray-core/transport/internet/stat"
|
||||||
"github.com/xtls/xray-core/transport/internet/tls"
|
|
||||||
"github.com/xtls/xray-core/transport/pipe"
|
|
||||||
"google.golang.org/protobuf/proto"
|
"google.golang.org/protobuf/proto"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -63,7 +60,6 @@ type Handler struct {
|
|||||||
streamSettings *internet.MemoryStreamConfig
|
streamSettings *internet.MemoryStreamConfig
|
||||||
proxyConfig proto.Message
|
proxyConfig proto.Message
|
||||||
proxy proxy.Outbound
|
proxy proxy.Outbound
|
||||||
outboundManager outbound.Manager
|
|
||||||
mux *mux.ClientManager
|
mux *mux.ClientManager
|
||||||
xudp *mux.ClientManager
|
xudp *mux.ClientManager
|
||||||
udp443 string
|
udp443 string
|
||||||
@@ -77,7 +73,6 @@ func NewHandler(ctx context.Context, config *core.OutboundHandlerConfig) (outbou
|
|||||||
uplinkCounter, downlinkCounter := getStatCounter(v, config.Tag)
|
uplinkCounter, downlinkCounter := getStatCounter(v, config.Tag)
|
||||||
h := &Handler{
|
h := &Handler{
|
||||||
tag: config.Tag,
|
tag: config.Tag,
|
||||||
outboundManager: v.GetFeature(outbound.ManagerType()).(outbound.Manager),
|
|
||||||
uplinkCounter: uplinkCounter,
|
uplinkCounter: uplinkCounter,
|
||||||
downlinkCounter: downlinkCounter,
|
downlinkCounter: downlinkCounter,
|
||||||
}
|
}
|
||||||
@@ -92,7 +87,7 @@ func NewHandler(ctx context.Context, config *core.OutboundHandlerConfig) (outbou
|
|||||||
h.senderSettings = s
|
h.senderSettings = s
|
||||||
mss, err := internet.ToMemoryStreamConfig(s.StreamSettings)
|
mss, err := internet.ToMemoryStreamConfig(s.StreamSettings)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, errors.New("failed to parse stream settings").Base(err).AtWarning()
|
return nil, errors.New("failed to parse stream settings").Base(err)
|
||||||
}
|
}
|
||||||
h.streamSettings = mss
|
h.streamSettings = mss
|
||||||
default:
|
default:
|
||||||
@@ -108,9 +103,11 @@ func NewHandler(ctx context.Context, config *core.OutboundHandlerConfig) (outbou
|
|||||||
|
|
||||||
ctx = session.ContextWithFullHandler(ctx, h)
|
ctx = session.ContextWithFullHandler(ctx, h)
|
||||||
|
|
||||||
newCtx := session.ContextWithStreamSettings(ctx, h.streamSettings)
|
if h.streamSettings != nil {
|
||||||
|
ctx = session.ContextWithStreamSettings(ctx, h.streamSettings)
|
||||||
|
}
|
||||||
|
|
||||||
rawProxyHandler, err := common.CreateObject(newCtx, proxyConfig)
|
rawProxyHandler, err := common.CreateObject(ctx, proxyConfig)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
@@ -197,7 +194,6 @@ func (h *Handler) Dispatch(ctx context.Context, link *transport.Link) {
|
|||||||
common.Interrupt(link.Reader)
|
common.Interrupt(link.Reader)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
} else {
|
} else {
|
||||||
unchangedDomain := ob.Target.Address.Domain()
|
unchangedDomain := ob.Target.Address.Domain()
|
||||||
ob.Target.Address = net.IPAddress(ips[dice.Roll(len(ips))])
|
ob.Target.Address = net.IPAddress(ips[dice.Roll(len(ips))])
|
||||||
@@ -221,7 +217,7 @@ func (h *Handler) Dispatch(ctx context.Context, link *transport.Link) {
|
|||||||
if ob.Target.Network == net.Network_UDP && ob.Target.Port == 443 {
|
if ob.Target.Network == net.Network_UDP && ob.Target.Port == 443 {
|
||||||
switch h.udp443 {
|
switch h.udp443 {
|
||||||
case "reject":
|
case "reject":
|
||||||
test(errors.New("XUDP rejected UDP/443 traffic").AtInfo())
|
test(errors.New("XUDP rejected UDP/443 traffic"))
|
||||||
return
|
return
|
||||||
case "skip":
|
case "skip":
|
||||||
goto out
|
goto out
|
||||||
@@ -270,66 +266,26 @@ func (h *Handler) DestIpAddress() net.IP {
|
|||||||
|
|
||||||
// Dial implements internet.Dialer.
|
// Dial implements internet.Dialer.
|
||||||
func (h *Handler) Dial(ctx context.Context, dest net.Destination) (stat.Connection, error) {
|
func (h *Handler) Dial(ctx context.Context, dest net.Destination) (stat.Connection, error) {
|
||||||
if h.senderSettings != nil {
|
if h.senderSettings != nil && h.senderSettings.Via != nil {
|
||||||
|
|
||||||
if h.senderSettings.ProxySettings.HasTag() {
|
|
||||||
|
|
||||||
tag := h.senderSettings.ProxySettings.Tag
|
|
||||||
handler := h.outboundManager.GetHandler(tag)
|
|
||||||
if handler != nil {
|
|
||||||
errors.LogDebug(ctx, "proxying to ", tag, " for dest ", dest)
|
|
||||||
outbounds := session.OutboundsFromContext(ctx)
|
|
||||||
ctx = session.ContextWithOutbounds(ctx, append(outbounds, &session.Outbound{
|
|
||||||
Target: dest,
|
|
||||||
Tag: tag,
|
|
||||||
})) // add another outbound in session ctx
|
|
||||||
opts := pipe.OptionsFromContext(ctx)
|
|
||||||
uplinkReader, uplinkWriter := pipe.New(opts...)
|
|
||||||
downlinkReader, downlinkWriter := pipe.New(opts...)
|
|
||||||
|
|
||||||
go handler.Dispatch(ctx, &transport.Link{Reader: uplinkReader, Writer: downlinkWriter})
|
|
||||||
conn := cnc.NewConnection(cnc.ConnectionInputMulti(uplinkWriter), cnc.ConnectionOutputMulti(downlinkReader))
|
|
||||||
|
|
||||||
if config := tls.ConfigFromStreamSettings(h.streamSettings); config != nil {
|
|
||||||
tlsConfig := config.GetTLSConfig(tls.WithDestination(dest))
|
|
||||||
conn = tls.Client(conn, tlsConfig)
|
|
||||||
}
|
|
||||||
|
|
||||||
return h.getStatCouterConnection(conn), nil
|
|
||||||
}
|
|
||||||
|
|
||||||
errors.LogError(ctx, "failed to get outbound handler with tag: ", tag)
|
|
||||||
return nil, errors.New("failed to get outbound handler with tag: " + tag)
|
|
||||||
}
|
|
||||||
|
|
||||||
if h.senderSettings.Via != nil {
|
|
||||||
outbounds := session.OutboundsFromContext(ctx)
|
outbounds := session.OutboundsFromContext(ctx)
|
||||||
ob := outbounds[len(outbounds)-1]
|
ob := outbounds[len(outbounds)-1]
|
||||||
h.SetOutboundGateway(ctx, ob)
|
h.SetOutboundGateway(ctx, ob)
|
||||||
}
|
}
|
||||||
}
|
|
||||||
|
|
||||||
conn, err := internet.Dial(ctx, dest, h.streamSettings)
|
conn, err := internet.Dial(ctx, dest, h.streamSettings)
|
||||||
conn = h.getStatCouterConnection(conn)
|
conn = h.getStatCouterConnection(conn)
|
||||||
outbounds := session.OutboundsFromContext(ctx)
|
|
||||||
if outbounds != nil {
|
|
||||||
ob := outbounds[len(outbounds)-1]
|
|
||||||
ob.Conn = conn
|
|
||||||
} else {
|
|
||||||
// for Vision's pre-connect
|
|
||||||
}
|
|
||||||
return conn, err
|
return conn, err
|
||||||
}
|
}
|
||||||
|
|
||||||
func (h *Handler) SetOutboundGateway(ctx context.Context, ob *session.Outbound) {
|
func (h *Handler) SetOutboundGateway(ctx context.Context, ob *session.Outbound) {
|
||||||
if ob.Gateway == nil && h.senderSettings != nil && h.senderSettings.Via != nil && !h.senderSettings.ProxySettings.HasTag() && (h.streamSettings.SocketSettings == nil || len(h.streamSettings.SocketSettings.DialerProxy) == 0) {
|
if ob.Gateway == nil && h.senderSettings != nil && h.senderSettings.Via != nil &&
|
||||||
|
(h.streamSettings.SocketSettings == nil || len(h.streamSettings.SocketSettings.DialerProxy) == 0) {
|
||||||
var domain string
|
var domain string
|
||||||
addr := h.senderSettings.Via.AsAddress()
|
addr := h.senderSettings.Via.AsAddress()
|
||||||
domain = h.senderSettings.Via.GetDomain()
|
domain = h.senderSettings.Via.GetDomain()
|
||||||
switch {
|
switch {
|
||||||
case h.senderSettings.ViaCidr != "":
|
case h.senderSettings.ViaCidr != "":
|
||||||
ob.Gateway = ParseRandomIP(addr, h.senderSettings.ViaCidr)
|
ob.Gateway = ParseRandomIP(addr, h.senderSettings.ViaCidr)
|
||||||
|
|
||||||
case domain == "origin":
|
case domain == "origin":
|
||||||
if inbound := session.InboundFromContext(ctx); inbound != nil {
|
if inbound := session.InboundFromContext(ctx); inbound != nil {
|
||||||
if inbound.Local.IsValid() && inbound.Local.Address.Family().IsIP() {
|
if inbound.Local.IsValid() && inbound.Local.Address.Family().IsIP() {
|
||||||
@@ -344,11 +300,9 @@ func (h *Handler) SetOutboundGateway(ctx context.Context, ob *session.Outbound)
|
|||||||
errors.LogDebug(ctx, "use inbound source ip as sendthrough: ", inbound.Source.Address.String())
|
errors.LogDebug(ctx, "use inbound source ip as sendthrough: ", inbound.Source.Address.String())
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
// case addr.Family().IsDomain():
|
default: // case addr.Family().IsDomain():
|
||||||
default:
|
|
||||||
ob.Gateway = addr
|
ob.Gateway = addr
|
||||||
}
|
}
|
||||||
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -68,13 +68,13 @@ func (p *Portal) HandleConnection(ctx context.Context, link *transport.Link) err
|
|||||||
outbounds := session.OutboundsFromContext(ctx)
|
outbounds := session.OutboundsFromContext(ctx)
|
||||||
ob := outbounds[len(outbounds)-1]
|
ob := outbounds[len(outbounds)-1]
|
||||||
if ob == nil {
|
if ob == nil {
|
||||||
return errors.New("outbound metadata not found").AtError()
|
return errors.New("outbound metadata not found")
|
||||||
}
|
}
|
||||||
|
|
||||||
if isDomain(ob.Target, p.domain) {
|
if isDomain(ob.Target, p.domain) {
|
||||||
muxClient, err := mux.NewClientWorker(*link, mux.ClientStrategy{})
|
muxClient, err := mux.NewClientWorker(*link, mux.ClientStrategy{})
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return errors.New("failed to create mux client worker").Base(err).AtWarning()
|
return errors.New("failed to create mux client worker").Base(err)
|
||||||
}
|
}
|
||||||
|
|
||||||
worker, err := NewPortalWorker(muxClient)
|
worker, err := NewPortalWorker(muxClient)
|
||||||
|
|||||||
@@ -115,7 +115,7 @@ func (rr *RoutingRule) BuildCondition() (Condition, error) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
if conds.Len() == 0 {
|
if conds.Len() == 0 {
|
||||||
return nil, errors.New("this rule has no effective fields").AtWarning()
|
return nil, errors.New("this rule has no effective fields")
|
||||||
}
|
}
|
||||||
|
|
||||||
return conds, nil
|
return conds, nil
|
||||||
@@ -145,7 +145,7 @@ func (br *BalancingRule) Build(ohm outbound.Manager, dispatcher routing.Dispatch
|
|||||||
}
|
}
|
||||||
s, ok := i.(*StrategyLeastLoadConfig)
|
s, ok := i.(*StrategyLeastLoadConfig)
|
||||||
if !ok {
|
if !ok {
|
||||||
return nil, errors.New("not a StrategyLeastLoadConfig").AtError()
|
return nil, errors.New("not a StrategyLeastLoadConfig")
|
||||||
}
|
}
|
||||||
leastLoadStrategy := NewLeastLoadStrategy(s)
|
leastLoadStrategy := NewLeastLoadStrategy(s)
|
||||||
return &Balancer{
|
return &Balancer{
|
||||||
|
|||||||
+12
-2
@@ -587,6 +587,8 @@ type Config struct {
|
|||||||
DomainStrategy Config_DomainStrategy `protobuf:"varint,1,opt,name=domain_strategy,json=domainStrategy,proto3,enum=xray.app.router.Config_DomainStrategy" json:"domain_strategy,omitempty"`
|
DomainStrategy Config_DomainStrategy `protobuf:"varint,1,opt,name=domain_strategy,json=domainStrategy,proto3,enum=xray.app.router.Config_DomainStrategy" json:"domain_strategy,omitempty"`
|
||||||
Rule []*RoutingRule `protobuf:"bytes,2,rep,name=rule,proto3" json:"rule,omitempty"`
|
Rule []*RoutingRule `protobuf:"bytes,2,rep,name=rule,proto3" json:"rule,omitempty"`
|
||||||
BalancingRule []*BalancingRule `protobuf:"bytes,3,rep,name=balancing_rule,json=balancingRule,proto3" json:"balancing_rule,omitempty"`
|
BalancingRule []*BalancingRule `protobuf:"bytes,3,rep,name=balancing_rule,json=balancingRule,proto3" json:"balancing_rule,omitempty"`
|
||||||
|
// Absolute path to the Lua routing script.
|
||||||
|
Script string `protobuf:"bytes,4,opt,name=script,proto3" json:"script,omitempty"`
|
||||||
unknownFields protoimpl.UnknownFields
|
unknownFields protoimpl.UnknownFields
|
||||||
sizeCache protoimpl.SizeCache
|
sizeCache protoimpl.SizeCache
|
||||||
}
|
}
|
||||||
@@ -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,167 @@
|
|||||||
|
package router
|
||||||
|
|
||||||
|
import (
|
||||||
|
"runtime"
|
||||||
|
"strings"
|
||||||
|
|
||||||
|
"github.com/xtls/xray-core/common/errors"
|
||||||
|
xlua "github.com/xtls/xray-core/common/lua"
|
||||||
|
"github.com/xtls/xray-core/common/net"
|
||||||
|
"github.com/xtls/xray-core/features/routing"
|
||||||
|
lua "github.com/yuin/gopher-lua"
|
||||||
|
)
|
||||||
|
|
||||||
|
const (
|
||||||
|
luaContextType = "xray.router.Context"
|
||||||
|
luaAttributesType = "xray.router.Attributes"
|
||||||
|
)
|
||||||
|
|
||||||
|
// RegisterLua makes xray.router available to routing scripts.
|
||||||
|
func (r *Router) RegisterLua(L *lua.LState) {
|
||||||
|
registerLuaContext(L)
|
||||||
|
|
||||||
|
L.PreloadModule("xray.router", func(L *lua.LState) int {
|
||||||
|
module := L.NewTable()
|
||||||
|
|
||||||
|
module.RawSetString("NetworkUnknown", lua.LNumber(net.Network_Unknown))
|
||||||
|
module.RawSetString("NetworkTCP", lua.LNumber(net.Network_TCP))
|
||||||
|
module.RawSetString("NetworkUDP", lua.LNumber(net.Network_UDP))
|
||||||
|
module.RawSetString("NetworkUNIX", lua.LNumber(net.Network_UNIX))
|
||||||
|
module.RawSetString("LocalOS", lua.LString(runtime.GOOS))
|
||||||
|
|
||||||
|
module.RawSetString("PickOutbound", L.NewFunction(func(L *lua.LState) int {
|
||||||
|
tag, ok := L.Get(2).(lua.LString)
|
||||||
|
if !ok {
|
||||||
|
L.ArgError(2, "balancer tag must be a string")
|
||||||
|
return 0
|
||||||
|
}
|
||||||
|
balancer, found := (*r.balancers.Load())[string(tag)]
|
||||||
|
if !found {
|
||||||
|
xlua.PushNil(L)
|
||||||
|
xlua.PushError(L, errors.New("balancer ", tag, " not found"))
|
||||||
|
return 2
|
||||||
|
}
|
||||||
|
outboundTag, err := balancer.PickOutbound()
|
||||||
|
xlua.PushString(L, outboundTag)
|
||||||
|
xlua.PushError(L, err)
|
||||||
|
return 2
|
||||||
|
}))
|
||||||
|
|
||||||
|
module.RawSetString("FindProcess", L.NewFunction(func(L *lua.LState) int {
|
||||||
|
pid, name, path, err := findProcess(checkLuaContext(L), net.FindProcess)
|
||||||
|
xlua.PushNumber(L, pid)
|
||||||
|
xlua.PushString(L, name)
|
||||||
|
xlua.PushString(L, path)
|
||||||
|
xlua.PushError(L, err)
|
||||||
|
return 4
|
||||||
|
}))
|
||||||
|
|
||||||
|
L.Push(module)
|
||||||
|
return 1
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func registerLuaContext(L *lua.LState) {
|
||||||
|
attributes := L.NewTypeMetatable(luaAttributesType)
|
||||||
|
L.SetField(attributes, "__index", L.NewFunction(func(L *lua.LState) int {
|
||||||
|
values := L.CheckUserData(1).Value.(map[string]string)
|
||||||
|
key := L.CheckString(2)
|
||||||
|
if value, found := values[key]; found {
|
||||||
|
xlua.PushString(L, value)
|
||||||
|
} else {
|
||||||
|
xlua.PushNil(L)
|
||||||
|
}
|
||||||
|
return 1
|
||||||
|
}))
|
||||||
|
methods := L.NewTable()
|
||||||
|
L.SetFuncs(methods, map[string]lua.LGFunction{
|
||||||
|
"GetSourceIPs": func(L *lua.LState) int {
|
||||||
|
xlua.PushUserData(L, checkLuaContext(L).GetSourceIPs())
|
||||||
|
return 1
|
||||||
|
},
|
||||||
|
"GetTargetIPs": func(L *lua.LState) int {
|
||||||
|
xlua.PushUserData(L, checkLuaContext(L).GetTargetIPs())
|
||||||
|
return 1
|
||||||
|
},
|
||||||
|
"GetLocalIPs": func(L *lua.LState) int {
|
||||||
|
xlua.PushUserData(L, checkLuaContext(L).GetLocalIPs())
|
||||||
|
return 1
|
||||||
|
},
|
||||||
|
"GetAttributes": func(L *lua.LState) int {
|
||||||
|
values := L.NewUserData()
|
||||||
|
values.Value = checkLuaContext(L).GetAttributes()
|
||||||
|
L.SetMetatable(values, attributes)
|
||||||
|
L.Push(values)
|
||||||
|
return 1
|
||||||
|
},
|
||||||
|
})
|
||||||
|
L.SetField(L.NewTypeMetatable(luaContextType), "__index", methods)
|
||||||
|
}
|
||||||
|
|
||||||
|
func checkLuaContext(L *lua.LState) routing.Context {
|
||||||
|
ctx, ok := L.CheckUserData(1).Value.(routing.Context)
|
||||||
|
if !ok {
|
||||||
|
L.ArgError(1, "routing context expected")
|
||||||
|
}
|
||||||
|
return ctx
|
||||||
|
}
|
||||||
|
|
||||||
|
// callLuaHook invokes HandleRoute in the supplied state.
|
||||||
|
func (r *Router) callLuaHook(L *lua.LState, routeCtx routing.Context) (string, string, error) {
|
||||||
|
top := L.GetTop()
|
||||||
|
defer L.SetTop(top)
|
||||||
|
fn := L.GetGlobal("HandleRoute")
|
||||||
|
if fn.Type() != lua.LTFunction {
|
||||||
|
return "", "", errors.New("routing script must define HandleRoute(...)")
|
||||||
|
}
|
||||||
|
value := L.NewUserData()
|
||||||
|
value.Value = routeCtx
|
||||||
|
L.SetMetatable(value, L.GetTypeMetatable(luaContextType))
|
||||||
|
if err := L.CallByParam(lua.P{Fn: fn, NRet: 3, Protect: true},
|
||||||
|
value, lua.LString(routeCtx.GetInboundTag()), lua.LNumber(routeCtx.GetSourcePort()),
|
||||||
|
lua.LNumber(routeCtx.GetTargetPort()), lua.LNumber(routeCtx.GetLocalPort()),
|
||||||
|
lua.LString(strings.ToLower(routeCtx.GetTargetDomain())), lua.LNumber(routeCtx.GetNetwork()),
|
||||||
|
lua.LString(routeCtx.GetProtocol()), lua.LString(routeCtx.GetUser()),
|
||||||
|
lua.LNumber(routeCtx.GetVlessRoute()), lua.LBool(routeCtx.GetSkipDNSResolve())); err != nil {
|
||||||
|
return "", "", err
|
||||||
|
}
|
||||||
|
return readLuaRouteResult(L.Get(-3), L.Get(-2), L.Get(-1))
|
||||||
|
}
|
||||||
|
|
||||||
|
func readLuaRouteResult(tagValue, ruleValue, errorValue lua.LValue) (string, string, error) {
|
||||||
|
if err := xlua.ReadError(errorValue, "routing script error must be an error or string"); err != nil {
|
||||||
|
return "", "", err
|
||||||
|
}
|
||||||
|
tag, err := xlua.ReadOptionalString(tagValue, "routing script outboundTag must be a string or nil")
|
||||||
|
if err != nil || tag == "" {
|
||||||
|
return "", "", err
|
||||||
|
}
|
||||||
|
ruleTag, err := xlua.ReadOptionalString(ruleValue, "routing script ruleTag must be a string")
|
||||||
|
if err != nil {
|
||||||
|
return "", "", err
|
||||||
|
}
|
||||||
|
return tag, ruleTag, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
type processFinder func(string, string, uint16, string, uint16) (int, string, string, error)
|
||||||
|
|
||||||
|
func findProcess(ctx routing.Context, finder processFinder) (int, string, string, error) {
|
||||||
|
sources := ctx.GetSourceIPs()
|
||||||
|
if len(sources) == 0 {
|
||||||
|
return 0, "", "", errors.New("process lookup requires a source IP")
|
||||||
|
}
|
||||||
|
var network string
|
||||||
|
switch ctx.GetNetwork() {
|
||||||
|
case net.Network_TCP:
|
||||||
|
network = "tcp"
|
||||||
|
case net.Network_UDP:
|
||||||
|
network = "udp"
|
||||||
|
default:
|
||||||
|
return 0, "", "", errors.New("process lookup requires TCP or UDP")
|
||||||
|
}
|
||||||
|
targetIP, targetPort := "", uint16(0)
|
||||||
|
if targets := ctx.GetTargetIPs(); len(targets) > 0 {
|
||||||
|
targetIP, targetPort = targets[0].String(), uint16(ctx.GetTargetPort())
|
||||||
|
}
|
||||||
|
return finder(network, sources[0].String(), uint16(ctx.GetSourcePort()), targetIP, targetPort)
|
||||||
|
}
|
||||||
@@ -0,0 +1,300 @@
|
|||||||
|
package router
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
go_errors "errors"
|
||||||
|
"runtime"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/xtls/xray-core/common/geodata"
|
||||||
|
"github.com/xtls/xray-core/common/net"
|
||||||
|
"github.com/xtls/xray-core/common/protocol"
|
||||||
|
"github.com/xtls/xray-core/common/session"
|
||||||
|
"github.com/xtls/xray-core/features/routing"
|
||||||
|
routing_session "github.com/xtls/xray-core/features/routing/session"
|
||||||
|
lua "github.com/yuin/gopher-lua"
|
||||||
|
)
|
||||||
|
|
||||||
|
type luaRouteTestContext struct {
|
||||||
|
*routing_session.Context
|
||||||
|
sourceIPs, targetIPs, localIPs []net.IP
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *luaRouteTestContext) GetSourceIPs() []net.IP { return c.sourceIPs }
|
||||||
|
func (c *luaRouteTestContext) GetTargetIPs() []net.IP { return c.targetIPs }
|
||||||
|
func (c *luaRouteTestContext) GetLocalIPs() []net.IP { return c.localIPs }
|
||||||
|
|
||||||
|
func newLuaRouteTestContext() *luaRouteTestContext {
|
||||||
|
return &luaRouteTestContext{
|
||||||
|
Context: &routing_session.Context{
|
||||||
|
Inbound: &session.Inbound{
|
||||||
|
Tag: "in", VlessRoute: 4321,
|
||||||
|
Source: net.TCPDestination(net.LocalHostIP, 1234),
|
||||||
|
Local: net.TCPDestination(net.LocalHostIP, 5678),
|
||||||
|
User: &protocol.MemoryUser{Email: "user@example.com"},
|
||||||
|
},
|
||||||
|
Outbound: &session.Outbound{
|
||||||
|
Target: net.TCPDestination(net.LocalHostIP, 443),
|
||||||
|
RouteTarget: net.TCPDestination(net.DomainAddress("MiXeD.Example."), 443),
|
||||||
|
},
|
||||||
|
Content: &session.Content{
|
||||||
|
Protocol: "tls", Attributes: map[string]string{"key": "value"}, SkipDNSResolve: true,
|
||||||
|
},
|
||||||
|
},
|
||||||
|
sourceIPs: []net.IP{{127, 0, 0, 2}},
|
||||||
|
targetIPs: []net.IP{{127, 0, 0, 3}},
|
||||||
|
localIPs: []net.IP{{127, 0, 0, 1}},
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func newLuaRouterState(t *testing.T, script string) (*Router, *lua.LState) {
|
||||||
|
t.Helper()
|
||||||
|
r := new(Router)
|
||||||
|
if err := r.Init(context.Background(), &Config{}, nil, nil, nil); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
L := lua.NewState()
|
||||||
|
t.Cleanup(L.Close)
|
||||||
|
r.RegisterLua(L)
|
||||||
|
geodata.RegisterLua(L)
|
||||||
|
if err := L.DoString(script); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
return r, L
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestLuaRouteBinding(t *testing.T) {
|
||||||
|
r, L := newLuaRouterState(t, `
|
||||||
|
local router = require("xray.router")
|
||||||
|
local matcher = require("xray.geodata").BuildIPMatcher("127.0.0.0/8")
|
||||||
|
assert(router.NetworkUnknown == 0 and router.NetworkTCP == 2)
|
||||||
|
assert(router.NetworkUDP == 3 and router.NetworkUNIX == 4)
|
||||||
|
assert(router.BuildIPMatcher == nil and router.BuildDomainMatcher == nil)
|
||||||
|
function HandleRoute(ctx, inboundTag, sourcePort, targetPort, localPort,
|
||||||
|
targetDomain, network, protocol, user, vlessRoute, skipDNSResolve, ...)
|
||||||
|
assert(select("#", ...) == 0)
|
||||||
|
assert(inboundTag == "in" and sourcePort == 1234 and targetPort == 443 and localPort == 5678)
|
||||||
|
assert(targetDomain == "mixed.example." and network == router.NetworkTCP)
|
||||||
|
assert(protocol == "tls" and user == "user@example.com" and vlessRoute == 4321 and skipDNSResolve)
|
||||||
|
assert(ctx.GetNetwork == nil and ctx.Context == nil)
|
||||||
|
savedContext = ctx
|
||||||
|
sourceIPs, targetIPs, localIPs = ctx:GetSourceIPs(), ctx:GetTargetIPs(), ctx:GetLocalIPs()
|
||||||
|
attributes = ctx:GetAttributes()
|
||||||
|
assert(matcher:AnyMatch(sourceIPs) and matcher:AnyMatch(targetIPs) and matcher:AnyMatch(localIPs))
|
||||||
|
assert(attributes.key == "value" and attributes.missing == nil)
|
||||||
|
assert(not pcall(function() attributes.key = "changed" end))
|
||||||
|
return "out", "rule"
|
||||||
|
end`)
|
||||||
|
|
||||||
|
ctx := newLuaRouteTestContext()
|
||||||
|
tag, rule, err := r.callLuaHook(L, ctx)
|
||||||
|
if err != nil || tag != "out" || rule != "rule" {
|
||||||
|
t.Fatalf("hook = %q, %q, %v", tag, rule, err)
|
||||||
|
}
|
||||||
|
if L.GetGlobal("savedContext").(*lua.LUserData).Value != ctx {
|
||||||
|
t.Fatal("routing context was copied")
|
||||||
|
}
|
||||||
|
for _, tc := range []struct {
|
||||||
|
name string
|
||||||
|
want []net.IP
|
||||||
|
}{
|
||||||
|
{"sourceIPs", ctx.sourceIPs},
|
||||||
|
{"targetIPs", ctx.targetIPs},
|
||||||
|
{"localIPs", ctx.localIPs},
|
||||||
|
} {
|
||||||
|
got := L.GetGlobal(tc.name).(*lua.LUserData).Value.([]net.IP)
|
||||||
|
if &got[0] != &tc.want[0] {
|
||||||
|
t.Fatalf("%s storage was copied", tc.name)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
ctx.Content.Attributes["key"] = "updated"
|
||||||
|
L.SetGlobal("expectedOS", lua.LString(runtime.GOOS))
|
||||||
|
if err := L.DoString(`
|
||||||
|
assert(attributes.key == "updated")
|
||||||
|
assert(require("xray.router").LocalOS == expectedOS)`); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestLuaRouteResult(t *testing.T) {
|
||||||
|
nativeErr := go_errors.New("native failure")
|
||||||
|
for _, tc := range []struct {
|
||||||
|
name, body, tag, rule, wantErr string
|
||||||
|
native bool
|
||||||
|
}{
|
||||||
|
{name: "route", body: `return "out", "rule"`, tag: "out", rule: "rule"},
|
||||||
|
{name: "no match", body: `return nil`},
|
||||||
|
{name: "empty tag", body: `return ""`},
|
||||||
|
{name: "no match ignores rule", body: `return nil, false`},
|
||||||
|
{name: "empty tag ignores rule", body: `return "", false`},
|
||||||
|
{name: "missing rule", body: `return "out"`, tag: "out"},
|
||||||
|
{name: "invalid tag", body: `return 1`, wantErr: "outboundTag"},
|
||||||
|
{name: "invalid rule", body: `return "out", false`, wantErr: "ruleTag"},
|
||||||
|
{name: "string error", body: `return nil, nil, "script failure"`, wantErr: "script failure"},
|
||||||
|
{name: "native error", body: `return nil, nil, nativeError`, native: true},
|
||||||
|
{name: "error overrides invalid tags", body: `return false, false, nativeError`, native: true},
|
||||||
|
{name: "invalid error", body: `return "out", "rule", false`, wantErr: "error or string"},
|
||||||
|
{name: "wrong error userdata", body: `return "out", "rule", wrongError`, wantErr: "error or string"},
|
||||||
|
{name: "runtime error", body: `error("runtime failure")`, wantErr: "runtime failure"},
|
||||||
|
} {
|
||||||
|
t.Run(tc.name, func(t *testing.T) {
|
||||||
|
r, L := newLuaRouterState(t, "function HandleRoute() "+tc.body+" end")
|
||||||
|
value := L.NewUserData()
|
||||||
|
value.Value = nativeErr
|
||||||
|
L.SetGlobal("nativeError", value)
|
||||||
|
wrong := L.NewUserData()
|
||||||
|
wrong.Value = "not a native error"
|
||||||
|
L.SetGlobal("wrongError", wrong)
|
||||||
|
L.Push(lua.LTrue)
|
||||||
|
|
||||||
|
tag, rule, err := r.callLuaHook(L, &routing_session.Context{})
|
||||||
|
if tag != tc.tag || rule != tc.rule {
|
||||||
|
t.Fatalf("result = %q, %q, %v", tag, rule, err)
|
||||||
|
}
|
||||||
|
switch {
|
||||||
|
case tc.native:
|
||||||
|
if err != nativeErr {
|
||||||
|
t.Fatalf("error = %v, want original error", err)
|
||||||
|
}
|
||||||
|
case tc.wantErr != "":
|
||||||
|
if err == nil || !strings.Contains(err.Error(), tc.wantErr) {
|
||||||
|
t.Fatalf("error = %v, want %q", err, tc.wantErr)
|
||||||
|
}
|
||||||
|
case err != nil:
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if L.GetTop() != 1 || L.Get(1) != lua.LTrue {
|
||||||
|
t.Fatal("hook did not restore the stack")
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestLuaRouteCancellation(t *testing.T) {
|
||||||
|
r, L := newLuaRouterState(t, `function HandleRoute() while true do end end`)
|
||||||
|
ctx, cancel := context.WithCancel(context.Background())
|
||||||
|
cancel()
|
||||||
|
L.SetContext(ctx)
|
||||||
|
if _, _, err := r.callLuaHook(L, &routing_session.Context{}); err == nil {
|
||||||
|
t.Fatal("CallLuaHook did not stop after context cancellation")
|
||||||
|
}
|
||||||
|
if L.Context() != ctx || L.GetTop() != 0 {
|
||||||
|
t.Fatal("CallLuaHook did not restore the Lua state")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestFindProcess(t *testing.T) {
|
||||||
|
for _, tc := range []struct {
|
||||||
|
name, network, target string
|
||||||
|
targetPort uint16
|
||||||
|
modify func(*luaRouteTestContext)
|
||||||
|
wantErr bool
|
||||||
|
}{
|
||||||
|
{name: "TCP", network: "tcp", target: "127.0.0.3", targetPort: 443},
|
||||||
|
{name: "UDP", network: "udp", target: "127.0.0.3", targetPort: 443, modify: func(c *luaRouteTestContext) {
|
||||||
|
c.Outbound.Target.Network = net.Network_UDP
|
||||||
|
}},
|
||||||
|
{name: "domain target", network: "tcp", modify: func(c *luaRouteTestContext) { c.targetIPs = nil }},
|
||||||
|
{name: "missing source", modify: func(c *luaRouteTestContext) { c.sourceIPs = nil }, wantErr: true},
|
||||||
|
{name: "unsupported network", modify: func(c *luaRouteTestContext) {
|
||||||
|
c.Outbound.Target.Network = net.Network_UNIX
|
||||||
|
}, wantErr: true},
|
||||||
|
} {
|
||||||
|
t.Run(tc.name, func(t *testing.T) {
|
||||||
|
ctx := newLuaRouteTestContext()
|
||||||
|
if tc.modify != nil {
|
||||||
|
tc.modify(ctx)
|
||||||
|
}
|
||||||
|
called := false
|
||||||
|
pid, name, path, err := findProcess(ctx, func(network, source string, sourcePort uint16, target string, targetPort uint16) (int, string, string, error) {
|
||||||
|
called = true
|
||||||
|
if network != tc.network || source != "127.0.0.2" || sourcePort != 1234 || target != tc.target || targetPort != tc.targetPort {
|
||||||
|
t.Fatalf("endpoints = %s %s:%d -> %s:%d", network, source, sourcePort, target, targetPort)
|
||||||
|
}
|
||||||
|
return 42, "process", "/path/process", nil
|
||||||
|
})
|
||||||
|
if tc.wantErr {
|
||||||
|
if err == nil || called {
|
||||||
|
t.Fatalf("findProcess = %d, %q, %q, %v", pid, name, path, err)
|
||||||
|
}
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if err != nil || !called || pid != 42 || name != "process" || path != "/path/process" {
|
||||||
|
t.Fatalf("findProcess = %d, %q, %q, %v", pid, name, path, err)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// BenchmarkLuaRouteHookCall isolates a preloaded Lua hook and its routing context bridge.
|
||||||
|
// The direct case runs an equivalent native routing rule.
|
||||||
|
func BenchmarkLuaRouteHookCall(b *testing.B) {
|
||||||
|
r := new(Router)
|
||||||
|
if err := r.Init(context.Background(), &Config{Rule: []*RoutingRule{{
|
||||||
|
TargetTag: &RoutingRule_Tag{Tag: "out"},
|
||||||
|
RuleTag: "rule",
|
||||||
|
InboundTag: []string{"in"},
|
||||||
|
Networks: []net.Network{net.Network_TCP},
|
||||||
|
Ip: []*geodata.IPRule{{
|
||||||
|
Value: &geodata.IPRule_Custom{Custom: &geodata.CIDRRule{
|
||||||
|
Cidr: &geodata.CIDR{Ip: []byte{127, 0, 0, 0}, Prefix: 8},
|
||||||
|
}},
|
||||||
|
}},
|
||||||
|
}}}, nil, nil, nil); err != nil {
|
||||||
|
b.Fatal(err)
|
||||||
|
}
|
||||||
|
L := lua.NewState()
|
||||||
|
defer L.Close()
|
||||||
|
r.RegisterLua(L)
|
||||||
|
geodata.RegisterLua(L)
|
||||||
|
if err := L.DoString(`
|
||||||
|
local router = require("xray.router")
|
||||||
|
local matcher = require("xray.geodata").BuildIPMatcher("127.0.0.0/8")
|
||||||
|
function HandleRoute(ctx, inboundTag, sourcePort, targetPort, localPort,
|
||||||
|
targetDomain, network, protocol, user, vlessRoute, skipDNSResolve)
|
||||||
|
if inboundTag == "in" and network == router.NetworkTCP and matcher:AnyMatch(ctx:GetTargetIPs()) then
|
||||||
|
return "out", "rule"
|
||||||
|
end
|
||||||
|
end
|
||||||
|
`); err != nil {
|
||||||
|
b.Fatal(err)
|
||||||
|
}
|
||||||
|
|
||||||
|
L.SetContext(context.Background())
|
||||||
|
routeCtx := newLuaRouteTestContext()
|
||||||
|
for _, benchmark := range []struct {
|
||||||
|
name string
|
||||||
|
route func() (string, string, error)
|
||||||
|
}{
|
||||||
|
{"direct", func() (string, string, error) {
|
||||||
|
route, err := r.PickRoute(routeCtx)
|
||||||
|
if err != nil {
|
||||||
|
return "", "", err
|
||||||
|
}
|
||||||
|
return route.GetOutboundTag(), route.GetRuleTag(), nil
|
||||||
|
}},
|
||||||
|
{"lua_hook", func() (string, string, error) {
|
||||||
|
return r.callLuaHook(L, routeCtx)
|
||||||
|
}},
|
||||||
|
} {
|
||||||
|
b.Run(benchmark.name, func(b *testing.B) {
|
||||||
|
b.ReportAllocs()
|
||||||
|
b.ResetTimer()
|
||||||
|
var tag, rule string
|
||||||
|
var err error
|
||||||
|
for i := 0; i < b.N; i++ {
|
||||||
|
tag, rule, err = benchmark.route()
|
||||||
|
if err != nil {
|
||||||
|
b.Fatal(err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
b.StopTimer()
|
||||||
|
if tag != "out" || rule != "rule" {
|
||||||
|
b.Fatalf("route() = %q, %q; want out, rule", tag, rule)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
var _ routing.Context = (*luaRouteTestContext)(nil)
|
||||||
@@ -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,68 @@
|
|||||||
|
package router
|
||||||
|
|
||||||
|
import (
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/xtls/xray-core/app/dns"
|
||||||
|
"github.com/xtls/xray-core/common"
|
||||||
|
"github.com/xtls/xray-core/common/errors"
|
||||||
|
"github.com/xtls/xray-core/common/geodata"
|
||||||
|
"github.com/xtls/xray-core/common/log"
|
||||||
|
xlua "github.com/xtls/xray-core/common/lua"
|
||||||
|
"github.com/xtls/xray-core/features/routing"
|
||||||
|
lua "github.com/yuin/gopher-lua"
|
||||||
|
)
|
||||||
|
|
||||||
|
const scriptExecutionTimeout = 6 * time.Second
|
||||||
|
|
||||||
|
type scriptEngine struct {
|
||||||
|
router *Router
|
||||||
|
pool *xlua.Pool
|
||||||
|
}
|
||||||
|
|
||||||
|
func newScriptEngine(path string, router *Router) (*scriptEngine, error) {
|
||||||
|
program, err := xlua.CompileFile(path)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
e := &scriptEngine{router: router}
|
||||||
|
e.pool, err = xlua.NewPool(router.ctx, scriptExecutionTimeout, program.NewStateFactory(
|
||||||
|
scriptExecutionTimeout*20,
|
||||||
|
func(L *lua.LState) {
|
||||||
|
geodata.RegisterLua(L)
|
||||||
|
log.RegisterLua(L)
|
||||||
|
router.RegisterLua(L)
|
||||||
|
dns.RegisterLua(L, router.dns)
|
||||||
|
},
|
||||||
|
func(L *lua.LState) error {
|
||||||
|
if L.GetGlobal("HandleRoute").Type() != lua.LTFunction {
|
||||||
|
return errors.New("routing script must define HandleRoute(...)")
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}))
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
errors.LogInfo(router.ctx, "routing script initialized from ", path)
|
||||||
|
return e, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (e *scriptEngine) close() {
|
||||||
|
e.pool.Close()
|
||||||
|
}
|
||||||
|
|
||||||
|
func (e *scriptEngine) pickRoute(ctx routing.Context) (routing.Route, error) {
|
||||||
|
var tag, ruleTag string
|
||||||
|
err := e.pool.WithState(nil, 0, func(L *lua.LState) error {
|
||||||
|
var hookErr error
|
||||||
|
tag, ruleTag, hookErr = e.router.callLuaHook(L, ctx)
|
||||||
|
return hookErr
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
if tag == "" {
|
||||||
|
return nil, common.ErrNoClue
|
||||||
|
}
|
||||||
|
return &Route{Context: ctx, outboundTag: tag, ruleTag: ruleTag}, nil
|
||||||
|
}
|
||||||
@@ -0,0 +1,372 @@
|
|||||||
|
package router
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
stdnet "net"
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"sync"
|
||||||
|
"sync/atomic"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
wireDNS "github.com/miekg/dns"
|
||||||
|
"github.com/xtls/xray-core/app/dispatcher"
|
||||||
|
appdns "github.com/xtls/xray-core/app/dns"
|
||||||
|
"github.com/xtls/xray-core/app/proxyman"
|
||||||
|
_ "github.com/xtls/xray-core/app/proxyman/outbound"
|
||||||
|
"github.com/xtls/xray-core/common"
|
||||||
|
"github.com/xtls/xray-core/common/net"
|
||||||
|
"github.com/xtls/xray-core/common/serial"
|
||||||
|
"github.com/xtls/xray-core/core"
|
||||||
|
featureDNS "github.com/xtls/xray-core/features/dns"
|
||||||
|
"github.com/xtls/xray-core/features/outbound"
|
||||||
|
"github.com/xtls/xray-core/features/routing"
|
||||||
|
routing_session "github.com/xtls/xray-core/features/routing/session"
|
||||||
|
"github.com/xtls/xray-core/proxy/blackhole"
|
||||||
|
"github.com/xtls/xray-core/proxy/freedom"
|
||||||
|
)
|
||||||
|
|
||||||
|
type luaRouteDNSClient struct {
|
||||||
|
featureDNS.Client
|
||||||
|
lookup func(string, featureDNS.IPOption) ([]net.IP, uint32, error)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (d *luaRouteDNSClient) LookupIP(domain string, option featureDNS.IPOption) ([]net.IP, uint32, error) {
|
||||||
|
return d.lookup(domain, option)
|
||||||
|
}
|
||||||
|
|
||||||
|
type luaRouteOutboundManager struct{ outbound.Manager }
|
||||||
|
|
||||||
|
func (*luaRouteOutboundManager) Select(selectors []string) []string { return selectors }
|
||||||
|
|
||||||
|
func writeRouteScript(t *testing.T, script string) string {
|
||||||
|
t.Helper()
|
||||||
|
path := filepath.Join(t.TempDir(), "route.lua")
|
||||||
|
if err := os.WriteFile(path, []byte(script), 0o600); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
return path
|
||||||
|
}
|
||||||
|
|
||||||
|
func startLuaRouter(t *testing.T, script string, d featureDNS.Client, config *Config) *Router {
|
||||||
|
t.Helper()
|
||||||
|
if config == nil {
|
||||||
|
config = &Config{}
|
||||||
|
}
|
||||||
|
config.Script = writeRouteScript(t, script)
|
||||||
|
r := new(Router)
|
||||||
|
if err := r.Init(context.Background(), config, d, &luaRouteOutboundManager{}, nil); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if err := r.Start(); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
t.Cleanup(func() {
|
||||||
|
if err := r.Close(); err != nil {
|
||||||
|
t.Error(err)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
return r
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRouterScriptStartup(t *testing.T) {
|
||||||
|
for _, tc := range []struct{ name, script string }{
|
||||||
|
{"syntax error", "function HandleRoute("},
|
||||||
|
{"missing hook", "value = 1"},
|
||||||
|
{"initialization error", `error("setup failed")`},
|
||||||
|
} {
|
||||||
|
t.Run(tc.name, func(t *testing.T) {
|
||||||
|
r := new(Router)
|
||||||
|
if err := r.Init(context.Background(), &Config{Script: writeRouteScript(t, tc.script)}, nil, nil, nil); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
defer r.Close()
|
||||||
|
if err := r.Start(); err == nil {
|
||||||
|
t.Fatal("Start accepted an invalid routing script")
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRouterScriptRouting(t *testing.T) {
|
||||||
|
var dnsCalls atomic.Int32
|
||||||
|
d := &luaRouteDNSClient{lookup: func(string, featureDNS.IPOption) ([]net.IP, uint32, error) {
|
||||||
|
dnsCalls.Add(1)
|
||||||
|
return []net.IP{{1, 2, 3, 4}}, 60, nil
|
||||||
|
}}
|
||||||
|
r := startLuaRouter(t, `
|
||||||
|
function HandleRoute(ctx, inbound)
|
||||||
|
if inbound == "miss" then return nil end
|
||||||
|
return "lua-out", "lua-rule"
|
||||||
|
end`, d, &Config{
|
||||||
|
DomainStrategy: Config_IpOnDemand,
|
||||||
|
Rule: []*RoutingRule{{
|
||||||
|
TargetTag: &RoutingRule_Tag{Tag: "json-out"},
|
||||||
|
Networks: []net.Network{net.Network_TCP},
|
||||||
|
}},
|
||||||
|
})
|
||||||
|
ctx := newLuaRouteTestContext()
|
||||||
|
ctx.Content.SkipDNSResolve = false
|
||||||
|
route, err := r.PickRoute(ctx)
|
||||||
|
if err != nil || route.GetOutboundTag() != "lua-out" || route.GetRuleTag() != "lua-rule" || route.(*Route).Context != ctx {
|
||||||
|
t.Fatalf("route = %v, %v", route, err)
|
||||||
|
}
|
||||||
|
ctx.Inbound.Tag = "miss"
|
||||||
|
if route, err := r.PickRoute(ctx); route != nil || err != common.ErrNoClue {
|
||||||
|
t.Fatalf("miss = %v, %v", route, err)
|
||||||
|
}
|
||||||
|
if dnsCalls.Load() != 0 {
|
||||||
|
t.Fatal("script routing implicitly resolved DNS")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRouterScriptModules(t *testing.T) {
|
||||||
|
ips := []net.IP{{127, 0, 0, 7}}
|
||||||
|
calls := 0
|
||||||
|
d := &luaRouteDNSClient{lookup: func(domain string, option featureDNS.IPOption) ([]net.IP, uint32, error) {
|
||||||
|
calls++
|
||||||
|
if domain != "mixed.example." || !option.IPv4Enable || option.IPv6Enable || !option.FakeEnable {
|
||||||
|
t.Fatalf("dns.Query arguments = %q, %+v", domain, option)
|
||||||
|
}
|
||||||
|
return ips, 17, nil
|
||||||
|
}}
|
||||||
|
r := startLuaRouter(t, `
|
||||||
|
local dns = require("xray.dns")
|
||||||
|
local matcher = require("xray.geodata").BuildIPMatcher("127.0.0.0/8")
|
||||||
|
assert(dns.Servers == nil and type(dns.Query) == "function")
|
||||||
|
assert(type(require("xray.log").Info) == "function")
|
||||||
|
function HandleRoute(ctx, inbound, sourcePort, targetPort, localPort, domain)
|
||||||
|
local ips, ttl, err = dns.Query(domain, true, false, true)
|
||||||
|
assert(not err and ttl == 17)
|
||||||
|
assert(matcher:AnyMatch(ips) and matcher:AnyMatch(ctx:GetTargetIPs()))
|
||||||
|
return "out"
|
||||||
|
end`, d, nil)
|
||||||
|
if _, err := r.PickRoute(newLuaRouteTestContext()); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if calls != 1 {
|
||||||
|
t.Fatalf("DNS calls = %d, want 1", calls)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRouterScriptBalancerReload(t *testing.T) {
|
||||||
|
config := func(tag string) *Config {
|
||||||
|
return &Config{BalancingRule: []*BalancingRule{{
|
||||||
|
Tag: "balance", Strategy: "roundrobin", OutboundSelector: []string{tag},
|
||||||
|
}}}
|
||||||
|
}
|
||||||
|
r := startLuaRouter(t, `
|
||||||
|
local router = require("xray.router")
|
||||||
|
function HandleRoute()
|
||||||
|
local tag, err = router:PickOutbound("balance")
|
||||||
|
return tag, "balanced", err
|
||||||
|
end`, nil, config("old"))
|
||||||
|
pick := func(want string) {
|
||||||
|
t.Helper()
|
||||||
|
route, err := r.PickRoute(&routing_session.Context{})
|
||||||
|
if err != nil || route.GetOutboundTag() != want || route.GetRuleTag() != "balanced" {
|
||||||
|
t.Fatalf("route = %v, %v, want %q", route, err, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
pick("old")
|
||||||
|
if err := r.SetOverrideTarget("balance", "override"); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
pick("override")
|
||||||
|
if err := r.SetOverrideTarget("balance", ""); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if err := r.ReloadRules(config("new"), false); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
pick("new")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRouterScriptConcurrentBalancerReload(t *testing.T) {
|
||||||
|
config := func(tag string) *Config {
|
||||||
|
return &Config{BalancingRule: []*BalancingRule{{
|
||||||
|
Tag: "balance", Strategy: "roundrobin", OutboundSelector: []string{tag},
|
||||||
|
}}}
|
||||||
|
}
|
||||||
|
r := startLuaRouter(t, `
|
||||||
|
local router = require("xray.router")
|
||||||
|
function HandleRoute()
|
||||||
|
local tag, err = router:PickOutbound("balance")
|
||||||
|
return tag, nil, err
|
||||||
|
end`, nil, config("a"))
|
||||||
|
|
||||||
|
var wg sync.WaitGroup
|
||||||
|
for range 4 {
|
||||||
|
wg.Go(func() {
|
||||||
|
for range 20 {
|
||||||
|
route, err := r.PickRoute(&routing_session.Context{})
|
||||||
|
if err != nil {
|
||||||
|
t.Errorf("PickRoute: %v", err)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if tag := route.GetOutboundTag(); tag != "a" && tag != "b" {
|
||||||
|
t.Errorf("unexpected tag %q", tag)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
wg.Go(func() {
|
||||||
|
for range 20 {
|
||||||
|
for _, tag := range []string{"a", "b"} {
|
||||||
|
if err := r.ReloadRules(config(tag), false); err != nil {
|
||||||
|
t.Error(err)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
})
|
||||||
|
wg.Wait()
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRouterScriptStateReuse(t *testing.T) {
|
||||||
|
r := startLuaRouter(t, `
|
||||||
|
local calls = 0
|
||||||
|
function HandleRoute(ctx, inbound)
|
||||||
|
calls = calls + 1
|
||||||
|
if inbound == "miss" then return nil end
|
||||||
|
if inbound == "fail" then error("failed") end
|
||||||
|
return tostring(calls)
|
||||||
|
end`, nil, nil)
|
||||||
|
ctx := newLuaRouteTestContext()
|
||||||
|
pick := func(want string) {
|
||||||
|
t.Helper()
|
||||||
|
route, err := r.PickRoute(ctx)
|
||||||
|
if err != nil || route.GetOutboundTag() != want {
|
||||||
|
t.Fatalf("route = %v, %v, want %q", route, err, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
pick("1")
|
||||||
|
ctx.Inbound.Tag = "miss"
|
||||||
|
if _, err := r.PickRoute(ctx); err != common.ErrNoClue {
|
||||||
|
t.Fatalf("miss = %v", err)
|
||||||
|
}
|
||||||
|
ctx.Inbound.Tag = "in"
|
||||||
|
pick("3")
|
||||||
|
ctx.Inbound.Tag = "fail"
|
||||||
|
if _, err := r.PickRoute(ctx); err == nil {
|
||||||
|
t.Fatal("script error was ignored")
|
||||||
|
}
|
||||||
|
ctx.Inbound.Tag = "in"
|
||||||
|
pick("1")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRouterScriptDNSDispatcherReentry(t *testing.T) {
|
||||||
|
conn, err := stdnet.ListenPacket("udp4", "127.0.0.1:0")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
port := conn.LocalAddr().(*stdnet.UDPAddr).Port
|
||||||
|
ready, stopped := make(chan struct{}), make(chan error, 1)
|
||||||
|
var queries atomic.Int32
|
||||||
|
server := &wireDNS.Server{
|
||||||
|
PacketConn: conn,
|
||||||
|
NotifyStartedFunc: func() {
|
||||||
|
close(ready)
|
||||||
|
},
|
||||||
|
Handler: wireDNS.HandlerFunc(func(w wireDNS.ResponseWriter, query *wireDNS.Msg) {
|
||||||
|
queries.Add(1)
|
||||||
|
response := new(wireDNS.Msg).SetReply(query)
|
||||||
|
for _, question := range query.Question {
|
||||||
|
if question.Name == "nested.example." && question.Qtype == wireDNS.TypeA {
|
||||||
|
response.Answer = append(response.Answer, &wireDNS.A{
|
||||||
|
Hdr: wireDNS.RR_Header{Name: question.Name, Rrtype: wireDNS.TypeA, Class: wireDNS.ClassINET, Ttl: 60},
|
||||||
|
A: stdnet.IP{127, 0, 0, 7},
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if err := w.WriteMsg(response); err != nil {
|
||||||
|
t.Error(err)
|
||||||
|
}
|
||||||
|
}),
|
||||||
|
}
|
||||||
|
go func() { stopped <- server.ActivateAndServe() }()
|
||||||
|
defer func() {
|
||||||
|
server.Shutdown()
|
||||||
|
select {
|
||||||
|
case err := <-stopped:
|
||||||
|
if err != nil {
|
||||||
|
t.Error(err)
|
||||||
|
}
|
||||||
|
case <-time.After(3 * time.Second):
|
||||||
|
t.Error("DNS server did not stop")
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
select {
|
||||||
|
case <-ready:
|
||||||
|
case err := <-stopped:
|
||||||
|
t.Fatalf("DNS server startup: %v", err)
|
||||||
|
case <-time.After(3 * time.Second):
|
||||||
|
t.Fatal("DNS server did not start")
|
||||||
|
}
|
||||||
|
|
||||||
|
dnsScript := writeRouteScript(t, `
|
||||||
|
local server = require("xray.dns").Servers[1]
|
||||||
|
function HandleDNSQuery(domain, ipv4, ipv6, fake)
|
||||||
|
return server:Query(domain, ipv4, ipv6, fake)
|
||||||
|
end`)
|
||||||
|
routerScript := writeRouteScript(t, `
|
||||||
|
local router = require("xray.router")
|
||||||
|
local dns = require("xray.dns")
|
||||||
|
local matcher = require("xray.geodata").BuildIPMatcher("127.0.0.7")
|
||||||
|
local active = false
|
||||||
|
function HandleRoute(ctx, inbound, sourcePort, targetPort, localPort, domain, network,
|
||||||
|
protocol, user, vlessRoute, skipDNSResolve)
|
||||||
|
assert(not active, "borrowed Router VM reentered")
|
||||||
|
if inbound == "dns" then
|
||||||
|
assert(network == router.NetworkUDP and skipDNSResolve == false)
|
||||||
|
return "direct", "dns-route"
|
||||||
|
end
|
||||||
|
active = true
|
||||||
|
local ips, ttl, err = dns.Query("nested.example", true, false, false)
|
||||||
|
assert(not err and matcher:AnyMatch(ips) and active)
|
||||||
|
active = false
|
||||||
|
return "direct", "outer-route"
|
||||||
|
end`)
|
||||||
|
instance, err := core.New(&core.Config{
|
||||||
|
App: []*serial.TypedMessage{
|
||||||
|
serial.ToTypedMessage(&appdns.Config{
|
||||||
|
Tag: "dns", Script: dnsScript, DisableCache: true,
|
||||||
|
NameServer: []*appdns.NameServer{{
|
||||||
|
Id: "upstream", TimeoutMs: 1000,
|
||||||
|
Address: &net.Endpoint{
|
||||||
|
Network: net.Network_UDP,
|
||||||
|
Address: &net.IPOrDomain{Address: &net.IPOrDomain_Ip{Ip: []byte{127, 0, 0, 1}}},
|
||||||
|
Port: uint32(port),
|
||||||
|
},
|
||||||
|
}},
|
||||||
|
}),
|
||||||
|
serial.ToTypedMessage(&Config{Script: routerScript}),
|
||||||
|
serial.ToTypedMessage(&dispatcher.Config{}),
|
||||||
|
serial.ToTypedMessage(&proxyman.OutboundConfig{}),
|
||||||
|
},
|
||||||
|
Outbound: []*core.OutboundHandlerConfig{
|
||||||
|
{Tag: "default", ProxySettings: serial.ToTypedMessage(&blackhole.Config{})},
|
||||||
|
{Tag: "direct", ProxySettings: serial.ToTypedMessage(&freedom.Config{
|
||||||
|
FinalRules: []*freedom.FinalRuleConfig{{Action: freedom.RuleAction_Allow}},
|
||||||
|
})},
|
||||||
|
},
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
defer instance.Close()
|
||||||
|
if err := instance.Start(); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
r := instance.GetFeature(routing.RouterType()).(*Router)
|
||||||
|
route, err := r.PickRoute(newLuaRouteTestContext())
|
||||||
|
if err != nil || route.GetOutboundTag() != "direct" || route.GetRuleTag() != "outer-route" {
|
||||||
|
t.Fatalf("nested DNS routing = %v, %v", route, err)
|
||||||
|
}
|
||||||
|
if queries.Load() == 0 {
|
||||||
|
t.Fatal("DNS query did not pass through the dispatcher")
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -6,6 +6,7 @@ import (
|
|||||||
|
|
||||||
type windowsReader struct {
|
type windowsReader struct {
|
||||||
bufs []syscall.WSABuf
|
bufs []syscall.WSABuf
|
||||||
|
ready bool
|
||||||
}
|
}
|
||||||
|
|
||||||
func (r *windowsReader) Init(bs []*Buffer) {
|
func (r *windowsReader) Init(bs []*Buffer) {
|
||||||
@@ -15,6 +16,7 @@ func (r *windowsReader) Init(bs []*Buffer) {
|
|||||||
for _, b := range bs {
|
for _, b := range bs {
|
||||||
r.bufs = append(r.bufs, syscall.WSABuf{Len: uint32(Size), Buf: &b.v[0]})
|
r.bufs = append(r.bufs, syscall.WSABuf{Len: uint32(Size), Buf: &b.v[0]})
|
||||||
}
|
}
|
||||||
|
r.ready = false
|
||||||
}
|
}
|
||||||
|
|
||||||
func (r *windowsReader) Clear() {
|
func (r *windowsReader) Clear() {
|
||||||
@@ -25,6 +27,14 @@ func (r *windowsReader) Clear() {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (r *windowsReader) Read(fd uintptr) int32 {
|
func (r *windowsReader) Read(fd uintptr) int32 {
|
||||||
|
// On the first invocation, we return -1 to indicate "not ready"
|
||||||
|
// to make rawConn.Read wait for readability using the runtime's own mechanism
|
||||||
|
// because syscall.WSARecv() is a blocking call when used with nil OVERLAPPED
|
||||||
|
if !r.ready {
|
||||||
|
r.ready = true
|
||||||
|
return -1
|
||||||
|
}
|
||||||
|
|
||||||
var nBytes uint32
|
var nBytes uint32
|
||||||
var flags uint32
|
var flags uint32
|
||||||
err := syscall.WSARecv(syscall.Handle(fd), &r.bufs[0], uint32(len(r.bufs)), &nBytes, &flags, nil, nil)
|
err := syscall.WSARecv(syscall.Handle(fd), &r.bufs[0], uint32(len(r.bufs)), &nBytes, &flags, nil, nil)
|
||||||
|
|||||||
@@ -118,7 +118,9 @@ func (w *BufferedWriter) Write(b []byte) (int, error) {
|
|||||||
|
|
||||||
nBytes, err := w.buffer.Write(b)
|
nBytes, err := w.buffer.Write(b)
|
||||||
totalBytes += nBytes
|
totalBytes += nBytes
|
||||||
if err != nil {
|
|
||||||
|
// ErrBufferFull means a partial write, so flush below and continue
|
||||||
|
if err != nil && err != ErrBufferFull {
|
||||||
return totalBytes, err
|
return totalBytes, err
|
||||||
}
|
}
|
||||||
if !w.buffered || w.buffer.IsFull() {
|
if !w.buffered || w.buffer.IsFull() {
|
||||||
|
|||||||
@@ -10,12 +10,12 @@ import (
|
|||||||
|
|
||||||
// [,)
|
// [,)
|
||||||
func RandBetween(from int64, to int64) int64 {
|
func RandBetween(from int64, to int64) int64 {
|
||||||
if from == to {
|
|
||||||
return from
|
|
||||||
}
|
|
||||||
if from > to {
|
if from > to {
|
||||||
from, to = to, from
|
from, to = to, from
|
||||||
}
|
}
|
||||||
|
if d := to - from; d == 0 || d == 1 {
|
||||||
|
return from
|
||||||
|
}
|
||||||
bigInt, _ := rand.Int(rand.Reader, big.NewInt(to-from))
|
bigInt, _ := rand.Int(rand.Reader, big.NewInt(to-from))
|
||||||
return from + bigInt.Int64()
|
return from + bigInt.Int64()
|
||||||
}
|
}
|
||||||
|
|||||||
+4
-56
@@ -18,17 +18,12 @@ type hasInnerError interface {
|
|||||||
Unwrap() error
|
Unwrap() error
|
||||||
}
|
}
|
||||||
|
|
||||||
type hasSeverity interface {
|
|
||||||
Severity() log.Severity
|
|
||||||
}
|
|
||||||
|
|
||||||
// Error is an error object with underlying error.
|
// Error is an error object with underlying error.
|
||||||
type Error struct {
|
type Error struct {
|
||||||
prefix []interface{}
|
prefix []interface{}
|
||||||
message []interface{}
|
message []interface{}
|
||||||
caller string
|
caller string
|
||||||
inner error
|
inner error
|
||||||
severity log.Severity
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// Error implements error.Error().
|
// Error implements error.Error().
|
||||||
@@ -69,46 +64,6 @@ func (err *Error) Base(e error) *Error {
|
|||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
func (err *Error) atSeverity(s log.Severity) *Error {
|
|
||||||
err.severity = s
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
|
|
||||||
func (err *Error) Severity() log.Severity {
|
|
||||||
if err.inner == nil {
|
|
||||||
return err.severity
|
|
||||||
}
|
|
||||||
|
|
||||||
if s, ok := err.inner.(hasSeverity); ok {
|
|
||||||
as := s.Severity()
|
|
||||||
if as < err.severity {
|
|
||||||
return as
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
return err.severity
|
|
||||||
}
|
|
||||||
|
|
||||||
// AtDebug sets the severity to debug.
|
|
||||||
func (err *Error) AtDebug() *Error {
|
|
||||||
return err.atSeverity(log.Severity_Debug)
|
|
||||||
}
|
|
||||||
|
|
||||||
// AtInfo sets the severity to info.
|
|
||||||
func (err *Error) AtInfo() *Error {
|
|
||||||
return err.atSeverity(log.Severity_Info)
|
|
||||||
}
|
|
||||||
|
|
||||||
// AtWarning sets the severity to warning.
|
|
||||||
func (err *Error) AtWarning() *Error {
|
|
||||||
return err.atSeverity(log.Severity_Warning)
|
|
||||||
}
|
|
||||||
|
|
||||||
// AtError sets the severity to error.
|
|
||||||
func (err *Error) AtError() *Error {
|
|
||||||
return err.atSeverity(log.Severity_Error)
|
|
||||||
}
|
|
||||||
|
|
||||||
// String returns the string representation of this error.
|
// String returns the string representation of this error.
|
||||||
func (err *Error) String() string {
|
func (err *Error) String() string {
|
||||||
return err.Error()
|
return err.Error()
|
||||||
@@ -133,7 +88,6 @@ func New(msg ...interface{}) *Error {
|
|||||||
}
|
}
|
||||||
return &Error{
|
return &Error{
|
||||||
message: msg,
|
message: msg,
|
||||||
severity: log.Severity_Info,
|
|
||||||
caller: details,
|
caller: details,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -171,6 +125,9 @@ func LogErrorInner(ctx context.Context, inner error, msg ...interface{}) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func doLog(ctx context.Context, inner error, severity log.Severity, msg ...interface{}) {
|
func doLog(ctx context.Context, inner error, severity log.Severity, msg ...interface{}) {
|
||||||
|
if log.GetSeverity() < severity {
|
||||||
|
return
|
||||||
|
}
|
||||||
pc, _, _, _ := runtime.Caller(2)
|
pc, _, _, _ := runtime.Caller(2)
|
||||||
details := runtime.FuncForPC(pc).Name()
|
details := runtime.FuncForPC(pc).Name()
|
||||||
if len(details) >= trim {
|
if len(details) >= trim {
|
||||||
@@ -182,7 +139,6 @@ func doLog(ctx context.Context, inner error, severity log.Severity, msg ...inter
|
|||||||
}
|
}
|
||||||
err := &Error{
|
err := &Error{
|
||||||
message: msg,
|
message: msg,
|
||||||
severity: severity,
|
|
||||||
caller: details,
|
caller: details,
|
||||||
inner: inner,
|
inner: inner,
|
||||||
}
|
}
|
||||||
@@ -193,7 +149,7 @@ func doLog(ctx context.Context, inner error, severity log.Severity, msg ...inter
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
log.Record(&log.GeneralMessage{
|
log.Record(&log.GeneralMessage{
|
||||||
Severity: GetSeverity(err),
|
Severity: severity,
|
||||||
Content: err,
|
Content: err,
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
@@ -217,11 +173,3 @@ L:
|
|||||||
}
|
}
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
// GetSeverity returns the actual severity of the error, including inner errors.
|
|
||||||
func GetSeverity(err error) log.Severity {
|
|
||||||
if s, ok := err.(hasSeverity); ok {
|
|
||||||
return s.Severity()
|
|
||||||
}
|
|
||||||
return log.Severity_Info
|
|
||||||
}
|
|
||||||
|
|||||||
@@ -7,30 +7,21 @@ import (
|
|||||||
|
|
||||||
"github.com/google/go-cmp/cmp"
|
"github.com/google/go-cmp/cmp"
|
||||||
. "github.com/xtls/xray-core/common/errors"
|
. "github.com/xtls/xray-core/common/errors"
|
||||||
"github.com/xtls/xray-core/common/log"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
func TestError(t *testing.T) {
|
func TestError(t *testing.T) {
|
||||||
err := New("TestError")
|
err := New("TestError")
|
||||||
if v := GetSeverity(err); v != log.Severity_Info {
|
if v := err.Error(); !strings.Contains(v, "TestError") {
|
||||||
t.Error("severity: ", v)
|
t.Error("error: ", v)
|
||||||
}
|
}
|
||||||
|
|
||||||
err = New("TestError2").Base(io.EOF)
|
err = New("TestError2").Base(io.EOF)
|
||||||
if v := GetSeverity(err); v != log.Severity_Info {
|
if v := err.Error(); !strings.Contains(v, "EOF") {
|
||||||
t.Error("severity: ", v)
|
t.Error("error: ", v)
|
||||||
}
|
}
|
||||||
|
|
||||||
err = New("TestError3").Base(io.EOF).AtWarning()
|
err = New("TestError3").Base(io.EOF)
|
||||||
if v := GetSeverity(err); v != log.Severity_Warning {
|
err = New("TestError4").Base(err)
|
||||||
t.Error("severity: ", v)
|
|
||||||
}
|
|
||||||
|
|
||||||
err = New("TestError4").Base(io.EOF).AtWarning()
|
|
||||||
err = New("TestError5").Base(err)
|
|
||||||
if v := GetSeverity(err); v != log.Severity_Warning {
|
|
||||||
t.Error("severity: ", v)
|
|
||||||
}
|
|
||||||
if v := err.Error(); !strings.Contains(v, "EOF") {
|
if v := err.Error(); !strings.Contains(v, "EOF") {
|
||||||
t.Error("error: ", v)
|
t.Error("error: ", v)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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()
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
+210
-59
@@ -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 {
|
||||||
|
if err == io.EOF && len(prefix) > 0 {
|
||||||
|
err = io.ErrUnexpectedEOF // as io.ReadFull
|
||||||
}
|
}
|
||||||
prefix := prefixBuf[:prefixL]
|
return 0, err
|
||||||
if _, err := io.ReadFull(br, prefix); err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
}
|
||||||
|
if bodyL >= need && len(prefix) >= need && int(prefix[1]) == codeL && bytes.Equal(prefix[2:], code) {
|
||||||
match := false
|
return bodyL, nil
|
||||||
if bodyL >= need {
|
|
||||||
if int(prefix[1]) == codeL && bytes.Equal(prefix[2:need], code) {
|
|
||||||
if !readBody {
|
|
||||||
return nil, nil
|
|
||||||
}
|
|
||||||
match = true
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
remain := bodyL - prefixL
|
|
||||||
if match {
|
|
||||||
out := make([]byte, bodyL)
|
|
||||||
copy(out, prefix)
|
|
||||||
if remain > 0 {
|
|
||||||
if _, err := io.ReadFull(br, out[prefixL:]); err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return out, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
if remain > 0 {
|
|
||||||
if _, err := br.Discard(remain); err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
}
|
||||||
|
if _, err := br.Discard(bodyL); err != nil {
|
||||||
|
return 0, 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)
|
|
||||||
|
type siteDecoder struct {
|
||||||
|
want []string
|
||||||
|
has []bool
|
||||||
|
fn func(Domain_Type, []byte)
|
||||||
|
}
|
||||||
|
|
||||||
|
func newSiteDecoder(attrs string, fn func(Domain_Type, []byte)) *siteDecoder {
|
||||||
|
d := &siteDecoder{fn: fn}
|
||||||
|
if attrs != "" {
|
||||||
|
d.want = strings.Split(attrs, "@")
|
||||||
|
d.has = make([]bool, len(d.want))
|
||||||
|
}
|
||||||
|
return d
|
||||||
|
}
|
||||||
|
|
||||||
|
// decode walks the whole fields at the start of b, a part of an encoded GeoSite (see geodat.proto),
|
||||||
|
// calls fn for every domain that has all attrs and returns how many bytes it used. A field cut off
|
||||||
|
// by the end of b is an error unless more is set. It accepts and rejects what proto.Unmarshal does.
|
||||||
|
func (d *siteDecoder) decode(b []byte, more bool) (int, error) {
|
||||||
|
used := 0
|
||||||
|
for used < len(b) {
|
||||||
|
f, n, err := consumeField(b[used:])
|
||||||
|
if err == io.ErrUnexpectedEOF && more {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
if err != nil {
|
||||||
|
return used, err
|
||||||
|
}
|
||||||
|
used += n
|
||||||
|
if f.typ != protowire.BytesType {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
switch f.num {
|
||||||
|
case 1: // code
|
||||||
|
if !utf8.Valid(f.v) {
|
||||||
|
return used, errInvalidUTF8
|
||||||
|
}
|
||||||
|
case 2: // domain
|
||||||
|
t, value, err := decodeDomain(f.v, d.want, d.has)
|
||||||
|
if err != nil {
|
||||||
|
return used, err
|
||||||
|
}
|
||||||
|
if !slices.Contains(d.has, false) {
|
||||||
|
d.fn(t, value)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return used, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// decodeDomain decodes an encoded Domain and sets has[i] if one of its attributes has the key want[i].
|
||||||
|
func decodeDomain(b []byte, want []string, has []bool) (t Domain_Type, value []byte, err error) {
|
||||||
|
clear(has)
|
||||||
|
for len(b) > 0 {
|
||||||
|
f, n, err := consumeField(b)
|
||||||
|
if err != nil {
|
||||||
|
return 0, nil, err
|
||||||
|
}
|
||||||
|
b = b[n:]
|
||||||
|
switch {
|
||||||
|
case f.num == 1 && f.typ == protowire.VarintType: // type
|
||||||
|
t = Domain_Type(f.x)
|
||||||
|
case f.num == 2 && f.typ == protowire.BytesType: // value
|
||||||
|
if !utf8.Valid(f.v) {
|
||||||
|
return 0, nil, errInvalidUTF8
|
||||||
|
}
|
||||||
|
value = f.v
|
||||||
|
case f.num == 3 && f.typ == protowire.BytesType: // attribute
|
||||||
|
key, err := decodeAttributeKey(f.v)
|
||||||
|
if err != nil {
|
||||||
|
return 0, nil, err
|
||||||
|
}
|
||||||
|
for i, w := range want {
|
||||||
|
if string(key) == w {
|
||||||
|
has[i] = true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return t, value, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// decodeAttributeKey returns the key of an encoded Domain.Attribute.
|
||||||
|
func decodeAttributeKey(b []byte) ([]byte, error) {
|
||||||
|
var key []byte
|
||||||
|
for len(b) > 0 {
|
||||||
|
f, n, err := consumeField(b)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
b = b[n:]
|
||||||
matcher := NewAllAttrsMatcher(attrs)
|
if f.num == 1 && f.typ == protowire.BytesType {
|
||||||
if matcher == nil {
|
if !utf8.Valid(f.v) {
|
||||||
return domains, nil
|
return nil, errInvalidUTF8
|
||||||
}
|
}
|
||||||
|
key = f.v
|
||||||
filtered := make([]*Domain, 0, len(domains))
|
|
||||||
for _, d := range domains {
|
|
||||||
if matcher.Match(d) {
|
|
||||||
filtered = append(filtered, d)
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
return key, nil
|
||||||
|
}
|
||||||
|
|
||||||
return filtered, 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,58 @@
|
|||||||
|
package geodata
|
||||||
|
|
||||||
|
import (
|
||||||
|
lua "github.com/yuin/gopher-lua"
|
||||||
|
luar "layeh.com/gopher-luar"
|
||||||
|
)
|
||||||
|
|
||||||
|
// RegisterLua makes xray.geodata available to require in an LState.
|
||||||
|
func RegisterLua(L *lua.LState) {
|
||||||
|
L.PreloadModule("xray.geodata", func(L *lua.LState) int {
|
||||||
|
module := L.NewTable()
|
||||||
|
|
||||||
|
module.RawSetString("BuildDomainMatcher", L.NewFunction(func(L *lua.LState) int {
|
||||||
|
parsed, err := ParseDomainRules(luaRules(L), Domain_Domain)
|
||||||
|
if err != nil {
|
||||||
|
L.RaiseError("%v", err)
|
||||||
|
return 0
|
||||||
|
}
|
||||||
|
matcher, err := DomainReg.BuildDomainMatcher(parsed)
|
||||||
|
if err != nil {
|
||||||
|
L.RaiseError("%v", err)
|
||||||
|
return 0
|
||||||
|
}
|
||||||
|
L.Push(luar.New(L, matcher))
|
||||||
|
return 1
|
||||||
|
}))
|
||||||
|
|
||||||
|
module.RawSetString("BuildIPMatcher", L.NewFunction(func(L *lua.LState) int {
|
||||||
|
parsed, err := ParseIPRules(luaRules(L))
|
||||||
|
if err != nil {
|
||||||
|
L.RaiseError("%v", err)
|
||||||
|
return 0
|
||||||
|
}
|
||||||
|
matcher, err := IPReg.BuildIPMatcher(parsed)
|
||||||
|
if err != nil {
|
||||||
|
L.RaiseError("%v", err)
|
||||||
|
return 0
|
||||||
|
}
|
||||||
|
L.Push(luar.New(L, matcher))
|
||||||
|
return 1
|
||||||
|
}))
|
||||||
|
L.Push(module)
|
||||||
|
return 1
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func luaRules(L *lua.LState) []string {
|
||||||
|
rules := make([]string, L.GetTop())
|
||||||
|
for i := range rules {
|
||||||
|
value, ok := L.Get(i + 1).(lua.LString)
|
||||||
|
if !ok {
|
||||||
|
L.RaiseError("geodata rules must be strings")
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
rules[i] = string(value)
|
||||||
|
}
|
||||||
|
return rules
|
||||||
|
}
|
||||||
@@ -0,0 +1,66 @@
|
|||||||
|
package geodata
|
||||||
|
|
||||||
|
import (
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/xtls/xray-core/common/net"
|
||||||
|
lua "github.com/yuin/gopher-lua"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestLuaIPMatcher(t *testing.T) {
|
||||||
|
L := lua.NewState()
|
||||||
|
defer L.Close()
|
||||||
|
RegisterLua(L)
|
||||||
|
ip := L.NewUserData()
|
||||||
|
ip.Value = net.ParseIP("127.0.0.1")
|
||||||
|
L.SetGlobal("ip", ip)
|
||||||
|
ips := L.NewUserData()
|
||||||
|
ips.Value = []net.IP{ip.Value.(net.IP), net.ParseIP("8.8.8.8")}
|
||||||
|
L.SetGlobal("ips", ips)
|
||||||
|
if err := L.DoString(`
|
||||||
|
local matcher = require("xray.geodata").BuildIPMatcher("127.0.0.0/8", "::1")
|
||||||
|
assert(matcher:Match(ip))
|
||||||
|
assert(matcher:AnyMatch(ips))
|
||||||
|
assert(not matcher:Matches(ips))
|
||||||
|
local matched, unmatched = matcher:FilterIPs(ips)
|
||||||
|
assert(type(matched) == "userdata" and type(unmatched) == "userdata")
|
||||||
|
assert(#matched == 1 and #unmatched == 1)
|
||||||
|
`); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestLuaDomainMatcher(t *testing.T) {
|
||||||
|
L := lua.NewState()
|
||||||
|
defer L.Close()
|
||||||
|
RegisterLua(L)
|
||||||
|
if err := L.DoString(`
|
||||||
|
local matcher = require("xray.geodata").BuildDomainMatcher("example.com", "full:other.com")
|
||||||
|
assert(matcher:MatchAny("example.com"))
|
||||||
|
assert(matcher:MatchAny("www.example.com"))
|
||||||
|
assert(matcher:MatchAny("other.com"))
|
||||||
|
assert(not matcher:MatchAny("www.other.com"))
|
||||||
|
assert(#(matcher:Match("www.example.com")) == 1)
|
||||||
|
`); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestLuaMatchersRejectInvalidRules(t *testing.T) {
|
||||||
|
for _, tc := range []struct {
|
||||||
|
name string
|
||||||
|
script string
|
||||||
|
}{
|
||||||
|
{"IP rule", `require("xray.geodata").BuildIPMatcher("not-an-ip")`},
|
||||||
|
{"non-string domain rule", `require("xray.geodata").BuildDomainMatcher("example.com", true)`},
|
||||||
|
} {
|
||||||
|
t.Run(tc.name, func(t *testing.T) {
|
||||||
|
L := lua.NewState()
|
||||||
|
defer L.Close()
|
||||||
|
RegisterLua(L)
|
||||||
|
if err := L.DoString(tc.script); err == nil {
|
||||||
|
t.Fatal("invalid geodata rule was accepted")
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -1,6 +1,7 @@
|
|||||||
package strmatcher_test
|
package strmatcher_test
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"regexp"
|
||||||
"strconv"
|
"strconv"
|
||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
@@ -72,6 +73,64 @@ func BenchmarkSubstrMatcher(b *testing.B) {
|
|||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func BenchmarkRegexMatcher(b *testing.B) {
|
||||||
|
patterns := []string{ // taken from geosite
|
||||||
|
`(^|\.)91porn\.(best|com|cool|fun|group|party|plus|site|tw|work)$`,
|
||||||
|
`(^|\.)91porn[0-9]{3}\.me$`,
|
||||||
|
`(^|\.)apiproxy-device-prod-nlb-.+\.amazonaws\.com$`,
|
||||||
|
`(^|\.)dualstack\.apiproxy-.+\.amazonaws\.com$`,
|
||||||
|
`(^|\.)aqdk[0-9]{3}\.com$`,
|
||||||
|
`(^|\.)bilibili3(0[1-9]|1[0-2])\.xyz$`,
|
||||||
|
`(^|\.)byyum([3589]|2[235689]|3[34]|4[1-9]|5[1-79]|6[0134679])?\.com$`,
|
||||||
|
`(^|\.)fiftymvapi\..+$`,
|
||||||
|
`(^|\.)gossipfuli[0-9]{3,4}\.xyz$`,
|
||||||
|
`(^|\.)kpkuang\.(bond|fun|info|one|us)$`,
|
||||||
|
`(^|\.)rule34\.(asia|us|world|xxx|xyz)$`,
|
||||||
|
`(^|\.)[a-z][1-9][0-9][a-z]\.com$`,
|
||||||
|
`.+\.awsdns-[0-9][0-9]\.(co\.uk|com|net|org)$`,
|
||||||
|
`.+\.dkr\.ecr\.[^\.]+\.amazonaws\.com$`,
|
||||||
|
`^(.+\.)*zh\.okaapps\.com$`,
|
||||||
|
`^cdn\d-epicgames-\d+\.file\.myqcloud\.com$`,
|
||||||
|
`^chatgpt-async-webps-prod-\S+-\d+\.webpubsub\.azure\.com$`,
|
||||||
|
`^r+[0-9]+(---|\.)sn-(2x3|ni5|j5o)\w{5}\.googlevideo\.com$`,
|
||||||
|
`^speed\.(coe|open)\.ad\.[a-z]{2,6}\.prod\.hosts\.ooklaserver\.net$`,
|
||||||
|
`javdb\d+\.com$`,
|
||||||
|
}
|
||||||
|
domains := []string{
|
||||||
|
"www.google.com", "rr3---sn-4g5edndy.googlevideo.com", "r1---sn-2x3abcde.googlevideo.com", "i.ytimg.com",
|
||||||
|
"graph.facebook.com", "api.twitter.com", "www.baidu.com", "github.com", "objects.githubusercontent.com",
|
||||||
|
"login.microsoftonline.com", "e1234.dscb.akamaiedge.net", "d1a2b3c4d5e6f7.cloudfront.net",
|
||||||
|
"s3.us-east-1.amazonaws.com", "123456789012.dkr.ecr.us-east-1.amazonaws.com", "www.wikipedia.org",
|
||||||
|
"discord.com", "telegram.org", "store.steampowered.com", "www.91porn.com", "ns-1234.awsdns-12.org",
|
||||||
|
}
|
||||||
|
bench := func(b *testing.B, ctor func(pattern string) func(string) bool) {
|
||||||
|
var matchers []func(string) bool
|
||||||
|
for _, p := range patterns {
|
||||||
|
matchers = append(matchers, ctor(p))
|
||||||
|
}
|
||||||
|
b.ResetTimer()
|
||||||
|
for i := 0; i < b.N; i++ {
|
||||||
|
for _, d := range domains {
|
||||||
|
for _, match := range matchers {
|
||||||
|
_ = match(d)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
b.Run("regexp", func(b *testing.B) {
|
||||||
|
bench(b, func(pattern string) func(string) bool {
|
||||||
|
return regexp.MustCompile(pattern).MatchString
|
||||||
|
})
|
||||||
|
})
|
||||||
|
b.Run("prefilter", func(b *testing.B) {
|
||||||
|
bench(b, func(pattern string) func(string) bool {
|
||||||
|
m, err := Regex.New(pattern)
|
||||||
|
common.Must(err)
|
||||||
|
return m.Match
|
||||||
|
})
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
// Utility functions for benchmark
|
// Utility functions for benchmark
|
||||||
|
|
||||||
func benchmarkMatcherType(b *testing.B, t Type, ctor func() MatcherGroup) {
|
func benchmarkMatcherType(b *testing.B, t Type, ctor func() MatcherGroup) {
|
||||||
|
|||||||
@@ -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 result
|
||||||
return CompositeMatches(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,198 +1,440 @@
|
|||||||
package strmatcher
|
package strmatcher
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"math/bits"
|
"bytes"
|
||||||
"runtime"
|
"cmp"
|
||||||
"sort"
|
"encoding/binary"
|
||||||
|
"errors"
|
||||||
|
"math"
|
||||||
|
"slices"
|
||||||
"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 {
|
||||||
rules []string // RuleIdx -> pattern string, index 0 reserved for failed lookup
|
arena string
|
||||||
values [][]uint32 // RuleIdx -> registered matcher values for the pattern (Full Matcher takes precedence)
|
level0 []uint16 // bucket -> seed
|
||||||
level0 []uint32 // RollingHash & Mask -> seed for Memhash
|
level1 []uint32 // slot -> flags | record offset
|
||||||
level0Mask uint32 // Mask restricting RollingHash to 0 ~ len(level0)
|
fp []uint8 // slot -> low byte of its pattern's hash, rejects most misses without reading arena
|
||||||
level1 []uint32 // Memhash<seed> & Mask -> stored index for rules
|
n0, n1 uint32
|
||||||
level1Mask uint32 // Mask for restricting Memhash<seed> to 0 ~ len(level1)
|
mul uint64 // multiplier of the suffix hash
|
||||||
ruleInfos *map[string]mphRuleInfo
|
single uint32 // the only value if !multi
|
||||||
|
multi bool
|
||||||
|
|
||||||
|
buf []byte // build only, patterns in Add order
|
||||||
|
entries []mphEntry
|
||||||
}
|
}
|
||||||
|
|
||||||
func NewMphMatcherGroup() *MphMatcherGroup {
|
func NewMphMatcherGroup() *MphMatcherGroup {
|
||||||
return &MphMatcherGroup{
|
return new(MphMatcherGroup)
|
||||||
rules: []string{""},
|
|
||||||
values: [][]uint32{nil},
|
|
||||||
level0: nil,
|
|
||||||
level0Mask: 0,
|
|
||||||
level1: nil,
|
|
||||||
level1Mask: 0,
|
|
||||||
ruleInfos: &map[string]mphRuleInfo{}, // Only used for building, destroyed after build complete
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// AddFullMatcher implements MatcherGroupForFull.
|
// 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.values = append(g.values, nil)
|
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)
|
|
||||||
|
|
||||||
// Create buckets based on all rule's rolling hash
|
|
||||||
buckets := make([][]uint32, len(g.level0))
|
|
||||||
for ruleIdx := 1; ruleIdx < len(g.rules); ruleIdx++ { // Traverse rules starting from index 1 (0 reserved for failed lookup)
|
|
||||||
ruleInfo := (*g.ruleInfos)[g.rules[ruleIdx]]
|
|
||||||
bucketIdx := ruleInfo.rollingHash & g.level0Mask
|
|
||||||
buckets[bucketIdx] = append(buckets[bucketIdx], uint32(ruleIdx))
|
|
||||||
g.values[ruleIdx] = append(ruleInfo.matchers[Full], ruleInfo.matchers[Domain]...) // nolint:gocritic
|
|
||||||
}
|
}
|
||||||
g.ruleInfos = nil // Set ruleInfos nil to release memory
|
if uint64(len(g.buf)) > math.MaxUint32 {
|
||||||
runtime.GC() // peak mem
|
return errors.New("too many rules for MphMatcherGroup")
|
||||||
|
|
||||||
// Sort buckets in descending order with respect to each bucket's size
|
|
||||||
bucketIdxs := make([]int, len(buckets))
|
|
||||||
for bucketIdx := range buckets {
|
|
||||||
bucketIdxs[bucketIdx] = bucketIdx
|
|
||||||
}
|
}
|
||||||
sort.Slice(bucketIdxs, func(i, j int) bool { return len(buckets[bucketIdxs[i]]) > len(buckets[bucketIdxs[j]]) })
|
recs := g.writeRecords()
|
||||||
|
if len(g.arena) > mphOffMask {
|
||||||
// Exercise Hash, Displace, and Compress algorithm to construct minimal perfect hash table
|
return errors.New("too many rules for MphMatcherGroup")
|
||||||
occupied := make([]bool, len(g.level1)) // Whether a second-level hash has been already used
|
|
||||||
hashedBucket := make([]uint32, 0, 4) // Second-level hashes for each rule in a specific bucket
|
|
||||||
for _, bucketIdx := range bucketIdxs {
|
|
||||||
bucket := buckets[bucketIdx]
|
|
||||||
hashedBucket = hashedBucket[:0]
|
|
||||||
seed := uint32(0)
|
|
||||||
for len(hashedBucket) != len(bucket) {
|
|
||||||
for _, ruleIdx := range bucket {
|
|
||||||
memHash := MemHash(seed, g.rules[ruleIdx]) & g.level1Mask
|
|
||||||
if occupied[memHash] { // Collision occurred with this seed
|
|
||||||
for _, hash := range hashedBucket { // Revert all values in this hashed bucket
|
|
||||||
occupied[hash] = false
|
|
||||||
g.level1[hash] = 0
|
|
||||||
}
|
}
|
||||||
hashedBucket = hashedBucket[:0]
|
hashes := make([]uint64, len(recs))
|
||||||
seed++ // Try next seed
|
for _, mul := range mphMultipliers {
|
||||||
|
for i, rec := range recs {
|
||||||
|
hashes[i] = mphMix(mphHash(mul, g.recKey(rec)))
|
||||||
|
}
|
||||||
|
g.mul = mul
|
||||||
|
if err := g.place(recs, hashes); err != errMphCollision {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return errMphCollision
|
||||||
|
}
|
||||||
|
|
||||||
|
// writeRecords writes one record per distinct pattern to arena and returns flags | offset of each.
|
||||||
|
func (g *MphMatcherGroup) writeRecords() []uint32 {
|
||||||
|
g.multi = false
|
||||||
|
if len(g.entries) > 0 {
|
||||||
|
g.single = g.entries[0].value
|
||||||
|
for _, e := range g.entries {
|
||||||
|
if e.value != g.single {
|
||||||
|
g.multi = true
|
||||||
break
|
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
|
||||||
}
|
}
|
||||||
|
|
||||||
// Lookup searches for input in minimal perfect hash table and returns its index. 0 indicates not found.
|
// mphHash is the suffix hash of s, taken from the right: the hash of s[i:] is the state after reading s[i].
|
||||||
func (g *MphMatcherGroup) Lookup(rollingHash uint32, input string) uint32 {
|
func mphHash(mul uint64, s string) uint64 {
|
||||||
i0 := rollingHash & g.level0Mask
|
h := uint64(0)
|
||||||
seed := g.level0[i0]
|
for i := len(s) - 1; i >= 0; i-- {
|
||||||
i1 := MemHash(seed, input) & g.level1Mask
|
h = h*mul + uint64(s[i])
|
||||||
if n := g.level1[i1]; g.rules[n] == input {
|
}
|
||||||
return n
|
return h
|
||||||
|
}
|
||||||
|
|
||||||
|
// mphMix spreads the weak low bits of a suffix hash.
|
||||||
|
func mphMix(h uint64) uint64 {
|
||||||
|
h ^= h >> 32
|
||||||
|
h *= 0xd6e8feb86659fd93
|
||||||
|
return h ^ h>>32
|
||||||
|
}
|
||||||
|
|
||||||
|
func (g *MphMatcherGroup) bucket(f uint64) uint32 {
|
||||||
|
return uint32(((f >> 32) * uint64(g.n0)) >> 32)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (g *MphMatcherGroup) slot(f uint64, seed uint16) uint32 {
|
||||||
|
x := ((f ^ uint64(seed)*0x9e3779b97f4a7c15) * 0xc4ceb9fe1a85ec53) >> 32
|
||||||
|
return uint32((x * uint64(g.n1)) >> 32)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (g *MphMatcherGroup) uvarint(p uint32) (x, next uint32) {
|
||||||
|
for shift := 0; ; shift += 7 {
|
||||||
|
c := g.arena[p]
|
||||||
|
p++
|
||||||
|
x |= uint32(c&0x7f) << shift
|
||||||
|
if c < 0x80 {
|
||||||
|
return x, p
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// recSpan returns where the pattern of the record at off starts and how long it is.
|
||||||
|
func (g *MphMatcherGroup) recSpan(off uint32) (p, n uint32) {
|
||||||
|
n, p = uint32(g.arena[off]), off+1
|
||||||
|
if n == 255 {
|
||||||
|
n, p = g.uvarint(p)
|
||||||
|
}
|
||||||
|
return p, n
|
||||||
|
}
|
||||||
|
|
||||||
|
func (g *MphMatcherGroup) recKey(rec uint32) string {
|
||||||
|
p, n := g.recSpan(rec & mphOffMask)
|
||||||
|
return g.arena[p : p+n]
|
||||||
|
}
|
||||||
|
|
||||||
|
// lookup returns the level1 entry of s, or 0 if s is not a pattern. h is the suffix hash of s.
|
||||||
|
func (g *MphMatcherGroup) lookup(h uint64, s string) uint32 {
|
||||||
|
f := mphMix(h)
|
||||||
|
// bucket < n0 == len(level0) and slot < n1 == len(level1) == len(fp), skip the bounds checks
|
||||||
|
seed := *(*uint16)(unsafe.Add(unsafe.Pointer(unsafe.SliceData(g.level0)), uintptr(g.bucket(f))*2))
|
||||||
|
slot := uintptr(g.slot(f, seed))
|
||||||
|
if *(*uint8)(unsafe.Add(unsafe.Pointer(unsafe.SliceData(g.fp)), slot)) != uint8(f) {
|
||||||
|
return 0
|
||||||
|
}
|
||||||
|
e := *(*uint32)(unsafe.Add(unsafe.Pointer(unsafe.SliceData(g.level1)), slot*4))
|
||||||
|
if len(s) < 255 {
|
||||||
|
// A record whose length byte is len(s) has len(s) pattern bytes after it
|
||||||
|
p := unsafe.Add(unsafe.Pointer(unsafe.StringData(g.arena)), e&mphOffMask)
|
||||||
|
if int(*(*byte)(p)) == len(s) && unsafe.String((*byte)(unsafe.Add(p, 1)), len(s)) == s {
|
||||||
|
return e
|
||||||
|
}
|
||||||
|
return 0
|
||||||
|
}
|
||||||
|
if g.recKey(e) == s {
|
||||||
|
return e
|
||||||
}
|
}
|
||||||
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) appendValues(dst []uint32, e, want uint32) []uint32 {
|
||||||
|
if !g.multi {
|
||||||
|
for _, flag := range mphKinds {
|
||||||
|
if e&want&flag != 0 {
|
||||||
|
dst = append(dst, g.single)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return dst
|
||||||
|
}
|
||||||
|
if e&want == 0 {
|
||||||
|
return dst
|
||||||
|
}
|
||||||
|
p, n := g.recSpan(e & mphOffMask)
|
||||||
|
p += n
|
||||||
|
for _, flag := range mphKinds {
|
||||||
|
if e&flag == 0 {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
var count, v uint32
|
||||||
|
for count, p = g.uvarint(p); count > 0; count-- {
|
||||||
|
v, p = g.uvarint(p)
|
||||||
|
if want&flag != 0 {
|
||||||
|
dst = append(dst, v)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return dst
|
||||||
|
}
|
||||||
|
|
||||||
|
// Match implements MatcherGroup.Match. Values of an exact match come first (Full, then Domain), then those of
|
||||||
|
// the parent domains, nearest first.
|
||||||
func (g *MphMatcherGroup) Match(input string) []uint32 {
|
func (g *MphMatcherGroup) Match(input string) []uint32 {
|
||||||
matches := make([][]uint32, 0, 5)
|
var stack [8]uint32
|
||||||
hash := uint32(0)
|
parents := stack[:0] // TLD side first
|
||||||
|
h, mul := uint64(0), g.mul
|
||||||
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 mphIdx := g.Lookup(hash, input[i:]); mphIdx != 0 {
|
if e := g.lookup(h, input[i+1:]); e&(mphDomain|mphParent) != 0 {
|
||||||
matches = append(matches, g.values[mphIdx])
|
parents = append(parents, e)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
h = h*mul + uint64(input[i])
|
||||||
}
|
}
|
||||||
if mphIdx := g.Lookup(hash, input); mphIdx != 0 {
|
exact := g.lookup(h, input)
|
||||||
matches = append(matches, g.values[mphIdx])
|
if exact&(mphFull|mphDomain) == 0 && len(parents) == 0 {
|
||||||
|
return nil
|
||||||
}
|
}
|
||||||
return CompositeMatchesReverse(matches)
|
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})
|
||||||
|
}
|
||||||
|
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 true
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
return g.lookup(h, input)&(mphFull|mphDomain) != 0
|
||||||
return g.Lookup(hash, input) != 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)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -1,7 +1,10 @@
|
|||||||
package strmatcher_test
|
package strmatcher_test
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"math/rand"
|
||||||
"reflect"
|
"reflect"
|
||||||
|
"slices"
|
||||||
|
"strings"
|
||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
"github.com/xtls/xray-core/common"
|
"github.com/xtls/xray-core/common"
|
||||||
@@ -276,3 +279,142 @@ func TestEmptyMphMatcherGroup(t *testing.T) {
|
|||||||
t.Error("Expect [], but ", r)
|
t.Error("Expect [], but ", r)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestMphMatcherGroupRandom(t *testing.T) {
|
||||||
|
inputs := []string{""} // All strings over "ab." up to 7 bytes
|
||||||
|
for i := 0; len(inputs[i]) < 7; i++ {
|
||||||
|
for _, c := range []string{"a", "b", "."} {
|
||||||
|
inputs = append(inputs, inputs[i]+c)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
for seed := int64(0); seed < 300; seed++ {
|
||||||
|
r := rand.New(rand.NewSource(seed))
|
||||||
|
g := NewMphMatcherGroup()
|
||||||
|
full, domain := map[string][]uint32{}, map[string][]uint32{} // Stored pattern -> values
|
||||||
|
for value := uint32(r.Intn(200)); value > 0; value-- {
|
||||||
|
pattern := make([]byte, r.Intn(8))
|
||||||
|
for i := range pattern {
|
||||||
|
pattern[i] = "ab."[r.Intn(3)]
|
||||||
|
}
|
||||||
|
if p := string(pattern); r.Intn(2) == 0 {
|
||||||
|
g.AddFullMatcher(FullMatcher(p), value)
|
||||||
|
full[p] = append(full[p], value)
|
||||||
|
} else {
|
||||||
|
g.AddDomainMatcher(DomainMatcher(p), value)
|
||||||
|
domain[p] = append(domain[p], value)
|
||||||
|
domain["."+p] = append(domain["."+p], value)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
common.Must(g.Build())
|
||||||
|
for _, input := range inputs {
|
||||||
|
keys := []string{input} // Whole input first, then "." suffixes from longest to shortest
|
||||||
|
for i := range len(input) {
|
||||||
|
if input[i] == '.' {
|
||||||
|
keys = append(keys, input[i:])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
var want []uint32
|
||||||
|
for _, k := range keys {
|
||||||
|
want = append(append(want, full[k]...), domain[k]...)
|
||||||
|
}
|
||||||
|
// Compared as sets: Match reports a value once per matching pattern, and orders them differently
|
||||||
|
// from want for patterns and inputs with a leading dot
|
||||||
|
m := g.Match(input)
|
||||||
|
if !slices.Equal(sortedSet(m), sortedSet(want)) {
|
||||||
|
t.Fatalf("seed %d: Match(%q) = %v, want %v", seed, input, m, want)
|
||||||
|
}
|
||||||
|
if m := g.MatchAny(input); m != (len(want) > 0) {
|
||||||
|
t.Fatalf("seed %d: MatchAny(%q) = %v", seed, input, m)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMphMatcherGroupAppend(t *testing.T) {
|
||||||
|
g := NewMphMatcherGroup()
|
||||||
|
g.AddFullMatcher(FullMatcher("a.com"), 1)
|
||||||
|
g.AddFullMatcher(FullMatcher("b.com"), 2)
|
||||||
|
g.Build()
|
||||||
|
if m := append(g.Match("a.com"), 3); !slices.Equal(m, []uint32{1, 3}) {
|
||||||
|
t.Error("expect [1 3], but ", m)
|
||||||
|
}
|
||||||
|
if m := g.Match("b.com"); !slices.Equal(m, []uint32{2}) {
|
||||||
|
t.Error("expect [2], but ", m)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func sortedSet(v []uint32) []uint32 {
|
||||||
|
v = slices.Clone(v)
|
||||||
|
slices.Sort(v)
|
||||||
|
return slices.Compact(v)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMphMatcherGroupLongPattern(t *testing.T) {
|
||||||
|
long := strings.Repeat("a", 300) + ".com"
|
||||||
|
for _, values := range [][4]uint32{{1, 2, 3, 4}, {7, 7, 7, 7}} {
|
||||||
|
g := NewMphMatcherGroup()
|
||||||
|
g.AddDomainMatcher(DomainMatcher(long), values[0])
|
||||||
|
g.AddFullMatcher(FullMatcher("x."+long), values[1])
|
||||||
|
g.AddFullMatcher(FullMatcher(long[:255]), values[2]) // the shortest pattern stored with a long length
|
||||||
|
g.AddFullMatcher(FullMatcher(long[:254]), values[3])
|
||||||
|
common.Must(g.Build())
|
||||||
|
cases := []struct {
|
||||||
|
input string
|
||||||
|
want []uint32
|
||||||
|
}{
|
||||||
|
{long, []uint32{values[0]}},
|
||||||
|
{"www." + long, []uint32{values[0]}},
|
||||||
|
{"x." + long, []uint32{values[1], values[0]}},
|
||||||
|
{long[1:], nil},
|
||||||
|
{"a" + long, nil},
|
||||||
|
{long[:255], []uint32{values[2]}},
|
||||||
|
{long[:254], []uint32{values[3]}},
|
||||||
|
{long[:256], nil},
|
||||||
|
{long[:253], nil},
|
||||||
|
}
|
||||||
|
for _, c := range cases {
|
||||||
|
if m := g.Match(c.input); !slices.Equal(m, c.want) {
|
||||||
|
t.Errorf("Match(%d bytes) = %v, want %v", len(c.input), m, c.want)
|
||||||
|
}
|
||||||
|
if m := g.MatchAny(c.input); m != (c.want != nil) {
|
||||||
|
t.Errorf("MatchAny(%d bytes) = %v", len(c.input), m)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// A pattern longer than 65535 bytes builds and matches: a record's length is a uvarint,
|
||||||
|
// so the only cap was the build-time length field, now widened to uint32.
|
||||||
|
huge := strings.Repeat("a", 70000)
|
||||||
|
g := NewMphMatcherGroup()
|
||||||
|
g.AddFullMatcher(FullMatcher(strings.Repeat("a", 65535)), 1)
|
||||||
|
g.AddDomainMatcher(DomainMatcher(huge+".com"), 2)
|
||||||
|
g.AddFullMatcher(FullMatcher("a.com"), 3)
|
||||||
|
common.Must(g.Build())
|
||||||
|
if !g.MatchAny(strings.Repeat("a", 65535)) || g.MatchAny(strings.Repeat("a", 65534)) {
|
||||||
|
t.Error("wrong answer for a 65535-byte pattern")
|
||||||
|
}
|
||||||
|
if m := g.Match(huge + ".com"); !slices.Equal(m, []uint32{2}) {
|
||||||
|
t.Errorf("Match(%d-byte input) = %v, want [2]", len(huge)+4, m)
|
||||||
|
}
|
||||||
|
if m := g.Match("x." + huge + ".com"); !slices.Equal(m, []uint32{2}) {
|
||||||
|
t.Errorf("Match(subdomain of a %d-byte pattern) = %v, want [2]", len(huge)+4, m)
|
||||||
|
}
|
||||||
|
if g.MatchAny(huge) { // the 70000-byte label on its own is not a rule
|
||||||
|
t.Error("unexpected match for the bare 70000-byte label")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMphMatcherGroupBuildOnce(t *testing.T) {
|
||||||
|
g := NewMphMatcherGroup()
|
||||||
|
g.AddFullMatcher(FullMatcher("a.com"), 1)
|
||||||
|
common.Must(g.Build())
|
||||||
|
if err := g.Build(); err == nil || !g.MatchAny("a.com") {
|
||||||
|
t.Errorf("second Build() = %v, MatchAny(a.com) = %v", err, g.MatchAny("a.com"))
|
||||||
|
}
|
||||||
|
defer func() {
|
||||||
|
if recover() == nil {
|
||||||
|
t.Error("Add after Build did not panic")
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
g.AddDomainMatcher(DomainMatcher("b.com"), 2)
|
||||||
|
}
|
||||||
|
|||||||
@@ -2,9 +2,12 @@ package strmatcher
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"errors"
|
"errors"
|
||||||
|
"math/bits"
|
||||||
"regexp"
|
"regexp"
|
||||||
|
"regexp/syntax"
|
||||||
"slices"
|
"slices"
|
||||||
"strings"
|
"strings"
|
||||||
|
"unicode"
|
||||||
"unicode/utf8"
|
"unicode/utf8"
|
||||||
|
|
||||||
"golang.org/x/net/idna"
|
"golang.org/x/net/idna"
|
||||||
@@ -74,6 +77,273 @@ 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
|
||||||
|
tail []byteSet // tail[i] holds the bytes a matching input can have i bytes before its end
|
||||||
|
rest *byteSet // the bytes it can have further before, nil if any
|
||||||
|
}
|
||||||
|
|
||||||
|
func newRegexMatcher(pattern string) (Matcher, error) {
|
||||||
|
regex, err := regexp.Compile(pattern)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
m := &RegexMatcher{pattern: regex}
|
||||||
|
if re, err := syntax.Parse(pattern, syntax.Perl); err == nil { // same flags as regexp.Compile
|
||||||
|
m.literals = requiredLiterals(re, nil)
|
||||||
|
slices.SortStableFunc(m.literals, func(a, b string) int { return len(b) - len(a) })
|
||||||
|
m.tail, m.rest = tailGuard(re)
|
||||||
|
}
|
||||||
|
return m, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// byteSet is a set of bytes. The bytes >= 0x80 share one bit with 0x7f.
|
||||||
|
type byteSet [4]uint32
|
||||||
|
|
||||||
|
func (s *byteSet) add(c byte) { c = min(c, 0x7f); s[c>>5] |= 1 << (c & 31) }
|
||||||
|
func (s *byteSet) has(c byte) bool { c = min(c, 0x7f); return s[c>>5]&(1<<(c&31)) != 0 }
|
||||||
|
func (s *byteSet) or(t *byteSet) {
|
||||||
|
for i := range s {
|
||||||
|
s[i] |= t[i]
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
var allBytes = byteSet{^uint32(0), ^uint32(0), ^uint32(0), ^uint32(0)}
|
||||||
|
|
||||||
|
// tailLen is how many positions before the end of the input tailGuard tells apart.
|
||||||
|
const tailLen = 8
|
||||||
|
|
||||||
|
// tailBudget caps how many repetition steps tailGuard walks. Only nested repeats can make the
|
||||||
|
// walk explode, so only they are charged: a flat pattern, however long, is walked once and keeps
|
||||||
|
// its guard.
|
||||||
|
const tailBudget = 100000
|
||||||
|
|
||||||
|
// tailWalk is a set of positions in the input, counted in bytes before its end.
|
||||||
|
type tailWalk struct {
|
||||||
|
at uint32 // bit i: exactly i bytes before the end, for i < tailLen
|
||||||
|
far bool // tailLen or more bytes before the end
|
||||||
|
free bool // not tied to the end of the input yet
|
||||||
|
}
|
||||||
|
|
||||||
|
func (w tailWalk) union(v tailWalk) tailWalk {
|
||||||
|
return tailWalk{w.at | v.at, w.far || v.far, w.free || v.free}
|
||||||
|
}
|
||||||
|
|
||||||
|
type tailBuilder struct {
|
||||||
|
tail [tailLen]byteSet
|
||||||
|
rest byteSet
|
||||||
|
void bool
|
||||||
|
work int
|
||||||
|
}
|
||||||
|
|
||||||
|
// tailGuard walks re backwards from the end of the input and collects the bytes an input
|
||||||
|
// matching re can have at each position before its end. It returns nil, nil when a branch
|
||||||
|
// of re does not end with $ or when nested repeats push the walk past tailBudget.
|
||||||
|
func tailGuard(re *syntax.Regexp) ([]byteSet, *byteSet) {
|
||||||
|
var b tailBuilder
|
||||||
|
w := b.walk(re, tailWalk{free: true})
|
||||||
|
b.stop(w)
|
||||||
|
if b.void {
|
||||||
|
return nil, nil
|
||||||
|
}
|
||||||
|
if w.at != 0 { // a match can start here, so any bytes can come before
|
||||||
|
for i := bits.TrailingZeros32(w.at); i < tailLen; i++ {
|
||||||
|
b.tail[i] = allBytes
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if w.at != 0 || w.far {
|
||||||
|
b.rest = allBytes
|
||||||
|
}
|
||||||
|
n := tailLen
|
||||||
|
for n > 0 && b.tail[n-1] == b.rest {
|
||||||
|
n--
|
||||||
|
}
|
||||||
|
var tail []byteSet
|
||||||
|
if n > 0 {
|
||||||
|
tail = slices.Clone(b.tail[:n])
|
||||||
|
}
|
||||||
|
if b.rest != allBytes {
|
||||||
|
rest := b.rest
|
||||||
|
return tail, &rest
|
||||||
|
}
|
||||||
|
return tail, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// stop ends the paths of w. One that never met $ lets its match be followed by anything.
|
||||||
|
func (b *tailBuilder) stop(w tailWalk) {
|
||||||
|
if w.free {
|
||||||
|
b.void = true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (b *tailBuilder) walk(re *syntax.Regexp, w tailWalk) tailWalk {
|
||||||
|
if w == (tailWalk{}) || b.void {
|
||||||
|
return w
|
||||||
|
}
|
||||||
|
switch re.Op {
|
||||||
|
case syntax.OpNoMatch:
|
||||||
|
return tailWalk{}
|
||||||
|
case syntax.OpLiteral:
|
||||||
|
for i := len(re.Rune) - 1; i >= 0; i-- {
|
||||||
|
var set byteSet
|
||||||
|
set.add(byte(min(re.Rune[i], utf8.RuneSelf)))
|
||||||
|
if re.Flags&syntax.FoldCase != 0 {
|
||||||
|
for f := unicode.SimpleFold(re.Rune[i]); f != re.Rune[i]; f = unicode.SimpleFold(f) {
|
||||||
|
set.add(byte(min(f, utf8.RuneSelf)))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
w = b.step(w, &set)
|
||||||
|
}
|
||||||
|
return w
|
||||||
|
case syntax.OpCharClass:
|
||||||
|
var set byteSet
|
||||||
|
for i := 0; i+1 < len(re.Rune); i += 2 {
|
||||||
|
for r := min(re.Rune[i], utf8.RuneSelf); r <= min(re.Rune[i+1], utf8.RuneSelf); r++ {
|
||||||
|
set.add(byte(r))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return b.step(w, &set)
|
||||||
|
case syntax.OpAnyChar, syntax.OpAnyCharNotNL: // a domain has no \n to reject
|
||||||
|
return b.step(w, &allBytes)
|
||||||
|
case syntax.OpBeginText: // nothing comes before
|
||||||
|
b.stop(w)
|
||||||
|
return tailWalk{}
|
||||||
|
case syntax.OpEndText:
|
||||||
|
out := tailWalk{at: w.at & 1}
|
||||||
|
if w.free {
|
||||||
|
out.at = 1
|
||||||
|
}
|
||||||
|
return out
|
||||||
|
case syntax.OpCapture:
|
||||||
|
return b.walk(re.Sub[0], w)
|
||||||
|
case syntax.OpConcat:
|
||||||
|
for i := len(re.Sub) - 1; i >= 0; i-- {
|
||||||
|
w = b.walk(re.Sub[i], w)
|
||||||
|
}
|
||||||
|
return w
|
||||||
|
case syntax.OpAlternate:
|
||||||
|
var out tailWalk
|
||||||
|
for _, sub := range re.Sub {
|
||||||
|
out = out.union(b.walk(sub, w))
|
||||||
|
}
|
||||||
|
return out
|
||||||
|
case syntax.OpQuest:
|
||||||
|
return b.repeat(re.Sub[0], w, 1)
|
||||||
|
case syntax.OpStar:
|
||||||
|
return b.repeat(re.Sub[0], w, -1)
|
||||||
|
case syntax.OpPlus:
|
||||||
|
return b.repeat(re.Sub[0], b.walk(re.Sub[0], w), -1)
|
||||||
|
case syntax.OpRepeat:
|
||||||
|
for i := 0; i < re.Min; i++ {
|
||||||
|
if b.charge() {
|
||||||
|
return w
|
||||||
|
}
|
||||||
|
w = b.walk(re.Sub[0], w)
|
||||||
|
}
|
||||||
|
if re.Max < 0 {
|
||||||
|
return b.repeat(re.Sub[0], w, -1)
|
||||||
|
}
|
||||||
|
return b.repeat(re.Sub[0], w, re.Max-re.Min)
|
||||||
|
}
|
||||||
|
return w // empty match, line and word boundaries: no constraint
|
||||||
|
}
|
||||||
|
|
||||||
|
// charge counts one repetition step and reports whether the walk has run out of budget. Only
|
||||||
|
// repeats re-walk their body, so charging them alone bounds the blow-up of nested repeats while
|
||||||
|
// leaving a single linear pass, of any length, free.
|
||||||
|
func (b *tailBuilder) charge() bool {
|
||||||
|
b.work++
|
||||||
|
if b.work > tailBudget {
|
||||||
|
b.void = true
|
||||||
|
}
|
||||||
|
return b.void
|
||||||
|
}
|
||||||
|
|
||||||
|
// repeat walks back over up to n more repetitions of re, any number if n < 0.
|
||||||
|
func (b *tailBuilder) repeat(re *syntax.Regexp, w tailWalk, n int) tailWalk {
|
||||||
|
for ; n != 0; n-- {
|
||||||
|
if b.charge() {
|
||||||
|
return w
|
||||||
|
}
|
||||||
|
next := w.union(b.walk(re, w))
|
||||||
|
if next == w {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
w = next
|
||||||
|
}
|
||||||
|
return w
|
||||||
|
}
|
||||||
|
|
||||||
|
// step walks back over one character whose last byte is in set. A character that can be
|
||||||
|
// non-ASCII can take up to 4 bytes, all >= 0x80; regexp matches an invalid byte as U+FFFD.
|
||||||
|
func (b *tailBuilder) step(w tailWalk, set *byteSet) tailWalk {
|
||||||
|
out := tailWalk{far: w.far, free: w.free}
|
||||||
|
if w.far {
|
||||||
|
b.rest.or(set)
|
||||||
|
}
|
||||||
|
width := 1
|
||||||
|
if set.has(0x80) {
|
||||||
|
width = utf8.UTFMax
|
||||||
|
}
|
||||||
|
for i := 0; i < tailLen; i++ {
|
||||||
|
if w.at&(1<<i) == 0 {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
b.tail[i].or(set)
|
||||||
|
for n := 1; n <= width; n++ {
|
||||||
|
if j := i + n; j < tailLen {
|
||||||
|
out.at |= 1 << j
|
||||||
|
if n < width {
|
||||||
|
b.tail[j].add(0x80)
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
out.far = true
|
||||||
|
if n < width {
|
||||||
|
b.rest.add(0x80)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return out
|
||||||
|
}
|
||||||
|
|
||||||
|
// mayMatch reports whether s passes the tail guard.
|
||||||
|
func (m *RegexMatcher) mayMatch(s string) bool {
|
||||||
|
n := len(s)
|
||||||
|
if m.rest == nil {
|
||||||
|
n = min(n, len(m.tail))
|
||||||
|
}
|
||||||
|
for i := 0; i < n; i++ {
|
||||||
|
set := m.rest
|
||||||
|
if i < len(m.tail) {
|
||||||
|
set = &m.tail[i]
|
||||||
|
}
|
||||||
|
if !set.has(s[len(s)-1-i]) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
// requiredLiterals appends to dst the case-sensitive strings that every match of re contains.
|
||||||
|
func requiredLiterals(re *syntax.Regexp, dst []string) []string {
|
||||||
|
switch re.Op {
|
||||||
|
case syntax.OpLiteral:
|
||||||
|
// regexp matches U+FFFD against invalid UTF-8 bytes, strings.Contains does not
|
||||||
|
if re.Flags&syntax.FoldCase == 0 && !slices.Contains(re.Rune, utf8.RuneError) {
|
||||||
|
dst = append(dst, string(re.Rune))
|
||||||
|
}
|
||||||
|
case syntax.OpCapture, syntax.OpPlus:
|
||||||
|
dst = requiredLiterals(re.Sub[0], dst)
|
||||||
|
case syntax.OpRepeat:
|
||||||
|
if re.Min > 0 {
|
||||||
|
dst = requiredLiterals(re.Sub[0], dst)
|
||||||
|
}
|
||||||
|
case syntax.OpConcat:
|
||||||
|
for _, sub := range re.Sub {
|
||||||
|
dst = requiredLiterals(sub, dst)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return dst
|
||||||
}
|
}
|
||||||
|
|
||||||
func (*RegexMatcher) Type() Type {
|
func (*RegexMatcher) Type() Type {
|
||||||
@@ -89,6 +359,14 @@ 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 {
|
||||||
|
if !strings.Contains(s, l) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
}
|
||||||
return m.pattern.MatchString(s)
|
return m.pattern.MatchString(s)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -102,11 +380,7 @@ func (t Type) New(pattern string) (Matcher, error) {
|
|||||||
case Domain:
|
case Domain:
|
||||||
return DomainMatcher(pattern), nil
|
return DomainMatcher(pattern), nil
|
||||||
case Regex: // 1. regex matching is case-sensitive
|
case Regex: // 1. regex matching is case-sensitive
|
||||||
regex, err := regexp.Compile(pattern)
|
return newRegexMatcher(pattern)
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
return &RegexMatcher{pattern: regex}, nil
|
|
||||||
default:
|
default:
|
||||||
return nil, errors.New("unknown matcher type")
|
return nil, errors.New("unknown matcher type")
|
||||||
}
|
}
|
||||||
@@ -135,11 +409,7 @@ func (t Type) NewDomainPattern(pattern string) (Matcher, error) {
|
|||||||
}
|
}
|
||||||
return DomainMatcher(pattern), nil
|
return DomainMatcher(pattern), nil
|
||||||
case Regex: // Regex's charset not in LDH subset
|
case Regex: // Regex's charset not in LDH subset
|
||||||
regex, err := regexp.Compile(pattern)
|
return newRegexMatcher(pattern)
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
return &RegexMatcher{pattern: regex}, nil
|
|
||||||
default:
|
default:
|
||||||
return nil, errors.New("unknown matcher type")
|
return nil, errors.New("unknown matcher type")
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -0,0 +1,233 @@
|
|||||||
|
package strmatcher
|
||||||
|
|
||||||
|
import (
|
||||||
|
"hash/fnv"
|
||||||
|
"math/rand/v2"
|
||||||
|
"regexp"
|
||||||
|
"regexp/syntax"
|
||||||
|
"slices"
|
||||||
|
"strconv"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
"unicode"
|
||||||
|
"unicode/utf8"
|
||||||
|
)
|
||||||
|
|
||||||
|
var regexLiteralCases = []struct {
|
||||||
|
pattern string
|
||||||
|
literals []string
|
||||||
|
}{
|
||||||
|
{`(^|\.)91porn\.(best|com)$`, []string{"91porn."}},
|
||||||
|
{`.+\.awsdns-cn-[0-9][0-9]\.(biz|com|net|top)$`, []string{".awsdns-cn-", "."}},
|
||||||
|
{`^r+[0-9]+(---|\.)sn-(2x3|ni5|j5o)\w{5}\.googlevideo\.com$`, []string{".googlevideo.com", "sn-", "r"}},
|
||||||
|
{`(?i)abc`, nil},
|
||||||
|
{`ab(?i:CD)ef`, []string{"ab", "ef"}},
|
||||||
|
{`(abc)?x`, []string{"x"}},
|
||||||
|
{`(abc)*x`, []string{"x"}},
|
||||||
|
{`x{0,3}yy`, []string{"yy"}},
|
||||||
|
{`(ab)+c{2}`, []string{"ab", "c"}},
|
||||||
|
{`abc|abd`, []string{"ab"}},
|
||||||
|
{`\Qa.b\E`, []string{"a.b"}},
|
||||||
|
{`a\x{FFFD}b`, nil},
|
||||||
|
{`^[^.]+$`, nil},
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRegexRequiredLiterals(t *testing.T) {
|
||||||
|
for _, test := range regexLiteralCases {
|
||||||
|
m, err := newRegexMatcher(test.pattern)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if got := m.(*RegexMatcher).literals; !slices.Equal(got, test.literals) {
|
||||||
|
t.Errorf("%s: got %q, want %q", test.pattern, got, test.literals)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
var regexTailCases = []struct {
|
||||||
|
pattern string
|
||||||
|
guard bool
|
||||||
|
match []string // inputs the pattern matches
|
||||||
|
reject []string // inputs the tail guard alone rejects
|
||||||
|
}{
|
||||||
|
{`^[a-z]([a-z0-9-]{0,61}[a-z0-9])?$`, true, []string{"a", "localhost", "x-1"}, []string{"www.example.com", "localhost.", "LOCALHOST", "a b"}},
|
||||||
|
{`(^|\.)[a-z][1-9][0-9][a-z]\.com$`, true, []string{"a12b.com", "x.q10z.com"}, []string{"google.com", "a12b.co", "a12b.com.", "ab12.com"}},
|
||||||
|
{`^hses[1-7]?\.akamaized\.net$`, true, []string{"hses.akamaized.net", "hses3.akamaized.net"}, []string{"xhses.akamaized.net", "www.hses.akamaized.net"}},
|
||||||
|
{`(?i)k\.net$`, true, []string{"k.net", "K.NET", "\u212a.net"}, []string{"x.net", "k.nex"}},
|
||||||
|
{`[^.]+\.cn$`, true, []string{"a.cn", "\xff.cn", "\u4e2d.cn"}, []string{"a.cnn", "a.c"}},
|
||||||
|
{`\x{FFFD}$`, true, []string{"\xff", "a\xc3", "\uFFFD"}, []string{"a", "\xff."}},
|
||||||
|
{`^.\.cn$`, true, []string{"a.cn", "\u4E2D.cn", "\xff.cn"}, []string{"ab.cn"}},
|
||||||
|
{`^$`, true, []string{""}, []string{"a"}},
|
||||||
|
{`(^|\.)youyuapi\..+$`, false, []string{"youyuapi.com"}, nil},
|
||||||
|
{`abc`, false, []string{"abc", "xabcx"}, nil},
|
||||||
|
{`^ab`, false, []string{"ab", "abc"}, nil},
|
||||||
|
{`a$|b`, false, []string{"a", "bx"}, nil},
|
||||||
|
{`(?m)a$`, false, []string{"a", "a\nb"}, nil},
|
||||||
|
{strings.Repeat(`(?:abcdefgh(?:a`, 20) + strings.Repeat(`)*)*`, 20) + `\.com$`, false, []string{".com", "abcdefgha.com"}, nil}, // over tailBudget
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRegexTailGuard(t *testing.T) {
|
||||||
|
for _, test := range regexTailCases {
|
||||||
|
m, err := newRegexMatcher(test.pattern)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
rm := m.(*RegexMatcher)
|
||||||
|
if guard := rm.tail != nil || rm.rest != nil; guard != test.guard {
|
||||||
|
t.Errorf("%s: guard %v, want %v", test.pattern, guard, test.guard)
|
||||||
|
}
|
||||||
|
for _, s := range test.match {
|
||||||
|
if !rm.pattern.MatchString(s) || !rm.Match(s) {
|
||||||
|
t.Errorf("%s: %q does not match", test.pattern, s)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
for _, s := range test.reject {
|
||||||
|
if rm.pattern.MatchString(s) || rm.mayMatch(s) {
|
||||||
|
t.Errorf("%s: %q passes the guard", test.pattern, s)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestRegexTailGuardFlatAlternation checks that a long but non-recursive pattern keeps its
|
||||||
|
// guard. Only nested repeats are charged against tailBudget, so a flat alternation of many
|
||||||
|
// names, however large, is walked once and guarded; its guard is checked against regexp.
|
||||||
|
func TestRegexTailGuardFlatAlternation(t *testing.T) {
|
||||||
|
var sb strings.Builder
|
||||||
|
sb.WriteString("(?:")
|
||||||
|
for i := 0; i < 20000; i++ {
|
||||||
|
if i > 0 {
|
||||||
|
sb.WriteByte('|')
|
||||||
|
}
|
||||||
|
sb.WriteString("name")
|
||||||
|
sb.WriteString(strconv.Itoa(i))
|
||||||
|
}
|
||||||
|
sb.WriteString(`)\.example\.com$`)
|
||||||
|
m, err := newRegexMatcher(sb.String())
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
rm := m.(*RegexMatcher)
|
||||||
|
if rm.tail == nil && rm.rest == nil {
|
||||||
|
t.Fatal("flat alternation of 20000 names lost its guard")
|
||||||
|
}
|
||||||
|
for _, s := range []string{"name0.example.com", "name19999.example.com", "x.name12345.example.com"} {
|
||||||
|
if !rm.pattern.MatchString(s) || !rm.Match(s) {
|
||||||
|
t.Errorf("%q should match", s)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
for _, s := range []string{"name0.example.org", "name0.example.com.", "name0.example.con", "google.com"} {
|
||||||
|
if rm.pattern.MatchString(s) {
|
||||||
|
t.Fatalf("test bug: %q matches the pattern", s)
|
||||||
|
}
|
||||||
|
if rm.mayMatch(s) {
|
||||||
|
t.Errorf("%q should be rejected by the guard", s)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// sampleMatch appends a string that re matches, assertions aside, unless it runs out of
|
||||||
|
// budget, which it spends one per call so that nested repeats stay cheap.
|
||||||
|
func sampleMatch(sb *strings.Builder, re *syntax.Regexp, rnd *rand.Rand, budget *int) {
|
||||||
|
if *budget <= 0 {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
*budget--
|
||||||
|
switch re.Op {
|
||||||
|
case syntax.OpLiteral:
|
||||||
|
for _, r := range re.Rune {
|
||||||
|
if re.Flags&syntax.FoldCase != 0 {
|
||||||
|
for n := rnd.IntN(4); n > 0; n-- {
|
||||||
|
r = unicode.SimpleFold(r)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
sampleRune(sb, r, rnd)
|
||||||
|
}
|
||||||
|
case syntax.OpCharClass:
|
||||||
|
if len(re.Rune) > 0 {
|
||||||
|
i := rnd.IntN(len(re.Rune)/2) * 2
|
||||||
|
sampleRune(sb, re.Rune[i]+rnd.Int32N(min(re.Rune[i+1]-re.Rune[i]+1, 300)), rnd)
|
||||||
|
}
|
||||||
|
case syntax.OpAnyChar, syntax.OpAnyCharNotNL:
|
||||||
|
sampleRune(sb, []rune{'a', '.', '\n', 0xe9, 0x212a, utf8.RuneError}[rnd.IntN(6)], rnd)
|
||||||
|
case syntax.OpCapture:
|
||||||
|
sampleMatch(sb, re.Sub[0], rnd, budget)
|
||||||
|
case syntax.OpConcat:
|
||||||
|
for _, sub := range re.Sub {
|
||||||
|
sampleMatch(sb, sub, rnd, budget)
|
||||||
|
}
|
||||||
|
case syntax.OpAlternate:
|
||||||
|
sampleMatch(sb, re.Sub[rnd.IntN(len(re.Sub))], rnd, budget)
|
||||||
|
case syntax.OpQuest, syntax.OpStar, syntax.OpPlus, syntax.OpRepeat:
|
||||||
|
lo, hi := 0, 3
|
||||||
|
switch re.Op {
|
||||||
|
case syntax.OpQuest:
|
||||||
|
hi = 1
|
||||||
|
case syntax.OpPlus:
|
||||||
|
lo = 1
|
||||||
|
case syntax.OpRepeat:
|
||||||
|
lo, hi = re.Min, re.Min+3
|
||||||
|
if re.Max >= 0 {
|
||||||
|
hi = min(hi, re.Max)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
for n := lo + rnd.IntN(hi-lo+1); n > 0; n-- {
|
||||||
|
sampleMatch(sb, re.Sub[0], rnd, budget)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func sampleRune(sb *strings.Builder, r rune, rnd *rand.Rand) {
|
||||||
|
if r == utf8.RuneError && rnd.IntN(2) == 0 {
|
||||||
|
sb.WriteByte(0x80 | byte(rnd.IntN(0x80))) // regexp matches an invalid byte as U+FFFD
|
||||||
|
return
|
||||||
|
}
|
||||||
|
sb.WriteRune(r)
|
||||||
|
}
|
||||||
|
|
||||||
|
func FuzzRegexMatcher(f *testing.F) {
|
||||||
|
inputs := []string{
|
||||||
|
"", "x", "yy", "abd", "ccc", "ABC", "abCDef", "abcdef", "abababcc", "a.b", "a\xffb", "a\uFFFDb",
|
||||||
|
"www.91porn.com", "ns1.awsdns-cn-01.top", "r1---sn-2x3abcde.googlevideo.com",
|
||||||
|
}
|
||||||
|
for _, test := range regexLiteralCases {
|
||||||
|
for _, s := range inputs {
|
||||||
|
f.Add(test.pattern, s)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
for _, test := range regexTailCases {
|
||||||
|
for _, s := range append(test.match, test.reject...) {
|
||||||
|
f.Add(test.pattern, s)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
f.Fuzz(func(t *testing.T, pattern, s string) {
|
||||||
|
re, err := regexp.Compile(pattern)
|
||||||
|
if err != nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
m, _ := newRegexMatcher(pattern)
|
||||||
|
check := func(s string) {
|
||||||
|
if got, want := m.Match(s), re.MatchString(s); got != want {
|
||||||
|
t.Errorf("pattern %q, input %q: got %v, want %v", pattern, s, got, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
check(s)
|
||||||
|
// random inputs seldom match, so also try strings built from the pattern
|
||||||
|
parsed, _ := syntax.Parse(pattern, syntax.Perl)
|
||||||
|
h := fnv.New64a()
|
||||||
|
h.Write([]byte(s))
|
||||||
|
rnd := rand.New(rand.NewPCG(h.Sum64(), 1))
|
||||||
|
for range 8 {
|
||||||
|
var sb strings.Builder
|
||||||
|
budget := 256
|
||||||
|
sampleMatch(&sb, parsed, rnd, &budget)
|
||||||
|
sample := sb.String()
|
||||||
|
check(sample)
|
||||||
|
check(s + sample)
|
||||||
|
if len(sample) > 0 && len(s) > 0 {
|
||||||
|
i := rnd.IntN(len(sample))
|
||||||
|
check(sample[:i] + s[:1] + sample[i+1:])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
@@ -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 result
|
||||||
return CompositeMatches(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
|
||||||
|
}
|
||||||
|
|||||||
+21
-25
@@ -1,7 +1,7 @@
|
|||||||
package log // import "github.com/xtls/xray-core/common/log"
|
package log // import "github.com/xtls/xray-core/common/log"
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"sync"
|
"sync/atomic"
|
||||||
|
|
||||||
"github.com/xtls/xray-core/common/serial"
|
"github.com/xtls/xray-core/common/serial"
|
||||||
)
|
)
|
||||||
@@ -29,36 +29,32 @@ func (m *GeneralMessage) String() string {
|
|||||||
|
|
||||||
// Record writes a message into log stream.
|
// Record writes a message into log stream.
|
||||||
func Record(msg Message) {
|
func Record(msg Message) {
|
||||||
logHandler.Handle(msg)
|
if h := logHandler.Load(); h != nil {
|
||||||
|
(*h).Handle(msg)
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
var logHandler syncHandler
|
type SeverityLogger interface {
|
||||||
|
Handler
|
||||||
|
Severity() Severity
|
||||||
|
}
|
||||||
|
|
||||||
|
func GetSeverity() Severity {
|
||||||
|
if h := logHandler.Load(); h != nil {
|
||||||
|
if sh, ok := (*h).(SeverityLogger); ok {
|
||||||
|
return sh.Severity()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
// log everything by default
|
||||||
|
return Severity_Debug
|
||||||
|
}
|
||||||
|
|
||||||
|
var logHandler atomic.Pointer[Handler]
|
||||||
|
|
||||||
// RegisterHandler registers a new handler as current log handler. Previous registered handler will be discarded.
|
// RegisterHandler registers a new handler as current log handler. Previous registered handler will be discarded.
|
||||||
func RegisterHandler(handler Handler) {
|
func RegisterHandler(handler Handler) {
|
||||||
if handler == nil {
|
if handler == nil {
|
||||||
panic("Log handler is nil")
|
panic("Log handler is nil")
|
||||||
}
|
}
|
||||||
logHandler.Set(handler)
|
logHandler.Store(&handler)
|
||||||
}
|
|
||||||
|
|
||||||
type syncHandler struct {
|
|
||||||
sync.RWMutex
|
|
||||||
Handler
|
|
||||||
}
|
|
||||||
|
|
||||||
func (h *syncHandler) Handle(msg Message) {
|
|
||||||
h.RLock()
|
|
||||||
defer h.RUnlock()
|
|
||||||
|
|
||||||
if h.Handler != nil {
|
|
||||||
h.Handler.Handle(msg)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func (h *syncHandler) Set(handler Handler) {
|
|
||||||
h.Lock()
|
|
||||||
defer h.Unlock()
|
|
||||||
|
|
||||||
h.Handler = handler
|
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -68,6 +68,10 @@ func (l *serverityLogger) Handle(msg Message) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (l *serverityLogger) Severity() Severity {
|
||||||
|
return l.logLevel
|
||||||
|
}
|
||||||
|
|
||||||
func (l *generalLogger) run() {
|
func (l *generalLogger) run() {
|
||||||
defer l.access.Signal()
|
defer l.access.Signal()
|
||||||
|
|
||||||
|
|||||||
@@ -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.NewTable()
|
||||||
|
var source, prefix string // cache
|
||||||
|
for name, severity := range map[string]Severity{
|
||||||
|
"Debug": Severity_Debug,
|
||||||
|
"Info": Severity_Info,
|
||||||
|
"Warning": Severity_Warning,
|
||||||
|
"Error": Severity_Error,
|
||||||
|
} {
|
||||||
|
module.RawSetString(name, L.NewFunction(func(L *lua.LState) int {
|
||||||
|
if GetSeverity() < severity {
|
||||||
|
return 0
|
||||||
|
}
|
||||||
|
var content strings.Builder
|
||||||
|
// Prefix with the calling script's filename.
|
||||||
|
if caller, ok := L.GetStack(1); ok {
|
||||||
|
if _, err := L.GetInfo("S", caller, lua.LNil); err == nil && caller.Source != "" {
|
||||||
|
if caller.Source != source {
|
||||||
|
source = caller.Source
|
||||||
|
prefix = filepath.Base(strings.TrimPrefix(source, "@")) + ": "
|
||||||
|
}
|
||||||
|
content.WriteString(prefix)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
for i := 1; i <= L.GetTop(); i++ {
|
||||||
|
content.WriteString(luaLogString(L, L.Get(i)))
|
||||||
|
}
|
||||||
|
Record(&GeneralMessage{
|
||||||
|
Severity: severity,
|
||||||
|
Content: content.String(),
|
||||||
|
})
|
||||||
|
return 0
|
||||||
|
}))
|
||||||
|
}
|
||||||
|
L.Push(module)
|
||||||
|
return 1
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func luaLogString(L *lua.LState, value lua.LValue) string {
|
||||||
|
if ud, ok := value.(*lua.LUserData); ok {
|
||||||
|
if err, ok := ud.Value.(error); ok {
|
||||||
|
return err.Error()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if _, ok := L.GetMetaField(value, "__tostring").(*lua.LFunction); ok {
|
||||||
|
return L.ToStringMeta(value).String()
|
||||||
|
}
|
||||||
|
return value.String()
|
||||||
|
}
|
||||||
@@ -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,150 @@
|
|||||||
|
package lua
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"errors"
|
||||||
|
"sync"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
glua "github.com/yuin/gopher-lua"
|
||||||
|
)
|
||||||
|
|
||||||
|
const maxIdleStates = 16
|
||||||
|
|
||||||
|
// Pool lends each state to one caller at a time. It grows on contention and
|
||||||
|
// keeps up to maxIdleStates idle states until Close. Acquire/Release callers
|
||||||
|
// decide reusability; WithState uses its callback's error.
|
||||||
|
type Pool struct {
|
||||||
|
ctx context.Context
|
||||||
|
cancel context.CancelFunc
|
||||||
|
timeout time.Duration
|
||||||
|
|
||||||
|
factory LStateFactory
|
||||||
|
idle []*glua.LState
|
||||||
|
top int
|
||||||
|
|
||||||
|
mu sync.Mutex
|
||||||
|
active sync.WaitGroup
|
||||||
|
closed bool
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewPool tests the factory by creating one state during initialization.
|
||||||
|
func NewPool(ctx context.Context, timeout time.Duration, factory LStateFactory) (*Pool, error) {
|
||||||
|
if timeout <= 0 {
|
||||||
|
return nil, errors.New("Lua pool timeout must be positive")
|
||||||
|
}
|
||||||
|
|
||||||
|
poolCtx, cancel := context.WithCancel(ctx)
|
||||||
|
|
||||||
|
state, err := factory(poolCtx)
|
||||||
|
if err != nil {
|
||||||
|
cancel()
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
return &Pool{ctx: poolCtx, cancel: cancel, timeout: timeout, factory: factory, idle: []*glua.LState{state}, top: state.GetTop()}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Acquire returns an initialized exclusive state, growing the pool if necessary.
|
||||||
|
// ctx is passed to the factory for state creation; nil uses the pool context.
|
||||||
|
func (p *Pool) Acquire(ctx context.Context) (*glua.LState, error) {
|
||||||
|
p.mu.Lock()
|
||||||
|
if p.closed {
|
||||||
|
p.mu.Unlock()
|
||||||
|
return nil, errors.New("Lua pool is closed")
|
||||||
|
}
|
||||||
|
if err := p.ctx.Err(); err != nil {
|
||||||
|
p.mu.Unlock()
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
if ctx == nil {
|
||||||
|
ctx = p.ctx
|
||||||
|
} else if err := ctx.Err(); err != nil {
|
||||||
|
p.mu.Unlock()
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
p.active.Add(1)
|
||||||
|
|
||||||
|
n := len(p.idle)
|
||||||
|
if n != 0 {
|
||||||
|
state := p.idle[n-1]
|
||||||
|
p.idle = p.idle[:n-1]
|
||||||
|
p.mu.Unlock()
|
||||||
|
return state, nil
|
||||||
|
}
|
||||||
|
p.mu.Unlock()
|
||||||
|
|
||||||
|
// TODO: Limit the total number of states. When the limit is reached, wait
|
||||||
|
// for a Release instead of creating another state; allow the wait to be
|
||||||
|
// cancelled by the caller or by Close.
|
||||||
|
state, err := p.factory(ctx)
|
||||||
|
if err != nil {
|
||||||
|
p.active.Done()
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
return state, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// WithState runs work on an exclusive state and releases it afterward.
|
||||||
|
// Nil ctx and zero timeout use pool defaults. The timeout starts after acquisition.
|
||||||
|
func (p *Pool) WithState(ctx context.Context, timeout time.Duration, work func(*glua.LState) error) error {
|
||||||
|
state, err := p.Acquire(ctx)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if ctx == nil {
|
||||||
|
ctx = p.ctx
|
||||||
|
}
|
||||||
|
if timeout == 0 {
|
||||||
|
timeout = p.timeout
|
||||||
|
}
|
||||||
|
ctx, cancel := context.WithTimeout(ctx, timeout)
|
||||||
|
state.SetContext(ctx)
|
||||||
|
reusable := false
|
||||||
|
defer func() {
|
||||||
|
cancel()
|
||||||
|
p.Release(state, reusable)
|
||||||
|
}()
|
||||||
|
err = work(state)
|
||||||
|
reusable = err == nil
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
// Release resets a state for reuse or closes it.
|
||||||
|
func (p *Pool) Release(state *glua.LState, reusable bool) {
|
||||||
|
if reusable {
|
||||||
|
state.RemoveContext()
|
||||||
|
state.SetTop(p.top)
|
||||||
|
p.mu.Lock()
|
||||||
|
if !p.closed && p.ctx.Err() == nil && len(p.idle) < maxIdleStates {
|
||||||
|
p.idle = append(p.idle, state)
|
||||||
|
} else {
|
||||||
|
reusable = false
|
||||||
|
}
|
||||||
|
p.mu.Unlock()
|
||||||
|
}
|
||||||
|
|
||||||
|
if !reusable {
|
||||||
|
state.Close()
|
||||||
|
}
|
||||||
|
|
||||||
|
p.active.Done()
|
||||||
|
}
|
||||||
|
|
||||||
|
// Close cancels the pool context, closes idle states, and waits for borrowed states.
|
||||||
|
func (p *Pool) Close() {
|
||||||
|
p.mu.Lock()
|
||||||
|
if !p.closed {
|
||||||
|
p.closed = true
|
||||||
|
p.cancel()
|
||||||
|
for _, state := range p.idle {
|
||||||
|
state.Close()
|
||||||
|
}
|
||||||
|
p.idle = nil
|
||||||
|
}
|
||||||
|
p.mu.Unlock()
|
||||||
|
|
||||||
|
p.active.Wait()
|
||||||
|
}
|
||||||
@@ -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)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -38,7 +38,7 @@ func (m *ClientManager) Dispatch(ctx context.Context, link *transport.Link) erro
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
return errors.New("unable to find an available mux client").AtWarning()
|
return errors.New("unable to find an available mux client")
|
||||||
}
|
}
|
||||||
|
|
||||||
type WorkerPicker interface {
|
type WorkerPicker interface {
|
||||||
|
|||||||
+1
-1
@@ -117,7 +117,7 @@ func (f *FrameMetadata) Unmarshal(reader io.Reader, readSourceAndLocal bool) err
|
|||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
if metaLen > 512 {
|
if metaLen > 512 {
|
||||||
return errors.New("invalid metalen ", metaLen).AtError()
|
return errors.New("invalid metalen ", metaLen)
|
||||||
}
|
}
|
||||||
|
|
||||||
b := buf.New()
|
b := buf.New()
|
||||||
|
|||||||
@@ -351,7 +351,7 @@ func (w *ServerWorker) handleFrame(ctx context.Context, reader *buf.BufferedRead
|
|||||||
err = w.handleStatusKeep(&meta, reader)
|
err = w.handleStatusKeep(&meta, reader)
|
||||||
default:
|
default:
|
||||||
status := meta.SessionStatus
|
status := meta.SessionStatus
|
||||||
return errors.New("unknown status: ", status).AtError()
|
return errors.New("unknown status: ", status)
|
||||||
}
|
}
|
||||||
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
|||||||
@@ -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,16 +1,13 @@
|
|||||||
package http
|
package http
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"bufio"
|
|
||||||
"bytes"
|
"bytes"
|
||||||
"context"
|
"context"
|
||||||
"errors"
|
"errors"
|
||||||
"io"
|
|
||||||
"net/http"
|
|
||||||
"strings"
|
"strings"
|
||||||
"unsafe"
|
|
||||||
|
|
||||||
"github.com/xtls/xray-core/common"
|
"github.com/xtls/xray-core/common"
|
||||||
|
"github.com/xtls/xray-core/common/net"
|
||||||
"github.com/xtls/xray-core/common/session"
|
"github.com/xtls/xray-core/common/session"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -42,67 +39,79 @@ func (h *SniffHeader) Domain() string {
|
|||||||
}
|
}
|
||||||
|
|
||||||
var (
|
var (
|
||||||
validMethods = map[string]bool{}
|
methods = [...]string{"get", "post", "head", "put", "delete", "options", "connect"}
|
||||||
errNotHTTP = errors.New("not an HTTP request")
|
|
||||||
|
errNotHTTPMethod = errors.New("not an HTTP method")
|
||||||
)
|
)
|
||||||
|
|
||||||
func init() {
|
func beginWithHTTPMethod(b []byte) error {
|
||||||
// https://www.iana.org/assignments/http-methods
|
for _, m := range &methods {
|
||||||
methods := []string{
|
if len(b) >= len(m) && strings.EqualFold(string(b[:len(m)]), m) {
|
||||||
"ACL", "BASELINE-CONTROL", "BIND", "CHECKIN", "CHECKOUT",
|
return nil
|
||||||
"CONNECT", "COPY", "DELETE", "GET", "HEAD",
|
|
||||||
"LABEL", "LINK", "LOCK", "MERGE", "MKACTIVITY",
|
|
||||||
"MKCALENDAR", "MKCOL", "MKREDIRECTREF", "MKWORKSPACE", "MOVE",
|
|
||||||
"OPTIONS", "ORDERPATCH", "PATCH", "POST", "PRI",
|
|
||||||
"PROPFIND", "PROPPATCH", "PUT", "QUERY", "REBIND",
|
|
||||||
"REPORT", "SEARCH", "TRACE", "UNBIND", "UNCHECKOUT",
|
|
||||||
"UNLINK", "UNLOCK", "UPDATE", "UPDATEREDIRECTREF", "VERSION-CONTROL",
|
|
||||||
}
|
}
|
||||||
for _, m := range methods {
|
|
||||||
validMethods[m] = true
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func isValidHTTPMethod(b []byte) bool {
|
if len(b) < len(m) {
|
||||||
if len(b) == 0 {
|
return common.ErrNoClue
|
||||||
return false
|
|
||||||
}
|
}
|
||||||
idx := bytes.IndexByte(b, ' ')
|
|
||||||
if idx == -1 {
|
|
||||||
return false
|
|
||||||
}
|
}
|
||||||
method := unsafe.String(unsafe.SliceData(b), idx)
|
|
||||||
return validMethods[method]
|
return errNotHTTPMethod
|
||||||
}
|
}
|
||||||
|
|
||||||
func SniffHTTP(b []byte, c context.Context) (*SniffHeader, error) {
|
func SniffHTTP(b []byte, c context.Context) (*SniffHeader, error) {
|
||||||
if !isValidHTTPMethod(b) {
|
|
||||||
return nil, errNotHTTP
|
|
||||||
}
|
|
||||||
content := session.ContentFromContext(c)
|
content := session.ContentFromContext(c)
|
||||||
r, err := http.ReadRequest(bufio.NewReader(bytes.NewReader(b)))
|
ShouldSniffAttr := true
|
||||||
if err != nil {
|
|
||||||
if err == io.ErrUnexpectedEOF {
|
|
||||||
return nil, common.ErrNoClue
|
|
||||||
}
|
|
||||||
return nil, errNotHTTP
|
|
||||||
}
|
|
||||||
if r.Host == "" {
|
|
||||||
return nil, common.ErrNoClue
|
|
||||||
}
|
|
||||||
sh := &SniffHeader{
|
|
||||||
version: HTTP1,
|
|
||||||
host: r.Host,
|
|
||||||
}
|
|
||||||
// If content.Attributes have information, that means it comes from HTTP inbound PlainHTTP mode.
|
// If content.Attributes have information, that means it comes from HTTP inbound PlainHTTP mode.
|
||||||
// It will set attributes, so skip it.
|
// It will set attributes, so skip it.
|
||||||
if content != nil && len(content.Attributes) == 0 {
|
if content == nil || len(content.Attributes) != 0 {
|
||||||
for key, h := range r.Header {
|
ShouldSniffAttr = false
|
||||||
content.Attributes[key] = strings.Join(h, ",")
|
|
||||||
}
|
}
|
||||||
content.Attributes[":method"] = r.Method
|
if err := beginWithHTTPMethod(b); err != nil {
|
||||||
content.Attributes[":path"] = r.URL.Path
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
|
sh := &SniffHeader{
|
||||||
|
version: HTTP1,
|
||||||
|
}
|
||||||
|
|
||||||
|
headers := bytes.Split(b, []byte{'\n'})
|
||||||
|
for i := 1; i < len(headers); i++ {
|
||||||
|
header := headers[i]
|
||||||
|
if len(header) == 0 {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
parts := bytes.SplitN(header, []byte{':'}, 2)
|
||||||
|
if len(parts) != 2 {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
key := strings.ToLower(string(parts[0]))
|
||||||
|
value := string(bytes.TrimSpace(parts[1]))
|
||||||
|
if ShouldSniffAttr {
|
||||||
|
content.SetAttribute(key, value) // Put header in attribute
|
||||||
|
}
|
||||||
|
if key == "host" {
|
||||||
|
rawHost := strings.ToLower(value)
|
||||||
|
dest, err := ParseHost(rawHost, net.Port(80))
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
sh.host = dest.Address.String()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
// Parse request line
|
||||||
|
// Request line is like this
|
||||||
|
// "GET /homo/114514 HTTP/1.1"
|
||||||
|
if len(headers) > 0 && ShouldSniffAttr {
|
||||||
|
RequestLineParts := bytes.Split(headers[0], []byte{' '})
|
||||||
|
if len(RequestLineParts) == 3 {
|
||||||
|
content.SetAttribute(":method", string(RequestLineParts[0]))
|
||||||
|
content.SetAttribute(":path", string(RequestLineParts[1]))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(sh.host) > 0 {
|
||||||
return sh, nil
|
return sh, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil, common.ErrNoClue
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -14,76 +14,75 @@ func TestHTTPHeaders(t *testing.T) {
|
|||||||
err bool
|
err bool
|
||||||
}{
|
}{
|
||||||
{
|
{
|
||||||
input: "GET /tutorials/other/top-20-mysql-best-practices/ HTTP/1.1\r\n" +
|
input: `GET /tutorials/other/top-20-mysql-best-practices/ HTTP/1.1
|
||||||
"Host: net.tutsplus.com\r\n" +
|
Host: net.tutsplus.com
|
||||||
"User-Agent: Mozilla/5.0 (Windows; U; Windows NT 6.1; en-US; rv:1.9.1.5) Gecko/20091102 Firefox/3.5.5 (.NET CLR 3.5.30729)\r\n" +
|
User-Agent: Mozilla/5.0 (Windows; U; Windows NT 6.1; en-US; rv:1.9.1.5) Gecko/20091102 Firefox/3.5.5 (.NET CLR 3.5.30729)
|
||||||
"Accept: text/html,application/xhtml+xml,application/xml;q=0.9,*/*;q=0.8\r\n" +
|
Accept: text/html,application/xhtml+xml,application/xml;q=0.9,*/*;q=0.8
|
||||||
"Accept-Language: en-us,en;q=0.5\r\n" +
|
Accept-Language: en-us,en;q=0.5
|
||||||
"Accept-Encoding: gzip,deflate\r\n" +
|
Accept-Encoding: gzip,deflate
|
||||||
"Accept-Charset: ISO-8859-1,utf-8;q=0.7,*;q=0.7\r\n" +
|
Accept-Charset: ISO-8859-1,utf-8;q=0.7,*;q=0.7
|
||||||
"Keep-Alive: 300\r\n" +
|
Keep-Alive: 300
|
||||||
"Connection: keep-alive\r\n" +
|
Connection: keep-alive
|
||||||
"Cookie: PHPSESSID=r2t5uvjq435r4q7ib3vtdjq120\r\n" +
|
Cookie: PHPSESSID=r2t5uvjq435r4q7ib3vtdjq120
|
||||||
"Pragma: no-cache\r\n" +
|
Pragma: no-cache
|
||||||
"Cache-Control: no-cache\r\n" +
|
Cache-Control: no-cache`,
|
||||||
"\r\n",
|
|
||||||
domain: "net.tutsplus.com",
|
domain: "net.tutsplus.com",
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
input: "POST /foo.php HTTP/1.1\r\n" +
|
input: `POST /foo.php HTTP/1.1
|
||||||
"Host: localhost\r\n" +
|
Host: localhost
|
||||||
"User-Agent: Mozilla/5.0 (Windows; U; Windows NT 6.1; en-US; rv:1.9.1.5) Gecko/20091102 Firefox/3.5.5 (.NET CLR 3.5.30729)\r\n" +
|
User-Agent: Mozilla/5.0 (Windows; U; Windows NT 6.1; en-US; rv:1.9.1.5) Gecko/20091102 Firefox/3.5.5 (.NET CLR 3.5.30729)
|
||||||
"Accept: text/html,application/xhtml+xml,application/xml;q=0.9,*/*;q=0.8\r\n" +
|
Accept: text/html,application/xhtml+xml,application/xml;q=0.9,*/*;q=0.8
|
||||||
"Accept-Language: en-us,en;q=0.5\r\n" +
|
Accept-Language: en-us,en;q=0.5
|
||||||
"Accept-Encoding: gzip,deflate\r\n" +
|
Accept-Encoding: gzip,deflate
|
||||||
"Accept-Charset: ISO-8859-1,utf-8;q=0.7,*;q=0.7\r\n" +
|
Accept-Charset: ISO-8859-1,utf-8;q=0.7,*;q=0.7
|
||||||
"Keep-Alive: 300\r\n" +
|
Keep-Alive: 300
|
||||||
"Connection: keep-alive\r\n" +
|
Connection: keep-alive
|
||||||
"Referer: http://localhost/test.php\r\n" +
|
Referer: http://localhost/test.php
|
||||||
"Content-Type: application/x-www-form-urlencoded\r\n" +
|
Content-Type: application/x-www-form-urlencoded
|
||||||
"Content-Length: 43\r\n" +
|
Content-Length: 43
|
||||||
"\r\n" +
|
|
||||||
"first_name=John&last_name=Doe&action=Submit",
|
first_name=John&last_name=Doe&action=Submit`,
|
||||||
domain: "localhost",
|
domain: "localhost",
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
input: "X /foo.php HTTP/1.1\r\n" +
|
input: `X /foo.php HTTP/1.1
|
||||||
"Host: localhost\r\n" +
|
Host: localhost
|
||||||
"User-Agent: Mozilla/5.0 (Windows; U; Windows NT 6.1; en-US; rv:1.9.1.5) Gecko/20091102 Firefox/3.5.5 (.NET CLR 3.5.30729)\r\n" +
|
User-Agent: Mozilla/5.0 (Windows; U; Windows NT 6.1; en-US; rv:1.9.1.5) Gecko/20091102 Firefox/3.5.5 (.NET CLR 3.5.30729)
|
||||||
"Accept: text/html,application/xhtml+xml,application/xml;q=0.9,*/*;q=0.8\r\n" +
|
Accept: text/html,application/xhtml+xml,application/xml;q=0.9,*/*;q=0.8
|
||||||
"Accept-Language: en-us,en;q=0.5\r\n" +
|
Accept-Language: en-us,en;q=0.5
|
||||||
"Accept-Encoding: gzip,deflate\r\n" +
|
Accept-Encoding: gzip,deflate
|
||||||
"Accept-Charset: ISO-8859-1,utf-8;q=0.7,*;q=0.7\r\n" +
|
Accept-Charset: ISO-8859-1,utf-8;q=0.7,*;q=0.7
|
||||||
"Keep-Alive: 300\r\n" +
|
Keep-Alive: 300
|
||||||
"Connection: keep-alive\r\n" +
|
Connection: keep-alive
|
||||||
"Referer: http://localhost/test.php\r\n" +
|
Referer: http://localhost/test.php
|
||||||
"Content-Type: application/x-www-form-urlencoded\r\n" +
|
Content-Type: application/x-www-form-urlencoded
|
||||||
"Content-Length: 43\r\n" +
|
Content-Length: 43
|
||||||
"\r\n" +
|
|
||||||
"first_name=John&last_name=Doe&action=Submit",
|
first_name=John&last_name=Doe&action=Submit`,
|
||||||
domain: "",
|
domain: "",
|
||||||
err: true,
|
err: true,
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
input: "GET /foo.php HTTP/1.1\r\n" +
|
input: `GET /foo.php HTTP/1.1
|
||||||
"User-Agent: Mozilla/5.0 (Windows; U; Windows NT 6.1; en-US; rv:1.9.1.5) Gecko/20091102 Firefox/3.5.5 (.NET CLR 3.5.30729)\r\n" +
|
User-Agent: Mozilla/5.0 (Windows; U; Windows NT 6.1; en-US; rv:1.9.1.5) Gecko/20091102 Firefox/3.5.5 (.NET CLR 3.5.30729)
|
||||||
"Accept: text/html,application/xhtml+xml,application/xml;q=0.9,*/*;q=0.8\r\n" +
|
Accept: text/html,application/xhtml+xml,application/xml;q=0.9,*/*;q=0.8
|
||||||
"Accept-Language: en-us,en;q=0.5\r\n" +
|
Accept-Language: en-us,en;q=0.5
|
||||||
"Accept-Encoding: gzip,deflate\r\n" +
|
Accept-Encoding: gzip,deflate
|
||||||
"Accept-Charset: ISO-8859-1,utf-8;q=0.7,*;q=0.7\r\n" +
|
Accept-Charset: ISO-8859-1,utf-8;q=0.7,*;q=0.7
|
||||||
"Keep-Alive: 300\r\n" +
|
Keep-Alive: 300
|
||||||
"Connection: keep-alive\r\n" +
|
Connection: keep-alive
|
||||||
"Referer: http://localhost/test.php\r\n" +
|
Referer: http://localhost/test.php
|
||||||
"Content-Type: application/x-www-form-urlencoded\r\n" +
|
Content-Type: application/x-www-form-urlencoded
|
||||||
"Content-Length: 43\r\n" +
|
Content-Length: 43
|
||||||
"\r\n" +
|
|
||||||
"Host: localhost\r\n" +
|
Host: localhost
|
||||||
"first_name=John&last_name=Doe&action=Submit",
|
first_name=John&last_name=Doe&action=Submit`,
|
||||||
domain: "",
|
domain: "",
|
||||||
err: true,
|
err: true,
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
input: "GET /tutorials/other/top-20-mysql-best-practices/ HTTP/1.1\r\n",
|
input: `GET /tutorials/other/top-20-mysql-best-practices/ HTTP/1.1`,
|
||||||
domain: "",
|
domain: "",
|
||||||
err: true,
|
err: true,
|
||||||
},
|
},
|
||||||
@@ -98,7 +97,6 @@ func TestHTTPHeaders(t *testing.T) {
|
|||||||
} else {
|
} else {
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Errorf("Expect no error but actually %s in test %v", err.Error(), test)
|
t.Errorf("Expect no error but actually %s in test %v", err.Error(), test)
|
||||||
continue
|
|
||||||
}
|
}
|
||||||
if header.Domain() != test.domain {
|
if header.Domain() != test.domain {
|
||||||
t.Error("expected domain ", test.domain, " but got ", header.Domain())
|
t.Error("expected domain ", test.domain, " but got ", header.Domain())
|
||||||
|
|||||||
@@ -7,7 +7,7 @@ import (
|
|||||||
|
|
||||||
func (u *User) GetTypedAccount() (Account, error) {
|
func (u *User) GetTypedAccount() (Account, error) {
|
||||||
if u.GetAccount() == nil {
|
if u.GetAccount() == nil {
|
||||||
return nil, errors.New("Account is missing").AtWarning()
|
return nil, errors.New("Account is missing")
|
||||||
}
|
}
|
||||||
|
|
||||||
rawAccount, err := u.Account.GetInstance()
|
rawAccount, err := u.Account.GetInstance()
|
||||||
|
|||||||
@@ -70,8 +70,6 @@ type Outbound struct {
|
|||||||
Tag string
|
Tag string
|
||||||
// Name of the outbound proxy that handles the connection.
|
// Name of the outbound proxy that handles the connection.
|
||||||
Name string
|
Name string
|
||||||
// Unused. Conn is actually internet.Connection. May be nil. It is currently nil for outbound with proxySettings
|
|
||||||
Conn net.Conn
|
|
||||||
// CanSpliceCopy is a property for this connection
|
// CanSpliceCopy is a property for this connection
|
||||||
// 1 = can, 2 = after processing protocol info should be able to, 3 = cannot
|
// 1 = can, 2 = after processing protocol info should be able to, 3 = cannot
|
||||||
CanSpliceCopy int
|
CanSpliceCopy int
|
||||||
|
|||||||
@@ -1,53 +0,0 @@
|
|||||||
package singbridge
|
|
||||||
|
|
||||||
import (
|
|
||||||
M "github.com/sagernet/sing/common/metadata"
|
|
||||||
N "github.com/sagernet/sing/common/network"
|
|
||||||
"github.com/xtls/xray-core/common/errors"
|
|
||||||
"github.com/xtls/xray-core/common/net"
|
|
||||||
)
|
|
||||||
|
|
||||||
func ToNetwork(network string) net.Network {
|
|
||||||
switch N.NetworkName(network) {
|
|
||||||
case N.NetworkTCP:
|
|
||||||
return net.Network_TCP
|
|
||||||
case N.NetworkUDP:
|
|
||||||
return net.Network_UDP
|
|
||||||
default:
|
|
||||||
return net.Network_Unknown
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func ToDestination(socksaddr M.Socksaddr, network net.Network) (net.Destination, error) {
|
|
||||||
// IsFqdn() implicitly checks if the domain name is valid
|
|
||||||
if socksaddr.IsFqdn() {
|
|
||||||
return net.Destination{
|
|
||||||
Network: network,
|
|
||||||
Address: net.DomainAddress(socksaddr.Fqdn),
|
|
||||||
Port: net.Port(socksaddr.Port),
|
|
||||||
}, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// IsIP() implicitly checks if the IP address is valid
|
|
||||||
if socksaddr.IsIP() {
|
|
||||||
return net.Destination{
|
|
||||||
Network: network,
|
|
||||||
Address: net.IPAddress(socksaddr.Addr.AsSlice()),
|
|
||||||
Port: net.Port(socksaddr.Port),
|
|
||||||
}, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
return net.Destination{}, errors.New("invalid socks address: ", socksaddr)
|
|
||||||
}
|
|
||||||
|
|
||||||
func ToSocksaddr(destination net.Destination) M.Socksaddr {
|
|
||||||
var addr M.Socksaddr
|
|
||||||
switch destination.Address.Family() {
|
|
||||||
case net.AddressFamilyDomain:
|
|
||||||
addr.Fqdn = destination.Address.Domain()
|
|
||||||
default:
|
|
||||||
addr.Addr = M.AddrFromIP(destination.Address.IP())
|
|
||||||
}
|
|
||||||
addr.Port = uint16(destination.Port)
|
|
||||||
return addr
|
|
||||||
}
|
|
||||||
@@ -1,72 +0,0 @@
|
|||||||
package singbridge
|
|
||||||
|
|
||||||
import (
|
|
||||||
"context"
|
|
||||||
"os"
|
|
||||||
|
|
||||||
M "github.com/sagernet/sing/common/metadata"
|
|
||||||
N "github.com/sagernet/sing/common/network"
|
|
||||||
"github.com/xtls/xray-core/common/net"
|
|
||||||
"github.com/xtls/xray-core/common/net/cnc"
|
|
||||||
"github.com/xtls/xray-core/common/session"
|
|
||||||
"github.com/xtls/xray-core/proxy"
|
|
||||||
"github.com/xtls/xray-core/transport"
|
|
||||||
"github.com/xtls/xray-core/transport/internet"
|
|
||||||
"github.com/xtls/xray-core/transport/pipe"
|
|
||||||
)
|
|
||||||
|
|
||||||
var _ N.Dialer = (*XrayDialer)(nil)
|
|
||||||
|
|
||||||
type XrayDialer struct {
|
|
||||||
internet.Dialer
|
|
||||||
}
|
|
||||||
|
|
||||||
func NewDialer(dialer internet.Dialer) *XrayDialer {
|
|
||||||
return &XrayDialer{dialer}
|
|
||||||
}
|
|
||||||
|
|
||||||
func (d *XrayDialer) DialContext(ctx context.Context, network string, destination M.Socksaddr) (net.Conn, error) {
|
|
||||||
dest, err := ToDestination(destination, ToNetwork(network))
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
return d.Dialer.Dial(ctx, dest)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (d *XrayDialer) ListenPacket(ctx context.Context, destination M.Socksaddr) (net.PacketConn, error) {
|
|
||||||
return nil, os.ErrInvalid
|
|
||||||
}
|
|
||||||
|
|
||||||
type XrayOutboundDialer struct {
|
|
||||||
outbound proxy.Outbound
|
|
||||||
dialer internet.Dialer
|
|
||||||
}
|
|
||||||
|
|
||||||
func NewOutboundDialer(outbound proxy.Outbound, dialer internet.Dialer) *XrayOutboundDialer {
|
|
||||||
return &XrayOutboundDialer{outbound, dialer}
|
|
||||||
}
|
|
||||||
|
|
||||||
func (d *XrayOutboundDialer) DialContext(ctx context.Context, network string, destination M.Socksaddr) (net.Conn, error) {
|
|
||||||
dest, err := ToDestination(destination, ToNetwork(network))
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
outbounds := session.OutboundsFromContext(ctx)
|
|
||||||
if len(outbounds) == 0 {
|
|
||||||
outbounds = []*session.Outbound{{}}
|
|
||||||
ctx = session.ContextWithOutbounds(ctx, outbounds)
|
|
||||||
}
|
|
||||||
ob := outbounds[len(outbounds)-1]
|
|
||||||
ob.Target = dest
|
|
||||||
|
|
||||||
opts := []pipe.Option{pipe.WithSizeLimit(64 * 1024)}
|
|
||||||
uplinkReader, uplinkWriter := pipe.New(opts...)
|
|
||||||
downlinkReader, downlinkWriter := pipe.New(opts...)
|
|
||||||
conn := cnc.NewConnection(cnc.ConnectionInputMulti(downlinkWriter), cnc.ConnectionOutputMulti(uplinkReader))
|
|
||||||
go d.outbound.Process(ctx, &transport.Link{Reader: downlinkReader, Writer: uplinkWriter}, d.dialer)
|
|
||||||
return conn, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (d *XrayOutboundDialer) ListenPacket(ctx context.Context, destination M.Socksaddr) (net.PacketConn, error) {
|
|
||||||
return nil, os.ErrInvalid
|
|
||||||
}
|
|
||||||
@@ -1,10 +0,0 @@
|
|||||||
package singbridge
|
|
||||||
|
|
||||||
import E "github.com/sagernet/sing/common/exceptions"
|
|
||||||
|
|
||||||
func ReturnError(err error) error {
|
|
||||||
if E.IsClosedOrCanceled(err) {
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
@@ -1,58 +0,0 @@
|
|||||||
package singbridge
|
|
||||||
|
|
||||||
import (
|
|
||||||
"context"
|
|
||||||
"io"
|
|
||||||
|
|
||||||
M "github.com/sagernet/sing/common/metadata"
|
|
||||||
N "github.com/sagernet/sing/common/network"
|
|
||||||
"github.com/xtls/xray-core/common/buf"
|
|
||||||
"github.com/xtls/xray-core/common/errors"
|
|
||||||
"github.com/xtls/xray-core/common/net"
|
|
||||||
"github.com/xtls/xray-core/features/routing"
|
|
||||||
"github.com/xtls/xray-core/transport"
|
|
||||||
)
|
|
||||||
|
|
||||||
var (
|
|
||||||
_ N.TCPConnectionHandler = (*Dispatcher)(nil)
|
|
||||||
_ N.UDPConnectionHandler = (*Dispatcher)(nil)
|
|
||||||
)
|
|
||||||
|
|
||||||
type Dispatcher struct {
|
|
||||||
upstream routing.Dispatcher
|
|
||||||
newErrorFunc func(values ...any) *errors.Error
|
|
||||||
}
|
|
||||||
|
|
||||||
func NewDispatcher(dispatcher routing.Dispatcher, newErrorFunc func(values ...any) *errors.Error) *Dispatcher {
|
|
||||||
return &Dispatcher{
|
|
||||||
upstream: dispatcher,
|
|
||||||
newErrorFunc: newErrorFunc,
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func (d *Dispatcher) NewConnection(ctx context.Context, conn net.Conn, metadata M.Metadata) error {
|
|
||||||
dest, err := ToDestination(metadata.Destination, net.Network_TCP)
|
|
||||||
if err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
xConn := NewConn(conn)
|
|
||||||
return d.upstream.DispatchLink(ctx, dest, &transport.Link{
|
|
||||||
Reader: xConn,
|
|
||||||
Writer: xConn,
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
func (d *Dispatcher) NewPacketConnection(ctx context.Context, conn N.PacketConn, metadata M.Metadata) error {
|
|
||||||
dest, err := ToDestination(metadata.Destination, net.Network_UDP)
|
|
||||||
if err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
return d.upstream.DispatchLink(ctx, dest, &transport.Link{
|
|
||||||
Reader: buf.NewPacketReader(conn.(io.Reader)),
|
|
||||||
Writer: buf.NewWriter(conn.(io.Writer)),
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
func (d *Dispatcher) NewError(ctx context.Context, err error) {
|
|
||||||
errors.LogInfo(ctx, err.Error())
|
|
||||||
}
|
|
||||||
@@ -1,70 +0,0 @@
|
|||||||
package singbridge
|
|
||||||
|
|
||||||
import (
|
|
||||||
"context"
|
|
||||||
|
|
||||||
"github.com/sagernet/sing/common/logger"
|
|
||||||
"github.com/xtls/xray-core/common/errors"
|
|
||||||
)
|
|
||||||
|
|
||||||
var _ logger.ContextLogger = (*XrayLogger)(nil)
|
|
||||||
|
|
||||||
type XrayLogger struct {
|
|
||||||
newError func(values ...any) *errors.Error
|
|
||||||
}
|
|
||||||
|
|
||||||
func NewLogger(newErrorFunc func(values ...any) *errors.Error) *XrayLogger {
|
|
||||||
return &XrayLogger{
|
|
||||||
newErrorFunc,
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func (l *XrayLogger) Trace(args ...any) {
|
|
||||||
}
|
|
||||||
|
|
||||||
func (l *XrayLogger) Debug(args ...any) {
|
|
||||||
errors.LogDebug(context.Background(), args...)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (l *XrayLogger) Info(args ...any) {
|
|
||||||
errors.LogInfo(context.Background(), args...)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (l *XrayLogger) Warn(args ...any) {
|
|
||||||
errors.LogWarning(context.Background(), args...)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (l *XrayLogger) Error(args ...any) {
|
|
||||||
errors.LogError(context.Background(), args...)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (l *XrayLogger) Fatal(args ...any) {
|
|
||||||
}
|
|
||||||
|
|
||||||
func (l *XrayLogger) Panic(args ...any) {
|
|
||||||
}
|
|
||||||
|
|
||||||
func (l *XrayLogger) TraceContext(ctx context.Context, args ...any) {
|
|
||||||
}
|
|
||||||
|
|
||||||
func (l *XrayLogger) DebugContext(ctx context.Context, args ...any) {
|
|
||||||
errors.LogDebug(ctx, args...)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (l *XrayLogger) InfoContext(ctx context.Context, args ...any) {
|
|
||||||
errors.LogInfo(ctx, args...)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (l *XrayLogger) WarnContext(ctx context.Context, args ...any) {
|
|
||||||
errors.LogWarning(ctx, args...)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (l *XrayLogger) ErrorContext(ctx context.Context, args ...any) {
|
|
||||||
errors.LogError(ctx, args...)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (l *XrayLogger) FatalContext(ctx context.Context, args ...any) {
|
|
||||||
}
|
|
||||||
|
|
||||||
func (l *XrayLogger) PanicContext(ctx context.Context, args ...any) {
|
|
||||||
}
|
|
||||||
@@ -1,107 +0,0 @@
|
|||||||
package singbridge
|
|
||||||
|
|
||||||
import (
|
|
||||||
"context"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
B "github.com/sagernet/sing/common/buf"
|
|
||||||
"github.com/sagernet/sing/common/bufio"
|
|
||||||
M "github.com/sagernet/sing/common/metadata"
|
|
||||||
"github.com/xtls/xray-core/common"
|
|
||||||
"github.com/xtls/xray-core/common/buf"
|
|
||||||
"github.com/xtls/xray-core/common/net"
|
|
||||||
"github.com/xtls/xray-core/common/signal"
|
|
||||||
"github.com/xtls/xray-core/transport"
|
|
||||||
)
|
|
||||||
|
|
||||||
func CopyPacketConn(ctx context.Context, inboundConn net.Conn, link *transport.Link, destination net.Destination, serverConn net.PacketConn) error {
|
|
||||||
cancel := func() {
|
|
||||||
common.Interrupt(link.Reader)
|
|
||||||
common.Interrupt(serverConn)
|
|
||||||
}
|
|
||||||
conn := &PacketConnWrapper{
|
|
||||||
Reader: link.Reader,
|
|
||||||
Writer: link.Writer,
|
|
||||||
Dest: destination,
|
|
||||||
Conn: inboundConn,
|
|
||||||
T: signal.CancelAfterInactivity(ctx, cancel, 300*time.Second),
|
|
||||||
}
|
|
||||||
return ReturnError(bufio.CopyPacketConn(ctx, conn, bufio.NewPacketConn(serverConn)))
|
|
||||||
}
|
|
||||||
|
|
||||||
type PacketConnWrapper struct {
|
|
||||||
buf.Reader
|
|
||||||
buf.Writer
|
|
||||||
net.Conn
|
|
||||||
Dest net.Destination
|
|
||||||
cached buf.MultiBuffer
|
|
||||||
|
|
||||||
// A simple patch to avoid goroutine leak since sing infra cannot awake read block by write err
|
|
||||||
T *signal.ActivityTimer
|
|
||||||
}
|
|
||||||
|
|
||||||
func (w *PacketConnWrapper) ReadPacket(buffer *B.Buffer) (addr M.Socksaddr, err error) {
|
|
||||||
w.T.Update()
|
|
||||||
defer func() {
|
|
||||||
if err != nil {
|
|
||||||
// uplinkonly
|
|
||||||
w.T.SetTimeout(2 * time.Second)
|
|
||||||
}
|
|
||||||
}()
|
|
||||||
if w.cached != nil {
|
|
||||||
mb, bb := buf.SplitFirst(w.cached)
|
|
||||||
if bb == nil {
|
|
||||||
w.cached = nil
|
|
||||||
} else {
|
|
||||||
buffer.Write(bb.Bytes())
|
|
||||||
w.cached = mb
|
|
||||||
var destination net.Destination
|
|
||||||
if bb.UDP != nil {
|
|
||||||
destination = *bb.UDP
|
|
||||||
} else {
|
|
||||||
destination = w.Dest
|
|
||||||
}
|
|
||||||
bb.Release()
|
|
||||||
return ToSocksaddr(destination), nil
|
|
||||||
}
|
|
||||||
}
|
|
||||||
mb, err := w.ReadMultiBuffer()
|
|
||||||
nb, bb := buf.SplitFirst(mb)
|
|
||||||
if bb == nil {
|
|
||||||
return M.Socksaddr{}, nil
|
|
||||||
} else {
|
|
||||||
buffer.Write(bb.Bytes())
|
|
||||||
w.cached = nb
|
|
||||||
var destination net.Destination
|
|
||||||
if bb.UDP != nil {
|
|
||||||
destination = *bb.UDP
|
|
||||||
} else {
|
|
||||||
destination = w.Dest
|
|
||||||
}
|
|
||||||
bb.Release()
|
|
||||||
return ToSocksaddr(destination), nil
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func (w *PacketConnWrapper) WritePacket(buffer *B.Buffer, destination M.Socksaddr) (err error) {
|
|
||||||
w.T.Update()
|
|
||||||
defer func() {
|
|
||||||
if err != nil {
|
|
||||||
// downlinkonly
|
|
||||||
w.T.SetTimeout(5 * time.Second)
|
|
||||||
}
|
|
||||||
}()
|
|
||||||
endpoint, err := ToDestination(destination, net.Network_UDP)
|
|
||||||
if err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
vBuf := buf.New()
|
|
||||||
vBuf.Write(buffer.Bytes())
|
|
||||||
vBuf.UDP = &endpoint
|
|
||||||
return w.WriteMultiBuffer(buf.MultiBuffer{vBuf})
|
|
||||||
}
|
|
||||||
|
|
||||||
func (w *PacketConnWrapper) Close() error {
|
|
||||||
buf.ReleaseMulti(w.cached)
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
@@ -1,81 +0,0 @@
|
|||||||
package singbridge
|
|
||||||
|
|
||||||
import (
|
|
||||||
"context"
|
|
||||||
"io"
|
|
||||||
"net"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
"github.com/sagernet/sing/common/bufio"
|
|
||||||
"github.com/xtls/xray-core/common"
|
|
||||||
"github.com/xtls/xray-core/common/buf"
|
|
||||||
"github.com/xtls/xray-core/common/signal"
|
|
||||||
"github.com/xtls/xray-core/transport"
|
|
||||||
)
|
|
||||||
|
|
||||||
func CopyConn(ctx context.Context, inboundConn net.Conn, link *transport.Link, serverConn net.Conn) error {
|
|
||||||
conn := &PipeConnWrapper{
|
|
||||||
W: link.Writer,
|
|
||||||
Conn: inboundConn,
|
|
||||||
}
|
|
||||||
if ir, ok := link.Reader.(io.Reader); ok {
|
|
||||||
conn.R = ir
|
|
||||||
} else {
|
|
||||||
conn.R = &buf.BufferedReader{Reader: link.Reader}
|
|
||||||
}
|
|
||||||
cancel := func() {
|
|
||||||
common.Interrupt(link.Reader)
|
|
||||||
common.Interrupt(serverConn)
|
|
||||||
}
|
|
||||||
conn.T = signal.CancelAfterInactivity(ctx, cancel, 300*time.Second)
|
|
||||||
return ReturnError(bufio.CopyConn(ctx, conn, serverConn))
|
|
||||||
}
|
|
||||||
|
|
||||||
type PipeConnWrapper struct {
|
|
||||||
R io.Reader
|
|
||||||
W buf.Writer
|
|
||||||
net.Conn
|
|
||||||
|
|
||||||
// A simple patch to avoid goroutine leak since sing infra cannot awake read block by write err
|
|
||||||
T *signal.ActivityTimer
|
|
||||||
}
|
|
||||||
|
|
||||||
func (w *PipeConnWrapper) Close() error {
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (w *PipeConnWrapper) Read(b []byte) (n int, err error) {
|
|
||||||
w.T.Update()
|
|
||||||
n, err = w.R.Read(b)
|
|
||||||
if err != nil {
|
|
||||||
// uplinkonly
|
|
||||||
w.T.SetTimeout(2 * time.Second)
|
|
||||||
}
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
func (w *PipeConnWrapper) Write(p []byte) (n int, err error) {
|
|
||||||
w.T.Update()
|
|
||||||
n = len(p)
|
|
||||||
var mb buf.MultiBuffer
|
|
||||||
pLen := len(p)
|
|
||||||
for pLen > 0 {
|
|
||||||
buffer := buf.New()
|
|
||||||
if pLen > buf.Size {
|
|
||||||
_, err = buffer.Write(p[:buf.Size])
|
|
||||||
p = p[buf.Size:]
|
|
||||||
} else {
|
|
||||||
buffer.Write(p)
|
|
||||||
}
|
|
||||||
pLen -= int(buffer.Len())
|
|
||||||
mb = append(mb, buffer)
|
|
||||||
}
|
|
||||||
err = w.W.WriteMultiBuffer(mb)
|
|
||||||
if err != nil {
|
|
||||||
n = 0
|
|
||||||
buf.ReleaseMulti(mb)
|
|
||||||
// downlinkonly
|
|
||||||
w.T.SetTimeout(5 * time.Second)
|
|
||||||
}
|
|
||||||
return
|
|
||||||
}
|
|
||||||
@@ -1,66 +0,0 @@
|
|||||||
package singbridge
|
|
||||||
|
|
||||||
import (
|
|
||||||
"time"
|
|
||||||
|
|
||||||
"github.com/sagernet/sing/common"
|
|
||||||
"github.com/sagernet/sing/common/bufio"
|
|
||||||
N "github.com/sagernet/sing/common/network"
|
|
||||||
"github.com/xtls/xray-core/common/buf"
|
|
||||||
"github.com/xtls/xray-core/common/net"
|
|
||||||
)
|
|
||||||
|
|
||||||
var (
|
|
||||||
_ buf.Reader = (*Conn)(nil)
|
|
||||||
_ buf.TimeoutReader = (*Conn)(nil)
|
|
||||||
_ buf.Writer = (*Conn)(nil)
|
|
||||||
)
|
|
||||||
|
|
||||||
type Conn struct {
|
|
||||||
net.Conn
|
|
||||||
writer N.VectorisedWriter
|
|
||||||
}
|
|
||||||
|
|
||||||
func NewConn(conn net.Conn) *Conn {
|
|
||||||
writer, _ := bufio.CreateVectorisedWriter(conn)
|
|
||||||
return &Conn{
|
|
||||||
Conn: conn,
|
|
||||||
writer: writer,
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *Conn) ReadMultiBuffer() (buf.MultiBuffer, error) {
|
|
||||||
buffer, err := buf.ReadBuffer(c.Conn)
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
return buf.MultiBuffer{buffer}, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *Conn) ReadMultiBufferTimeout(duration time.Duration) (buf.MultiBuffer, error) {
|
|
||||||
err := c.SetReadDeadline(time.Now().Add(duration))
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
defer c.SetReadDeadline(time.Time{})
|
|
||||||
return c.ReadMultiBuffer()
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *Conn) WriteMultiBuffer(bufferList buf.MultiBuffer) error {
|
|
||||||
defer buf.ReleaseMulti(bufferList)
|
|
||||||
if c.writer != nil {
|
|
||||||
bytesList := make([][]byte, len(bufferList))
|
|
||||||
for i, buffer := range bufferList {
|
|
||||||
bytesList[i] = buffer.Bytes()
|
|
||||||
}
|
|
||||||
return common.Error(bufio.WriteVectorised(c.writer, bytesList))
|
|
||||||
}
|
|
||||||
// Since this conn is only used by tun, we don't force buffer writes to merge.
|
|
||||||
for _, buffer := range bufferList {
|
|
||||||
_, err := c.Conn.Write(buffer.Bytes())
|
|
||||||
if err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
+2
-2
@@ -16,7 +16,7 @@ var typeCreatorRegistry = make(map[reflect.Type]ConfigCreator)
|
|||||||
func RegisterConfig(config interface{}, configCreator ConfigCreator) error {
|
func RegisterConfig(config interface{}, configCreator ConfigCreator) error {
|
||||||
configType := reflect.TypeOf(config)
|
configType := reflect.TypeOf(config)
|
||||||
if _, found := typeCreatorRegistry[configType]; found {
|
if _, found := typeCreatorRegistry[configType]; found {
|
||||||
return errors.New(configType.Name() + " is already registered").AtError()
|
return errors.New(configType.Name() + " is already registered")
|
||||||
}
|
}
|
||||||
typeCreatorRegistry[configType] = configCreator
|
typeCreatorRegistry[configType] = configCreator
|
||||||
return nil
|
return nil
|
||||||
@@ -27,7 +27,7 @@ func CreateObject(ctx context.Context, config interface{}) (interface{}, error)
|
|||||||
configType := reflect.TypeOf(config)
|
configType := reflect.TypeOf(config)
|
||||||
creator, found := typeCreatorRegistry[configType]
|
creator, found := typeCreatorRegistry[configType]
|
||||||
if !found {
|
if !found {
|
||||||
return nil, errors.New(configType.String() + " is not registered").AtError()
|
return nil, errors.New(configType.String() + " is not registered")
|
||||||
}
|
}
|
||||||
return creator(ctx, config)
|
return creator(ctx, config)
|
||||||
}
|
}
|
||||||
|
|||||||
+4
-4
@@ -125,7 +125,7 @@ func LoadConfig(formatName string, input interface{}) (*Config, error) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
if f == "" {
|
if f == "" {
|
||||||
return nil, errors.New("Failed to get format of ", file).AtWarning()
|
return nil, errors.New("Failed to get format of ", file)
|
||||||
}
|
}
|
||||||
|
|
||||||
if f == "protobuf" {
|
if f == "protobuf" {
|
||||||
@@ -142,7 +142,7 @@ func LoadConfig(formatName string, input interface{}) (*Config, error) {
|
|||||||
if len(v) == 1 {
|
if len(v) == 1 {
|
||||||
return configLoaderByName["protobuf"].Loader(v)
|
return configLoaderByName["protobuf"].Loader(v)
|
||||||
} else {
|
} else {
|
||||||
return nil, errors.New("Only one protobuf config file is allowed").AtWarning()
|
return nil, errors.New("Only one protobuf config file is allowed")
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -152,11 +152,11 @@ func LoadConfig(formatName string, input interface{}) (*Config, error) {
|
|||||||
if f, found := configLoaderByName[formatName]; found {
|
if f, found := configLoaderByName[formatName]; found {
|
||||||
return f.Loader(v)
|
return f.Loader(v)
|
||||||
} else {
|
} else {
|
||||||
return nil, errors.New("Unable to load config in", formatName).AtWarning()
|
return nil, errors.New("Unable to load config in", formatName)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
return nil, errors.New("Unable to load config").AtWarning()
|
return nil, errors.New("Unable to load config")
|
||||||
}
|
}
|
||||||
|
|
||||||
func loadProtobufConfig(data []byte) (*Config, error) {
|
func loadProtobufConfig(data []byte) (*Config, error) {
|
||||||
|
|||||||
+2
-2
@@ -19,8 +19,8 @@ import (
|
|||||||
|
|
||||||
var (
|
var (
|
||||||
Version_x byte = 26
|
Version_x byte = 26
|
||||||
Version_y byte = 7
|
Version_y byte = 9
|
||||||
Version_z byte = 28
|
Version_z byte = 30
|
||||||
)
|
)
|
||||||
|
|
||||||
var (
|
var (
|
||||||
|
|||||||
@@ -13,7 +13,7 @@ type FakeDNSEngine interface {
|
|||||||
|
|
||||||
var (
|
var (
|
||||||
FakeIPv4Pool = "198.18.0.0/15"
|
FakeIPv4Pool = "198.18.0.0/15"
|
||||||
FakeIPv6Pool = "fc00::/18"
|
FakeIPv6Pool = "2001:2::/48"
|
||||||
)
|
)
|
||||||
|
|
||||||
type FakeDNSEngineRev0 interface {
|
type FakeDNSEngineRev0 interface {
|
||||||
|
|||||||
@@ -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()
|
||||||
|
}
|
||||||
@@ -1,6 +1,6 @@
|
|||||||
module github.com/xtls/xray-core
|
module github.com/xtls/xray-core
|
||||||
|
|
||||||
go 1.26
|
go 1.27
|
||||||
|
|
||||||
require (
|
require (
|
||||||
github.com/apernet/quic-go v0.61.1-0.20260806010916-184d081eef3e
|
github.com/apernet/quic-go v0.61.1-0.20260806010916-184d081eef3e
|
||||||
@@ -18,25 +18,26 @@ require (
|
|||||||
github.com/pires/go-proxyproto v0.15.0
|
github.com/pires/go-proxyproto v0.15.0
|
||||||
github.com/refraction-networking/utls v1.8.3-0.20260301010127-aa6edf4b11af
|
github.com/refraction-networking/utls v1.8.3-0.20260301010127-aa6edf4b11af
|
||||||
github.com/robfig/cron/v3 v3.0.1
|
github.com/robfig/cron/v3 v3.0.1
|
||||||
github.com/sagernet/sing v0.5.1
|
|
||||||
github.com/sagernet/sing-shadowsocks v0.2.7
|
|
||||||
github.com/stretchr/testify v1.12.1
|
github.com/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-20260322125925-9234c772ba8f
|
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.55.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
|
||||||
golang.org/x/net v0.58.0
|
golang.org/x/net v0.59.0
|
||||||
golang.org/x/sync v0.22.0
|
golang.org/x/sync v0.23.0
|
||||||
golang.org/x/sys v0.47.0
|
golang.org/x/sys v0.48.0
|
||||||
golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2
|
golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2
|
||||||
golang.zx2c4.com/wireguard v0.0.0-20250521234502-f333402bd9cb
|
golang.zx2c4.com/wireguard v0.0.0-20250521234502-f333402bd9cb
|
||||||
golang.zx2c4.com/wireguard/windows v1.0.1
|
golang.zx2c4.com/wireguard/windows v1.1.1
|
||||||
google.golang.org/grpc v1.83.1
|
google.golang.org/grpc v1.84.0
|
||||||
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
|
||||||
)
|
)
|
||||||
|
|
||||||
require (
|
require (
|
||||||
@@ -48,7 +49,6 @@ require (
|
|||||||
github.com/juju/ratelimit v1.0.2 // indirect
|
github.com/juju/ratelimit v1.0.2 // indirect
|
||||||
github.com/klauspost/compress v1.17.4 // indirect
|
github.com/klauspost/compress v1.17.4 // indirect
|
||||||
github.com/koron/go-ssdp v0.0.4 // indirect
|
github.com/koron/go-ssdp v0.0.4 // indirect
|
||||||
github.com/kr/text v0.2.0 // indirect
|
|
||||||
github.com/libp2p/go-netroute v0.2.1 // indirect
|
github.com/libp2p/go-netroute v0.2.1 // indirect
|
||||||
github.com/pion/dtls/v3 v3.1.5 // indirect
|
github.com/pion/dtls/v3 v3.1.5 // indirect
|
||||||
github.com/pion/logging v0.2.4 // indirect
|
github.com/pion/logging v0.2.4 // indirect
|
||||||
@@ -57,8 +57,9 @@ require (
|
|||||||
github.com/vishvananda/netns v0.0.5 // indirect
|
github.com/vishvananda/netns v0.0.5 // indirect
|
||||||
github.com/wlynxg/anet v0.0.5 // indirect
|
github.com/wlynxg/anet v0.0.5 // indirect
|
||||||
go.yaml.in/yaml/v3 v3.0.5 // indirect
|
go.yaml.in/yaml/v3 v3.0.5 // indirect
|
||||||
golang.org/x/text v0.41.0 // indirect
|
golang.org/x/text v0.42.0 // indirect
|
||||||
golang.org/x/time v0.14.0 // indirect
|
golang.org/x/time v0.14.0 // indirect
|
||||||
google.golang.org/genproto/googleapis/rpc v0.0.0-20260526163538-3dc84a4a5aaa // indirect
|
golang.org/x/tools v0.49.0 // indirect
|
||||||
|
google.golang.org/genproto/googleapis/rpc v0.0.0-20260706201446-f0a921348800 // indirect
|
||||||
gopkg.in/yaml.v2 v2.4.0 // indirect
|
gopkg.in/yaml.v2 v2.4.0 // indirect
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -2,17 +2,15 @@ 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/cespare/xxhash/v2 v2.3.0 h1:UL815xU9SqsFlibzuggzjXhog7bL6oX9BbNZnL2UFvs=
|
github.com/chzyer/logex v1.1.10/go.mod h1:+Ywpsq7O8HXn0nuIou7OrIPyXbp3wmkHB+jjWRnGsAI=
|
||||||
github.com/cespare/xxhash/v2 v2.3.0/go.mod h1:VGX0DQ3Q6kWi7AoAeZDth3/j3BFtOZR5XLFGgcrjCOs=
|
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/creack/pty v1.1.9/go.mod h1:oKZEueFk5CKHvIhNR5MUki03XCEU+Q6VDXinZuGJ33E=
|
|
||||||
github.com/ghodss/yaml v1.0.1-0.20220118164431-d8423dcdf344 h1:Arcl6UOIS/kgO2nW3A65HN+7CMjSDP/gofXL4CZt1V4=
|
github.com/ghodss/yaml v1.0.1-0.20220118164431-d8423dcdf344 h1:Arcl6UOIS/kgO2nW3A65HN+7CMjSDP/gofXL4CZt1V4=
|
||||||
github.com/ghodss/yaml v1.0.1-0.20220118164431-d8423dcdf344/go.mod h1:GIjDIg/heH5DOkXY3YJ/wNhfHsQHoXGjl8G8amsYQ1I=
|
github.com/ghodss/yaml v1.0.1-0.20220118164431-d8423dcdf344/go.mod h1:GIjDIg/heH5DOkXY3YJ/wNhfHsQHoXGjl8G8amsYQ1I=
|
||||||
github.com/go-logr/logr v1.4.3 h1:CjnDlHq8ikf6E492q6eKboGOC0T8CDaOvkHCIg8idEI=
|
github.com/go-quicktest/qt v1.102.0 h1:HSQxCeh5YZH3EL3W39ixjtyaEhcWSXQHtHnMBzSs474=
|
||||||
github.com/go-logr/logr v1.4.3/go.mod h1:9T104GzyrTigFIr8wt5mBrctHMim0Nb2HLGrmQ40KvY=
|
github.com/go-quicktest/qt v1.102.0/go.mod h1:p4lGIVX+8Wa6ZPNDvqcxq36XpUDLh42FLetFU7odllI=
|
||||||
github.com/go-logr/stdr v1.2.2 h1:hSWxHoqTgW2S2qGc0LTAI563KZ5YKYRhT3MFKZMbjag=
|
|
||||||
github.com/go-logr/stdr v1.2.2/go.mod h1:mMo/vtBO5dYbehREoey6XUKy/eSumjCCveDpRre4VKE=
|
|
||||||
github.com/golang/mock v1.7.0-rc.1 h1:YojYx61/OLFsiv6Rw1Z96LpldJIy31o+UHmwAUMJ6/U=
|
github.com/golang/mock v1.7.0-rc.1 h1:YojYx61/OLFsiv6Rw1Z96LpldJIy31o+UHmwAUMJ6/U=
|
||||||
github.com/golang/mock v1.7.0-rc.1/go.mod h1:s42URUywIqd+OcERslBJvOjepvNymP31m3q8d/GkuRs=
|
github.com/golang/mock v1.7.0-rc.1/go.mod h1:s42URUywIqd+OcERslBJvOjepvNymP31m3q8d/GkuRs=
|
||||||
github.com/golang/protobuf v1.5.4 h1:i7eJL8qZTpSEXOPTxNKhASYpMn+8e5Q6AdndVa1dWek=
|
github.com/golang/protobuf v1.5.4 h1:i7eJL8qZTpSEXOPTxNKhASYpMn+8e5Q6AdndVa1dWek=
|
||||||
@@ -73,12 +71,8 @@ github.com/refraction-networking/utls v1.8.3-0.20260301010127-aa6edf4b11af h1:er
|
|||||||
github.com/refraction-networking/utls v1.8.3-0.20260301010127-aa6edf4b11af/go.mod h1:jkSOEkLqn+S/jtpEHPOsVv/4V4EVnelwbMQl4vCWXAM=
|
github.com/refraction-networking/utls v1.8.3-0.20260301010127-aa6edf4b11af/go.mod h1:jkSOEkLqn+S/jtpEHPOsVv/4V4EVnelwbMQl4vCWXAM=
|
||||||
github.com/robfig/cron/v3 v3.0.1 h1:WdRxkvbJztn8LMz/QEvLN5sBU+xKpSqwwUO1Pjr4qDs=
|
github.com/robfig/cron/v3 v3.0.1 h1:WdRxkvbJztn8LMz/QEvLN5sBU+xKpSqwwUO1Pjr4qDs=
|
||||||
github.com/robfig/cron/v3 v3.0.1/go.mod h1:eQICP3HwyT7UooqI/z+Ov+PtYAWygg1TEWWzGIFLtro=
|
github.com/robfig/cron/v3 v3.0.1/go.mod h1:eQICP3HwyT7UooqI/z+Ov+PtYAWygg1TEWWzGIFLtro=
|
||||||
github.com/rogpeppe/go-internal v1.10.0 h1:TMyTOH3F/DB16zRVcYyreMH6GnZZrwQVAoYjRBZyWFQ=
|
github.com/rogpeppe/go-internal v1.16.0 h1:O9DK+vNMDVGLr2BeZqmpLeMjiMNkuXfcqntWbZV6S5g=
|
||||||
github.com/rogpeppe/go-internal v1.10.0/go.mod h1:UQnix2H7Ngw/k4C5ijL5+65zddjncjaFoBhdsK/akog=
|
github.com/rogpeppe/go-internal v1.16.0/go.mod h1:DrUVZyrJU+txYW5/1kwtXQSMFio52ZOxX7yM1VHvnxs=
|
||||||
github.com/sagernet/sing v0.5.1 h1:mhL/MZVq0TjuvHcpYcFtmSD1BFOxZ/+8ofbNZcg1k1Y=
|
|
||||||
github.com/sagernet/sing v0.5.1/go.mod h1:ARkL0gM13/Iv5VCZmci/NuoOlePoIsW0m7BWfln/Hak=
|
|
||||||
github.com/sagernet/sing-shadowsocks v0.2.7 h1:zaopR1tbHEw5Nk6FAkM05wCslV6ahVegEZaKMv9ipx8=
|
|
||||||
github.com/sagernet/sing-shadowsocks v0.2.7/go.mod h1:0rIKJZBR65Qi0zwdKezt4s57y/Tl1ofkaq6NlkzVuyE=
|
|
||||||
github.com/stretchr/testify v1.12.1 h1:EuwCh5fleGS7H32xRwO3wRGT7DxrDhLAT6FF8MpWDWE=
|
github.com/stretchr/testify v1.12.1 h1:EuwCh5fleGS7H32xRwO3wRGT7DxrDhLAT6FF8MpWDWE=
|
||||||
github.com/stretchr/testify v1.12.1/go.mod h1:MDEgiDPPsNp5cuIrHPPCyornHKgEVbtFUmoNlxoYthg=
|
github.com/stretchr/testify v1.12.1/go.mod h1:MDEgiDPPsNp5cuIrHPPCyornHKgEVbtFUmoNlxoYthg=
|
||||||
github.com/vishvananda/netlink v1.3.1 h1:3AEMt62VKqz90r0tmNhog0r/PpWKmrEShJU0wJW6bV0=
|
github.com/vishvananda/netlink v1.3.1 h1:3AEMt62VKqz90r0tmNhog0r/PpWKmrEShJU0wJW6bV0=
|
||||||
@@ -87,21 +81,12 @@ github.com/vishvananda/netns v0.0.5 h1:DfiHV+j8bA32MFM7bfEunvT8IAqQ/NzSJHtcmW5zd
|
|||||||
github.com/vishvananda/netns v0.0.5/go.mod h1:SpkAiCQRtJ6TvvxPnOSyH3BMl6unz3xZlaprSwhNNJM=
|
github.com/vishvananda/netns v0.0.5/go.mod h1:SpkAiCQRtJ6TvvxPnOSyH3BMl6unz3xZlaprSwhNNJM=
|
||||||
github.com/wlynxg/anet v0.0.5 h1:J3VJGi1gvo0JwZ/P1/Yc/8p63SoW98B5dHkYDmpgvvU=
|
github.com/wlynxg/anet v0.0.5 h1:J3VJGi1gvo0JwZ/P1/Yc/8p63SoW98B5dHkYDmpgvvU=
|
||||||
github.com/wlynxg/anet v0.0.5/go.mod h1:eay5PRQr7fIVAMbTbchTnO9gG65Hg/uYGdc7mguHxoA=
|
github.com/wlynxg/anet v0.0.5/go.mod h1:eay5PRQr7fIVAMbTbchTnO9gG65Hg/uYGdc7mguHxoA=
|
||||||
github.com/xtls/reality v0.0.0-20260322125925-9234c772ba8f h1:iy2JRioxmUpoJ3SzbFPyTxHZMbR/rSHP7dOOgYaq1O8=
|
github.com/xtls/reality v0.0.0-20260908062103-8cdf7bf9c7f0 h1:rb+fKQFhz+5I2PPuQsNYxI5mUU840XWYtRF0ZBjvkws=
|
||||||
github.com/xtls/reality v0.0.0-20260322125925-9234c772ba8f/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=
|
||||||
go.opentelemetry.io/auto/sdk v1.2.1 h1:jXsnJ4Lmnqd11kwkBV2LgLoFMZKizbCi5fNZ/ipaZ64=
|
github.com/yuin/gopher-lua v0.0.0-20190206043414-8bfc7677f583/go.mod h1:gqRgreBUhTSL0GeU64rtZ3Uq3wtjOa/TB2YfrtkCbVQ=
|
||||||
go.opentelemetry.io/auto/sdk v1.2.1/go.mod h1:KRTj+aOaElaLi+wW1kO/DZRXwkF4C5xPbEe3ZiIhN7Y=
|
github.com/yuin/gopher-lua v1.1.2 h1:yF/FjE3hD65tBbt0VXLE13HWS9h34fdzJmrWRXwobGA=
|
||||||
go.opentelemetry.io/otel v1.44.0 h1:JjwHmHpA4iZ3wBxluu2fbbE7j4kqlE8jXyAyPXH7HqU=
|
github.com/yuin/gopher-lua v1.1.2/go.mod h1:7aRmXIWl37SqRf0koeyylBEzJ+aPt8A+mmkQ4f1ntR8=
|
||||||
go.opentelemetry.io/otel v1.44.0/go.mod h1:BMgjTHL9WPRlRjL2oZCBTL4whCGtXch2H4BhOPIAyYc=
|
|
||||||
go.opentelemetry.io/otel/metric v1.44.0 h1:1w0gILTcHdr3YI+ixLyjemwrVnsMURbTZFrSYCdDdmc=
|
|
||||||
go.opentelemetry.io/otel/metric v1.44.0/go.mod h1:8O7hanEPBNgEMmybD3s2VBKcgWOCsA6tzHBPODAiquo=
|
|
||||||
go.opentelemetry.io/otel/sdk v1.44.0 h1:nHYwb9lK+fJPU/dnT6s7W7Z8itMWyqrnVfbheVYrZ58=
|
|
||||||
go.opentelemetry.io/otel/sdk v1.44.0/go.mod h1:Osuydd3Se74nqjAKxid74N5eC+jfEqfTegHRnq58oK0=
|
|
||||||
go.opentelemetry.io/otel/sdk/metric v1.44.0 h1:3LlKgI+VjbVsjNRFZJZAJ30WjXC5VkNRks6si09iEfI=
|
|
||||||
go.opentelemetry.io/otel/sdk/metric v1.44.0/go.mod h1:5B5pMARnXxKhltooO4xUuCBorl65a4EpnTalObqOigA=
|
|
||||||
go.opentelemetry.io/otel/trace v1.44.0 h1:jxF5CsGYCe74MCRx2X4g7WsY/VBKRqqpNvXlX/6gtIk=
|
|
||||||
go.opentelemetry.io/otel/trace v1.44.0/go.mod h1:oLl1jrMQAVo6v3GAggN+1VH9VIz9iUSvW53sW1Q8PIE=
|
|
||||||
go.uber.org/mock v0.5.2 h1:LbtPTcP8A5k9WPXj54PPPbjcI4Y6lhyOZXn+VS7wNko=
|
go.uber.org/mock v0.5.2 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=
|
||||||
@@ -110,8 +95,8 @@ go4.org/netipx v0.0.0-20231129151722-fdeea329fbba h1:0b9z3AuHCjxk0x/opv64kcgZLBs
|
|||||||
go4.org/netipx v0.0.0-20231129151722-fdeea329fbba/go.mod h1:PLyyIXexvUFg3Owu6p/WfdlivPbZJsZdgWZlrGope/Y=
|
go4.org/netipx v0.0.0-20231129151722-fdeea329fbba/go.mod h1:PLyyIXexvUFg3Owu6p/WfdlivPbZJsZdgWZlrGope/Y=
|
||||||
golang.org/x/crypto v0.0.0-20190308221718-c2843e01d9a2/go.mod h1:djNgcEr1/C05ACkg1iLfiJU5Ep61QUkGW8qpdssI0+w=
|
golang.org/x/crypto v0.0.0-20190308221718-c2843e01d9a2/go.mod h1:djNgcEr1/C05ACkg1iLfiJU5Ep61QUkGW8qpdssI0+w=
|
||||||
golang.org/x/crypto v0.0.0-20191011191535-87dc89f01550/go.mod h1:yigFU9vqHzYiE8UmvKecakEJjdnWj3jj499lnFckfCI=
|
golang.org/x/crypto v0.0.0-20191011191535-87dc89f01550/go.mod h1:yigFU9vqHzYiE8UmvKecakEJjdnWj3jj499lnFckfCI=
|
||||||
golang.org/x/crypto v0.55.0 h1:+KWHjbgOaAQ66dh/YlkZKHlz9ZUlq61AFirAR9ntP8M=
|
golang.org/x/crypto v0.57.0 h1:3ZVCjf8Ggz7zneR/EHRVx68Ctf+2pmIMP2UFhh9cC6M=
|
||||||
golang.org/x/crypto v0.55.0/go.mod h1:uq0V9dE/fzQuJtbnL+2EhWOE63vo164FY8xqEnV9xis=
|
golang.org/x/crypto v0.57.0/go.mod h1:Fdz0i5U6CoizGwLda9DttjSk6qlZo25zYNtR+ycvuZA=
|
||||||
golang.org/x/exp v0.0.0-20240506185415-9bf2ced13842 h1:vr/HnozRka3pE4EsMEg1lgkXJkTFJCVUX+S/ZT6wYzM=
|
golang.org/x/exp v0.0.0-20240506185415-9bf2ced13842 h1:vr/HnozRka3pE4EsMEg1lgkXJkTFJCVUX+S/ZT6wYzM=
|
||||||
golang.org/x/exp v0.0.0-20240506185415-9bf2ced13842/go.mod h1:XtvwrStGgqGPLc4cjQfWqZHG1YFdYs6swckp8vpsjnc=
|
golang.org/x/exp v0.0.0-20240506185415-9bf2ced13842/go.mod h1:XtvwrStGgqGPLc4cjQfWqZHG1YFdYs6swckp8vpsjnc=
|
||||||
golang.org/x/lint v0.0.0-20200302205851-738671d3881b/go.mod h1:3xt1FjdF8hUf6vQPIChWIBhFzV8gjjsPE/fR3IyQdNY=
|
golang.org/x/lint v0.0.0-20200302205851-738671d3881b/go.mod h1:3xt1FjdF8hUf6vQPIChWIBhFzV8gjjsPE/fR3IyQdNY=
|
||||||
@@ -120,12 +105,13 @@ golang.org/x/mod v0.5.1/go.mod h1:5OXOZSfqPIIbmVBIIKWRFfZjPR0E5r58TLhUjH0a2Ro=
|
|||||||
golang.org/x/net v0.0.0-20190404232315-eb5bcb51f2a3/go.mod h1:t9HGtf8HONx5eT2rtn7q6eTqICYqUVnKs3thJo3Qplg=
|
golang.org/x/net v0.0.0-20190404232315-eb5bcb51f2a3/go.mod h1:t9HGtf8HONx5eT2rtn7q6eTqICYqUVnKs3thJo3Qplg=
|
||||||
golang.org/x/net v0.0.0-20190620200207-3b0461eec859/go.mod h1:z5CRVTTTmAJ677TzLLGU+0bjPO0LkuOLi4/5GtJWs/s=
|
golang.org/x/net v0.0.0-20190620200207-3b0461eec859/go.mod h1:z5CRVTTTmAJ677TzLLGU+0bjPO0LkuOLi4/5GtJWs/s=
|
||||||
golang.org/x/net v0.0.0-20211015210444-4f30a5c0130f/go.mod h1:9nx3DQGgdP8bBQD5qxJ1jj9UTztislL4KSBs9R2vV5Y=
|
golang.org/x/net v0.0.0-20211015210444-4f30a5c0130f/go.mod h1:9nx3DQGgdP8bBQD5qxJ1jj9UTztislL4KSBs9R2vV5Y=
|
||||||
golang.org/x/net v0.58.0 h1:ynWG7rqYi4ccpTEuPZ2QGWHktVEM9DMCj9yzDE0Q7To=
|
golang.org/x/net v0.59.0 h1:5zfYln+w5XCxwrnMMJPufRgNoXEaGxl0wo5GqPXyues=
|
||||||
golang.org/x/net v0.58.0/go.mod h1:YwCddHnFlT7eLQqVprV19OnhLGtc5xOKgE0RyqgfWAU=
|
golang.org/x/net v0.59.0/go.mod h1:2DA/G1UfVbCpQPeWTmMPGY7Cs2PkBkwu743bVX5PIVg=
|
||||||
golang.org/x/sync v0.0.0-20190423024810-112230192c58/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
golang.org/x/sync v0.0.0-20190423024810-112230192c58/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
||||||
golang.org/x/sync v0.0.0-20210220032951-036812b2e83c/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
golang.org/x/sync v0.0.0-20210220032951-036812b2e83c/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
||||||
golang.org/x/sync v0.22.0 h1:SZjpbeLmrCk4xhRSZFNZW5gFUeCeFgjekvI/+gfScek=
|
golang.org/x/sync v0.23.0 h1:KameEIfc1IkluZyXWLn39Wd4tURc6GbCiISGiZm2bQk=
|
||||||
golang.org/x/sync v0.22.0/go.mod h1:9xrNwdLfx4jkKbNva9FpL6vEN7evnE43NNNJQ2LF3+0=
|
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=
|
||||||
@@ -133,20 +119,22 @@ golang.org/x/sys v0.0.0-20210423082822-04245dca01da/go.mod h1:h1NjWce9XRLGQEsW7w
|
|||||||
golang.org/x/sys v0.0.0-20211019181941-9d821ace8654/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
|
golang.org/x/sys v0.0.0-20211019181941-9d821ace8654/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
|
||||||
golang.org/x/sys v0.2.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
|
golang.org/x/sys v0.2.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
|
||||||
golang.org/x/sys v0.10.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
|
golang.org/x/sys v0.10.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
|
||||||
golang.org/x/sys v0.47.0 h1:o7XGOvZQCADBQQ4Y7VNq2dRWQR7JmOUW8Kxx4ZsNgWs=
|
golang.org/x/sys v0.48.0 h1:bbX/i/6MgT9BVLM9RT1thmxL04yeTAhbEz4SyadbXoo=
|
||||||
golang.org/x/sys v0.47.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw=
|
golang.org/x/sys v0.48.0/go.mod h1:hNLxWAXmnKAxqDtdwIYC4bM9oQPEecfsnNMuSxOs3og=
|
||||||
golang.org/x/term v0.0.0-20201126162022-7de9c90e9dd1/go.mod h1:bj7SfCRtBDWHUb9snDiAeCFNEtKQo2Wmx5Cou7ajbmo=
|
golang.org/x/term v0.0.0-20201126162022-7de9c90e9dd1/go.mod h1:bj7SfCRtBDWHUb9snDiAeCFNEtKQo2Wmx5Cou7ajbmo=
|
||||||
golang.org/x/text v0.3.0/go.mod h1:NqM8EUOU14njkJ3fqMW+pc6Ldnwhi/IjpwHt7yyuwOQ=
|
golang.org/x/text v0.3.0/go.mod h1:NqM8EUOU14njkJ3fqMW+pc6Ldnwhi/IjpwHt7yyuwOQ=
|
||||||
golang.org/x/text v0.3.6/go.mod h1:5Zoc/QRtKVWzQhOtBMvqHzDpF6irO9z98xDceosuGiQ=
|
golang.org/x/text v0.3.6/go.mod h1:5Zoc/QRtKVWzQhOtBMvqHzDpF6irO9z98xDceosuGiQ=
|
||||||
golang.org/x/text v0.3.7/go.mod h1:u+2+/6zg+i71rQMx5EYifcz6MCKuco9NR6JIITiCfzQ=
|
golang.org/x/text v0.3.7/go.mod h1:u+2+/6zg+i71rQMx5EYifcz6MCKuco9NR6JIITiCfzQ=
|
||||||
golang.org/x/text v0.41.0 h1:vz/seA0lnX87Othu2f/0L24RcgrXD9/YFTSuGjj3rH8=
|
golang.org/x/text v0.42.0 h1:JbOZXgfeCPU9gacVtYliJqOhD+zhrEqK4LfdpmlUZqI=
|
||||||
golang.org/x/text v0.41.0/go.mod h1:jvf1O8ajNzZqhSrQBPbutR/EB83Cc0CFrezNQIwbb5M=
|
golang.org/x/text v0.42.0/go.mod h1:ojzP1Z+2QtioaF8DTtO8K5q7JWVVYwZKenzujK0Zd0E=
|
||||||
golang.org/x/time v0.14.0 h1:MRx4UaLrDotUKUdCIqzPC48t1Y9hANFKIRpNx+Te8PI=
|
golang.org/x/time v0.14.0 h1:MRx4UaLrDotUKUdCIqzPC48t1Y9hANFKIRpNx+Te8PI=
|
||||||
golang.org/x/time v0.14.0/go.mod h1:eL/Oa2bBBK0TkX57Fyni+NgnyQQN4LitPmob2Hjnqw4=
|
golang.org/x/time v0.14.0/go.mod h1:eL/Oa2bBBK0TkX57Fyni+NgnyQQN4LitPmob2Hjnqw4=
|
||||||
golang.org/x/tools v0.0.0-20180917221912-90fa682c2a6e/go.mod h1:n7NCudcB/nEzxVGmLbDWY5pfWTLqBcC2KZ6jyYvM4mQ=
|
golang.org/x/tools v0.0.0-20180917221912-90fa682c2a6e/go.mod h1:n7NCudcB/nEzxVGmLbDWY5pfWTLqBcC2KZ6jyYvM4mQ=
|
||||||
golang.org/x/tools v0.0.0-20191119224855-298f0cb1881e/go.mod h1:b+2E5dAYhXwXZwtnZ6UAqBI28+e2cm9otk0dWdXHAEo=
|
golang.org/x/tools v0.0.0-20191119224855-298f0cb1881e/go.mod h1:b+2E5dAYhXwXZwtnZ6UAqBI28+e2cm9otk0dWdXHAEo=
|
||||||
golang.org/x/tools v0.0.0-20200130002326-2f3ba24bd6e7/go.mod h1:TB2adYChydJhpapKDTa4BR/hXlZSLoq2Wpct/0txZ28=
|
golang.org/x/tools v0.0.0-20200130002326-2f3ba24bd6e7/go.mod h1:TB2adYChydJhpapKDTa4BR/hXlZSLoq2Wpct/0txZ28=
|
||||||
golang.org/x/tools v0.1.8/go.mod h1:nABZi5QlRsZVlzPpHl034qft6wpY4eDcsTt5AaioBiU=
|
golang.org/x/tools v0.1.8/go.mod h1:nABZi5QlRsZVlzPpHl034qft6wpY4eDcsTt5AaioBiU=
|
||||||
|
golang.org/x/tools v0.49.0 h1:3NI7VXzL9+1WZD52Dx2ttoPwD5DWrFGpl9mFZDlmisI=
|
||||||
|
golang.org/x/tools v0.49.0/go.mod h1:SJNXV9DBKT0UbdttsQjbfJlAE/q+y36++zo3uL3N0Oo=
|
||||||
golang.org/x/xerrors v0.0.0-20190717185122-a985d3407aa7/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0=
|
golang.org/x/xerrors v0.0.0-20190717185122-a985d3407aa7/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0=
|
||||||
golang.org/x/xerrors v0.0.0-20191011141410-1b5146add898/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0=
|
golang.org/x/xerrors v0.0.0-20191011141410-1b5146add898/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0=
|
||||||
golang.org/x/xerrors v0.0.0-20200804184101-5ec99f83aff1/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0=
|
golang.org/x/xerrors v0.0.0-20200804184101-5ec99f83aff1/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0=
|
||||||
@@ -154,14 +142,14 @@ golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2 h1:B82qJJgjvYKsXS9jeu
|
|||||||
golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2/go.mod h1:deeaetjYA+DHMHg+sMSMI58GrEteJUUzzw7en6TJQcI=
|
golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2/go.mod h1:deeaetjYA+DHMHg+sMSMI58GrEteJUUzzw7en6TJQcI=
|
||||||
golang.zx2c4.com/wireguard v0.0.0-20250521234502-f333402bd9cb h1:whnFRlWMcXI9d+ZbWg+4sHnLp52d5yiIPUxMBSt4X9A=
|
golang.zx2c4.com/wireguard v0.0.0-20250521234502-f333402bd9cb h1:whnFRlWMcXI9d+ZbWg+4sHnLp52d5yiIPUxMBSt4X9A=
|
||||||
golang.zx2c4.com/wireguard v0.0.0-20250521234502-f333402bd9cb/go.mod h1:rpwXGsirqLqN2L0JDJQlwOboGHmptD5ZD6T2VmcqhTw=
|
golang.zx2c4.com/wireguard v0.0.0-20250521234502-f333402bd9cb/go.mod h1:rpwXGsirqLqN2L0JDJQlwOboGHmptD5ZD6T2VmcqhTw=
|
||||||
golang.zx2c4.com/wireguard/windows v1.0.1 h1:eOxiDVbywPC+ZQqvdCK7x+ZwWXKbYv50TtH8ysFIbw8=
|
golang.zx2c4.com/wireguard/windows v1.1.1 h1:8/H97U1v1PNDNcBsMZgU3KFuND9MQdTsU2NOwmCXArE=
|
||||||
golang.zx2c4.com/wireguard/windows v1.0.1/go.mod h1:+fbT3FFdX4zzYDLwJh5+HPEcNN/3HyNdzhNSVsQM+zs=
|
golang.zx2c4.com/wireguard/windows v1.1.1/go.mod h1:+fbT3FFdX4zzYDLwJh5+HPEcNN/3HyNdzhNSVsQM+zs=
|
||||||
gonum.org/v1/gonum v0.17.0 h1:VbpOemQlsSMrYmn7T2OUvQ4dqxQXU+ouZFQsZOx50z4=
|
gonum.org/v1/gonum v0.17.0 h1:VbpOemQlsSMrYmn7T2OUvQ4dqxQXU+ouZFQsZOx50z4=
|
||||||
gonum.org/v1/gonum v0.17.0/go.mod h1:El3tOrEuMpv2UdMrbNlKEh9vd86bmQ6vqIcDwxEOc1E=
|
gonum.org/v1/gonum v0.17.0/go.mod h1:El3tOrEuMpv2UdMrbNlKEh9vd86bmQ6vqIcDwxEOc1E=
|
||||||
google.golang.org/genproto/googleapis/rpc v0.0.0-20260526163538-3dc84a4a5aaa h1:mZHHdPZl0dbGHCflZgAq/Q468DWVFcU2whhB2KAo8fk=
|
google.golang.org/genproto/googleapis/rpc v0.0.0-20260706201446-f0a921348800 h1:qEHAMpSaUhtD0p3NbEEI83HwNGFxEwaSJ1G9PLnCBZE=
|
||||||
google.golang.org/genproto/googleapis/rpc v0.0.0-20260526163538-3dc84a4a5aaa/go.mod h1:4Hqkh8ycfw05ld/3BWL7rJOSfebL2Q+DVDeRgYgxUU8=
|
google.golang.org/genproto/googleapis/rpc v0.0.0-20260706201446-f0a921348800/go.mod h1:4Hqkh8ycfw05ld/3BWL7rJOSfebL2Q+DVDeRgYgxUU8=
|
||||||
google.golang.org/grpc v1.83.1 h1:HIO0+BEtBP6soyqvqC8sNUjZ7bTs+0hFQuFF+RAy++Y=
|
google.golang.org/grpc v1.84.0 h1:soMyaPJ8pAak5PIQ0DGBUir0XRo2fRoMqhNWMLlLxO0=
|
||||||
google.golang.org/grpc v1.83.1/go.mod h1:kDyl6SKsiHKt0uylY5gtn5cEjkrIOhQOGDgIc4JGwzQ=
|
google.golang.org/grpc v1.84.0/go.mod h1:ljCht0DrxQrXBDRTZp52Qxh3Ffk8CdYm2sj4O2QN2C0=
|
||||||
google.golang.org/protobuf v1.36.12 h1:pJOKDDOyeXErUroCihFAd5LQuwXBSpVnKGrj5o/fwxc=
|
google.golang.org/protobuf v1.36.12 h1:pJOKDDOyeXErUroCihFAd5LQuwXBSpVnKGrj5o/fwxc=
|
||||||
google.golang.org/protobuf v1.36.12/go.mod h1:HTf+CrKn2C3g5S8VImy6tdcUvCska2kB7j23XfzDpco=
|
google.golang.org/protobuf v1.36.12/go.mod h1:HTf+CrKn2C3g5S8VImy6tdcUvCska2kB7j23XfzDpco=
|
||||||
gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0=
|
gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0=
|
||||||
@@ -174,5 +162,9 @@ 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/go.mod h1:SmBHHrljiZu/uoypeKup3rFzP6eoC9UwCp2iH5E3jZA=
|
||||||
|
|||||||
+18
-28
@@ -1,52 +1,42 @@
|
|||||||
package conf
|
package conf
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"encoding/json"
|
"encoding/base64"
|
||||||
|
"strings"
|
||||||
|
|
||||||
"github.com/xtls/xray-core/common/errors"
|
"github.com/xtls/xray-core/common/errors"
|
||||||
"github.com/xtls/xray-core/common/serial"
|
|
||||||
"github.com/xtls/xray-core/proxy/blackhole"
|
"github.com/xtls/xray-core/proxy/blackhole"
|
||||||
"google.golang.org/protobuf/proto"
|
"google.golang.org/protobuf/proto"
|
||||||
)
|
)
|
||||||
|
|
||||||
type NoneResponse struct{}
|
type ResponseConfig struct {
|
||||||
|
Type string `json:"type"`
|
||||||
func (*NoneResponse) Build() (proto.Message, error) {
|
CustomResponseData string `json:"customResponseData"`
|
||||||
return new(blackhole.NoneResponse), nil
|
|
||||||
}
|
|
||||||
|
|
||||||
type HTTPResponse struct{}
|
|
||||||
|
|
||||||
func (*HTTPResponse) Build() (proto.Message, error) {
|
|
||||||
return new(blackhole.HTTPResponse), nil
|
|
||||||
}
|
}
|
||||||
|
|
||||||
type BlackholeConfig struct {
|
type BlackholeConfig struct {
|
||||||
Response json.RawMessage `json:"response"`
|
Response *ResponseConfig `json:"response"`
|
||||||
}
|
}
|
||||||
|
|
||||||
func (v *BlackholeConfig) Build() (proto.Message, error) {
|
func (v *BlackholeConfig) Build() (proto.Message, error) {
|
||||||
config := new(blackhole.Config)
|
config := new(blackhole.Config)
|
||||||
if v.Response != nil {
|
if v.Response != nil {
|
||||||
response, _, err := configLoader.Load(v.Response)
|
responseName := strings.ToLower(v.Response.Type)
|
||||||
|
switch responseName {
|
||||||
|
case "none", "":
|
||||||
|
config.Response = &blackhole.Response{Type: "none"}
|
||||||
|
case "http":
|
||||||
|
config.Response = &blackhole.Response{Type: "http"}
|
||||||
|
case "custom":
|
||||||
|
data, err := base64.StdEncoding.DecodeString(v.Response.CustomResponseData)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, errors.New("Config: Failed to parse Blackhole response config.").Base(err)
|
return nil, errors.New("failed to decode custom response data: " + err.Error())
|
||||||
}
|
}
|
||||||
responseSettings, err := response.(Buildable).Build()
|
config.Response = &blackhole.Response{Type: "custom", CustomResponseData: data}
|
||||||
if err != nil {
|
default:
|
||||||
return nil, err
|
return nil, errors.New("unknown blackhole response: " + responseName)
|
||||||
}
|
}
|
||||||
config.Response = serial.ToTypedMessage(responseSettings)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
return config, nil
|
return config, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
var configLoader = NewJSONConfigLoader(
|
|
||||||
ConfigCreatorCache{
|
|
||||||
"none": func() interface{} { return new(NoneResponse) },
|
|
||||||
"http": func() interface{} { return new(HTTPResponse) },
|
|
||||||
},
|
|
||||||
"type",
|
|
||||||
"",
|
|
||||||
)
|
|
||||||
|
|||||||
@@ -3,7 +3,6 @@ package conf_test
|
|||||||
import (
|
import (
|
||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
"github.com/xtls/xray-core/common/serial"
|
|
||||||
. "github.com/xtls/xray-core/infra/conf"
|
. "github.com/xtls/xray-core/infra/conf"
|
||||||
"github.com/xtls/xray-core/proxy/blackhole"
|
"github.com/xtls/xray-core/proxy/blackhole"
|
||||||
)
|
)
|
||||||
@@ -22,7 +21,7 @@ func TestHTTPResponseJSON(t *testing.T) {
|
|||||||
}`,
|
}`,
|
||||||
Parser: loadJSON(creator),
|
Parser: loadJSON(creator),
|
||||||
Output: &blackhole.Config{
|
Output: &blackhole.Config{
|
||||||
Response: serial.ToTypedMessage(&blackhole.HTTPResponse{}),
|
Response: &blackhole.Response{Type: "http"},
|
||||||
},
|
},
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
@@ -32,3 +31,27 @@ func TestHTTPResponseJSON(t *testing.T) {
|
|||||||
},
|
},
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestCustomResponseJSON(t *testing.T) {
|
||||||
|
creator := func() Buildable {
|
||||||
|
return new(BlackholeConfig)
|
||||||
|
}
|
||||||
|
|
||||||
|
runMultiTestCase(t, []TestCase{
|
||||||
|
{
|
||||||
|
Input: `{
|
||||||
|
"response": {
|
||||||
|
"type": "custom",
|
||||||
|
"customResponseData": "Y3VzdG9tIHJlc3BvbnNl"
|
||||||
|
}
|
||||||
|
}`,
|
||||||
|
Parser: loadJSON(creator),
|
||||||
|
Output: &blackhole.Config{
|
||||||
|
Response: &blackhole.Response{
|
||||||
|
Type: "custom",
|
||||||
|
CustomResponseData: []byte("custom response"),
|
||||||
|
},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|||||||
@@ -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)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
+2
-2
@@ -97,7 +97,7 @@ func (v *HTTPClientConfig) Build() (proto.Message, error) {
|
|||||||
user.Email = v.Email
|
user.Email = v.Email
|
||||||
} else {
|
} else {
|
||||||
if err := json.Unmarshal(rawUser, user); err != nil {
|
if err := json.Unmarshal(rawUser, user); err != nil {
|
||||||
return nil, errors.New("failed to parse HTTP user").Base(err).AtError()
|
return nil, errors.New("failed to parse HTTP user").Base(err)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
account := new(HTTPAccount)
|
account := new(HTTPAccount)
|
||||||
@@ -106,7 +106,7 @@ func (v *HTTPClientConfig) Build() (proto.Message, error) {
|
|||||||
account.Password = v.Password
|
account.Password = v.Password
|
||||||
} else {
|
} else {
|
||||||
if err := json.Unmarshal(rawUser, account); err != nil {
|
if err := json.Unmarshal(rawUser, account); err != nil {
|
||||||
return nil, errors.New("failed to parse HTTP account").Base(err).AtError()
|
return nil, errors.New("failed to parse HTTP account").Base(err)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
user.Account = serial.ToTypedMessage(account.Build())
|
user.Account = serial.ToTypedMessage(account.Build())
|
||||||
|
|||||||
+1
-1
@@ -18,7 +18,7 @@ func RegisterConfigureFilePostProcessingStage(name string, stage ConfigureFilePo
|
|||||||
func PostProcessConfigureFile(conf *Config) error {
|
func PostProcessConfigureFile(conf *Config) error {
|
||||||
for k, v := range configureFilePostProcessingStages {
|
for k, v := range configureFilePostProcessingStages {
|
||||||
if err := v.Process(conf); err != nil {
|
if err := v.Process(conf); err != nil {
|
||||||
return errors.New("Rejected by Postprocessing Stage ", k).AtError().Base(err)
|
return errors.New("Rejected by Postprocessing Stage ", k).Base(err)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
return nil
|
return nil
|
||||||
|
|||||||
@@ -13,7 +13,7 @@ type ConfigCreatorCache map[string]ConfigCreator
|
|||||||
|
|
||||||
func (v ConfigCreatorCache) RegisterCreator(id string, creator ConfigCreator) error {
|
func (v ConfigCreatorCache) RegisterCreator(id string, creator ConfigCreator) error {
|
||||||
if _, found := v[id]; found {
|
if _, found := v[id]; found {
|
||||||
return errors.New(id, " already registered.").AtError()
|
return errors.New(id, " already registered.")
|
||||||
}
|
}
|
||||||
|
|
||||||
v[id] = creator
|
v[id] = creator
|
||||||
@@ -61,7 +61,7 @@ func (v *JSONConfigLoader) Load(raw []byte) (interface{}, string, error) {
|
|||||||
}
|
}
|
||||||
rawID, found := obj[v.idKey]
|
rawID, found := obj[v.idKey]
|
||||||
if !found {
|
if !found {
|
||||||
return nil, "", errors.New(v.idKey, " not found in JSON context").AtError()
|
return nil, "", errors.New(v.idKey, " not found in JSON context")
|
||||||
}
|
}
|
||||||
var id string
|
var id string
|
||||||
if err := json.Unmarshal(rawID, &id); err != nil {
|
if err := json.Unmarshal(rawID, &id); err != nil {
|
||||||
|
|||||||
@@ -0,0 +1,103 @@
|
|||||||
|
package conf
|
||||||
|
|
||||||
|
import (
|
||||||
|
"net/netip"
|
||||||
|
"strings"
|
||||||
|
|
||||||
|
"github.com/xtls/xray-core/common/errors"
|
||||||
|
"github.com/xtls/xray-core/common/protocol"
|
||||||
|
"github.com/xtls/xray-core/common/serial"
|
||||||
|
"github.com/xtls/xray-core/proxy/masque"
|
||||||
|
"google.golang.org/protobuf/proto"
|
||||||
|
)
|
||||||
|
|
||||||
|
type MasqueClientConfig struct {
|
||||||
|
Address *Address `json:"address"`
|
||||||
|
Port uint16 `json:"port"`
|
||||||
|
RemoteDNS []string `json:"remoteDNS"`
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *MasqueClientConfig) Build() (proto.Message, error) {
|
||||||
|
if c.Address == nil {
|
||||||
|
return nil, errors.New(`MASQUE: "address" is not set`)
|
||||||
|
}
|
||||||
|
if c.Port == 0 {
|
||||||
|
return nil, errors.New(`MASQUE: "port" is not set`)
|
||||||
|
}
|
||||||
|
for _, s := range c.RemoteDNS {
|
||||||
|
if _, err := netip.ParseAddr(s); err != nil {
|
||||||
|
return nil, errors.New(`MASQUE: invalid "remoteDNS" `, s).Base(err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return &masque.ClientConfig{
|
||||||
|
Server: &protocol.ServerEndpoint{
|
||||||
|
Address: c.Address.Build(),
|
||||||
|
Port: uint32(c.Port),
|
||||||
|
},
|
||||||
|
RemoteDns: c.RemoteDNS,
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
type MasqueUserConfig struct {
|
||||||
|
Pass string `json:"pass"`
|
||||||
|
Level uint32 `json:"level"`
|
||||||
|
Email string `json:"email"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type MasqueServerConfig struct {
|
||||||
|
Users []*MasqueUserConfig `json:"users"`
|
||||||
|
Clients []*MasqueUserConfig `json:"clients"`
|
||||||
|
Address []string `json:"address"`
|
||||||
|
MTU uint32 `json:"mtu"`
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *MasqueServerConfig) Build() (proto.Message, error) {
|
||||||
|
if c.Clients != nil {
|
||||||
|
c.Users = c.Clients
|
||||||
|
}
|
||||||
|
config := &masque.ServerConfig{
|
||||||
|
Address: c.Address,
|
||||||
|
Mtu: c.MTU,
|
||||||
|
}
|
||||||
|
emails := make(map[string]bool)
|
||||||
|
for _, user := range c.Users {
|
||||||
|
if user.Email == "" {
|
||||||
|
return nil, errors.New(`MASQUE: "email" is empty`)
|
||||||
|
}
|
||||||
|
if strings.Contains(user.Email, ":") {
|
||||||
|
return nil, errors.New(`MASQUE: invalid "email" `, user.Email)
|
||||||
|
}
|
||||||
|
if user.Pass == "" {
|
||||||
|
return nil, errors.New(`MASQUE: "pass" of `, user.Email, ` is empty`)
|
||||||
|
}
|
||||||
|
email := strings.ToLower(user.Email)
|
||||||
|
if emails[email] {
|
||||||
|
return nil, errors.New(`MASQUE: duplicate "email" `, user.Email)
|
||||||
|
}
|
||||||
|
emails[email] = true
|
||||||
|
config.Users = append(config.Users, &protocol.User{
|
||||||
|
Email: user.Email,
|
||||||
|
Level: user.Level,
|
||||||
|
Account: serial.ToTypedMessage(&masque.Account{Password: user.Pass}),
|
||||||
|
})
|
||||||
|
}
|
||||||
|
if len(c.Address) == 0 {
|
||||||
|
return nil, errors.New(`MASQUE: "address" is not set`)
|
||||||
|
}
|
||||||
|
var v4, v6 bool
|
||||||
|
for _, s := range c.Address {
|
||||||
|
prefix, err := netip.ParsePrefix(s)
|
||||||
|
if err != nil {
|
||||||
|
return nil, errors.New(`MASQUE: invalid "address" `, s).Base(err)
|
||||||
|
}
|
||||||
|
if prefix.Addr().Is4() && v4 || prefix.Addr().Is6() && v6 {
|
||||||
|
return nil, errors.New(`MASQUE: "address" takes at most one IPv4 and one IPv6 prefix`)
|
||||||
|
}
|
||||||
|
v4 = v4 || prefix.Addr().Is4()
|
||||||
|
v6 = v6 || prefix.Addr().Is6()
|
||||||
|
}
|
||||||
|
if c.MTU != 0 && (c.MTU < 1280 || c.MTU > 65535) {
|
||||||
|
return nil, errors.New(`MASQUE: "mtu" must be between 1280 and 65535`)
|
||||||
|
}
|
||||||
|
return config, nil
|
||||||
|
}
|
||||||
@@ -0,0 +1,188 @@
|
|||||||
|
package conf_test
|
||||||
|
|
||||||
|
import (
|
||||||
|
"encoding/json"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/xtls/xray-core/common/protocol"
|
||||||
|
"github.com/xtls/xray-core/common/serial"
|
||||||
|
. "github.com/xtls/xray-core/infra/conf"
|
||||||
|
masqueproxy "github.com/xtls/xray-core/proxy/masque"
|
||||||
|
"github.com/xtls/xray-core/transport/internet/masque"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestMasqueConfig(t *testing.T) {
|
||||||
|
creator := func() Buildable {
|
||||||
|
return new(MasqueConfig)
|
||||||
|
}
|
||||||
|
|
||||||
|
runMultiTestCase(t, []TestCase{
|
||||||
|
{
|
||||||
|
Input: `{}`,
|
||||||
|
Parser: loadJSON(creator),
|
||||||
|
Output: &masque.Config{Path: "/.well-known/masque/ip/*/*/"},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
Input: `{
|
||||||
|
"host": "example.com:8443",
|
||||||
|
"path": "/.well-known/masque/ip/{target}/{ipproto}/",
|
||||||
|
"headers": {"Authorization": "Basic dTpw"}
|
||||||
|
}`,
|
||||||
|
Parser: loadJSON(creator),
|
||||||
|
Output: &masque.Config{
|
||||||
|
Host: "example.com:8443",
|
||||||
|
Path: "/.well-known/masque/ip/*/*/",
|
||||||
|
Headers: map[string]string{"Authorization": "Basic dTpw"},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
Input: `{"path": "/masque/ip{?target,ipproto}"}`,
|
||||||
|
Parser: loadJSON(creator),
|
||||||
|
Output: &masque.Config{Path: "/masque/ip?target=*&ipproto=*"},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
Input: `{"user": "u", "pass": "p:q", "headers": {"X-Token": "a"}}`,
|
||||||
|
Parser: loadJSON(creator),
|
||||||
|
Output: &masque.Config{
|
||||||
|
Path: "/.well-known/masque/ip/*/*/",
|
||||||
|
Headers: map[string]string{"Authorization": "Basic dTpwOnE=", "X-Token": "a"},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
})
|
||||||
|
|
||||||
|
for _, input := range []string{
|
||||||
|
`{"path": "/masque/{target}/{ipproto}/{dns}"}`,
|
||||||
|
`{"path": "masque"}`,
|
||||||
|
`{"host": "example.com/path"}`,
|
||||||
|
`{"headers": {"host": "example.com"}}`,
|
||||||
|
`{"headers": {"Capsule-Protocol": "?0"}}`,
|
||||||
|
`{"headers": {"X Token": "a"}}`,
|
||||||
|
`{"headers": {"X-Token": "a\r\nb"}}`,
|
||||||
|
`{"user": "u:v", "pass": "p"}`,
|
||||||
|
`{"user": "u", "pass": "p", "headers": {"authorization": "Basic dTpw"}}`,
|
||||||
|
} {
|
||||||
|
if _, err := loadJSON(creator)(input); err == nil {
|
||||||
|
t.Errorf("expected an error for %s", input)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMasqueOutboundConfig(t *testing.T) {
|
||||||
|
build := func(s string) error {
|
||||||
|
c := new(OutboundDetourConfig)
|
||||||
|
if err := json.Unmarshal([]byte(s), c); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
_, err := c.Build()
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := build(`{
|
||||||
|
"protocol": "masque",
|
||||||
|
"settings": {"address": "example.com", "port": 443},
|
||||||
|
"streamSettings": {"network": "masque", "security": "tls"},
|
||||||
|
"mux": {"enabled": false, "concurrency": -1}
|
||||||
|
}`); err != nil {
|
||||||
|
t.Error(err)
|
||||||
|
}
|
||||||
|
for _, input := range []string{
|
||||||
|
`{"protocol": "masque", "settings": {"address": "example.com"}, "streamSettings": {"network": "masque", "security": "tls"}}`,
|
||||||
|
`{"protocol": "masque", "settings": {"address": "example.com", "port": 443}, "streamSettings": {"network": "masque", "security": "tls"}, "mux": {"enabled": true}}`,
|
||||||
|
`{"protocol": "masque", "settings": {"address": "example.com", "port": 443}, "streamSettings": {"network": "masque", "security": "tls"}, "mux": {"enabled": true, "concurrency": -1}}`,
|
||||||
|
`{"protocol": "freedom", "streamSettings": {"network": "masque", "security": "tls"}}`,
|
||||||
|
} {
|
||||||
|
if err := build(input); err == nil {
|
||||||
|
t.Errorf("expected an error for %s", input)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMasqueServerConfig(t *testing.T) {
|
||||||
|
creator := func() Buildable {
|
||||||
|
return new(MasqueServerConfig)
|
||||||
|
}
|
||||||
|
|
||||||
|
runMultiTestCase(t, []TestCase{
|
||||||
|
{
|
||||||
|
Input: `{
|
||||||
|
"users": [{"email": "u@example.com", "pass": "p", "level": 1}],
|
||||||
|
"address": ["10.13.0.1/24", "fd13::1/64"],
|
||||||
|
"mtu": 1400
|
||||||
|
}`,
|
||||||
|
Parser: loadJSON(creator),
|
||||||
|
Output: &masqueproxy.ServerConfig{
|
||||||
|
Users: []*protocol.User{{
|
||||||
|
Email: "u@example.com",
|
||||||
|
Level: 1,
|
||||||
|
Account: serial.ToTypedMessage(&masqueproxy.Account{Password: "p"}),
|
||||||
|
}},
|
||||||
|
Address: []string{"10.13.0.1/24", "fd13::1/64"},
|
||||||
|
Mtu: 1400,
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
Input: `{"clients": [{"email": "u", "pass": "p:q"}], "address": ["10.13.0.1/24"]}`,
|
||||||
|
Parser: loadJSON(creator),
|
||||||
|
Output: &masqueproxy.ServerConfig{
|
||||||
|
Users: []*protocol.User{{
|
||||||
|
Email: "u",
|
||||||
|
Account: serial.ToTypedMessage(&masqueproxy.Account{Password: "p:q"}),
|
||||||
|
}},
|
||||||
|
Address: []string{"10.13.0.1/24"},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
Input: `{"address": ["10.13.0.1/24"]}`,
|
||||||
|
Parser: loadJSON(creator),
|
||||||
|
Output: &masqueproxy.ServerConfig{
|
||||||
|
Address: []string{"10.13.0.1/24"},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
})
|
||||||
|
|
||||||
|
for _, input := range []string{
|
||||||
|
`{"users": [{"email": "u:v", "pass": "p"}], "address": ["10.13.0.1/24"]}`,
|
||||||
|
`{"users": [{"email": "", "pass": "p"}], "address": ["10.13.0.1/24"]}`,
|
||||||
|
`{"users": [{"pass": "p"}], "address": ["10.13.0.1/24"]}`,
|
||||||
|
`{"users": [{"email": "u", "pass": ""}], "address": ["10.13.0.1/24"]}`,
|
||||||
|
`{"users": [{"email": "u", "pass": "p"}, {"email": "U", "pass": "q"}], "address": ["10.13.0.1/24"]}`,
|
||||||
|
`{"users": [{"email": "u", "pass": "p"}]}`,
|
||||||
|
`{"users": [{"email": "u", "pass": "p"}], "address": ["10.13.0.1"]}`,
|
||||||
|
`{"users": [{"email": "u", "pass": "p"}], "address": ["10.13.0.1/24", "10.14.0.1/24"]}`,
|
||||||
|
`{"users": [{"email": "u", "pass": "p"}], "address": ["fd13::1/64", "fd14::1/64"]}`,
|
||||||
|
`{"users": [{"email": "u", "pass": "p"}], "address": ["10.13.0.1/24"], "mtu": 1000}`,
|
||||||
|
`{"users": [{"email": "u", "pass": "p"}], "address": ["10.13.0.1/24"], "mtu": 70000}`,
|
||||||
|
} {
|
||||||
|
if _, err := loadJSON(creator)(input); err == nil {
|
||||||
|
t.Errorf("expected an error for %s", input)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMasqueInboundConfig(t *testing.T) {
|
||||||
|
build := func(s string) error {
|
||||||
|
c := new(InboundDetourConfig)
|
||||||
|
if err := json.Unmarshal([]byte(s), c); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
_, err := c.Build()
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := build(`{
|
||||||
|
"protocol": "masque",
|
||||||
|
"port": 443,
|
||||||
|
"settings": {"users": [{"email": "u@example.com", "pass": "p"}], "address": ["10.13.0.1/24"]},
|
||||||
|
"streamSettings": {"network": "masque", "security": "tls"}
|
||||||
|
}`); err != nil {
|
||||||
|
t.Error(err)
|
||||||
|
}
|
||||||
|
if err := build(`{
|
||||||
|
"protocol": "vless",
|
||||||
|
"port": 443,
|
||||||
|
"settings": {"users": [{"id": "27848739-7e62-4138-9fd3-098a63964b6b"}], "decryption": "none"},
|
||||||
|
"streamSettings": {"network": "masque", "security": "tls"}
|
||||||
|
}`); err == nil {
|
||||||
|
t.Error("expected an error for the masque transport on a vless inbound")
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -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
|
||||||
|
|||||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user