Compare commits

..
Author SHA1 Message Date
patternihaandClaude Opus 5.5 747b153333 TUN inbound: Error out when autoSystemWfpBlockLeak or autoSystemDnsToGateway cannot apply
As asked in review, rather than run without them:

- The config is rejected, also by xray -test, for autoSystemWfpBlockLeak
  without autoSystemRoutingTable, or with "dns" but without dns, on
  Windows, and for autoSystemDnsToGateway without gateway on Linux.
- Xray does not start when the filters cannot be added, now on every
  Windows version, or when the system DNS cannot be set on Linux,
  instead of logging it and running without them.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
2026-09-30 08:46:45 +03:30
patternihaandClaude Opus 5.5 94cd83ccb4 TUN inbound: Rename "misconfig" to "misconfigtun"
The autoSystemWfpBlockLeak value that blocks an IP version not routed to
the TUN, as asked in review; "misconfig" was never released.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
2026-09-30 07:59:39 +03:30
patternihaandClaude Opus 5.5 1c225a8041 TUN inbound: Make autoSystemWfpBlockLeak a list of the leaks to block
autoSystemWfpBlockLeak now takes ["dns", "misconfig"] instead of true:
"dns" keeps DNS inside the TUN, and "misconfig" blocks an IP version
that no route leads to the TUN, the leak of a configuration that routes
only one of them. Either can be used alone, e.g. ["dns"] to block DNS
leaks while an IP version stays out of the TUN on purpose. Unknown
values are rejected. The config field becomes a repeated string with
the same number; it was never released.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
2026-09-30 07:46:51 +03:30
patternihaandClaude Opus 5.5 3a6bdb19ba TUN inbound: autoSystemDnsToGateway falls back to an IPv6 gateway on Linux
Without an IPv4 address in gateway, the system DNS now points at the
first IPv6 gateway plus one (e.g. fc00::1/64 -> fc00::2) instead of
nothing, and the routing check before the takeover accepts IPv6
addresses for it.

The README also says what each system does without gateway: Xray
assigns no address on Linux, Windows gives the TUN link-local ones
itself, and macOS and FreeBSD use 169.254.10.1/30.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
2026-09-30 07:18:55 +03:30
patternihaandClaude Opus 5.5 de02da553a TUN inbound: Block IPv4 too when it is not routed to the TUN on Windows
autoSystemWfpBlockLeak blocked IPv6 when the TUN had no IPv6 address or
no IPv6 route, and IPv4 never: with only IPv6 routed to the TUN, IPv4
went around it. Now each IP version is blocked when no route of it leads
to the TUN, except for loopback, DHCP, IPv6 neighbor and multicast
listener discovery, and Xray itself.

Addresses no longer count: without one of a version in gateway, Windows
gives the TUN a link-local one itself (fe80:: at once, 169.254.x.x after
some seconds), and what is routed to the TUN goes through it with that,
so a TUN with IPv6 routes but no IPv6 address had its IPv6 blocked for
nothing.

Tested on Windows 11, elevated: with only IPv6 routed, other programs'
IPv4 is denied, while loopback, a DHCP renew of Wi-Fi and Xray still
work; without gateway, IPv4 and IPv6 routed to the TUN enter it from
169.254.x.x and fe80::.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
2026-09-30 03:07:14 +03:30
patternihaandClaude Opus 5.5 4ec4fb8aab TUN inbound: Rename to autoSystemWfpBlockLeak and autoSystemDnsToGateway
autoSystemWFP becomes autoSystemWfpBlockLeak, saying that the WFP filters
block leaks, and autoSystemDNS becomes autoSystemDnsToGateway, saying
where it points the system DNS, so that pointing the system DNS at the
gateway on Windows later would fit the same name. Their config fields
keep their numbers; neither was released.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
2026-09-30 02:31:00 +03:30
patternihaandClaude Opus 5.5 63de6135cb TUN inbound: Rename strictRoute to autoSystemWFP
It only turns on the Windows Filtering Platform filters, along with Xray
resolving its own lookups while they restrict DNS, so it is named after
what it sets up in the system, like autoSystemRoutingTable and
autoSystemDNS. The config field becomes auto_system_wfp, with the same
number; strictRoute was never released.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
2026-09-30 01:22:22 +03:30
patternihaandClaude Opus 5.5 edd916b08e TUN inbound: Keep Windows' DNS Client from sending DoH/DoT outside the TUN
With strictRoute, only port 53 was kept inside the TUN. But Windows' DNS
Client service sends the queries for an interface's DNS servers out
through that interface, whatever the routes say, and since Windows 11
and Server 2022 it may send them over HTTPS or TLS, when that is set up
for the interface (as Windows Settings does) or for the server. Those
left through the physical link.

On those versions, the DNS Client service may now only connect through
the TUN, except for its mDNS and LLMNR. The filters recognize the
service by its SID in the token of its process, as Windows Firewall's
own rules for it do. Earlier versions only query port 53, and may run
the service in one process with others, so they get no such filters.
The port 53 rule stays, for the programs that query a resolver on the
local network themselves, and for those earlier versions.

Tested on Windows 11, elevated: with DoH set on Wi-Fi per adapter, per
network profile or by global auto-upgrade, none of the DNS Client's
connections left through Wi-Fi (WFP logged the drops by the new filter),
names still resolved through the TUN, mDNS and LLMNR still went out, and
other programs were unaffected, also in a real Xray run.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
2026-09-30 01:09:52 +03:30
patternihaandClaude Opus 5.5 eb29a4e3de TUN inbound: Make strictRoute false by default
Like sing-box's strict_route, strictRoute is now false by default, so the
Windows Filtering Platform filters are only added when it is set to true
(together with autoSystemRoutingTable). With unset meaning false,
strict_route becomes a plain bool field.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
2026-09-29 14:51:17 +03:30
patternihaandClaude Opus 5.5 2db099b34b TUN inbound: Block DNS and IPv6 leaks outside the TUN on Windows; Add strictRoute
Windows sends name queries to the DNS servers of all interfaces, and a
resolver on the local network (e.g. 192.168.1.1 from DHCP) is reached
through its more specific LAN route instead of the TUN, so DNS leaks
past it. IPv6 bypasses a TUN that cannot carry it.

With autoSystemRoutingTable set, the Windows TUN now adds Windows
Filtering Platform filters, all in one transaction and in a dynamic
session, so that they are removed when Xray exits, even if it crashes:
- DNS (port 53) only goes through the TUN, in both directions: its local
  address, and the interface it leaves or arrives by, must be the TUN's.
- IPv6 is blocked in both directions when the TUN has no IPv6 address or
  no IPv6 route, except loopback, neighbor and multicast listener
  discovery, and DHCPv6.
- Xray's own traffic is exempt: its connections out with a hard permit,
  which Windows Firewall rules do not override (like sing-box's
  strict_route), connections to its inbounds with an ordinary one.
If the filters cannot be added, the TUN does not start on Windows 10 and
later (only a warning on 7/8). The new `strictRoute` option (true by
default) turns them off.

Also on Windows:
- A warning for `dns` servers outside gateway and autoSystemRoutingTable,
  as queries to them cannot go through the TUN and are blocked.
- While DNS is restricted and autoOutboundsInterface is in use, Xray
  resolves the names it would ask Windows for itself (Go's resolver on
  its own sockets). Those lookups and the `localhost` DNS server skip the
  TUN's DNS servers, unless another interface uses them too, instead of
  looping back into the TUN.
- The DNS cache is flushed when the TUN starts and stops, and DNS
  registration is turned off on the TUN (through netsh before Windows 10
  1809).
- Close no longer panics when registering the route or interface change
  callbacks failed.
The README's Windows section describes all of it.

Tested on Windows 11, elevated, amd64 and 386: the filters, DNS arriving
through a real Wintun adapter and blocked outside it, the IPv6 block,
Windows Firewall rules, and a real Xray run. Windows 7/8 and Windows 10
before 1809 are untested.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
2026-09-28 08:15:40 +03:30
153 changed files with 4114 additions and 12124 deletions
-3
View File
@@ -470,9 +470,6 @@ func (d *DefaultDispatcher) routedDispatch(ctx context.Context, link *transport.
return // DO NOT CHANGE: the traffic shouldn't be processed by default outbound if the specified outbound tag doesn't exist (yet), e.g., VLESS Reverse Proxy
}
} else {
if err != common.ErrNoClue {
errors.LogErrorInner(ctx, err, "failed to pick route for ", destination)
}
errors.LogInfo(ctx, "default route for ", destination)
}
}
+6 -25
View File
@@ -93,7 +93,6 @@ type NameServer struct {
UnexpectedIp []*geodata.IPRule `protobuf:"bytes,13,rep,name=unexpected_ip,json=unexpectedIp,proto3" json:"unexpected_ip,omitempty"`
ActUnprior bool `protobuf:"varint,14,opt,name=actUnprior,proto3" json:"actUnprior,omitempty"`
PolicyID uint32 `protobuf:"varint,17,opt,name=policyID,proto3" json:"policyID,omitempty"`
Id string `protobuf:"bytes,18,opt,name=id,proto3" json:"id,omitempty"`
unknownFields protoimpl.UnknownFields
sizeCache protoimpl.SizeCache
}
@@ -240,13 +239,6 @@ func (x *NameServer) GetPolicyID() uint32 {
return 0
}
func (x *NameServer) GetId() string {
if x != nil {
return x.Id
}
return ""
}
type Config struct {
state protoimpl.MessageState `protogen:"open.v1"`
// NameServer list used by this DNS client.
@@ -266,10 +258,8 @@ type Config struct {
DisableFallback bool `protobuf:"varint,10,opt,name=disableFallback,proto3" json:"disableFallback,omitempty"`
DisableFallbackIfMatch bool `protobuf:"varint,11,opt,name=disableFallbackIfMatch,proto3" json:"disableFallbackIfMatch,omitempty"`
EnableParallelQuery bool `protobuf:"varint,14,opt,name=enableParallelQuery,proto3" json:"enableParallelQuery,omitempty"`
// Absolute path to the Lua DNS query script.
Script string `protobuf:"bytes,15,opt,name=script,proto3" json:"script,omitempty"`
unknownFields protoimpl.UnknownFields
sizeCache protoimpl.SizeCache
unknownFields protoimpl.UnknownFields
sizeCache protoimpl.SizeCache
}
func (x *Config) Reset() {
@@ -379,13 +369,6 @@ func (x *Config) GetEnableParallelQuery() bool {
return false
}
func (x *Config) GetScript() string {
if x != nil {
return x.Script
}
return ""
}
type Config_HostMapping struct {
state protoimpl.MessageState `protogen:"open.v1"`
Domain *geodata.DomainRule `protobuf:"bytes,2,opt,name=domain,proto3" json:"domain,omitempty"`
@@ -452,7 +435,7 @@ var File_app_dns_config_proto protoreflect.FileDescriptor
const file_app_dns_config_proto_rawDesc = "" +
"\n" +
"\x14app/dns/config.proto\x12\fxray.app.dns\x1a\x1ccommon/net/destination.proto\x1a\x1bcommon/geodata/geodat.proto\"\xee\x05\n" +
"\x14app/dns/config.proto\x12\fxray.app.dns\x1a\x1ccommon/net/destination.proto\x1a\x1bcommon/geodata/geodat.proto\"\xde\x05\n" +
"\n" +
"NameServer\x123\n" +
"\aaddress\x18\x01 \x01(\v2\x19.xray.common.net.EndpointR\aaddress\x12\x1b\n" +
@@ -478,11 +461,10 @@ const file_app_dns_config_proto_rawDesc = "" +
"\n" +
"actUnprior\x18\x0e \x01(\bR\n" +
"actUnprior\x12\x1a\n" +
"\bpolicyID\x18\x11 \x01(\rR\bpolicyID\x12\x0e\n" +
"\x02id\x18\x12 \x01(\tR\x02idB\x0f\n" +
"\bpolicyID\x18\x11 \x01(\rR\bpolicyIDB\x0f\n" +
"\r_disableCacheB\r\n" +
"\v_serveStaleB\x12\n" +
"\x10_serveExpiredTTLJ\x04\b\x04\x10\x05\"\x9a\x05\n" +
"\x10_serveExpiredTTLJ\x04\b\x04\x10\x05\"\x82\x05\n" +
"\x06Config\x129\n" +
"\vname_server\x18\x05 \x03(\v2\x18.xray.app.dns.NameServerR\n" +
"nameServer\x12\x1b\n" +
@@ -498,8 +480,7 @@ const file_app_dns_config_proto_rawDesc = "" +
"\x0fdisableFallback\x18\n" +
" \x01(\bR\x0fdisableFallback\x126\n" +
"\x16disableFallbackIfMatch\x18\v \x01(\bR\x16disableFallbackIfMatch\x120\n" +
"\x13enableParallelQuery\x18\x0e \x01(\bR\x13enableParallelQuery\x12\x16\n" +
"\x06script\x18\x0f \x01(\tR\x06script\x1a}\n" +
"\x13enableParallelQuery\x18\x0e \x01(\bR\x13enableParallelQuery\x1a}\n" +
"\vHostMapping\x127\n" +
"\x06domain\x18\x02 \x01(\v2\x1f.xray.common.geodata.DomainRuleR\x06domain\x12\x0e\n" +
"\x02ip\x18\x03 \x03(\fR\x02ip\x12%\n" +
-4
View File
@@ -27,7 +27,6 @@ message NameServer {
repeated xray.common.geodata.IPRule unexpected_ip = 13;
bool actUnprior = 14;
uint32 policyID = 17;
string id = 18;
}
enum QueryStrategy {
@@ -74,7 +73,4 @@ message Config {
bool disableFallbackIfMatch = 11;
bool enableParallelQuery = 14;
// Absolute path to the Lua DNS query script.
string script = 15;
}
-16
View File
@@ -31,8 +31,6 @@ type DNS struct {
domainMatcher geodata.DomainMatcher
matcherInfos []*DomainMatcherInfo
checkSystem bool
script *scriptEngine
scriptPath string
}
// DomainMatcherInfo contains information attached to index returned by Server.domainMatcher.
@@ -182,7 +180,6 @@ func New(ctx context.Context, config *Config) (*DNS, error) {
disableFallbackIfMatch: config.DisableFallbackIfMatch,
enableParallelQuery: config.EnableParallelQuery,
checkSystem: checkSystem,
scriptPath: config.Script,
}, nil
}
@@ -193,21 +190,11 @@ func (*DNS) Type() interface{} {
// Start implements common.Runnable.
func (s *DNS) Start() error {
if s.scriptPath != "" {
engine, err := newScriptEngine(s.scriptPath, s)
if err != nil {
return errors.New("failed to initialize DNS script").Base(err)
}
s.script = engine
}
return nil
}
// Close implements common.Closable.
func (s *DNS) Close() error {
if s.script != nil {
s.script.close()
}
return nil
}
@@ -292,9 +279,6 @@ func (s *DNS) LookupIP(domain string, option dns.IPOption) ([]net.IP, uint32, er
}
// Name servers lookup
if s.script != nil {
return s.script.query(domain, option)
}
if s.enableParallelQuery {
return s.parallelQuery(domain, option)
} else {
-169
View File
@@ -1,169 +0,0 @@
package dns
import (
"context"
"strings"
"github.com/xtls/xray-core/common/errors"
xlua "github.com/xtls/xray-core/common/lua"
"github.com/xtls/xray-core/common/net"
featureDNS "github.com/xtls/xray-core/features/dns"
"github.com/xtls/xray-core/features/dns/localdns"
lua "github.com/yuin/gopher-lua"
)
// luaDNSServer adapts configured and local DNS to the same Lua API.
type luaDNSServer struct {
id string
name string
query func(context.Context, string, featureDNS.IPOption) ([]net.IP, uint32, error)
}
// RegisterLua makes xray.dns available to scripts backed by client.
func RegisterLua(L *lua.LState, client featureDNS.Client) {
var servers []luaDNSServer
switch client := client.(type) {
case *DNS:
servers = luaServers(client)
case *localdns.Client:
servers = []luaDNSServer{{
id: "localhost",
name: "localhost",
query: func(_ context.Context, domain string, option featureDNS.IPOption) ([]net.IP, uint32, error) {
return client.LookupIP(domain, option)
},
}}
}
registerLua(L, servers, client)
}
// registerLua makes xray.dns available to DNS scripts.
func (s *DNS) registerLua(L *lua.LState) {
registerLua(L, luaServers(s), nil)
}
func luaServers(s *DNS) []luaDNSServer {
servers := make([]luaDNSServer, len(s.clients))
for i, client := range s.clients {
servers[i] = luaDNSServer{id: client.id, name: client.Name(), query: client.QueryIP}
}
return servers
}
func registerLua(L *lua.LState, servers []luaDNSServer, client featureDNS.Client) {
L.PreloadModule("xray.dns", func(L *lua.LState) int {
pushIPs := xlua.NewSlicePusher[net.IP](L)
serverList := L.CreateTable(len(servers), 0)
for i, client := range servers {
server := L.CreateTable(0, 2)
server.RawSetString("ID", lua.LString(client.id))
server.RawSetString("Query", L.NewFunction(func(L *lua.LState) int {
domain, ok := L.Get(2).(lua.LString)
if !ok {
L.RaiseError("server:Query requires a domain")
return 0
}
option := featureDNS.IPOption{
IPv4Enable: L.CheckBool(3),
IPv6Enable: L.CheckBool(4),
FakeEnable: L.CheckBool(5),
}
ctx := L.Context()
if ctx == nil {
L.RaiseError("server:Query requires an active DNS query")
return 0
}
var ips []net.IP
var ttl uint32
var err error
if !option.FakeEnable && strings.EqualFold(client.name, "FakeDNS") {
err = featureDNS.ErrEmptyResponse
} else {
ips, ttl, err = client.query(ctx, string(domain), option)
}
pushIPs(L, ips)
xlua.PushNumber(L, ttl)
xlua.PushError(L, err)
return 3
}))
serverList.RawSetInt(i+1, server)
}
module := L.CreateTable(0, 2)
if servers != nil {
module.RawSetString("Servers", serverList)
}
if client != nil {
module.RawSetString("Query", newLuaClientQuery(L, client, pushIPs))
}
L.Push(module)
return 1
})
}
func newLuaClientQuery(L *lua.LState, client featureDNS.Client, pushIPs func(*lua.LState, []net.IP)) *lua.LFunction {
return L.NewFunction(func(L *lua.LState) int {
domain, ok := L.Get(1).(lua.LString)
if !ok {
L.RaiseError("dns.Query requires a domain")
return 0
}
option := featureDNS.IPOption{
IPv4Enable: L.CheckBool(2),
IPv6Enable: L.CheckBool(3),
FakeEnable: L.CheckBool(4),
}
if L.Context() == nil {
L.RaiseError("dns.Query requires an active DNS query")
return 0
}
ips, ttl, err := client.LookupIP(string(domain), option)
pushIPs(L, ips)
xlua.PushNumber(L, ttl)
xlua.PushError(L, err)
return 3
})
}
// callLuaQuery runs HandleDNSQuery and leaves (ips, ttl, err) on the stack.
func callLuaQuery(L *lua.LState, domain string, option featureDNS.IPOption) error {
fn := L.GetGlobal("HandleDNSQuery")
if fn.Type() != lua.LTFunction {
return errors.New("DNS script must define HandleDNSQuery(...)")
}
return L.CallByParam(lua.P{Fn: fn, NRet: 3, Protect: true},
lua.LString(strings.ToLower(domain)),
lua.LBool(option.IPv4Enable),
lua.LBool(option.IPv6Enable),
lua.LBool(option.FakeEnable))
}
// readLuaQueryResult reads (ips, ttl, err) from the stack without copying the IPs.
func readLuaQueryResult(L *lua.LState) ([]net.IP, uint32, error) {
if err := xlua.ReadError(L.Get(-1), "DNS script error must be an error or string"); err != nil {
return nil, 0, err
}
ttl, err := xlua.ReadUint32(L.Get(-2), "DNS script returned invalid TTL")
if err != nil {
return nil, 0, err
}
addresses := L.Get(-3)
if addresses == lua.LNil {
return nil, 0, featureDNS.ErrEmptyResponse
}
ips, err := xlua.ReadUserData[[]net.IP](addresses, "DNS script IPs must be native IP slice userdata")
if err != nil {
return nil, 0, err
}
if len(ips) == 0 {
return nil, 0, featureDNS.ErrEmptyResponse
}
return ips, ttl, nil
}
-118
View File
@@ -1,118 +0,0 @@
package dns
import (
"context"
"os"
"path/filepath"
"testing"
"time"
"github.com/xtls/xray-core/common/net"
featureDNS "github.com/xtls/xray-core/features/dns"
lua "github.com/yuin/gopher-lua"
)
// BenchmarkLuaDNSHook isolates scalar argument bridging and a fixed return.
// It excludes upstream queries, result decoding, and state pool management.
func BenchmarkLuaDNSHook(b *testing.B) {
L := lua.NewState()
b.Cleanup(L.Close)
if err := L.DoString(`
function HandleDNSQuery(domain, ipv4, ipv6, fake)
return true
end
`); err != nil {
b.Fatal(err)
}
L.SetContext(context.Background())
option := featureDNS.IPOption{IPv4Enable: true}
if err := callLuaQuery(L, "example.com", option); err != nil {
b.Fatal(err)
}
if L.Get(-3) != lua.LTrue {
b.Fatal("hook did not return true")
}
L.Pop(3)
b.ReportAllocs()
b.ResetTimer()
for i := 0; i < b.N; i++ {
if err := callLuaQuery(L, "example.com", option); err != nil {
b.Fatal(err)
}
L.Pop(3)
}
}
// BenchmarkLuaDNSQuery queries the same preselected, in-memory upstream.
// client_query compares Client.QueryIP to a preloaded server:Query hook.
// script_query additionally measures production pool and timeout management.
// These cases do not measure DNS.LookupIP server selection or network latency.
func BenchmarkLuaDNSQuery(b *testing.B) {
ctx := context.Background()
option := featureDNS.IPOption{IPv4Enable: true}
ip := net.ParseIP("127.0.0.1")
upstream := &benchmarkLuaNameServer{ips: []net.IP{ip}}
client := &Client{server: upstream, ipOption: &option, timeoutMs: time.Second}
server := &DNS{ctx: ctx, clients: []*Client{client}}
const script = `
local server = require("xray.dns").Servers[1]
function HandleDNSQuery(domain, ipv4, ipv6, fake)
return server:Query(domain, ipv4, ipv6, fake)
end
`
L := lua.NewState()
b.Cleanup(L.Close)
server.registerLua(L)
if err := L.DoString(script); err != nil {
b.Fatal(err)
}
L.SetContext(ctx)
path := filepath.Join(b.TempDir(), "query.lua")
if err := os.WriteFile(path, []byte(script), 0o600); err != nil {
b.Fatal(err)
}
engine, err := newScriptEngine(path, server)
if err != nil {
b.Fatal(err)
}
b.Cleanup(engine.close)
for _, bench := range []struct {
name string
query func() ([]net.IP, uint32, error)
}{
{"client_query/native", func() ([]net.IP, uint32, error) {
return client.QueryIP(ctx, "example.com", option)
}},
{"client_query/lua", func() ([]net.IP, uint32, error) {
if err := callLuaQuery(L, "example.com", option); err != nil {
return nil, 0, err
}
ips, ttl, err := readLuaQueryResult(L)
L.Pop(3)
return ips, ttl, err
}},
{"script_query/lua", func() ([]net.IP, uint32, error) {
return engine.query("example.com", option)
}},
} {
b.Run(bench.name, func(b *testing.B) {
ips, ttl, err := bench.query()
if err != nil || ttl != 60 || len(ips) != 1 || !ips[0].Equal(ip) {
b.Fatalf("query() = %v, TTL %d, %v; want %v, TTL 60", ips, ttl, err, ip)
}
b.ReportAllocs()
b.ResetTimer()
for i := 0; i < b.N; i++ {
ips, ttl, err = bench.query()
if err != nil {
b.Fatal(err)
}
}
b.StopTimer()
if ttl != 60 || len(ips) != 1 || !ips[0].Equal(ip) {
b.Fatalf("query() = %v, TTL %d; want %v, TTL 60", ips, ttl, ip)
}
})
}
}
-272
View File
@@ -1,272 +0,0 @@
package dns
import (
"context"
go_errors "errors"
"strings"
"testing"
"time"
"github.com/xtls/xray-core/common/geodata"
"github.com/xtls/xray-core/common/net"
featureDNS "github.com/xtls/xray-core/features/dns"
"github.com/xtls/xray-core/features/dns/localdns"
lua "github.com/yuin/gopher-lua"
)
func TestReadLuaQueryResult(t *testing.T) {
wantIPs := []net.IP{net.ParseIP("8.8.8.8"), {127, 0, 0, 1}, net.ParseIP("::1")}
nativeErr := go_errors.New("upstream failed")
for _, tc := range []struct {
name, values string
wantIPs []net.IP
wantTTL uint32
wantErr error
wantMessage string
}{
{name: "IPs", values: `ips, 45`, wantIPs: wantIPs, wantTTL: 45},
{name: "nil IPs", values: `nil, 0`, wantErr: featureDNS.ErrEmptyResponse},
{name: "empty IPs", values: `emptyIPs, 0`, wantErr: featureDNS.ErrEmptyResponse},
{name: "native error", values: `nil, nil, nativeError`, wantErr: nativeErr},
{name: "string error", values: `nil, nil, "blocked"`, wantMessage: "blocked"},
{name: "fractional TTL", values: `ips, 1.5`, wantMessage: "invalid TTL"},
{name: "oversized TTL", values: `ips, 4294967296`, wantMessage: "invalid TTL"},
{name: "negative TTL", values: `ips, -1`, wantMessage: "invalid TTL"},
{name: "NaN TTL", values: `ips, 0/0`, wantMessage: "invalid TTL"},
{name: "missing TTL", values: `ips`, wantMessage: "invalid TTL"},
{name: "string IPs", values: `"127.0.0.1", 60`, wantMessage: "native IP slice"},
{name: "wrong userdata", values: `ip, 60`, wantMessage: "native IP slice"},
{name: "invalid error", values: `ips, 60, false`, wantMessage: "error or string"},
} {
t.Run(tc.name, func(t *testing.T) {
L := lua.NewState()
defer L.Close()
for name, value := range map[string]any{"ips": wantIPs, "ip": wantIPs[0], "emptyIPs": []net.IP(nil), "nativeError": nativeErr} {
ud := L.NewUserData()
ud.Value = value
L.SetGlobal(name, ud)
}
fn, err := L.LoadString("return " + tc.values)
if err != nil {
t.Fatal(err)
}
if err := L.CallByParam(lua.P{Fn: fn, NRet: 3, Protect: true}); err != nil {
t.Fatal(err)
}
ips, ttl, err := readLuaQueryResult(L)
switch {
case tc.wantErr != nil:
if err != tc.wantErr {
t.Fatalf("error = %v, want original error %v", err, tc.wantErr)
}
case tc.wantMessage != "":
if err == nil || !strings.Contains(err.Error(), tc.wantMessage) {
t.Fatalf("error = %v, want %q", err, tc.wantMessage)
}
case err != nil:
t.Fatal(err)
}
if ttl != tc.wantTTL || len(ips) != len(tc.wantIPs) {
t.Fatalf("result = %v, TTL %d; want %v, TTL %d", ips, ttl, tc.wantIPs, tc.wantTTL)
}
for i := range ips {
if !ips[i].Equal(tc.wantIPs[i]) {
t.Fatalf("IP %d = %v, want %v", i, ips[i], tc.wantIPs[i])
}
}
if len(ips) != 0 && &ips[0] != &tc.wantIPs[0] {
t.Fatal("result copied the IP slice")
}
})
}
}
func TestCallLuaQueryCancellation(t *testing.T) {
L := lua.NewState()
defer L.Close()
if err := L.DoString(`function HandleDNSQuery(domain, ipv4, ipv6, fake) while true do end end`); err != nil {
t.Fatal(err)
}
ctx, cancel := context.WithCancel(context.Background())
cancel()
L.SetContext(ctx)
err := callLuaQuery(L, "example.com", featureDNS.IPOption{IPv4Enable: true})
if err == nil {
t.Fatal("callLuaQuery did not stop after context cancellation")
}
if L.Context() != ctx {
t.Fatal("callLuaQuery changed the Lua state's context")
}
}
func TestCallLuaQuery(t *testing.T) {
L := lua.NewState()
defer L.Close()
addresses := L.NewUserData()
addresses.Value = []net.IP{net.ParseIP("127.0.0.1")}
L.SetGlobal("ips", addresses)
if err := L.DoString(`
function HandleDNSQuery(domain, ipv4, ipv6, fake)
assert(domain == "example.com")
assert(ipv4 and not ipv6 and not fake)
return ips, 60, nil
end
`); err != nil {
t.Fatal(err)
}
if err := callLuaQuery(L, "ExAmPlE.CoM", featureDNS.IPOption{IPv4Enable: true}); err != nil {
t.Fatal(err)
}
if L.GetTop() != 3 || L.Get(1) != addresses || L.Get(2) != lua.LNumber(60) || L.Get(3) != lua.LNil {
t.Fatal("callLuaQuery did not leave the three query results on the stack")
}
}
func TestLuaDNSServerQuery(t *testing.T) {
L := lua.NewState()
defer L.Close()
geodata.RegisterLua(L)
option := featureDNS.IPOption{IPv4Enable: true}
ips := []net.IP{net.ParseIP("127.0.0.1"), net.ParseIP("8.8.8.8")}
server := &DNS{clients: []*Client{{server: &benchmarkLuaNameServer{ips: ips}, ipOption: &option, timeoutMs: time.Second}}}
server.registerLua(L)
if err := L.DoString(`
local server = require("xray.dns").Servers[1]
local matcher = require("xray.geodata").BuildIPMatcher("127.0.0.0/8")
function HandleDNSQuery(domain, ipv4, ipv6, fake)
local ips, ttl, err = server:Query(domain, ipv4, ipv6, fake)
assert(type(ips) == "userdata" and not err)
assert(#ips == 2 and ips[1]:String() == "127.0.0.1" and ips[2]:String() == "8.8.8.8")
assert(matcher:Match(ips[1]) and not matcher:Match(ips[2]))
assert(matcher:AnyMatch(ips))
local matched, unmatched = matcher:FilterIPs(ips)
assert(#matched == 1 and #unmatched == 1)
assert(matched[1]:Equal(ips[1]) and unmatched[1]:Equal(ips[2]))
return matched, ttl, err
end
`); err != nil {
t.Fatal(err)
}
L.SetContext(context.Background())
if err := callLuaQuery(L, "example.com", option); err != nil {
t.Fatal(err)
}
got, ttl, err := readLuaQueryResult(L)
if err != nil || ttl != 60 || len(got) != 1 || !got[0].Equal(ips[0]) {
t.Fatalf("server query = %v, TTL %d, %v", got, ttl, err)
}
}
type luaDNSClient struct {
featureDNS.Client
lookup func(string, featureDNS.IPOption) ([]net.IP, uint32, error)
}
func (c *luaDNSClient) LookupIP(domain string, option featureDNS.IPOption) ([]net.IP, uint32, error) {
return c.lookup(domain, option)
}
func TestLuaDNSClientQuery(t *testing.T) {
L := lua.NewState()
defer L.Close()
L.SetContext(context.Background())
geodata.RegisterLua(L)
want := []net.IP{{127, 0, 0, 1}, net.ParseIP("::1")}
client := &luaDNSClient{lookup: func(domain string, option featureDNS.IPOption) ([]net.IP, uint32, error) {
if domain != "MiXeD.Example." || !option.IPv4Enable || option.IPv6Enable || !option.FakeEnable {
t.Fatalf("dns.Query arguments = %q, %+v", domain, option)
}
return want, 42, nil
}}
RegisterLua(L, client)
if err := L.DoString(`
local dns = require("xray.dns")
local matcher = require("xray.geodata").BuildIPMatcher("127.0.0.1")
assert(dns.Servers == nil)
ips, ttl, err = dns.Query("MiXeD.Example.", true, false, true)
assert(not err and ttl == 42 and matcher:AnyMatch(ips))
assert(#ips == 2 and ips[1]:String() == "127.0.0.1" and ips[2]:String() == "::1")
assert(matcher:Match(ips[1]) and not matcher:Match(ips[2]))
`); err != nil {
t.Fatal(err)
}
got := L.GetGlobal("ips").(*lua.LUserData).Value.([]net.IP)
if &got[0] != &want[0] {
t.Fatal("dns.Query copied the IP slice")
}
}
func TestLuaDNSLocalClient(t *testing.T) {
L := lua.NewState()
defer L.Close()
L.SetContext(context.Background())
RegisterLua(L, localdns.New())
if err := L.DoString(`
local dns = require("xray.dns")
assert(dns.Servers[1].ID == "localhost")
serverIPs, _, serverErr = dns.Servers[1]:Query("127.0.0.1", true, false, false)
clientIPs, _, clientErr = dns.Query("127.0.0.1", true, false, false)
assert(not serverErr and not clientErr)
assert(#serverIPs == 1 and #clientIPs == 1)
assert(serverIPs[1]:String() == "127.0.0.1" and serverIPs[1]:Equal(clientIPs[1]))
`); err != nil {
t.Fatal(err)
}
for _, name := range []string{"serverIPs", "clientIPs"} {
ips := L.GetGlobal(name).(*lua.LUserData).Value.([]net.IP)
if len(ips) != 1 || !ips[0].Equal(net.ParseIP("127.0.0.1")) {
t.Fatalf("%s = %v", name, ips)
}
}
}
func TestLuaDNSQueryEmptyIPs(t *testing.T) {
for _, tc := range []struct {
name string
ips []net.IP
}{
{"nil", nil},
{"empty", []net.IP{}},
} {
t.Run(tc.name, func(t *testing.T) {
L := lua.NewState()
defer L.Close()
L.SetContext(context.Background())
L.SetGlobal("expectNil", lua.LBool(tc.ips == nil))
client := &luaDNSClient{lookup: func(string, featureDNS.IPOption) ([]net.IP, uint32, error) {
return tc.ips, 0, featureDNS.ErrEmptyResponse
}}
registerLua(L, []luaDNSServer{{query: func(_ context.Context, domain string, option featureDNS.IPOption) ([]net.IP, uint32, error) {
return client.LookupIP(domain, option)
}}}, client)
if err := L.DoString(`
local dns = require("xray.dns")
for _, query in ipairs({
function() return dns.Servers[1]:Query("empty.example", true, false, false) end,
function() return dns.Query("empty.example", true, false, false) end,
}) do
local ips, ttl, err = query()
assert(ttl == 0 and err)
if expectNil then
assert(ips == nil)
else
assert(type(ips) == "userdata" and #ips == 0)
assert(not pcall(function() return ips[1] end))
end
end
`); err != nil {
t.Fatal(err)
}
})
}
}
type benchmarkLuaNameServer struct {
ips []net.IP
}
func (*benchmarkLuaNameServer) Name() string { return "benchmark" }
func (*benchmarkLuaNameServer) IsDisableCache() bool { return true }
func (s *benchmarkLuaNameServer) QueryIP(context.Context, string, featureDNS.IPOption) ([]net.IP, uint32, error) {
return s.ips, 60, nil
}
+1 -2
View File
@@ -29,7 +29,6 @@ type Server interface {
// Client is the interface for DNS client.
type Client struct {
id string
server Server
skipFallback bool
expectedIPs geodata.IPMatcher
@@ -98,7 +97,7 @@ func NewClient(
ipOption dns.IPOption,
updateRules func(bool),
) (*Client, error) {
client := &Client{id: ns.Id}
client := &Client{}
err := core.RequireFeatures(ctx, func(dispatcher routing.Dispatcher) error {
// Create a new server for each client for now
server, err := NewServer(ctx, ns.Address.AsDestination(), dispatcher, disableCache, serveStale, serveExpiredTTL, clientIP)
+1 -1
View File
@@ -49,5 +49,5 @@ func NewLocalNameServer() *LocalNameServer {
// NewLocalDNSClient creates localdns client object for directly lookup in system DNS.
func NewLocalDNSClient(ipOption dns.IPOption) *Client {
return &Client{id: "localhost", server: NewLocalNameServer(), ipOption: &ipOption}
return &Client{server: NewLocalNameServer(), ipOption: &ipOption}
}
-63
View File
@@ -1,63 +0,0 @@
package dns
import (
"time"
"github.com/xtls/xray-core/common/errors"
"github.com/xtls/xray-core/common/geodata"
"github.com/xtls/xray-core/common/log"
xlua "github.com/xtls/xray-core/common/lua"
"github.com/xtls/xray-core/common/net"
"github.com/xtls/xray-core/features/dns"
lua "github.com/yuin/gopher-lua"
)
const scriptExecutionTimeout = 6 * time.Second
type scriptEngine struct {
pool *xlua.Pool
}
func newScriptEngine(path string, server *DNS) (*scriptEngine, error) {
program, err := xlua.CompileFile(path)
if err != nil {
return nil, err
}
pool, err := xlua.NewPool(server.ctx, scriptExecutionTimeout, program.NewStateFactory(
scriptExecutionTimeout*20,
func(L *lua.LState) {
geodata.RegisterLua(L)
log.RegisterLua(L)
server.registerLua(L)
},
func(L *lua.LState) error {
if L.GetGlobal("HandleDNSQuery").Type() != lua.LTFunction {
return errors.New("DNS script must define HandleDNSQuery(...)")
}
return nil
}))
if err != nil {
return nil, err
}
errors.LogInfo(server.ctx, "DNS script initialized from ", path)
return &scriptEngine{pool: pool}, nil
}
func (e *scriptEngine) close() {
e.pool.Close()
}
func (e *scriptEngine) query(domain string, option dns.IPOption) (ips []net.IP, ttl uint32, queryErr error) {
if err := e.pool.WithState(nil, 0, func(L *lua.LState) error {
if err := callLuaQuery(L, domain, option); err != nil {
return err
}
ips, ttl, queryErr = readLuaQueryResult(L)
return nil
}); err != nil {
return nil, 0, err
}
return ips, ttl, queryErr
}
-280
View File
@@ -1,280 +0,0 @@
package dns
import (
"context"
go_errors "errors"
"os"
"path/filepath"
"strings"
"testing"
"time"
"github.com/xtls/xray-core/common/net"
featureDNS "github.com/xtls/xray-core/features/dns"
)
type scriptNameServer struct {
name string
answers map[string]net.IP
errors map[string]error
ttl uint32
calls int
}
func (s *scriptNameServer) Name() string { return s.name }
func (s *scriptNameServer) IsDisableCache() bool { return true }
func (s *scriptNameServer) QueryIP(ctx context.Context, domain string, _ featureDNS.IPOption) ([]net.IP, uint32, error) {
if err := ctx.Err(); err != nil {
return nil, 0, err
}
s.calls++
if err := s.errors[domain]; err != nil {
return nil, 0, err
}
ip, ok := s.answers[domain]
if !ok {
return nil, 0, featureDNS.ErrEmptyResponse
}
return []net.IP{ip}, s.ttl, nil
}
func TestDNSScriptQuery(t *testing.T) {
wantIP := net.ParseIP("127.0.0.1")
upstreamErr := go_errors.New("upstream failed")
for _, tc := range []struct {
name, body string
wantIPs []net.IP
wantTTL uint32
wantErr error
wantMessage string
wantCalls uint32
}{
{name: "IPs", body: `return server:Query(domain, ipv4, ipv6, fake)`, wantIPs: []net.IP{wantIP}, wantTTL: 60, wantCalls: 2},
{name: "empty result", body: `return nil, 0`, wantErr: featureDNS.ErrEmptyResponse, wantCalls: 2},
{name: "upstream error", body: `return server:Query("failed.example", ipv4, ipv6, fake)`, wantErr: upstreamErr, wantCalls: 2},
{name: "string error", body: `return nil, nil, "blocked"`, wantMessage: "blocked", wantCalls: 2},
{name: "invalid result", body: `return false, 0`, wantMessage: "native IP slice", wantCalls: 2},
{name: "execution error", body: `error("execution failed")`, wantMessage: "execution failed", wantCalls: 1},
} {
t.Run(tc.name, func(t *testing.T) {
script := `
local server = require("xray.dns").Servers[1]
local calls = 0
function HandleDNSQuery(domain, ipv4, ipv6, fake)
calls = calls + 1
if domain == "count.example" then
local ips, _, err = server:Query("good.example", ipv4, ipv6, fake)
return ips, calls, err
end
` + tc.body + `
end
`
path := filepath.Join(t.TempDir(), "query.lua")
if err := os.WriteFile(path, []byte(script), 0o600); err != nil {
t.Fatal(err)
}
option := featureDNS.IPOption{IPv4Enable: true}
upstream := &scriptNameServer{
name: "test",
answers: map[string]net.IP{"good.example": wantIP},
errors: map[string]error{"failed.example": upstreamErr},
ttl: 60,
}
server := &DNS{
ctx: context.Background(),
clients: []*Client{{server: upstream, ipOption: &option, timeoutMs: time.Second}},
}
engine, err := newScriptEngine(path, server)
if err != nil {
t.Fatal(err)
}
defer engine.close()
ips, ttl, err := engine.query("good.example", option)
switch {
case tc.wantErr != nil:
if err != tc.wantErr {
t.Fatalf("query error = %v, want original error %v", err, tc.wantErr)
}
case tc.wantMessage != "":
if err == nil || !strings.Contains(err.Error(), tc.wantMessage) {
t.Fatalf("query error = %v, want %q", err, tc.wantMessage)
}
case err != nil:
t.Fatal(err)
}
if ttl != tc.wantTTL || len(ips) != len(tc.wantIPs) {
t.Fatalf("query = %v, TTL %d; want %v, TTL %d", ips, ttl, tc.wantIPs, tc.wantTTL)
}
for i := range ips {
if !ips[i].Equal(tc.wantIPs[i]) {
t.Fatalf("IP %d = %v, want %v", i, ips[i], tc.wantIPs[i])
}
}
ips, calls, err := engine.query("count.example", option)
if err != nil || calls != tc.wantCalls || len(ips) != 1 || !ips[0].Equal(wantIP) {
t.Fatalf("next query = %v, calls %d, %v; want %v, calls %d", ips, calls, err, wantIP, tc.wantCalls)
}
})
}
}
func TestDNSScriptGeoIPFallback(t *testing.T) {
t.Setenv("xray.location.asset", filepath.Join("..", "..", "resources"))
script := `
local servers = require("xray.dns").Servers
local us_ips = require("xray.geodata").BuildIPMatcher("geoip:us")
local by_id = {}
for _, server in ipairs(servers) do
by_id[server.ID] = server
end
assert(by_id.primary and by_id.fallback, "primary and fallback DNS servers are required")
function HandleDNSQuery(domain, ipv4, ipv6, fake)
local ips, ttl, err = by_id.primary:Query(domain, ipv4, ipv6, fake)
if not err and us_ips:AnyMatch(ips) then
return ips, ttl, nil
end
return by_id.fallback:Query(domain, ipv4, ipv6, fake)
end
`
scriptPath := filepath.Join(t.TempDir(), "geoip_fallback.lua")
if err := os.WriteFile(scriptPath, []byte(script), 0o600); err != nil {
t.Fatal(err)
}
primary := &scriptNameServer{
name: "primary",
answers: map[string]net.IP{
"us.example": net.ParseIP("2001:4860:4860::8888"),
"other.example": net.ParseIP("127.0.0.1"),
},
ttl: 30,
}
fallback := &scriptNameServer{
name: "fallback",
answers: map[string]net.IP{"other.example": net.ParseIP("9.9.9.9")},
ttl: 60,
}
option := featureDNS.IPOption{IPv4Enable: true, IPv6Enable: true}
hosts, err := NewStaticHosts(nil)
if err != nil {
t.Fatal(err)
}
server := &DNS{
ctx: context.Background(),
hosts: hosts,
ipOption: &option,
scriptPath: scriptPath,
clients: []*Client{
{id: "primary", server: primary, ipOption: &option, timeoutMs: 2 * time.Second},
{id: "fallback", server: fallback, ipOption: &option, timeoutMs: 2 * time.Second},
},
}
if err := server.Start(); err != nil {
t.Fatal(err)
}
defer server.Close()
for _, tc := range []struct {
domain string
ip net.IP
ttl uint32
}{
{"Us.Example.", net.ParseIP("2001:4860:4860::8888"), 30},
{"other.example", net.ParseIP("9.9.9.9"), 60},
} {
ips, ttl, err := server.LookupIP(tc.domain, option)
if err != nil {
t.Fatalf("LookupIP(%q): %v", tc.domain, err)
}
if ttl != tc.ttl || len(ips) != 1 || !ips[0].Equal(tc.ip) {
t.Fatalf("LookupIP(%q) = %v, TTL %d; want %v, TTL %d", tc.domain, ips, ttl, tc.ip, tc.ttl)
}
}
if primary.calls != 2 || fallback.calls != 1 {
t.Fatalf("upstream calls: primary %d, fallback %d; want 2 and 1", primary.calls, fallback.calls)
}
}
func TestDNSScriptRejectsInvalidStartup(t *testing.T) {
for _, tc := range []struct {
name string
script string
}{
{"syntax", "function HandleDNSQuery("},
{"missing hook", "value = 1"},
{"top-level error", `error("setup failed")`},
} {
t.Run(tc.name, func(t *testing.T) {
path := filepath.Join(t.TempDir(), "script.lua")
if err := os.WriteFile(path, []byte(tc.script), 0o600); err != nil {
t.Fatal(err)
}
server := &DNS{ctx: context.Background(), scriptPath: path}
if err := server.Start(); err == nil {
t.Fatal("Start accepted an invalid DNS script")
}
if server.script != nil {
t.Fatal("Start retained a script engine after failure")
}
})
}
}
func TestDNSScriptFakeDNSOption(t *testing.T) {
path := filepath.Join(t.TempDir(), "script.lua")
script := `
local server = require("xray.dns").Servers[1]
local log = require("xray.log")
log.Info("DNS script loaded")
function HandleDNSQuery(domain, ipv4, ipv6, fake)
log.Debug("DNS query: ", domain)
local ips, ttl, err = server:Query(domain, ipv4, ipv6, fake)
if err then log.Error("DNS failed: ", err) end
return ips, ttl, err
end
`
if err := os.WriteFile(path, []byte(script), 0o600); err != nil {
t.Fatal(err)
}
option := featureDNS.IPOption{IPv4Enable: true}
hosts, err := NewStaticHosts(nil)
if err != nil {
t.Fatal(err)
}
upstream := &scriptNameServer{
name: "FakeDNS",
answers: map[string]net.IP{"good.example": net.ParseIP("198.18.0.1")},
ttl: 30,
}
server := &DNS{
ctx: context.Background(),
hosts: hosts,
ipOption: &option,
scriptPath: path,
clients: []*Client{{id: "fake", server: upstream, ipOption: &option, timeoutMs: time.Second}},
}
if err := server.Start(); err != nil {
t.Fatal(err)
}
defer server.Close()
if _, _, err := server.LookupIP("good.example", option); err != featureDNS.ErrEmptyResponse {
t.Fatalf("FakeDNS without FakeEnable = %v, want ErrEmptyResponse", err)
}
if upstream.calls != 0 {
t.Fatalf("FakeDNS was queried without FakeEnable: %d calls", upstream.calls)
}
withFake := featureDNS.IPOption{IPv4Enable: true, FakeEnable: true}
ips, ttl, err := server.LookupIP("good.example", withFake)
if err != nil || ttl != 30 || len(ips) != 1 || !ips[0].Equal(net.ParseIP("198.18.0.1")) {
t.Fatalf("FakeDNS with FakeEnable = %v, TTL %d, %v", ips, ttl, err)
}
if upstream.calls != 1 {
t.Fatalf("FakeDNS query count = %d, want 1", upstream.calls)
}
}
+4 -14
View File
@@ -587,10 +587,8 @@ type Config struct {
DomainStrategy Config_DomainStrategy `protobuf:"varint,1,opt,name=domain_strategy,json=domainStrategy,proto3,enum=xray.app.router.Config_DomainStrategy" json:"domain_strategy,omitempty"`
Rule []*RoutingRule `protobuf:"bytes,2,rep,name=rule,proto3" json:"rule,omitempty"`
BalancingRule []*BalancingRule `protobuf:"bytes,3,rep,name=balancing_rule,json=balancingRule,proto3" json:"balancing_rule,omitempty"`
// Absolute path to the Lua routing script.
Script string `protobuf:"bytes,4,opt,name=script,proto3" json:"script,omitempty"`
unknownFields protoimpl.UnknownFields
sizeCache protoimpl.SizeCache
unknownFields protoimpl.UnknownFields
sizeCache protoimpl.SizeCache
}
func (x *Config) Reset() {
@@ -644,13 +642,6 @@ func (x *Config) GetBalancingRule() []*BalancingRule {
return nil
}
func (x *Config) GetScript() string {
if x != nil {
return x.Script
}
return ""
}
var File_app_router_config_proto protoreflect.FileDescriptor
const file_app_router_config_proto_rawDesc = "" +
@@ -708,12 +699,11 @@ const file_app_router_config_proto_rawDesc = "" +
"\tbaselines\x18\x03 \x03(\x03R\tbaselines\x12\x1a\n" +
"\bexpected\x18\x04 \x01(\x05R\bexpected\x12\x16\n" +
"\x06maxRTT\x18\x05 \x01(\x03R\x06maxRTT\x12\x1c\n" +
"\ttolerance\x18\x06 \x01(\x02R\ttolerance\"\xae\x02\n" +
"\ttolerance\x18\x06 \x01(\x02R\ttolerance\"\x96\x02\n" +
"\x06Config\x12O\n" +
"\x0fdomain_strategy\x18\x01 \x01(\x0e2&.xray.app.router.Config.DomainStrategyR\x0edomainStrategy\x120\n" +
"\x04rule\x18\x02 \x03(\v2\x1c.xray.app.router.RoutingRuleR\x04rule\x12E\n" +
"\x0ebalancing_rule\x18\x03 \x03(\v2\x1e.xray.app.router.BalancingRuleR\rbalancingRule\x12\x16\n" +
"\x06script\x18\x04 \x01(\tR\x06script\"B\n" +
"\x0ebalancing_rule\x18\x03 \x03(\v2\x1e.xray.app.router.BalancingRuleR\rbalancingRule\"B\n" +
"\x0eDomainStrategy\x12\b\n" +
"\x04AsIs\x10\x00\x12\x10\n" +
"\fIpIfNonMatch\x10\x02\x12\x0e\n" +
-2
View File
@@ -110,6 +110,4 @@ message Config {
DomainStrategy domain_strategy = 1;
repeated RoutingRule rule = 2;
repeated BalancingRule balancing_rule = 3;
// Absolute path to the Lua routing script.
string script = 4;
}
-175
View File
@@ -1,175 +0,0 @@
package router
import (
"runtime"
"strings"
"github.com/xtls/xray-core/common/errors"
xlua "github.com/xtls/xray-core/common/lua"
"github.com/xtls/xray-core/common/net"
"github.com/xtls/xray-core/features/routing"
lua "github.com/yuin/gopher-lua"
)
const (
luaContextType = "xray.router.Context"
luaAttributesType = "xray.router.Attributes"
)
// RegisterLua makes xray.router available to routing scripts.
func (r *Router) RegisterLua(L *lua.LState) {
registerLuaContext(L)
L.PreloadModule("xray.router", func(L *lua.LState) int {
module := L.CreateTable(0, 7)
module.RawSetString("NetworkUnknown", lua.LNumber(net.Network_Unknown))
module.RawSetString("NetworkTCP", lua.LNumber(net.Network_TCP))
module.RawSetString("NetworkUDP", lua.LNumber(net.Network_UDP))
module.RawSetString("NetworkUNIX", lua.LNumber(net.Network_UNIX))
module.RawSetString("LocalOS", lua.LString(runtime.GOOS))
module.RawSetString("PickOutbound", L.NewFunction(func(L *lua.LState) int {
tag, ok := L.Get(2).(lua.LString)
if !ok {
L.ArgError(2, "balancer tag must be a string")
return 0
}
balancer, found := (*r.balancers.Load())[string(tag)]
if !found {
xlua.PushNil(L)
xlua.PushError(L, errors.New("balancer ", tag, " not found"))
return 2
}
outboundTag, err := balancer.PickOutbound()
xlua.PushString(L, outboundTag)
xlua.PushError(L, err)
return 2
}))
module.RawSetString("FindProcess", L.NewFunction(func(L *lua.LState) int {
pid, name, path, err := findProcess(checkLuaContext(L), net.FindProcess)
xlua.PushNumber(L, pid)
xlua.PushString(L, name)
xlua.PushString(L, path)
xlua.PushError(L, err)
return 4
}))
L.Push(module)
return 1
})
}
func registerLuaContext(L *lua.LState) {
pushIPs := xlua.NewSlicePusher[net.IP](L)
attributes := L.NewTypeMetatable(luaAttributesType)
L.SetField(attributes, "__index", L.NewFunction(func(L *lua.LState) int {
values := L.CheckUserData(1).Value.(map[string]string)
key := L.CheckString(2)
if value, found := values[key]; found {
xlua.PushString(L, value)
} else {
xlua.PushNil(L)
}
return 1
}))
methods := L.CreateTable(0, 4)
L.SetFuncs(methods, map[string]lua.LGFunction{
"GetSourceIPs": func(L *lua.LState) int {
pushIPs(L, checkLuaContext(L).GetSourceIPs())
return 1
},
"GetTargetIPs": func(L *lua.LState) int {
pushIPs(L, checkLuaContext(L).GetTargetIPs())
return 1
},
"GetLocalIPs": func(L *lua.LState) int {
pushIPs(L, checkLuaContext(L).GetLocalIPs())
return 1
},
"GetAttributes": func(L *lua.LState) int {
values := L.NewUserData()
values.Value = checkLuaContext(L).GetAttributes()
L.SetMetatable(values, attributes)
L.Push(values)
return 1
},
})
L.SetField(L.NewTypeMetatable(luaContextType), "__index", methods)
}
func checkLuaContext(L *lua.LState) routing.Context {
ctx, ok := L.CheckUserData(1).Value.(routing.Context)
if !ok {
L.ArgError(1, "routing context expected")
}
return ctx
}
// callLuaRoute runs HandleRoute and leaves (outboundTag, ruleTag, err) on the stack.
func callLuaRoute(L *lua.LState, ctx routing.Context) error {
fn := L.GetGlobal("HandleRoute")
if fn.Type() != lua.LTFunction {
return errors.New("routing script must define HandleRoute(...)")
}
value := L.NewUserData()
value.Value = ctx
L.SetMetatable(value, L.GetTypeMetatable(luaContextType))
return L.CallByParam(lua.P{Fn: fn, NRet: 3, Protect: true},
value,
lua.LString(ctx.GetInboundTag()),
lua.LNumber(ctx.GetSourcePort()),
lua.LNumber(ctx.GetTargetPort()),
lua.LNumber(ctx.GetLocalPort()),
lua.LString(strings.ToLower(ctx.GetTargetDomain())),
lua.LNumber(ctx.GetNetwork()),
lua.LString(ctx.GetProtocol()),
lua.LString(ctx.GetUser()),
lua.LNumber(ctx.GetVlessRoute()),
lua.LBool(ctx.GetSkipDNSResolve()))
}
// readLuaRouteResult reads (outboundTag, ruleTag, err) from the stack.
func readLuaRouteResult(L *lua.LState) (string, string, error) {
if err := xlua.ReadError(L.Get(-1), "routing script error must be an error or string"); err != nil {
return "", "", err
}
outboundTag, err := xlua.ReadOptionalString(L.Get(-3), "routing script outboundTag must be a string or nil")
if err != nil || outboundTag == "" {
return "", "", err
}
ruleTag, err := xlua.ReadOptionalString(L.Get(-2), "routing script ruleTag must be a string")
if err != nil {
return "", "", err
}
return outboundTag, ruleTag, nil
}
type processFinder func(string, string, uint16, string, uint16) (int, string, string, error)
func findProcess(ctx routing.Context, finder processFinder) (int, string, string, error) {
sources := ctx.GetSourceIPs()
if len(sources) == 0 {
return 0, "", "", errors.New("process lookup requires a source IP")
}
var network string
switch ctx.GetNetwork() {
case net.Network_TCP:
network = "tcp"
case net.Network_UDP:
network = "udp"
default:
return 0, "", "", errors.New("process lookup requires TCP or UDP")
}
targetIP, targetPort := "", uint16(0)
if targets := ctx.GetTargetIPs(); len(targets) > 0 {
targetIP, targetPort = targets[0].String(), uint16(ctx.GetTargetPort())
}
return finder(network, sources[0].String(), uint16(ctx.GetSourcePort()), targetIP, targetPort)
}
-223
View File
@@ -1,223 +0,0 @@
package router
import (
"context"
"fmt"
"os"
"path/filepath"
"strings"
"testing"
"github.com/xtls/xray-core/common/geodata"
"github.com/xtls/xray-core/common/net"
"github.com/xtls/xray-core/common/session"
"github.com/xtls/xray-core/features/routing"
routing_session "github.com/xtls/xray-core/features/routing/session"
lua "github.com/yuin/gopher-lua"
)
func benchmarkRouteContext(target net.Destination) *routing_session.Context {
// Use the production context: its IP getters construct a slice per call.
// The cached IP slices in luaRouteTestContext would undercount this cost.
return &routing_session.Context{
Inbound: &session.Inbound{
Tag: "in",
Source: net.TCPDestination(net.LocalHostIP, 1234),
Local: net.TCPDestination(net.LocalHostIP, 5678),
},
Outbound: &session.Outbound{Target: target},
Content: &session.Content{Protocol: "tls"},
}
}
func benchmarkRouteState(b *testing.B, r *Router, script string) *lua.LState {
b.Helper()
L := lua.NewState()
b.Cleanup(L.Close)
r.RegisterLua(L)
geodata.RegisterLua(L)
if err := L.DoString(script); err != nil {
b.Fatal(err)
}
L.SetContext(context.Background())
return L
}
// BenchmarkLuaRouteHook isolates argument bridging and a fixed-return hook.
// It excludes rules, result decoding, the state pool, and Route construction.
func BenchmarkLuaRouteHook(b *testing.B) {
L := benchmarkRouteState(b, new(Router), `
function HandleRoute(ctx, inboundTag, sourcePort, targetPort, localPort,
targetDomain, network, protocol, user, vlessRoute, skipDNSResolve)
return "out", "rule"
end
`)
ctx := benchmarkRouteContext(net.TCPDestination(net.LocalHostIP, 443))
if err := callLuaRoute(L, ctx); err != nil {
b.Fatal(err)
}
outboundTag, ruleTag, err := readLuaRouteResult(L)
L.Pop(3)
if err != nil || outboundTag != "out" || ruleTag != "rule" {
b.Fatalf("hook() = %q, %q, %v", outboundTag, ruleTag, err)
}
b.ReportAllocs()
b.ResetTimer()
for i := 0; i < b.N; i++ {
if err := callLuaRoute(L, ctx); err != nil {
b.Fatal(err)
}
L.Pop(3)
}
}
// BenchmarkLuaRoute compares equivalent ordered rules on the same session.
// rules returns tags only; pick_route uses Router.PickRoute on both sides.
// All compilation, matcher construction, and pool startup are outside timing.
func BenchmarkLuaRoute(b *testing.B) {
for _, name := range []string{"scalar", "ip", "domain", "domain_32_last"} {
b.Run(name, func(b *testing.B) {
config, script, ctx, wantTag, wantRule := benchmarkRouteFixture(b, name)
native := new(Router)
if err := native.Init(context.Background(), config, nil, nil, nil); err != nil {
b.Fatal(err)
}
L := benchmarkRouteState(b, native, script)
path := filepath.Join(b.TempDir(), "route.lua")
if err := os.WriteFile(path, []byte(script), 0o600); err != nil {
b.Fatal(err)
}
scripted := new(Router)
if err := scripted.Init(context.Background(), &Config{Script: path}, nil, nil, nil); err != nil {
b.Fatal(err)
}
if err := scripted.Start(); err != nil {
b.Fatal(err)
}
b.Cleanup(func() {
if err := scripted.Close(); err != nil {
b.Error(err)
}
})
for _, bench := range []struct {
name string
route func() (string, string, error)
}{
{"rules/native", func() (string, string, error) {
rule, _, err := native.pickRouteInternal(ctx)
if err != nil {
return "", "", err
}
tag, err := rule.GetTag()
return tag, rule.RuleTag, err
}},
{"rules/lua", func() (string, string, error) {
if err := callLuaRoute(L, ctx); err != nil {
return "", "", err
}
tag, ruleTag, err := readLuaRouteResult(L)
L.Pop(3)
return tag, ruleTag, err
}},
{"pick_route/native", func() (string, string, error) {
return benchmarkPickRoute(native, ctx)
}},
{"pick_route/lua", func() (string, string, error) {
return benchmarkPickRoute(scripted, ctx)
}},
} {
b.Run(bench.name, func(b *testing.B) {
// Validate and warm both paths before measuring steady state.
tag, ruleTag, err := bench.route()
if err != nil || tag != wantTag || ruleTag != wantRule {
b.Fatalf("route() = %q, %q, %v; want %q, %q", tag, ruleTag, err, wantTag, wantRule)
}
b.ReportAllocs()
b.ResetTimer()
for i := 0; i < b.N; i++ {
tag, ruleTag, err = bench.route()
if err != nil {
b.Fatal(err)
}
}
b.StopTimer()
if tag != wantTag || ruleTag != wantRule {
b.Fatalf("route() = %q, %q; want %q, %q", tag, ruleTag, wantTag, wantRule)
}
})
}
})
}
}
func benchmarkPickRoute(r *Router, ctx routing.Context) (string, string, error) {
route, err := r.PickRoute(ctx)
if err != nil {
return "", "", err
}
return route.GetOutboundTag(), route.GetRuleTag(), nil
}
func benchmarkRouteFixture(b *testing.B, name string) (*Config, string, routing.Context, string, string) {
b.Helper()
config := new(Config)
ctx := benchmarkRouteContext(net.TCPDestination(net.LocalHostIP, 443))
prelude := `local router = require("xray.router")
local geodata = require("xray.geodata")
`
body := `if inboundTag == "in" and network == router.NetworkTCP then return "out", "rule" end`
wantTag, wantRule := "out", "rule"
if name == "scalar" || name == "ip" {
rule := &RoutingRule{
TargetTag: &RoutingRule_Tag{Tag: wantTag},
RuleTag: wantRule,
InboundTag: []string{"in"},
Networks: []net.Network{net.Network_TCP},
}
if name == "ip" {
var err error
rule.Ip, err = geodata.ParseIPRules([]string{"127.0.0.0/8"})
if err != nil {
b.Fatal(err)
}
prelude += `local matcher = geodata.BuildIPMatcher("127.0.0.0/8")` + "\n"
body = `if inboundTag == "in" and network == router.NetworkTCP and matcher:AnyMatch(ctx:GetTargetIPs()) then return "out", "rule" end`
}
config.Rule = []*RoutingRule{rule}
} else {
count := 1
if name == "domain_32_last" {
count = 32
}
var rules strings.Builder
rules.WriteString("local rules = {\n")
for i := 0; i < count; i++ {
domain := fmt.Sprintf("route-%d.example.com", i)
tag, ruleTag := fmt.Sprintf("out-%d", i), fmt.Sprintf("rule-%d", i)
domains, err := geodata.ParseDomainRules([]string{"full:" + domain}, geodata.Domain_Domain)
if err != nil {
b.Fatal(err)
}
config.Rule = append(config.Rule, &RoutingRule{
TargetTag: &RoutingRule_Tag{Tag: tag}, RuleTag: ruleTag, Domain: domains,
})
fmt.Fprintf(&rules, "{geodata.BuildDomainMatcher(%q), %q, %q},\n", "full:"+domain, tag, ruleTag)
if i == count-1 {
ctx.Outbound.Target = net.TCPDestination(net.DomainAddress(domain), 443)
wantTag, wantRule = tag, ruleTag
}
}
rules.WriteString("}\n")
prelude += rules.String()
body = `for i = 1, #rules do
local rule = rules[i]
if rule[1]:MatchAny(targetDomain) then return rule[2], rule[3] end
end`
}
script := prelude + `function HandleRoute(ctx, inboundTag, sourcePort, targetPort, localPort,
targetDomain, network, protocol, user, vlessRoute, skipDNSResolve)
` + body + "\nend\n"
return config, script, ctx, wantTag, wantRule
}
-274
View File
@@ -1,274 +0,0 @@
package router
import (
"context"
go_errors "errors"
"runtime"
"strings"
"testing"
"github.com/xtls/xray-core/common/geodata"
"github.com/xtls/xray-core/common/net"
"github.com/xtls/xray-core/common/protocol"
"github.com/xtls/xray-core/common/session"
"github.com/xtls/xray-core/features/routing"
routing_session "github.com/xtls/xray-core/features/routing/session"
lua "github.com/yuin/gopher-lua"
)
type luaRouteTestContext struct {
*routing_session.Context
sourceIPs, targetIPs, localIPs []net.IP
}
func (c *luaRouteTestContext) GetSourceIPs() []net.IP { return c.sourceIPs }
func (c *luaRouteTestContext) GetTargetIPs() []net.IP { return c.targetIPs }
func (c *luaRouteTestContext) GetLocalIPs() []net.IP { return c.localIPs }
func newLuaRouteTestContext() *luaRouteTestContext {
return &luaRouteTestContext{
Context: &routing_session.Context{
Inbound: &session.Inbound{
Tag: "in", VlessRoute: 4321,
Source: net.TCPDestination(net.LocalHostIP, 1234),
Local: net.TCPDestination(net.LocalHostIP, 5678),
User: &protocol.MemoryUser{Email: "user@example.com"},
},
Outbound: &session.Outbound{
Target: net.TCPDestination(net.LocalHostIP, 443),
RouteTarget: net.TCPDestination(net.DomainAddress("MiXeD.Example."), 443),
},
Content: &session.Content{
Protocol: "tls", Attributes: map[string]string{"key": "value"}, SkipDNSResolve: true,
},
},
sourceIPs: []net.IP{{127, 0, 0, 2}},
targetIPs: []net.IP{{127, 0, 0, 3}},
localIPs: []net.IP{{127, 0, 0, 1}},
}
}
func newLuaRouterState(t *testing.T, script string) *lua.LState {
t.Helper()
r := new(Router)
if err := r.Init(context.Background(), &Config{}, nil, nil, nil); err != nil {
t.Fatal(err)
}
L := lua.NewState()
t.Cleanup(L.Close)
r.RegisterLua(L)
geodata.RegisterLua(L)
if err := L.DoString(script); err != nil {
t.Fatal(err)
}
return L
}
func TestLuaRouteBinding(t *testing.T) {
L := newLuaRouterState(t, `
local router = require("xray.router")
local matcher = require("xray.geodata").BuildIPMatcher("127.0.0.0/8")
assert(router.NetworkUnknown == 0 and router.NetworkTCP == 2)
assert(router.NetworkUDP == 3 and router.NetworkUNIX == 4)
assert(router.BuildIPMatcher == nil and router.BuildDomainMatcher == nil)
function HandleRoute(ctx, inboundTag, sourcePort, targetPort, localPort,
targetDomain, network, protocol, user, vlessRoute, skipDNSResolve, ...)
assert(select("#", ...) == 0)
assert(inboundTag == "in" and sourcePort == 1234 and targetPort == 443 and localPort == 5678)
assert(targetDomain == "mixed.example." and network == router.NetworkTCP)
assert(protocol == "tls" and user == "user@example.com" and vlessRoute == 4321 and skipDNSResolve)
assert(ctx.GetNetwork == nil and ctx.Context == nil)
savedContext = ctx
sourceIPs, targetIPs, localIPs = ctx:GetSourceIPs(), ctx:GetTargetIPs(), ctx:GetLocalIPs()
attributes = ctx:GetAttributes()
assert(#sourceIPs == 1 and #targetIPs == 1 and #localIPs == 1)
assert(sourceIPs[1]:String() == "127.0.0.2" and targetIPs[1]:String() == "127.0.0.3")
assert(localIPs[1]:String() == "127.0.0.1")
assert(matcher:Match(sourceIPs[1]) and matcher:Match(targetIPs[1]) and matcher:Match(localIPs[1]))
assert(matcher:AnyMatch(sourceIPs) and matcher:AnyMatch(targetIPs) and matcher:AnyMatch(localIPs))
local matched = matcher:FilterIPs(targetIPs)
assert(#matched == 1 and matched[1]:Equal(targetIPs[1]))
assert(attributes.key == "value" and attributes.missing == nil)
assert(not pcall(function() attributes.key = "changed" end))
return "out", "rule"
end`)
ctx := newLuaRouteTestContext()
if err := callLuaRoute(L, ctx); err != nil {
t.Fatal(err)
}
if L.GetTop() != 3 || L.Get(1) != lua.LString("out") || L.Get(2) != lua.LString("rule") || L.Get(3) != lua.LNil {
t.Fatal("callLuaRoute did not leave the three route results on the stack")
}
if L.GetGlobal("savedContext").(*lua.LUserData).Value != ctx {
t.Fatal("routing context was copied")
}
for _, tc := range []struct {
name string
want []net.IP
}{
{"sourceIPs", ctx.sourceIPs},
{"targetIPs", ctx.targetIPs},
{"localIPs", ctx.localIPs},
} {
got := L.GetGlobal(tc.name).(*lua.LUserData).Value.([]net.IP)
if &got[0] != &tc.want[0] {
t.Fatalf("%s storage was copied", tc.name)
}
}
ctx.Content.Attributes["key"] = "updated"
L.SetGlobal("expectedOS", lua.LString(runtime.GOOS))
if err := L.DoString(`
assert(attributes.key == "updated")
assert(require("xray.router").LocalOS == expectedOS)`); err != nil {
t.Fatal(err)
}
}
func TestLuaRouteEmptyIPs(t *testing.T) {
for _, tc := range []struct {
name string
ips []net.IP
}{
{"nil", nil},
{"empty", []net.IP{}},
} {
t.Run(tc.name, func(t *testing.T) {
L := newLuaRouterState(t, `
function HandleRoute(ctx)
for _, name in ipairs({"GetSourceIPs", "GetTargetIPs", "GetLocalIPs"}) do
local ips = ctx[name](ctx)
if expectNil then
assert(ips == nil)
else
assert(type(ips) == "userdata" and #ips == 0)
assert(not pcall(function() return ips[1] end))
end
end
return "out"
end
`)
L.SetGlobal("expectNil", lua.LBool(tc.ips == nil))
ctx := newLuaRouteTestContext()
ctx.sourceIPs, ctx.targetIPs, ctx.localIPs = tc.ips, tc.ips, tc.ips
if err := callLuaRoute(L, ctx); err != nil {
t.Fatal(err)
}
})
}
}
func TestReadLuaRouteResult(t *testing.T) {
nativeErr := go_errors.New("native failure")
for _, tc := range []struct {
name, values string
wantTag, wantRule string
wantErr error
wantMessage string
}{
{name: "route", values: `"out", "rule"`, wantTag: "out", wantRule: "rule"},
{name: "no match", values: `nil`},
{name: "empty tag", values: `""`},
{name: "no match ignores rule", values: `nil, false`},
{name: "empty tag ignores rule", values: `"", false`},
{name: "missing rule", values: `"out"`, wantTag: "out"},
{name: "invalid tag", values: `1`, wantMessage: "outboundTag"},
{name: "invalid rule", values: `"out", false`, wantMessage: "ruleTag"},
{name: "string error", values: `nil, nil, "script failure"`, wantMessage: "script failure"},
{name: "native error", values: `nil, nil, nativeError`, wantErr: nativeErr},
{name: "error overrides invalid tags", values: `false, false, nativeError`, wantErr: nativeErr},
{name: "invalid error", values: `"out", "rule", false`, wantMessage: "error or string"},
{name: "wrong error userdata", values: `"out", "rule", wrongError`, wantMessage: "error or string"},
} {
t.Run(tc.name, func(t *testing.T) {
L := lua.NewState()
defer L.Close()
for name, value := range map[string]any{"nativeError": nativeErr, "wrongError": "not a native error"} {
ud := L.NewUserData()
ud.Value = value
L.SetGlobal(name, ud)
}
fn, err := L.LoadString("return " + tc.values)
if err != nil {
t.Fatal(err)
}
if err := L.CallByParam(lua.P{Fn: fn, NRet: 3, Protect: true}); err != nil {
t.Fatal(err)
}
outboundTag, ruleTag, err := readLuaRouteResult(L)
if outboundTag != tc.wantTag || ruleTag != tc.wantRule {
t.Fatalf("result = %q, %q, %v; want %q, %q", outboundTag, ruleTag, err, tc.wantTag, tc.wantRule)
}
switch {
case tc.wantErr != nil:
if err != tc.wantErr {
t.Fatalf("error = %v, want original error", err)
}
case tc.wantMessage != "":
if err == nil || !strings.Contains(err.Error(), tc.wantMessage) {
t.Fatalf("error = %v, want %q", err, tc.wantMessage)
}
case err != nil:
t.Fatal(err)
}
})
}
}
func TestCallLuaRouteCancellation(t *testing.T) {
L := newLuaRouterState(t, `function HandleRoute() while true do end end`)
ctx, cancel := context.WithCancel(context.Background())
cancel()
L.SetContext(ctx)
if err := callLuaRoute(L, &routing_session.Context{}); err == nil {
t.Fatal("callLuaRoute did not stop after context cancellation")
}
if L.Context() != ctx {
t.Fatal("callLuaRoute changed the Lua state's context")
}
}
func TestFindProcess(t *testing.T) {
for _, tc := range []struct {
name, network, target string
targetPort uint16
modify func(*luaRouteTestContext)
wantErr bool
}{
{name: "TCP", network: "tcp", target: "127.0.0.3", targetPort: 443},
{name: "UDP", network: "udp", target: "127.0.0.3", targetPort: 443, modify: func(c *luaRouteTestContext) {
c.Outbound.Target.Network = net.Network_UDP
}},
{name: "domain target", network: "tcp", modify: func(c *luaRouteTestContext) { c.targetIPs = nil }},
{name: "missing source", modify: func(c *luaRouteTestContext) { c.sourceIPs = nil }, wantErr: true},
{name: "unsupported network", modify: func(c *luaRouteTestContext) {
c.Outbound.Target.Network = net.Network_UNIX
}, wantErr: true},
} {
t.Run(tc.name, func(t *testing.T) {
ctx := newLuaRouteTestContext()
if tc.modify != nil {
tc.modify(ctx)
}
called := false
pid, name, path, err := findProcess(ctx, func(network, source string, sourcePort uint16, target string, targetPort uint16) (int, string, string, error) {
called = true
if network != tc.network || source != "127.0.0.2" || sourcePort != 1234 || target != tc.target || targetPort != tc.targetPort {
t.Fatalf("endpoints = %s %s:%d -> %s:%d", network, source, sourcePort, target, targetPort)
}
return 42, "process", "/path/process", nil
})
if tc.wantErr {
if err == nil || called {
t.Fatalf("findProcess = %d, %q, %q, %v", pid, name, path, err)
}
return
}
if err != nil || !called || pid != 42 || name != "process" || path != "/path/process" {
t.Fatalf("findProcess = %d, %q, %q, %v", pid, name, path, err)
}
})
}
}
var _ routing.Context = (*luaRouteTestContext)(nil)
-17
View File
@@ -20,8 +20,6 @@ import (
type Router struct {
domainStrategy Config_DomainStrategy
rules atomic.Pointer[[]*Rule]
scriptPath string
script *scriptEngine
balancers atomic.Pointer[map[string]*Balancer]
dns dns.Client
@@ -42,7 +40,6 @@ type Route struct {
// Init initializes the Router.
func (r *Router) Init(ctx context.Context, config *Config, d dns.Client, ohm outbound.Manager, dispatcher routing.Dispatcher) error {
r.domainStrategy = config.DomainStrategy
r.scriptPath = config.Script
r.dns = d
r.ctx = ctx
r.ohm = ohm
@@ -55,10 +52,6 @@ func (r *Router) Init(ctx context.Context, config *Config, d dns.Client, ohm out
// PickRoute implements routing.Router.
func (r *Router) PickRoute(ctx routing.Context) (routing.Route, error) {
if r.script != nil {
return r.script.pickRoute(ctx)
}
originalCtx := ctx
rule, ctx, err := r.pickRouteInternal(ctx)
if err != nil {
@@ -228,13 +221,6 @@ func (r *Router) pickRouteInternal(ctx routing.Context) (*Rule, routing.Context,
// Start implements common.Runnable.
func (r *Router) Start() error {
if r.scriptPath != "" {
engine, err := newScriptEngine(r.scriptPath, r)
if err != nil {
return errors.New("failed to initialize routing script").Base(err)
}
r.script = engine
}
return nil
}
@@ -249,9 +235,6 @@ func closeWebhooks(rules []*Rule) {
// Close implements common.Closable.
func (r *Router) Close() error {
if r.script != nil {
r.script.close()
}
r.mu.Lock()
defer r.mu.Unlock()
closeWebhooks(*r.rules.Load())
-76
View File
@@ -1,76 +0,0 @@
package router
import (
"time"
"github.com/xtls/xray-core/app/dns"
"github.com/xtls/xray-core/common"
"github.com/xtls/xray-core/common/errors"
"github.com/xtls/xray-core/common/geodata"
"github.com/xtls/xray-core/common/log"
xlua "github.com/xtls/xray-core/common/lua"
"github.com/xtls/xray-core/features/routing"
lua "github.com/yuin/gopher-lua"
)
const scriptExecutionTimeout = 6 * time.Second
type scriptEngine struct {
pool *xlua.Pool
}
func newScriptEngine(path string, router *Router) (*scriptEngine, error) {
program, err := xlua.CompileFile(path)
if err != nil {
return nil, err
}
pool, err := xlua.NewPool(router.ctx, scriptExecutionTimeout, program.NewStateFactory(
scriptExecutionTimeout*20,
func(L *lua.LState) {
geodata.RegisterLua(L)
log.RegisterLua(L)
router.RegisterLua(L)
dns.RegisterLua(L, router.dns)
},
func(L *lua.LState) error {
if L.GetGlobal("HandleRoute").Type() != lua.LTFunction {
return errors.New("routing script must define HandleRoute(...)")
}
return nil
}))
if err != nil {
return nil, err
}
errors.LogInfo(router.ctx, "routing script initialized from ", path)
return &scriptEngine{pool: pool}, nil
}
func (e *scriptEngine) close() {
e.pool.Close()
}
func (e *scriptEngine) pickRoute(ctx routing.Context) (routing.Route, error) {
var outboundTag, ruleTag string
var routeErr error
if err := e.pool.WithState(nil, 0, func(L *lua.LState) error {
if err := callLuaRoute(L, ctx); err != nil {
return err
}
outboundTag, ruleTag, routeErr = readLuaRouteResult(L)
return nil
}); err != nil {
return nil, err
}
if routeErr != nil {
return nil, routeErr
}
if outboundTag == "" {
return nil, common.ErrNoClue
}
return &Route{Context: ctx, outboundTag: outboundTag, ruleTag: ruleTag}, nil
}
-382
View File
@@ -1,382 +0,0 @@
package router
import (
"context"
stdnet "net"
"os"
"path/filepath"
"strings"
"sync"
"sync/atomic"
"testing"
"time"
wireDNS "github.com/miekg/dns"
"github.com/xtls/xray-core/app/dispatcher"
appdns "github.com/xtls/xray-core/app/dns"
"github.com/xtls/xray-core/app/proxyman"
_ "github.com/xtls/xray-core/app/proxyman/outbound"
"github.com/xtls/xray-core/common"
"github.com/xtls/xray-core/common/net"
"github.com/xtls/xray-core/common/serial"
"github.com/xtls/xray-core/core"
featureDNS "github.com/xtls/xray-core/features/dns"
"github.com/xtls/xray-core/features/outbound"
"github.com/xtls/xray-core/features/routing"
routing_session "github.com/xtls/xray-core/features/routing/session"
"github.com/xtls/xray-core/proxy/blackhole"
"github.com/xtls/xray-core/proxy/freedom"
)
type luaRouteDNSClient struct {
featureDNS.Client
lookup func(string, featureDNS.IPOption) ([]net.IP, uint32, error)
}
func (d *luaRouteDNSClient) LookupIP(domain string, option featureDNS.IPOption) ([]net.IP, uint32, error) {
return d.lookup(domain, option)
}
type luaRouteOutboundManager struct{ outbound.Manager }
func (*luaRouteOutboundManager) Select(selectors []string) []string { return selectors }
func writeRouteScript(t *testing.T, script string) string {
t.Helper()
path := filepath.Join(t.TempDir(), "route.lua")
if err := os.WriteFile(path, []byte(script), 0o600); err != nil {
t.Fatal(err)
}
return path
}
func startLuaRouter(t *testing.T, script string, d featureDNS.Client, config *Config) *Router {
t.Helper()
if config == nil {
config = &Config{}
}
config.Script = writeRouteScript(t, script)
r := new(Router)
if err := r.Init(context.Background(), config, d, &luaRouteOutboundManager{}, nil); err != nil {
t.Fatal(err)
}
if err := r.Start(); err != nil {
t.Fatal(err)
}
t.Cleanup(func() {
if err := r.Close(); err != nil {
t.Error(err)
}
})
return r
}
func TestRouterScriptStartup(t *testing.T) {
for _, tc := range []struct{ name, script string }{
{"syntax error", "function HandleRoute("},
{"missing hook", "value = 1"},
{"initialization error", `error("setup failed")`},
} {
t.Run(tc.name, func(t *testing.T) {
r := new(Router)
if err := r.Init(context.Background(), &Config{Script: writeRouteScript(t, tc.script)}, nil, nil, nil); err != nil {
t.Fatal(err)
}
defer r.Close()
if err := r.Start(); err == nil {
t.Fatal("Start accepted an invalid routing script")
}
})
}
}
func TestRouterScriptRouting(t *testing.T) {
for _, tc := range []struct {
name, body string
wantTag, wantRule string
wantErr error
wantMessage string
wantCalls string
}{
{name: "route", body: `return "lua-out", "lua-rule"`, wantTag: "lua-out", wantRule: "lua-rule", wantCalls: "2"},
{name: "no match", body: `return nil`, wantErr: common.ErrNoClue, wantCalls: "2"},
{name: "empty tag", body: `return ""`, wantErr: common.ErrNoClue, wantCalls: "2"},
{name: "balancer error", body: `local tag, err = router:PickOutbound("missing"); return tag, nil, err`, wantMessage: "not found", wantCalls: "2"},
{name: "string error", body: `return nil, nil, "blocked"`, wantMessage: "blocked", wantCalls: "2"},
{name: "invalid tag", body: `return false`, wantMessage: "outboundTag", wantCalls: "2"},
{name: "invalid rule", body: `return "lua-out", false`, wantMessage: "ruleTag", wantCalls: "2"},
{name: "execution error", body: `error("execution failed")`, wantMessage: "execution failed", wantCalls: "1"},
} {
t.Run(tc.name, func(t *testing.T) {
var dnsCalls atomic.Int32
d := &luaRouteDNSClient{lookup: func(string, featureDNS.IPOption) ([]net.IP, uint32, error) {
dnsCalls.Add(1)
return []net.IP{{1, 2, 3, 4}}, 60, nil
}}
script := `
local router = require("xray.router")
local calls = 0
function HandleRoute(ctx, inbound)
calls = calls + 1
if inbound == "count" then return "lua-out", tostring(calls) end
` + tc.body + `
end
`
r := startLuaRouter(t, script, d, &Config{
DomainStrategy: Config_IpOnDemand,
Rule: []*RoutingRule{{
TargetTag: &RoutingRule_Tag{Tag: "json-out"},
Networks: []net.Network{net.Network_TCP},
}},
})
ctx := newLuaRouteTestContext()
ctx.Content.SkipDNSResolve = false
route, err := r.PickRoute(ctx)
switch {
case tc.wantErr != nil:
if err != tc.wantErr {
t.Fatalf("route error = %v, want %v", err, tc.wantErr)
}
case tc.wantMessage != "":
if err == nil || !strings.Contains(err.Error(), tc.wantMessage) {
t.Fatalf("route error = %v, want %q", err, tc.wantMessage)
}
case err != nil:
t.Fatal(err)
}
if tc.wantTag == "" {
if route != nil {
t.Fatalf("route = %v, want nil", route)
}
} else if route == nil || route.GetOutboundTag() != tc.wantTag || route.GetRuleTag() != tc.wantRule || route.(*Route).Context != ctx {
t.Fatalf("route = %v; want %q, %q and original context", route, tc.wantTag, tc.wantRule)
}
ctx.Inbound.Tag = "count"
route, err = r.PickRoute(ctx)
if err != nil || route == nil || route.GetOutboundTag() != "lua-out" || route.GetRuleTag() != tc.wantCalls {
t.Fatalf("next route = %v, %v; want lua-out, calls %s", route, err, tc.wantCalls)
}
if dnsCalls.Load() != 0 {
t.Fatal("script routing implicitly resolved DNS")
}
})
}
}
func TestRouterScriptModules(t *testing.T) {
ips := []net.IP{{127, 0, 0, 7}}
calls := 0
d := &luaRouteDNSClient{lookup: func(domain string, option featureDNS.IPOption) ([]net.IP, uint32, error) {
calls++
if domain != "mixed.example." || !option.IPv4Enable || option.IPv6Enable || !option.FakeEnable {
t.Fatalf("dns.Query arguments = %q, %+v", domain, option)
}
return ips, 17, nil
}}
r := startLuaRouter(t, `
local dns = require("xray.dns")
local matcher = require("xray.geodata").BuildIPMatcher("127.0.0.0/8")
assert(dns.Servers == nil and type(dns.Query) == "function")
assert(type(require("xray.log").Info) == "function")
function HandleRoute(ctx, inbound, sourcePort, targetPort, localPort, domain)
local ips, ttl, err = dns.Query(domain, true, false, true)
assert(not err and ttl == 17)
assert(matcher:AnyMatch(ips) and matcher:AnyMatch(ctx:GetTargetIPs()))
return "out"
end`, d, nil)
if _, err := r.PickRoute(newLuaRouteTestContext()); err != nil {
t.Fatal(err)
}
if calls != 1 {
t.Fatalf("DNS calls = %d, want 1", calls)
}
}
func TestRouterScriptBalancerReload(t *testing.T) {
config := func(tag string) *Config {
return &Config{BalancingRule: []*BalancingRule{{
Tag: "balance", Strategy: "roundrobin", OutboundSelector: []string{tag},
}}}
}
r := startLuaRouter(t, `
local router = require("xray.router")
function HandleRoute()
local tag, err = router:PickOutbound("balance")
return tag, "balanced", err
end`, nil, config("old"))
pick := func(want string) {
t.Helper()
route, err := r.PickRoute(&routing_session.Context{})
if err != nil || route.GetOutboundTag() != want || route.GetRuleTag() != "balanced" {
t.Fatalf("route = %v, %v, want %q", route, err, want)
}
}
pick("old")
if err := r.SetOverrideTarget("balance", "override"); err != nil {
t.Fatal(err)
}
pick("override")
if err := r.SetOverrideTarget("balance", ""); err != nil {
t.Fatal(err)
}
if err := r.ReloadRules(config("new"), false); err != nil {
t.Fatal(err)
}
pick("new")
}
func TestRouterScriptConcurrentBalancerReload(t *testing.T) {
config := func(tag string) *Config {
return &Config{BalancingRule: []*BalancingRule{{
Tag: "balance", Strategy: "roundrobin", OutboundSelector: []string{tag},
}}}
}
r := startLuaRouter(t, `
local router = require("xray.router")
function HandleRoute()
local tag, err = router:PickOutbound("balance")
return tag, nil, err
end`, nil, config("a"))
var wg sync.WaitGroup
for range 4 {
wg.Go(func() {
for range 20 {
route, err := r.PickRoute(&routing_session.Context{})
if err != nil {
t.Errorf("PickRoute: %v", err)
return
}
if tag := route.GetOutboundTag(); tag != "a" && tag != "b" {
t.Errorf("unexpected tag %q", tag)
}
}
})
}
wg.Go(func() {
for range 20 {
for _, tag := range []string{"a", "b"} {
if err := r.ReloadRules(config(tag), false); err != nil {
t.Error(err)
return
}
}
}
})
wg.Wait()
}
func TestRouterScriptDNSDispatcherReentry(t *testing.T) {
conn, err := stdnet.ListenPacket("udp4", "127.0.0.1:0")
if err != nil {
t.Fatal(err)
}
port := conn.LocalAddr().(*stdnet.UDPAddr).Port
ready, stopped := make(chan struct{}), make(chan error, 1)
var queries atomic.Int32
server := &wireDNS.Server{
PacketConn: conn,
NotifyStartedFunc: func() {
close(ready)
},
Handler: wireDNS.HandlerFunc(func(w wireDNS.ResponseWriter, query *wireDNS.Msg) {
queries.Add(1)
response := new(wireDNS.Msg).SetReply(query)
for _, question := range query.Question {
if question.Name == "nested.example." && question.Qtype == wireDNS.TypeA {
response.Answer = append(response.Answer, &wireDNS.A{
Hdr: wireDNS.RR_Header{Name: question.Name, Rrtype: wireDNS.TypeA, Class: wireDNS.ClassINET, Ttl: 60},
A: stdnet.IP{127, 0, 0, 7},
})
}
}
if err := w.WriteMsg(response); err != nil {
t.Error(err)
}
}),
}
go func() { stopped <- server.ActivateAndServe() }()
defer func() {
server.Shutdown()
select {
case err := <-stopped:
if err != nil {
t.Error(err)
}
case <-time.After(3 * time.Second):
t.Error("DNS server did not stop")
}
}()
select {
case <-ready:
case err := <-stopped:
t.Fatalf("DNS server startup: %v", err)
case <-time.After(3 * time.Second):
t.Fatal("DNS server did not start")
}
dnsScript := writeRouteScript(t, `
local server = require("xray.dns").Servers[1]
function HandleDNSQuery(domain, ipv4, ipv6, fake)
return server:Query(domain, ipv4, ipv6, fake)
end`)
routerScript := writeRouteScript(t, `
local router = require("xray.router")
local dns = require("xray.dns")
local matcher = require("xray.geodata").BuildIPMatcher("127.0.0.7")
local active = false
function HandleRoute(ctx, inbound, sourcePort, targetPort, localPort, domain, network,
protocol, user, vlessRoute, skipDNSResolve)
assert(not active, "borrowed Router VM reentered")
if inbound == "dns" then
assert(network == router.NetworkUDP and skipDNSResolve == false)
return "direct", "dns-route"
end
active = true
local ips, ttl, err = dns.Query("nested.example", true, false, false)
assert(not err and matcher:AnyMatch(ips) and active)
active = false
return "direct", "outer-route"
end`)
instance, err := core.New(&core.Config{
App: []*serial.TypedMessage{
serial.ToTypedMessage(&appdns.Config{
Tag: "dns", Script: dnsScript, DisableCache: true,
NameServer: []*appdns.NameServer{{
Id: "upstream", TimeoutMs: 1000,
Address: &net.Endpoint{
Network: net.Network_UDP,
Address: &net.IPOrDomain{Address: &net.IPOrDomain_Ip{Ip: []byte{127, 0, 0, 1}}},
Port: uint32(port),
},
}},
}),
serial.ToTypedMessage(&Config{Script: routerScript}),
serial.ToTypedMessage(&dispatcher.Config{}),
serial.ToTypedMessage(&proxyman.OutboundConfig{}),
},
Outbound: []*core.OutboundHandlerConfig{
{Tag: "default", ProxySettings: serial.ToTypedMessage(&blackhole.Config{})},
{Tag: "direct", ProxySettings: serial.ToTypedMessage(&freedom.Config{
FinalRules: []*freedom.FinalRuleConfig{{Action: freedom.RuleAction_Allow}},
})},
},
})
if err != nil {
t.Fatal(err)
}
defer instance.Close()
if err := instance.Start(); err != nil {
t.Fatal(err)
}
r := instance.GetFeature(routing.RouterType()).(*Router)
route, err := r.PickRoute(newLuaRouteTestContext())
if err != nil || route.GetOutboundTag() != "direct" || route.GetRuleTag() != "outer-route" {
t.Fatalf("nested DNS routing = %v, %v", route, err)
}
if queries.Load() == 0 {
t.Fatal("DNS query did not pass through the dispatcher")
}
}
+50 -33
View File
@@ -82,10 +82,19 @@ func (f *MphDomainMatcherFactory) BuildMatcher(rules []*DomainRule) (DomainMatch
}
g.Add(m, uint32(i))
case *DomainRule_Geosite:
err := loadSiteMatchers(v.Geosite, func(m strmatcher.Matcher) { g.Add(m, uint32(i)) })
domains, err := loadSiteWithAttrs(v.Geosite.File, v.Geosite.Code, v.Geosite.Attrs)
if err != nil {
return nil, err
}
for j, d := range domains {
domains[j] = nil // peak mem
m, err := parseDomain(d)
if err != nil {
errors.LogError(context.Background(), "ignore invalid geosite entry in ", v.Geosite.File, ":", v.Geosite.Code, " at index ", j, ", ", err)
continue
}
g.Add(m, uint32(i))
}
default:
panic("unknown domain rule type")
}
@@ -99,12 +108,12 @@ func (f *MphDomainMatcherFactory) BuildMatcher(rules []*DomainRule) (DomainMatch
return g, nil
}
type CompactMphDomainMatcherFactory struct {
type CompactDomainMatcherFactory struct {
sync.Mutex
shared *utils.WeakCacheMap[string, strmatcher.MphValueMatcher]
shared *utils.WeakCacheMap[string, strmatcher.LinearAnyMatcher]
}
func (f *CompactMphDomainMatcherFactory) getOrCreateFrom(rule *GeoSiteRule) (*strmatcher.MphValueMatcher, error) {
func (f *CompactDomainMatcherFactory) getOrCreateFrom(rule *GeoSiteRule) (strmatcher.MatcherSet, error) {
key := rule.File + ":" + rule.Code + "@" + rule.Attrs
f.Lock()
@@ -116,23 +125,33 @@ func (f *CompactMphDomainMatcherFactory) getOrCreateFrom(rule *GeoSiteRule) (*st
}
errors.LogDebug(context.Background(), "geodata geosite matcher cache MISS ", key)
s := strmatcher.NewMphValueMatcher()
if err := loadSiteMatchers(rule, func(m strmatcher.Matcher) { s.Add(m, 0) }); err != nil {
s := strmatcher.NewLinearAnyMatcher()
domains, err := loadSiteWithAttrs(rule.File, rule.Code, rule.Attrs)
if err != nil {
return nil, err
}
if err := s.Build(); err != nil {
return nil, err
for i, d := range domains {
domains[i] = nil // peak mem
m, err := parseDomain(d)
if err != nil {
errors.LogError(context.Background(), "ignore invalid geosite entry in ", rule.File, ":", rule.Code, " at index ", i, ", ", err)
continue
}
s.Add(m)
}
f.shared.Store(key, s)
return s, nil
return s, err
}
// BuildMatcher implements DomainMatcherFactory.
func (f *CompactMphDomainMatcherFactory) BuildMatcher(rules []*DomainRule) (DomainMatcher, error) {
func (f *CompactDomainMatcherFactory) BuildMatcher(rules []*DomainRule) (DomainMatcher, error) {
if len(rules) == 0 {
return nil, errors.New("empty domain rule list")
}
compact := new(CompactMphDomainMatcher)
compact := &CompactDomainMatcher{
matchers: make([]strmatcher.MatcherSet, 0, len(rules)),
values: make([]uint32, 0, len(rules)),
}
for i, r := range rules {
switch v := r.Value.(type) {
case *DomainRule_Custom:
@@ -149,7 +168,8 @@ func (f *CompactMphDomainMatcherFactory) BuildMatcher(rules []*DomainRule) (Doma
if err != nil {
return nil, err
}
compact.combiner.Add(m, uint32(i))
compact.matchers = append(compact.matchers, m)
compact.values = append(compact.values, uint32(i))
default:
panic("unknown domain rule type")
}
@@ -157,40 +177,37 @@ func (f *CompactMphDomainMatcherFactory) BuildMatcher(rules []*DomainRule) (Doma
return compact, nil
}
type CompactMphDomainMatcher struct {
type CompactDomainMatcher struct {
custom strmatcher.ValueMatcher
combiner strmatcher.MphValueMatcherCombiner
matchers []strmatcher.MatcherSet
values []uint32
}
// Match implements DomainMatcher.
func (c *CompactMphDomainMatcher) Match(input string) []uint32 {
result := c.combiner.Match(input)
func (c *CompactDomainMatcher) Match(input string) []uint32 {
var result []uint32
if c.custom != nil {
result = append(c.custom.Match(input), result...)
result = append(result, c.custom.Match(input)...)
}
for i, m := range c.matchers {
if m.MatchAny(input) {
result = append(result, c.values[i])
}
}
return result
}
// MatchAny implements DomainMatcher.
func (c *CompactMphDomainMatcher) MatchAny(input string) bool {
func (c *CompactDomainMatcher) MatchAny(input string) bool {
if c.custom != nil && c.custom.MatchAny(input) {
return true
}
return c.combiner.MatchAny(input)
}
// loadSiteMatchers calls add with a matcher for every domain of the geosite rule and logs the invalid ones.
func loadSiteMatchers(rule *GeoSiteRule, add func(strmatcher.Matcher)) error {
i := 0
return loadSite(rule.File, rule.Code, rule.Attrs, func(t Domain_Type, value []byte) {
m, err := parseDomain(&Domain{Type: t, Value: string(value)})
if err != nil {
errors.LogError(context.Background(), "ignore invalid geosite entry in ", rule.File, ":", rule.Code, " at index ", i, ", ", err)
} else {
add(m)
for _, m := range c.matchers {
if m.MatchAny(input) {
return true
}
i++
})
}
return false
}
func parseDomain(d *Domain) (strmatcher.Matcher, error) {
@@ -214,7 +231,7 @@ func parseDomain(d *Domain) (strmatcher.Matcher, error) {
func newDomainMatcherFactory() DomainMatcherFactory {
switch runtime.GOOS {
case "ios", "android":
return &CompactMphDomainMatcherFactory{shared: utils.NewWeakCacheMap[string, strmatcher.MphValueMatcher]()}
return &CompactDomainMatcherFactory{shared: utils.NewWeakCacheMap[string, strmatcher.LinearAnyMatcher]()}
default:
return &MphDomainMatcherFactory{shared: utils.NewWeakCacheMap[string, strmatcher.MphValueMatcher]()}
}
+2 -76
View File
@@ -4,7 +4,6 @@ import (
"path/filepath"
"reflect"
"slices"
"sync"
"testing"
"github.com/xtls/xray-core/common/geodata/strmatcher"
@@ -12,7 +11,7 @@ import (
)
func TestCompactDomainMatcher_PreservesCustomRuleIndices(t *testing.T) {
factory := &CompactMphDomainMatcherFactory{shared: utils.NewWeakCacheMap[string, strmatcher.MphValueMatcher]()}
factory := &CompactDomainMatcherFactory{shared: utils.NewWeakCacheMap[string, strmatcher.LinearAnyMatcher]()}
matcher, err := factory.BuildMatcher([]*DomainRule{
{Value: &DomainRule_Custom{Custom: &Domain{Type: Domain_Full, Value: "example.com"}}},
{Value: &DomainRule_Custom{Custom: &Domain{Type: Domain_Domain, Value: "example.com"}}},
@@ -33,7 +32,7 @@ func TestCompactDomainMatcher_PreservesCustomRuleIndices(t *testing.T) {
func TestCompactDomainMatcher_PreservesMixedRuleIndices(t *testing.T) {
t.Setenv("xray.location.asset", filepath.Join("..", "..", "resources"))
factory := &CompactMphDomainMatcherFactory{shared: utils.NewWeakCacheMap[string, strmatcher.MphValueMatcher]()}
factory := &CompactDomainMatcherFactory{shared: utils.NewWeakCacheMap[string, strmatcher.LinearAnyMatcher]()}
matcher, err := factory.BuildMatcher([]*DomainRule{
{Value: &DomainRule_Geosite{Geosite: &GeoSiteRule{File: DefaultGeoSiteDat, Code: "CN"}}},
{Value: &DomainRule_Custom{Custom: &Domain{Type: Domain_Full, Value: "163.com"}}},
@@ -73,76 +72,3 @@ func TestMphDomainMatcher_MatchReturnsDetachedSlice(t *testing.T) {
t.Fatalf("Match() after caller mutation = %v, want %v", gotAgain, []uint32{0, 1})
}
}
// DNS sorts every Match result in place, so a matcher must never hand out a
// slice it keeps, also when only its keyword or regex part matches.
func TestDomainMatcher_MatchResultsCanBeSortedConcurrently(t *testing.T) {
t.Setenv("xray.location.asset", filepath.Join("..", "..", "resources"))
rules := []*DomainRule{
{Value: &DomainRule_Custom{Custom: &Domain{Type: Domain_Full, Value: "example.com"}}},
{Value: &DomainRule_Custom{Custom: &Domain{Type: Domain_Domain, Value: "example.com"}}},
{Value: &DomainRule_Custom{Custom: &Domain{Type: Domain_Substr, Value: "exam"}}},
{Value: &DomainRule_Custom{Custom: &Domain{Type: Domain_Regex, Value: `^ex.*\.org$`}}},
{Value: &DomainRule_Custom{Custom: &Domain{Type: Domain_Substr, Value: "exam"}}},
{Value: &DomainRule_Geosite{Geosite: &GeoSiteRule{File: DefaultGeoSiteDat, Code: "CN"}}},
{Value: &DomainRule_Custom{Custom: &Domain{Type: Domain_Full, Value: "only.full.test"}}},
}
cases := []struct {
input string
want []uint32
}{
{"example.com", []uint32{0, 1, 2, 4}},
{"www.example.com", []uint32{1, 2, 4}},
{"exam.net", []uint32{2, 4}}, // keyword part only
{"example.org", []uint32{2, 3, 4}},
{"163.com", []uint32{5}},
{"www.163.com", []uint32{5}},
{"only.full.test", []uint32{6}}, // full part only
{"nomatch.test", nil},
}
factories := map[string]DomainMatcherFactory{
"mph": &MphDomainMatcherFactory{shared: utils.NewWeakCacheMap[string, strmatcher.MphValueMatcher]()},
"compact": &CompactMphDomainMatcherFactory{shared: utils.NewWeakCacheMap[string, strmatcher.MphValueMatcher]()},
}
for name, factory := range factories {
t.Run(name, func(t *testing.T) {
matcher, err := factory.BuildMatcher(rules)
if err != nil {
t.Fatalf("BuildMatcher() failed: %v", err)
}
for _, c := range cases {
got := matcher.Match(c.input)
if sorted := slices.Sorted(slices.Values(got)); !slices.Equal(sorted, c.want) {
t.Fatalf("Match(%q) = %v, want %v", c.input, sorted, c.want)
}
got = got[:cap(got)]
for j := range got {
got[j] = ^uint32(0)
}
if again := slices.Sorted(slices.Values(matcher.Match(c.input))); !slices.Equal(again, c.want) {
t.Fatalf("Match(%q) after caller mutation = %v, want %v", c.input, again, c.want)
}
}
var wg sync.WaitGroup
for range 8 {
wg.Add(1)
go func() {
defer wg.Done()
for range 500 {
for _, c := range cases {
got := matcher.Match(c.input)
slices.Sort(got)
if !slices.Equal(got, c.want) {
t.Errorf("Match(%q) = %v, want %v", c.input, got, c.want)
return
}
}
}
}()
}
wg.Wait()
})
}
}
+62 -213
View File
@@ -5,14 +5,11 @@ import (
"bytes"
"io"
"runtime"
"slices"
"strings"
"unicode/utf8"
"github.com/xtls/xray-core/common/errors"
"github.com/xtls/xray-core/common/platform/filesystem"
"google.golang.org/protobuf/encoding/protowire"
"google.golang.org/protobuf/proto"
)
@@ -55,56 +52,17 @@ func loadIP(file, code string) ([]*CIDR, error) {
return geoip.Cidr, nil
}
// loadSite calls fn, in file order, with the type and value of every domain of the geosite code
// that has all the "@"-separated attrs. It decodes the entry while reading the file instead of
// unmarshalling it into a []*Domain, so value is only valid during fn.
func loadSite(file, code, attrs string, fn func(Domain_Type, []byte)) error {
runtime.GC() // peak mem
r, err := filesystem.OpenAsset(file)
func loadSite(file, code string) ([]*Domain, error) {
bs, err := loadFile(file, code)
if err != nil {
return errors.New("failed to open ", file).Base(err)
return nil, err
}
defer r.Close()
br := bufio.NewReaderSize(r, 64*1024)
n, err := seek(br, []byte(code))
if err != nil {
return errors.New("failed to load code ", code, " from ", file).Base(err)
defer runtime.GC() // peak mem
var geosite GeoSite
if err := proto.Unmarshal(bs, &geosite); err != nil {
return nil, errors.New("error unmarshal Site in ", file, ":", code).Base(err)
}
loadErr := func(err error) error {
if err == io.EOF {
err = io.ErrUnexpectedEOF
}
return errors.New("failed to load code ", code, " from ", file).Base(err)
}
unmarshalErr := func(err error) error {
return errors.New("error unmarshal Site in ", file, ":", code).Base(err)
}
d := newSiteDecoder(attrs, fn)
for n > 0 {
w, err := br.Peek(min(n, br.Size()))
if err != nil {
return loadErr(err)
}
used, err := d.decode(w, len(w) < n)
if err != nil {
return unmarshalErr(err)
}
if used == 0 {
break // a field longer than the buffer
}
br.Discard(used)
n -= used
}
if n > 0 {
w := make([]byte, n)
if _, err := io.ReadFull(br, w); err != nil {
return loadErr(err)
}
if _, err := d.decode(w, false); err != nil {
return unmarshalErr(err)
}
}
return nil
return geosite.Domain, nil
}
func decodeVarint(br *bufio.Reader) (uint64, error) {
@@ -124,63 +82,68 @@ func decodeVarint(br *bufio.Reader) (uint64, error) {
}
func find(r io.Reader, code []byte, readBody bool) ([]byte, error) {
br := bufio.NewReaderSize(r, 64*1024)
bodyL, err := seek(br, code)
if err != nil || !readBody {
return nil, err
}
out := make([]byte, bodyL)
if _, err := io.ReadFull(br, out); err != nil {
return nil, err
}
return out, nil
}
// seek advances br to the body of the entry for code and returns the body length.
func seek(br *bufio.Reader, code []byte) (int, error) {
codeL := len(code)
if codeL == 0 {
return 0, errors.New("empty code")
return nil, errors.New("empty code")
}
br := bufio.NewReaderSize(r, 64*1024)
need := 2 + codeL // TODO: if code too long
prefixBuf := make([]byte, need)
for {
if _, err := br.ReadByte(); err != nil {
return 0, err
return nil, err
}
x, err := decodeVarint(br)
if err != nil {
return 0, err
return nil, err
}
bodyL := int(x)
if bodyL <= 0 {
return 0, errors.New("invalid body length: ", bodyL)
return nil, errors.New("invalid body length: ", bodyL)
}
// Peek no more than the buffer holds: a code longer than the buffer cannot match a single
// length byte anyway, so a short peek only skips it, as base find (io.ReadFull) does.
prefix, err := br.Peek(min(bodyL, need, br.Size()))
if err != nil {
if err == io.EOF && len(prefix) > 0 {
err = io.ErrUnexpectedEOF // as io.ReadFull
prefixL := bodyL
if prefixL > need {
prefixL = need
}
prefix := prefixBuf[:prefixL]
if _, err := io.ReadFull(br, prefix); err != nil {
return nil, err
}
match := false
if bodyL >= need {
if int(prefix[1]) == codeL && bytes.Equal(prefix[2:need], code) {
if !readBody {
return nil, nil
}
match = true
}
return 0, err
}
if bodyL >= need && len(prefix) >= need && int(prefix[1]) == codeL && bytes.Equal(prefix[2:], code) {
return bodyL, nil
remain := bodyL - prefixL
if match {
out := make([]byte, bodyL)
copy(out, prefix)
if remain > 0 {
if _, err := io.ReadFull(br, out[prefixL:]); err != nil {
return nil, err
}
}
return out, nil
}
if _, err := br.Discard(bodyL); err != nil {
return 0, err
if remain > 0 {
if _, err := br.Discard(remain); err != nil {
return nil, err
}
}
}
}
// AttributeMatcher, HasAttrMatcher, AllAttrsMatcher and NewAllAttrsMatcher are the exported
// attribute helpers that have been part of this package's API since #5814. The streaming loader
// above filters attributes itself without building a *Domain, so it does not use them, but they
// are kept for external callers. Their behaviour is unchanged.
type AttributeMatcher interface {
Match(*Domain) bool
}
@@ -222,137 +185,23 @@ func NewAllAttrsMatcher(attrs string) AttributeMatcher {
return m
}
var errInvalidUTF8 = errors.New("string field contains invalid UTF-8")
func loadSiteWithAttrs(file, code, attrs string) ([]*Domain, error) {
domains, err := loadSite(file, code)
if err != nil {
return nil, err
}
type siteDecoder struct {
want []string
has []bool
fn func(Domain_Type, []byte)
}
matcher := NewAllAttrsMatcher(attrs)
if matcher == nil {
return domains, nil
}
func newSiteDecoder(attrs string, fn func(Domain_Type, []byte)) *siteDecoder {
d := &siteDecoder{fn: fn}
if attrs != "" {
d.want = strings.Split(attrs, "@")
d.has = make([]bool, len(d.want))
filtered := make([]*Domain, 0, len(domains))
for _, d := range domains {
if matcher.Match(d) {
filtered = append(filtered, d)
}
}
return d
}
// decode walks the whole fields at the start of b, a part of an encoded GeoSite (see geodat.proto),
// calls fn for every domain that has all attrs and returns how many bytes it used. A field cut off
// by the end of b is an error unless more is set. It accepts and rejects what proto.Unmarshal does.
func (d *siteDecoder) decode(b []byte, more bool) (int, error) {
used := 0
for used < len(b) {
f, n, err := consumeField(b[used:])
if err == io.ErrUnexpectedEOF && more {
break
}
if err != nil {
return used, err
}
used += n
if f.typ != protowire.BytesType {
continue
}
switch f.num {
case 1: // code
if !utf8.Valid(f.v) {
return used, errInvalidUTF8
}
case 2: // domain
t, value, err := decodeDomain(f.v, d.want, d.has)
if err != nil {
return used, err
}
if !slices.Contains(d.has, false) {
d.fn(t, value)
}
}
}
return used, nil
}
// decodeDomain decodes an encoded Domain and sets has[i] if one of its attributes has the key want[i].
func decodeDomain(b []byte, want []string, has []bool) (t Domain_Type, value []byte, err error) {
clear(has)
for len(b) > 0 {
f, n, err := consumeField(b)
if err != nil {
return 0, nil, err
}
b = b[n:]
switch {
case f.num == 1 && f.typ == protowire.VarintType: // type
t = Domain_Type(f.x)
case f.num == 2 && f.typ == protowire.BytesType: // value
if !utf8.Valid(f.v) {
return 0, nil, errInvalidUTF8
}
value = f.v
case f.num == 3 && f.typ == protowire.BytesType: // attribute
key, err := decodeAttributeKey(f.v)
if err != nil {
return 0, nil, err
}
for i, w := range want {
if string(key) == w {
has[i] = true
}
}
}
}
return t, value, nil
}
// decodeAttributeKey returns the key of an encoded Domain.Attribute.
func decodeAttributeKey(b []byte) ([]byte, error) {
var key []byte
for len(b) > 0 {
f, n, err := consumeField(b)
if err != nil {
return nil, err
}
b = b[n:]
if f.num == 1 && f.typ == protowire.BytesType {
if !utf8.Valid(f.v) {
return nil, errInvalidUTF8
}
key = f.v
}
}
return key, nil
}
type protoField struct {
num protowire.Number
typ protowire.Type
v []byte // payload of a length-delimited field
x uint64 // value of a varint field
}
// consumeField parses the first field of an encoded message and returns it with its length.
func consumeField(b []byte) (protoField, int, error) {
num, typ, n := protowire.ConsumeTag(b)
if n < 0 {
return protoField{}, 0, protowire.ParseError(n)
}
if num > protowire.MaxValidNumber {
return protoField{}, 0, errors.New("invalid field number ", num)
}
f := protoField{num: num, typ: typ}
var m int
switch typ {
case protowire.BytesType:
f.v, m = protowire.ConsumeBytes(b[n:])
case protowire.VarintType:
f.x, m = protowire.ConsumeVarint(b[n:])
default:
m = protowire.ConsumeFieldValue(num, typ, b[n:])
}
if m < 0 {
return protoField{}, 0, protowire.ParseError(m)
}
return f, n + m, nil
return filtered, nil
}
-283
View File
@@ -1,283 +0,0 @@
package geodata
import (
"fmt"
"os"
"path/filepath"
"slices"
"strings"
"testing"
"google.golang.org/protobuf/encoding/protowire"
"google.golang.org/protobuf/proto"
)
type siteEntry struct {
Type Domain_Type
Value string
}
// unmarshalSite is what loadSite used to do: proto.Unmarshal, then keep the domains that have all attrs.
func unmarshalSite(b []byte, attrs string) ([]siteEntry, error) {
var site GeoSite
if err := proto.Unmarshal(b, &site); err != nil {
return nil, err
}
var entries []siteEntry
for _, d := range site.Domain {
ok := true
for _, key := range strings.Split(attrs, "@") {
ok = ok && (attrs == "" || slices.ContainsFunc(d.Attribute, func(a *Domain_Attribute) bool { return a.Key == key }))
}
if ok {
entries = append(entries, siteEntry{d.Type, d.Value})
}
}
return entries, nil
}
func checkDecodeSite(t *testing.T, name string, b []byte, attrs string) {
t.Helper()
want, wantErr := unmarshalSite(b, attrs)
var got []siteEntry
_, err := newSiteDecoder(attrs, func(typ Domain_Type, value []byte) {
got = append(got, siteEntry{typ, string(value)})
}).decode(b, false)
if (err == nil) != (wantErr == nil) {
t.Fatalf("%s@%s: error %v, proto.Unmarshal: %v", name, attrs, err, wantErr)
}
if err == nil && !slices.Equal(got, want) {
t.Fatalf("%s@%s: got %v, want %v", name, attrs, got, want)
}
}
func TestDecodeSiteMatchesUnmarshal(t *testing.T) {
bs, err := os.ReadFile(filepath.Join("..", "..", "resources", DefaultGeoSiteDat))
if err != nil {
t.Fatal(err)
}
for len(bs) > 0 {
num, typ, n := protowire.ConsumeTag(bs)
if n < 0 || num != 1 || typ != protowire.BytesType {
t.Fatal("unexpected GeoSiteList field")
}
entry, m := protowire.ConsumeBytes(bs[n:])
if m < 0 {
t.Fatal(protowire.ParseError(m))
}
bs = bs[n+m:]
var site GeoSite
if err := proto.Unmarshal(entry, &site); err != nil {
t.Fatal(err)
}
queries := []string{"", "none"}
for _, d := range site.Domain {
for _, a := range d.Attribute {
if !slices.Contains(queries, a.Key) {
queries = append(queries, a.Key, a.Key+"@none")
}
}
}
for _, attrs := range queries {
checkDecodeSite(t, site.Code, entry, attrs)
}
}
}
func TestDecodeSiteUnusualEncodings(t *testing.T) {
field := func(num protowire.Number, v []byte) []byte {
return protowire.AppendBytes(protowire.AppendTag(nil, num, protowire.BytesType), v)
}
typ := func(v Domain_Type) []byte {
return protowire.AppendVarint(protowire.AppendTag(nil, 1, protowire.VarintType), uint64(v))
}
value := func(s string) []byte { return field(2, []byte(s)) }
attr := func(keys ...string) []byte {
var b []byte
for _, k := range keys {
b = append(b, field(1, []byte(k))...)
}
return field(3, b)
}
domain := func(fields ...[]byte) []byte { return field(2, slices.Concat(fields...)) }
unknown := protowire.AppendFixed32(protowire.AppendTag(nil, 9, protowire.Fixed32Type), 1)
for name, b := range map[string][]byte{
"unknown field": domain(typ(Domain_Full), unknown, value("example.com")),
"repeated value": domain(value("a.com"), typ(Domain_Full), value("b.com")),
"repeated type": domain(typ(Domain_Full), value("a.com"), typ(Domain_Regex)),
"repeated key": domain(value("a.com"), attr("cn", "ads")),
"type as bytes": domain(field(1, []byte("x")), value("a.com")),
"no value": domain(typ(Domain_Domain), attr("cn")),
"truncated": domain(typ(Domain_Full), value("example.com"))[:10],
"invalid utf8": domain(value("example.\xff")),
"invalid key": domain(value("a.com"), attr("\xff")),
"bad field": protowire.AppendVarint(protowire.AppendTag(nil, protowire.MaxValidNumber+1, protowire.VarintType), 1),
"stray end group": protowire.AppendTag(nil, 5, protowire.EndGroupType),
} {
for _, attrs := range []string{"", "cn", "ads", "cn@ads"} {
checkDecodeSite(t, name, b, attrs)
}
}
}
// TestLoadSiteReadsInPieces covers what real lists never do: an entry far longer than the read
// buffer, with a field longer than the buffer in the middle, and a file cut short.
func TestLoadSiteReadsInPieces(t *testing.T) {
site := &GeoSite{Code: "BIG"}
for i := range 5000 {
d := &Domain{Type: Domain_Domain, Value: strings.Repeat("x", i%40) + ".example.com"}
if i%3 == 0 {
d.Attribute = []*Domain_Attribute{{Key: "cn"}}
}
if i == 2500 {
d = &Domain{Type: Domain_Regex, Value: strings.Repeat("a", 100_000)}
}
site.Domain = append(site.Domain, d)
}
list := &GeoSiteList{Entry: []*GeoSite{{Code: "SMALL", Domain: []*Domain{{Type: Domain_Full, Value: "a.com"}}}, site}}
bs, err := proto.Marshal(list)
if err != nil {
t.Fatal(err)
}
entry, err := proto.Marshal(site)
if err != nil {
t.Fatal(err)
}
dir := t.TempDir()
t.Setenv("xray.location.asset", dir)
write := func(b []byte) {
if err := os.WriteFile(filepath.Join(dir, "big.dat"), b, 0o644); err != nil {
t.Fatal(err)
}
}
for _, attrs := range []string{"", "cn"} {
want, _ := unmarshalSite(entry, attrs)
var got []siteEntry
write(bs)
err := loadSite("big.dat", "BIG", attrs, func(typ Domain_Type, value []byte) {
got = append(got, siteEntry{typ, string(value)})
})
if err != nil || !slices.Equal(got, want) {
t.Fatalf("attrs %q: %d entries, want %d, error %v", attrs, len(got), len(want), err)
}
for _, cut := range []int{30_000, len(bs) - 150_000, len(bs) - 1} {
write(bs[:cut])
if err := loadSite("big.dat", "BIG", attrs, func(Domain_Type, []byte) {}); err == nil {
t.Fatalf("file cut at %d of %d: no error", cut, len(bs))
}
}
}
}
// oneEntryGeoSiteFile wraps an encoded GeoSite as a one-entry GeoSiteList, the file loadSite reads.
func oneEntryGeoSiteFile(entry []byte) []byte {
return protowire.AppendBytes(protowire.AppendTag(nil, 1, protowire.BytesType), entry)
}
// TestLoadSiteWindowedMatchesSingleShot checks that the windowed reader in loadSite (its Peek/Discard
// loop, the more-break when a field is cut by a window edge, the used==0 fallback for a field longer
// than the buffer, and the tail path) reaches exactly the same result as decoding the whole entry at
// once, for a category several 64 KiB windows long, valid and then mutated near a window edge and
// early in the file: same error-or-not, and the same emitted (type, value) sequence when both accept.
func TestLoadSiteWindowedMatchesSingleShot(t *testing.T) {
const window = 64 * 1024
site := &GeoSite{Code: "BIG"}
for i := range 12000 { // ~250 KiB, four windows
d := &Domain{Type: Domain_Domain, Value: fmt.Sprintf("host%d.%s.example.com", i, strings.Repeat("y", i%30))}
if i%3 == 0 {
d.Attribute = []*Domain_Attribute{{Key: "cn"}}
}
site.Domain = append(site.Domain, d)
}
// a field longer than the buffer, straddling the third window, to force the used==0 fallback
site.Domain = slices.Insert(site.Domain, 8000, &Domain{Type: Domain_Regex, Value: strings.Repeat("a", 90_000)})
entry, err := proto.Marshal(site)
if err != nil {
t.Fatal(err)
}
dir := t.TempDir()
t.Setenv("xray.location.asset", dir)
// mutations of the encoded entry: unchanged, a byte flipped at several offsets (early windows and
// either side of a window edge), and truncations at the same places.
type mut struct {
name string
make func([]byte) []byte
}
muts := []mut{{"valid", func(b []byte) []byte { return b }}}
for _, off := range []int{3, 40, 4000, window - 1, window, window + 1, 2*window - 2, 2 * window} {
if off < len(entry) {
off := off
muts = append(muts, mut{fmt.Sprintf("flip@%d", off), func(b []byte) []byte {
c := slices.Clone(b)
c[off] ^= 0xff
return c
}})
muts = append(muts, mut{fmt.Sprintf("cut@%d", off), func(b []byte) []byte { return slices.Clone(b[:off]) }})
}
}
for _, attrs := range []string{"", "cn"} {
for _, m := range muts {
e := m.make(entry)
// single-shot reference: decode the whole entry in one call
var want []siteEntry
_, wantErr := newSiteDecoder(attrs, func(typ Domain_Type, value []byte) {
want = append(want, siteEntry{typ, string(value)})
}).decode(e, false)
// windowed: loadSite reads the file 64 KiB at a time
if err := os.WriteFile(filepath.Join(dir, "w.dat"), oneEntryGeoSiteFile(e), 0o644); err != nil {
t.Fatal(err)
}
var got []siteEntry
gotErr := loadSite("w.dat", "BIG", attrs, func(typ Domain_Type, value []byte) {
got = append(got, siteEntry{typ, string(value)})
})
if (gotErr == nil) != (wantErr == nil) {
t.Fatalf("%s attrs=%q: windowed err %v, single-shot err %v", m.name, attrs, gotErr, wantErr)
}
if gotErr == nil && !slices.Equal(got, want) {
t.Fatalf("%s attrs=%q: windowed got %d entries, single-shot %d", m.name, attrs, len(got), len(want))
}
}
}
}
// TestLoadSiteLongCode covers a geosite entry whose code is longer than the 64 KiB read buffer. seek
// must skip it (find compares a single length byte, so it never matches such a code) and still find a
// later entry, and looking the long code up must fail cleanly, like a missing code, not panic.
func TestLoadSiteLongCode(t *testing.T) {
longCode := strings.Repeat("Z", 70000)
list := &GeoSiteList{Entry: []*GeoSite{
{Code: "FIRST", Domain: []*Domain{{Type: Domain_Full, Value: "first.com"}}},
{Code: longCode, Domain: []*Domain{{Type: Domain_Full, Value: "huge.com"}}},
{Code: "AFTER", Domain: []*Domain{{Type: Domain_Domain, Value: "after.com"}}},
}}
bs, err := proto.Marshal(list)
if err != nil {
t.Fatal(err)
}
dir := t.TempDir()
t.Setenv("xray.location.asset", dir)
if err := os.WriteFile(filepath.Join(dir, "lc.dat"), bs, 0o644); err != nil {
t.Fatal(err)
}
collect := func(code string) ([]siteEntry, error) {
var got []siteEntry
err := loadSite("lc.dat", code, "", func(typ Domain_Type, value []byte) {
got = append(got, siteEntry{typ, string(value)})
})
return got, err
}
if got, err := collect("FIRST"); err != nil || !slices.Equal(got, []siteEntry{{Domain_Full, "first.com"}}) {
t.Fatalf("FIRST: %v %v", got, err)
}
if got, err := collect("AFTER"); err != nil || !slices.Equal(got, []siteEntry{{Domain_Domain, "after.com"}}) {
t.Fatalf("AFTER (past the oversized entry): %v %v", got, err)
}
if _, err := collect(longCode); err == nil {
t.Fatal("oversized code: expected a not-found error, got nil")
}
}
-163
View File
@@ -1,163 +0,0 @@
package geodata
import (
xlua "github.com/xtls/xray-core/common/lua"
"github.com/xtls/xray-core/common/net"
lua "github.com/yuin/gopher-lua"
)
// RegisterLua makes xray.geodata available to require in an LState.
func RegisterLua(L *lua.LState) {
L.PreloadModule("xray.geodata", func(L *lua.LState) int {
module := L.CreateTable(0, 2)
module.RawSetString("BuildDomainMatcher", L.NewFunction(func(L *lua.LState) int {
parsed, err := ParseDomainRules(luaRules(L), Domain_Domain)
if err != nil {
L.RaiseError("%v", err)
return 0
}
matcher, err := DomainReg.BuildDomainMatcher(parsed)
if err != nil {
L.RaiseError("%v", err)
return 0
}
xlua.PushWithDirectMethods(L, matcher, map[string]xlua.DirectMethod{
"Match": newLuaDomainMatch(xlua.NewSlicePusher[uint32](L)),
"MatchAny": luaDomainMatchAny,
})
return 1
}))
module.RawSetString("BuildIPMatcher", L.NewFunction(func(L *lua.LState) int {
parsed, err := ParseIPRules(luaRules(L))
if err != nil {
L.RaiseError("%v", err)
return 0
}
matcher, err := IPReg.BuildIPMatcher(parsed)
if err != nil {
L.RaiseError("%v", err)
return 0
}
xlua.PushWithDirectMethods(L, matcher, map[string]xlua.DirectMethod{
"Match": luaIPMatch,
"AnyMatch": luaIPAnyMatch,
"Matches": luaIPMatches,
"FilterIPs": newLuaIPFilterIPs(xlua.NewSlicePusher[net.IP](L)),
})
return 1
}))
L.Push(module)
return 1
})
}
// Read native Go values by type assertion; slices keep their original storage.
func readLuaIPMatcherArgs[T any](L *lua.LState) (IPMatcher, T, bool) {
var input T
if L.GetTop() != 2 {
return nil, input, false
}
value, ok := L.Get(1).(*lua.LUserData)
if !ok {
return nil, input, false
}
matcher, ok := value.Value.(IPMatcher)
if !ok {
return nil, input, false
}
if L.Get(2) == lua.LNil {
return matcher, input, true
}
value, ok = L.Get(2).(*lua.LUserData)
if !ok {
return nil, input, false
}
input, ok = value.Value.(T)
return matcher, input, ok
}
func luaIPMatch(L *lua.LState) (int, bool) {
matcher, ip, ok := readLuaIPMatcherArgs[net.IP](L)
if !ok {
return 0, false
}
L.Push(lua.LBool(matcher.Match(ip)))
return 1, true
}
func luaIPAnyMatch(L *lua.LState) (int, bool) {
matcher, ips, ok := readLuaIPMatcherArgs[[]net.IP](L)
if !ok {
return 0, false
}
L.Push(lua.LBool(matcher.AnyMatch(ips)))
return 1, true
}
func luaIPMatches(L *lua.LState) (int, bool) {
matcher, ips, ok := readLuaIPMatcherArgs[[]net.IP](L)
if !ok {
return 0, false
}
L.Push(lua.LBool(matcher.Matches(ips)))
return 1, true
}
func newLuaIPFilterIPs(pushIPs func(*lua.LState, []net.IP)) xlua.DirectMethod {
return func(L *lua.LState) (int, bool) {
matcher, ips, ok := readLuaIPMatcherArgs[[]net.IP](L)
if !ok {
return 0, false
}
matched, unmatched := matcher.FilterIPs(ips)
pushIPs(L, matched)
pushIPs(L, unmatched)
return 2, true
}
}
func newLuaDomainMatch(pushMatches func(*lua.LState, []uint32)) xlua.DirectMethod {
return func(L *lua.LState) (int, bool) {
if L.GetTop() == 2 {
if value, ok := L.Get(1).(*lua.LUserData); ok {
matcher, validMatcher := value.Value.(DomainMatcher)
domain, validDomain := L.Get(2).(lua.LString)
if validMatcher && validDomain {
pushMatches(L, matcher.Match(string(domain)))
return 1, true
}
}
}
return 0, false
}
}
func luaDomainMatchAny(L *lua.LState) (int, bool) {
if L.GetTop() == 2 {
if value, ok := L.Get(1).(*lua.LUserData); ok {
matcher, validMatcher := value.Value.(DomainMatcher)
domain, validDomain := L.Get(2).(lua.LString)
if validMatcher && validDomain {
L.Push(lua.LBool(matcher.MatchAny(string(domain))))
return 1, true
}
}
}
return 0, false
}
func luaRules(L *lua.LState) []string {
rules := make([]string, L.GetTop())
for i := range rules {
value, ok := L.Get(i + 1).(lua.LString)
if !ok {
L.RaiseError("geodata rules must be strings")
return nil
}
rules[i] = string(value)
}
return rules
}
-172
View File
@@ -1,172 +0,0 @@
package geodata
import (
"fmt"
"testing"
"github.com/xtls/xray-core/common/net"
lua "github.com/yuin/gopher-lua"
)
func TestLuaIPMatcher(t *testing.T) {
L := lua.NewState()
defer L.Close()
RegisterLua(L)
ip := L.NewUserData()
ip.Value = net.ParseIP("127.0.0.1")
L.SetGlobal("ip", ip)
ips := L.NewUserData()
ips.Value = []net.IP{ip.Value.(net.IP), net.ParseIP("8.8.8.8")}
L.SetGlobal("ips", ips)
if err := L.DoString(`
local matcher = require("xray.geodata").BuildIPMatcher("127.0.0.0/8", "::1")
assert(matcher:Match(ip))
assert(matcher:AnyMatch(ips))
assert(not matcher:Matches(ips))
local matched, unmatched = matcher:FilterIPs(ips)
assert(type(matched) == "userdata" and type(unmatched) == "userdata")
assert(#matched == 1 and #unmatched == 1)
`); err != nil {
t.Fatal(err)
}
}
func TestLuaDomainMatcher(t *testing.T) {
L := lua.NewState()
defer L.Close()
RegisterLua(L)
if err := L.DoString(`
local matcher = require("xray.geodata").BuildDomainMatcher("example.com", "full:other.com")
assert(matcher:MatchAny("example.com"))
assert(matcher:MatchAny("www.example.com"))
assert(matcher:MatchAny("other.com"))
assert(not matcher:MatchAny("www.other.com"))
assert(#(matcher:Match("www.example.com")) == 1)
`); err != nil {
t.Fatal(err)
}
}
func TestLuaMatchersRejectInvalidRules(t *testing.T) {
for _, tc := range []struct {
name string
script string
}{
{"IP rule", `require("xray.geodata").BuildIPMatcher("not-an-ip")`},
{"non-string domain rule", `require("xray.geodata").BuildDomainMatcher("example.com", true)`},
} {
t.Run(tc.name, func(t *testing.T) {
L := lua.NewState()
defer L.Close()
RegisterLua(L)
if err := L.DoString(tc.script); err == nil {
t.Fatal("invalid geodata rule was accepted")
}
})
}
}
func TestLuaMatcherArgumentsAndAliases(t *testing.T) {
L := lua.NewState()
defer L.Close()
RegisterLua(L)
ip := L.NewUserData()
ip.Value = net.ParseIP("127.0.0.1")
L.SetGlobal("ip", ip)
if err := L.DoString(`
local geodata = require("xray.geodata")
local matcher = geodata.BuildIPMatcher("127.0.0.0/8")
assert(matcher.Match == matcher.match and matcher.AnyMatch == matcher.anyMatch)
assert(matcher.Matches == matcher.matches and matcher.FilterIPs == matcher.filterIPs)
assert(matcher:match(ip))
assert(matcher:anyMatch({ip}) and matcher:matches({ip}))
assert(not matcher:AnyMatch(nil))
assert(matcher:Matches(nil) == matcher:Matches({}))
local matched, unmatched = matcher:FilterIPs({ip})
assert(#matched == 1 and matched[1]:Equal(ip))
assert(matcher:AnyMatch(matched) and matcher:Matches(matched))
local filtered, excluded = matcher:filterIPs(matched)
assert(#filtered == 1 and #excluded == 0 and filtered[1]:Equal(ip))
local emptyMatched, emptyUnmatched = matcher:FilterIPs(nil)
assert(#emptyMatched == 0 and #emptyUnmatched == 0)
matcher:SetReverse(true)
assert(not matcher:Match(ip) and not matcher:AnyMatch(matched))
matcher:ToggleReverse()
assert(matcher:Match(ip) and matcher:AnyMatch(matched))
assert(matcher.missing == nil)
local domain = geodata.BuildDomainMatcher("full:example.com")
assert(domain.Match == domain.match and domain.MatchAny == domain.matchAny)
assert(domain:matchAny("example.com"))
assert(#domain:Match("example.com") == 1)
assert(domain:match("example.com")[1] == 0)
assert(not pcall(function() matcher:AnyMatch() end))
assert(not pcall(function() matcher:AnyMatch(matched, true) end))
assert(not pcall(function() matcher.AnyMatch(ip, matched) end))
assert(not pcall(function() matcher:Match(true) end))
assert(not pcall(function() domain:MatchAny(123) end))
assert(not pcall(function() domain:MatchAny("example.com", true) end))
assert(not pcall(function() matcher:FilterIPs(true) end))
assert(not pcall(function() matcher:FilterIPs(matched, true) end))
assert(not pcall(function() domain:Match(123) end))
assert(not pcall(function() domain:Match("example.com", true) end))
`); err != nil {
t.Fatal(err)
}
}
// BenchmarkLuaMatcherCall measures repeated calls with prebuilt matchers and inputs.
func BenchmarkLuaMatcherCall(b *testing.B) {
L := lua.NewState()
defer L.Close()
RegisterLua(L)
ip := net.ParseIP("127.0.0.1")
for name, value := range map[string]any{"ip": ip, "ips": []net.IP{ip}} {
ud := L.NewUserData()
ud.Value = value
L.SetGlobal(name, ud)
}
if err := L.DoString(`
local geodata = require("xray.geodata")
ipMatcher = geodata.BuildIPMatcher("127.0.0.0/8")
domainMatcher = geodata.BuildDomainMatcher("full:example.com")
`); err != nil {
b.Fatal(err)
}
for _, benchmark := range []struct {
name, expression string
}{
{"ip_match", "ipMatcher:Match(ip)"},
{"ip_match_lower", "ipMatcher:match(ip)"},
{"ip_any_match", "ipMatcher:AnyMatch(ips)"},
{"ip_any_match_lower", "ipMatcher:anyMatch(ips)"},
{"ip_matches", "ipMatcher:Matches(ips)"},
{"ip_matches_lower", "ipMatcher:matches(ips)"},
{"domain_match_any", `domainMatcher:MatchAny("example.com")`},
{"domain_match_any_lower", `domainMatcher:matchAny("example.com")`},
{"ip_filter", "select(1, ipMatcher:FilterIPs(ips)) ~= nil"},
{"ip_filter_lower", "select(1, ipMatcher:filterIPs(ips)) ~= nil"},
{"domain_match", `#domainMatcher:Match("example.com") == 1`},
{"domain_match_lower", `#domainMatcher:match("example.com") == 1`},
{"ip_lua_table", "ipMatcher:AnyMatch({ip})"},
{"ip_lua_table_lower", "ipMatcher:anyMatch({ip})"},
} {
b.Run(benchmark.name, func(b *testing.B) {
if err := L.DoString(fmt.Sprintf("function benchmarkMatch() return %s end", benchmark.expression)); err != nil {
b.Fatal(err)
}
fn := L.GetGlobal("benchmarkMatch")
b.ReportAllocs()
b.ResetTimer()
for i := 0; i < b.N; i++ {
if err := L.CallByParam(lua.P{Fn: fn, NRet: 1, Protect: true}); err != nil {
b.Fatal(err)
}
if L.Get(-1) != lua.LTrue {
b.Fatal("matcher returned false")
}
L.Pop(1)
}
})
}
}
+12 -8
View File
@@ -52,9 +52,7 @@ func (g *MphIndexMatcher) Add(matcher Matcher) uint32 {
func (g *MphIndexMatcher) Build() error {
if g.mph != nil {
runtime.GC() // peak mem
if err := g.mph.Build(); err != nil {
return err
}
g.mph.Build()
}
runtime.GC() // peak mem
if g.ac != nil {
@@ -66,17 +64,23 @@ func (g *MphIndexMatcher) Build() error {
// Match implements IndexMatcher.Match.
func (g *MphIndexMatcher) Match(input string) []uint32 {
var result []uint32
result := make([][]uint32, 0, 5)
if g.mph != nil {
result = g.mph.Match(input) // a new slice, returned without another copy
if matches := g.mph.Match(input); len(matches) > 0 {
result = append(result, matches)
}
}
if g.ac != nil {
result = append(result, g.ac.Match(input)...)
if matches := g.ac.Match(input); len(matches) > 0 {
result = append(result, matches)
}
}
if g.regex != nil {
result = append(result, g.regex.Match(input)...)
if matches := g.regex.Match(input); len(matches) > 0 {
result = append(result, matches)
}
}
return result
return CompositeMatches(result)
}
// MatchAny implements IndexMatcher.MatchAny.
@@ -78,10 +78,6 @@ func TestMphIndexMatcher(t *testing.T) {
Input: "example.com",
Output: []uint32{10, 4},
},
{
Input: "apis.org",
Output: []uint32{2, 6},
},
}
matcherGroup := NewMphIndexMatcher()
for _, rule := range rules {
@@ -91,13 +87,8 @@ func TestMphIndexMatcher(t *testing.T) {
}
matcherGroup.Build()
for _, test := range cases {
m := matcherGroup.Match(test.Input)
if !reflect.DeepEqual(m, test.Output) {
t.Error("unexpected output: ", m, " for test case ", test)
}
clear(m) // the caller owns the result, so this must not change the next one
if m := matcherGroup.Match(test.Input); !reflect.DeepEqual(m, test.Output) {
t.Error("unexpected output after clearing the previous one: ", m, " for test case ", test)
t.Error("unexpected output: ", m, " for test case ", test)
}
}
}
+160 -369
View File
@@ -1,440 +1,231 @@
package strmatcher
import (
"bytes"
"cmp"
"encoding/binary"
"errors"
"math"
"slices"
"math/bits"
"runtime"
"sort"
"strings"
"unsafe"
)
// Flags of a level1 slot, stored above the record offset.
const (
mphDomain = 1 << 31 // matches the pattern and its subdomains
mphFull = 1 << 30 // matches the pattern only
mphParent = 1 << 29 // matches subdomains only, from a pattern with a leading dot
mphOffMask = mphParent - 1
)
// PrimeRK is the prime base used in Rabin-Karp algorithm.
const PrimeRK = 16777619
// Kinds of an added pattern, indexes of mphKinds.
const (
mphKindFull = iota
mphKindParent
mphKindDomain
)
// mphKinds are the slot flags in the order Match reports their values.
var mphKinds = [...]uint32{mphFull, mphParent, mphDomain}
// mphMultipliers are odd multipliers for the suffix hash. Build moves to the next one if two patterns collide.
var mphMultipliers = [...]uint64{0x9e3779b97f4a7c15, 0xc2b2ae3d27d4eb4f, 0x165667b19e3779f9, 0x27d4eb2f165667c5}
var (
errMphCollision = errors.New("strmatcher: suffix hash collision in MphMatcherGroup")
errMphBuilt = errors.New("strmatcher: MphMatcherGroup is already built")
)
type mphEntry struct {
off uint32 // pattern start in buf
value uint32
n uint32 // pattern length
kind uint8
// RollingHash calculates the rolling murmurHash of given string based on a provided suffix hash.
func RollingHash(hash uint32, input string) uint32 {
for i := len(input) - 1; i >= 0; i-- {
hash = hash*PrimeRK + uint32(input[i])
}
return hash
}
// MphMatcherGroup is an implementation of MatcherGroup for Full and Domain matchers.
// Each distinct pattern is stored once as a record in arena: its length (255 means a uvarint length follows),
// its bytes and, if the group holds more than one distinct value, its values. A minimal perfect hash table
// built with hash, displace and compress (http://cmph.sourceforge.net/papers/esa09.pdf) maps a pattern to its
// record. Patterns are hashed from the right, so one pass over the input hashes all its parent domains.
type MphMatcherGroup struct {
arena string
level0 []uint16 // bucket -> seed
level1 []uint32 // slot -> flags | record offset
fp []uint8 // slot -> low byte of its pattern's hash, rejects most misses without reading arena
n0, n1 uint32
mul uint64 // multiplier of the suffix hash
single uint32 // the only value if !multi
multi bool
// MemHash is the hash function used by go map, it utilizes available hardware instructions(behaves
// as aeshash if aes instruction is available).
// With different seed, each MemHash<seed> performs as distinct hash functions.
func MemHash(seed uint32, input string) uint32 {
return uint32(strhash(unsafe.Pointer(&input), uintptr(seed))) // nosemgrep
}
buf []byte // build only, patterns in Add order
entries []mphEntry
const (
mphMatchTypeCount = 2 // Full and Domain
)
type mphRuleInfo struct {
rollingHash uint32
matchers [mphMatchTypeCount][]uint32
}
// MphMatcherGroup is an implementation of MatcherGroup.
// It implements Rabin-Karp algorithm and minimal perfect hash table for Full and Domain matcher.
type MphMatcherGroup struct {
patterns string // All rule patterns concatenated
patternOffs []uint32 // RuleIdx -> patterns[patternOffs[i]:patternOffs[i+1]], index 0 reserved for failed lookup
values []uint32 // All registered matcher values concatenated
valueOffs []uint32 // RuleIdx -> values[valueOffs[i]:valueOffs[i+1]] (Full Matcher takes precedence)
level0 []uint32 // RollingHash & Mask -> seed for Memhash
level0Mask uint32 // Mask restricting RollingHash to 0 ~ len(level0)
level1 []uint32 // Memhash<seed> & Mask -> stored index for rules
level1Mask uint32 // Mask for restricting Memhash<seed> to 0 ~ len(level1)
rules []string // RuleIdx -> pattern string, only used for building
ruleInfos *map[string]mphRuleInfo
}
func NewMphMatcherGroup() *MphMatcherGroup {
return new(MphMatcherGroup)
return &MphMatcherGroup{
rules: []string{""},
level0: nil,
level0Mask: 0,
level1: nil,
level1Mask: 0,
ruleInfos: &map[string]mphRuleInfo{}, // Only used for building, destroyed after build complete
}
}
// AddFullMatcher implements MatcherGroupForFull.
func (g *MphMatcherGroup) AddFullMatcher(matcher FullMatcher, value uint32) {
g.add(matcher.Pattern(), mphKindFull, value)
pattern := strings.ToLower(matcher.Pattern())
g.addPattern(0, "", pattern, matcher.Type(), value)
}
// AddDomainMatcher implements MatcherGroupForDomain.
func (g *MphMatcherGroup) AddDomainMatcher(matcher DomainMatcher, value uint32) {
g.add(matcher.Pattern(), mphKindDomain, value)
pattern := strings.ToLower(matcher.Pattern())
hash := g.addPattern(0, "", pattern, matcher.Type(), value) // For full domain match
g.addPattern(hash, pattern, ".", matcher.Type(), value) // For partial domain match
}
func (g *MphMatcherGroup) add(pattern string, kind uint8, value uint32) {
if g.arena != "" {
panic(errMphBuilt)
}
pattern = strings.ToLower(pattern)
off := uint32(len(g.buf))
g.buf = append(g.buf, pattern...)
g.entries = append(g.entries, mphEntry{off: off, value: value, n: uint32(len(pattern)), kind: kind})
if len(pattern) > 0 && pattern[0] == '.' {
// ".x" has always matched "*.x" as well, so it also gets a parent-only record for "x"
g.entries = append(g.entries, mphEntry{off: off + 1, value: value, n: uint32(len(pattern) - 1), kind: mphKindParent})
func (g *MphMatcherGroup) addPattern(suffixHash uint32, suffixPattern string, pattern string, matcherType Type, value uint32) uint32 {
fullPattern := pattern + suffixPattern
info, found := (*g.ruleInfos)[fullPattern]
if !found {
info = mphRuleInfo{rollingHash: RollingHash(suffixHash, pattern)}
g.rules = append(g.rules, fullPattern)
}
info.matchers[matcherType] = append(info.matchers[matcherType], value)
(*g.ruleInfos)[fullPattern] = info
return info.rollingHash
}
func (g *MphMatcherGroup) key(i uint32) []byte {
e := &g.entries[i]
return g.buf[e.off : e.off+e.n]
}
// Build builds the hash table. It must be called once, after the last Add.
// Build builds a minimal perfect hash table for insert rules.
// Algorithm used: Hash, displace, and compress. See http://cmph.sourceforge.net/papers/esa09.pdf
func (g *MphMatcherGroup) Build() error {
if g.arena != "" {
return errMphBuilt
ruleCount := len(*g.ruleInfos)
g.level0 = make([]uint32, nextPow2(ruleCount/4))
g.level0Mask = uint32(len(g.level0) - 1)
g.level1 = make([]uint32, nextPow2(ruleCount))
g.level1Mask = uint32(len(g.level1) - 1)
// Flatten patterns and values so the built group has no per-rule objects
valueCount := 0
for _, ruleInfo := range *g.ruleInfos {
valueCount += len(ruleInfo.matchers[Full]) + len(ruleInfo.matchers[Domain])
}
if uint64(len(g.buf)) > math.MaxUint32 {
g.patterns = strings.Join(g.rules, "")
if uint64(len(g.patterns)) > math.MaxUint32 || uint64(valueCount) > math.MaxUint32 {
return errors.New("too many rules for MphMatcherGroup")
}
recs := g.writeRecords()
if len(g.arena) > mphOffMask {
return errors.New("too many rules for MphMatcherGroup")
}
hashes := make([]uint64, len(recs))
for _, mul := range mphMultipliers {
for i, rec := range recs {
hashes[i] = mphMix(mphHash(mul, g.recKey(rec)))
}
g.mul = mul
if err := g.place(recs, hashes); err != errMphCollision {
return err
}
}
return errMphCollision
}
g.patternOffs = make([]uint32, len(g.rules)+1)
g.values = make([]uint32, 0, valueCount)
g.valueOffs = make([]uint32, len(g.rules)+1)
// writeRecords writes one record per distinct pattern to arena and returns flags | offset of each.
func (g *MphMatcherGroup) writeRecords() []uint32 {
g.multi = false
if len(g.entries) > 0 {
g.single = g.entries[0].value
for _, e := range g.entries {
if e.value != g.single {
g.multi = true
break
}
}
// Create buckets based on all rule's rolling hash
buckets := make([][]uint32, len(g.level0))
for ruleIdx := 1; ruleIdx < len(g.rules); ruleIdx++ { // Traverse rules starting from index 1 (0 reserved for failed lookup)
ruleInfo := (*g.ruleInfos)[g.rules[ruleIdx]]
bucketIdx := ruleInfo.rollingHash & g.level0Mask
buckets[bucketIdx] = append(buckets[bucketIdx], uint32(ruleIdx))
g.patternOffs[ruleIdx+1] = g.patternOffs[ruleIdx] + uint32(len(g.rules[ruleIdx]))
g.values = append(append(g.values, ruleInfo.matchers[Full]...), ruleInfo.matchers[Domain]...)
g.valueOffs[ruleIdx+1] = uint32(len(g.values))
}
// Equal patterns become neighbours in Add order, so their values keep their priority
order := make([]uint32, len(g.entries))
for i := range order {
order[i] = uint32(i)
}
slices.SortFunc(order, func(a, b uint32) int {
return cmp.Or(bytes.Compare(g.key(a), g.key(b)), cmp.Compare(a, b))
})
g.rules = nil
g.ruleInfos = nil // Set ruleInfos nil to release memory
runtime.GC() // peak mem
size := len(g.buf) + len(g.entries) + 2
if g.multi {
size += 3 * len(g.entries)
// Sort buckets in descending order with respect to each bucket's size
bucketIdxs := make([]int, len(buckets))
for bucketIdx := range buckets {
bucketIdxs[bucketIdx] = bucketIdx
}
arena := make([]byte, 0, size)
recs := make([]uint32, 0, len(order))
var vals [len(mphKinds)][]uint32
for i := 0; i < len(order); {
k := g.key(order[i])
for t := range vals {
vals[t] = vals[t][:0]
}
for ; i < len(order) && bytes.Equal(g.key(order[i]), k); i++ {
e := &g.entries[order[i]]
if !slices.Contains(vals[e.kind], e.value) {
vals[e.kind] = append(vals[e.kind], e.value)
}
}
rec := uint32(len(arena))
if len(k) < 255 {
arena = append(arena, byte(len(k)))
} else {
arena = binary.AppendUvarint(append(arena, 255), uint64(len(k)))
}
arena = append(arena, k...)
for t, v := range vals {
if len(v) == 0 {
continue
}
rec |= mphKinds[t]
if g.multi {
arena = binary.AppendUvarint(arena, uint64(len(v)))
for _, x := range v {
arena = binary.AppendUvarint(arena, uint64(x))
sort.Slice(bucketIdxs, func(i, j int) bool { return len(buckets[bucketIdxs[i]]) > len(buckets[bucketIdxs[j]]) })
// Exercise Hash, Displace, and Compress algorithm to construct minimal perfect hash table
occupied := make([]bool, len(g.level1)) // Whether a second-level hash has been already used
hashedBucket := make([]uint32, 0, 4) // Second-level hashes for each rule in a specific bucket
for _, bucketIdx := range bucketIdxs {
bucket := buckets[bucketIdx]
hashedBucket = hashedBucket[:0]
seed := uint32(0)
for len(hashedBucket) != len(bucket) {
for _, ruleIdx := range bucket {
memHash := MemHash(seed, g.pattern(ruleIdx)) & g.level1Mask
if occupied[memHash] { // Collision occurred with this seed
for _, hash := range hashedBucket { // Revert all values in this hashed bucket
occupied[hash] = false
g.level1[hash] = 0
}
hashedBucket = hashedBucket[:0]
seed++ // Try next seed
break
}
occupied[memHash] = true
g.level1[memHash] = ruleIdx // The final value in the hash table
hashedBucket = append(hashedBucket, memHash)
}
}
recs = append(recs, rec)
}
// Lookups may point one byte past a pattern, and an empty group needs a record at offset 0 for empty slots
arena = append(arena, 0)
if len(recs) == 0 {
arena = append(arena, 0)
}
g.buf, g.entries = nil, nil
if cap(arena)-len(arena) > len(arena)/32 {
arena = slices.Clone(arena)
}
g.arena = unsafe.String(unsafe.SliceData(arena), len(arena)) // arena is not written after this
return recs
}
// place fills level0, level1 and fp: records are bucketed by hash, and each bucket, largest first, gets
// the first seed that puts all its records in free slots.
func (g *MphMatcherGroup) place(recs []uint32, hashes []uint64) error {
r := len(recs)
n0, n1 := max(1, r/3), max(1, r+r/99)
g.n0, g.n1 = uint32(n0), uint32(n1)
g.level0 = make([]uint16, n0)
g.level1 = make([]uint32, n1)
g.fp = make([]uint8, n1)
start := make([]uint32, n0+1)
for _, h := range hashes {
start[g.bucket(h)+1]++
}
for b := range n0 {
start[b+1] += start[b]
}
members := make([]uint32, r)
fill := slices.Clone(start[:n0])
for i, h := range hashes {
b := g.bucket(h)
members[fill[b]] = uint32(i)
fill[b]++
}
fill = nil
buckets := make([]uint32, n0)
for b := range buckets {
buckets[b] = uint32(b)
}
slices.SortStableFunc(buckets, func(a, b uint32) int {
return cmp.Compare(start[b+1]-start[b], start[a+1]-start[a])
})
occupied := make([]uint64, (n1+63)/64)
var slots []uint32
next:
for _, b := range buckets {
m := members[start[b]:start[b+1]]
if len(m) == 0 {
break
}
for i := range m {
for j := range i {
if hashes[m[i]] == hashes[m[j]] {
return errMphCollision // no seed can separate them
}
}
}
search:
for seed := range math.MaxUint16 + 1 {
slots = slots[:0]
for _, ri := range m {
s := g.slot(hashes[ri], uint16(seed))
if occupied[s/64]&(1<<(s%64)) != 0 || slices.Contains(slots, s) {
continue search
}
slots = append(slots, s)
}
for k, ri := range m {
s := slots[k]
occupied[s/64] |= 1 << (s % 64)
g.level1[s] = recs[ri]
g.fp[s] = uint8(hashes[ri])
}
g.level0[b] = uint16(seed)
continue next
}
return errors.New("strmatcher: no seed found for a bucket in MphMatcherGroup")
g.level0[bucketIdx] = seed // Displacement value for this bucket
}
return nil
}
// mphHash is the suffix hash of s, taken from the right: the hash of s[i:] is the state after reading s[i].
func mphHash(mul uint64, s string) uint64 {
h := uint64(0)
for i := len(s) - 1; i >= 0; i-- {
h = h*mul + uint64(s[i])
}
return h
func (g *MphMatcherGroup) pattern(ruleIdx uint32) string {
return g.patterns[g.patternOffs[ruleIdx]:g.patternOffs[ruleIdx+1]]
}
// mphMix spreads the weak low bits of a suffix hash.
func mphMix(h uint64) uint64 {
h ^= h >> 32
h *= 0xd6e8feb86659fd93
return h ^ h>>32
// valuesOf caps the capacity, so appending to a Match result can't overwrite the next rule's values.
func (g *MphMatcherGroup) valuesOf(ruleIdx uint32) []uint32 {
start, end := g.valueOffs[ruleIdx], g.valueOffs[ruleIdx+1]
return g.values[start:end:end]
}
func (g *MphMatcherGroup) bucket(f uint64) uint32 {
return uint32(((f >> 32) * uint64(g.n0)) >> 32)
}
func (g *MphMatcherGroup) slot(f uint64, seed uint16) uint32 {
x := ((f ^ uint64(seed)*0x9e3779b97f4a7c15) * 0xc4ceb9fe1a85ec53) >> 32
return uint32((x * uint64(g.n1)) >> 32)
}
func (g *MphMatcherGroup) uvarint(p uint32) (x, next uint32) {
for shift := 0; ; shift += 7 {
c := g.arena[p]
p++
x |= uint32(c&0x7f) << shift
if c < 0x80 {
return x, p
}
}
}
// recSpan returns where the pattern of the record at off starts and how long it is.
func (g *MphMatcherGroup) recSpan(off uint32) (p, n uint32) {
n, p = uint32(g.arena[off]), off+1
if n == 255 {
n, p = g.uvarint(p)
}
return p, n
}
func (g *MphMatcherGroup) recKey(rec uint32) string {
p, n := g.recSpan(rec & mphOffMask)
return g.arena[p : p+n]
}
// lookup returns the level1 entry of s, or 0 if s is not a pattern. h is the suffix hash of s.
func (g *MphMatcherGroup) lookup(h uint64, s string) uint32 {
f := mphMix(h)
// bucket < n0 == len(level0) and slot < n1 == len(level1) == len(fp), skip the bounds checks
seed := *(*uint16)(unsafe.Add(unsafe.Pointer(unsafe.SliceData(g.level0)), uintptr(g.bucket(f))*2))
slot := uintptr(g.slot(f, seed))
if *(*uint8)(unsafe.Add(unsafe.Pointer(unsafe.SliceData(g.fp)), slot)) != uint8(f) {
return 0
}
e := *(*uint32)(unsafe.Add(unsafe.Pointer(unsafe.SliceData(g.level1)), slot*4))
if len(s) < 255 {
// A record whose length byte is len(s) has len(s) pattern bytes after it
p := unsafe.Add(unsafe.Pointer(unsafe.StringData(g.arena)), e&mphOffMask)
if int(*(*byte)(p)) == len(s) && unsafe.String((*byte)(unsafe.Add(p, 1)), len(s)) == s {
return e
}
return 0
}
if g.recKey(e) == s {
return e
// Lookup searches for input in minimal perfect hash table and returns its index. 0 indicates not found.
func (g *MphMatcherGroup) Lookup(rollingHash uint32, input string) uint32 {
i0 := rollingHash & g.level0Mask
seed := g.level0[i0]
i1 := MemHash(seed, input) & g.level1Mask
n := g.level1[i1]
// Build only puts valid rule indices in level1, so n+1 < len(patternOffs) and the span is inside patterns.
// Skip the bounds checks, they made this hot path measurably slower than indexing a []string
offs := (*[2]uint32)(unsafe.Add(unsafe.Pointer(unsafe.SliceData(g.patternOffs)), uintptr(n)*4))
if start := offs[0]; int(offs[1]-start) == len(input) && unsafe.String((*byte)(unsafe.Add(unsafe.Pointer(unsafe.StringData(g.patterns)), start)), len(input)) == input {
return n
}
return 0
}
// appendValues appends the values of record e for the flags in want, in mphKinds order.
func (g *MphMatcherGroup) appendValues(dst []uint32, e, want uint32) []uint32 {
if !g.multi {
for _, flag := range mphKinds {
if e&want&flag != 0 {
dst = append(dst, g.single)
}
}
return dst
}
if e&want == 0 {
return dst
}
p, n := g.recSpan(e & mphOffMask)
p += n
for _, flag := range mphKinds {
if e&flag == 0 {
continue
}
var count, v uint32
for count, p = g.uvarint(p); count > 0; count-- {
v, p = g.uvarint(p)
if want&flag != 0 {
dst = append(dst, v)
}
}
}
return dst
}
// Match implements MatcherGroup.Match. Values of an exact match come first (Full, then Domain), then those of
// the parent domains, nearest first.
// Match implements MatcherGroup.Match.
func (g *MphMatcherGroup) Match(input string) []uint32 {
var stack [8]uint32
parents := stack[:0] // TLD side first
h, mul := uint64(0), g.mul
matches := make([][]uint32, 0, 5)
hash := uint32(0)
for i := len(input) - 1; i >= 0; i-- {
hash = hash*PrimeRK + uint32(input[i])
if input[i] == '.' {
if e := g.lookup(h, input[i+1:]); e&(mphDomain|mphParent) != 0 {
parents = append(parents, e)
if mphIdx := g.Lookup(hash, input[i:]); mphIdx != 0 {
matches = append(matches, g.valuesOf(mphIdx))
}
}
h = h*mul + uint64(input[i])
}
exact := g.lookup(h, input)
if exact&(mphFull|mphDomain) == 0 && len(parents) == 0 {
return nil
if mphIdx := g.Lookup(hash, input); mphIdx != 0 {
matches = append(matches, g.valuesOf(mphIdx))
}
result := g.appendValues(make([]uint32, 0, len(parents)+1), exact, mphFull|mphDomain)
for k := len(parents) - 1; k >= 0; k-- {
result = g.appendValues(result, parents[k], mphParent|mphDomain)
}
return result
return CompositeMatchesReverse(matches)
}
// MatchAny implements MatcherGroup.MatchAny.
func (g *MphMatcherGroup) MatchAny(input string) bool {
h, mul := uint64(0), g.mul
for i := len(input) - 1; i >= 0; i-- {
if input[i] == '.' && g.lookup(h, input[i+1:])&(mphDomain|mphParent) != 0 {
return true
}
h = h*mul + uint64(input[i])
}
return g.lookup(h, input)&(mphFull|mphDomain) != 0
}
// mphSuffix is the suffix hash of input[off:], a parent domain of the input.
type mphSuffix struct {
h uint64
off int
}
// mphSuffixes appends the suffix hashes of the parent domains of input to dst, TLD side first, and returns them
// with the hash of input itself: what MatchAny computes, computed once for several groups.
func mphSuffixes(dst []mphSuffix, mul uint64, input string) ([]mphSuffix, uint64) {
h := uint64(0)
hash := uint32(0)
for i := len(input) - 1; i >= 0; i-- {
hash = hash*PrimeRK + uint32(input[i])
if input[i] == '.' {
dst = append(dst, mphSuffix{h, i + 1})
if g.Lookup(hash, input[i:]) != 0 {
return true
}
}
h = h*mul + uint64(input[i])
}
return dst, h
return g.Lookup(hash, input) != 0
}
// matchAnyHashed is MatchAny with parents and h from mphSuffixes(_, mul, input).
func (g *MphMatcherGroup) matchAnyHashed(input string, parents []mphSuffix, h, mul uint64) bool {
if g.mul != mul {
return g.MatchAny(input) // built with a later multiplier after a collision
func nextPow2(v int) int {
if v <= 1 {
return 1
}
for _, p := range parents {
if g.lookup(p.h, input[p.off:])&(mphDomain|mphParent) != 0 {
return true
}
}
return g.lookup(h, input)&(mphFull|mphDomain) != 0
const MaxUInt = ^uint(0)
n := (MaxUInt >> bits.LeadingZeros(uint(v))) + 1
return int(n)
}
//go:noescape
//go:linkname strhash runtime.strhash
func strhash(p unsafe.Pointer, h uintptr) uintptr
@@ -1,108 +0,0 @@
package strmatcher
import (
"slices"
"testing"
)
func TestMphMatcherGroupHashCollision(t *testing.T) {
saved := mphMultipliers
defer func() { mphMultipliers = saved }()
mphMultipliers[0] = 1 // anagrams collide
g := NewMphMatcherGroup()
g.AddFullMatcher(FullMatcher("ab.com"), 1)
g.AddDomainMatcher(DomainMatcher("ba.com"), 2)
g.AddDomainMatcher(DomainMatcher("com"), 3)
if err := g.Build(); err != nil {
t.Fatal(err)
}
if g.mul != saved[1] {
t.Errorf("multiplier %#x, want the second one %#x", g.mul, saved[1])
}
for input, want := range map[string][]uint32{"ab.com": {1, 3}, "x.ba.com": {2, 3}, "x.ab.com": {3}, "ba.com": {2, 3}} {
if m := g.Match(input); !slices.Equal(m, want) {
t.Errorf("Match(%q) = %v, want %v", input, m, want)
}
}
// Thue-Morse strings of 2048 bytes and their complements collide for every odd multiplier
mphMultipliers = saved
a, b := make([]byte, 2048), make([]byte, 2048)
for i := range a {
a[i], b[i] = "ab"[bitsOnes(i)%2], "ba"[bitsOnes(i)%2]
}
g = NewMphMatcherGroup()
g.AddFullMatcher(FullMatcher(a), 1)
g.AddFullMatcher(FullMatcher(b), 1)
if err := g.Build(); err != errMphCollision {
t.Errorf("Build() = %v, want %v", err, errMphCollision)
}
}
func bitsOnes(i int) int {
n := 0
for ; i > 0; i &= i - 1 {
n++
}
return n
}
func TestMphValueMatcherCombiner(t *testing.T) {
build := func(matchers ...Matcher) *MphValueMatcher {
m := NewMphValueMatcher()
for _, x := range matchers {
m.Add(x, 0)
}
if err := m.Build(); err != nil {
t.Fatal(err)
}
return m
}
regex, err := Regex.New(`^a\d+\.net$`)
if err != nil {
t.Fatal(err)
}
saved := mphMultipliers
t.Cleanup(func() { mphMultipliers = saved })
mphMultipliers[0] = 1 // anagrams collide, so this one falls back to its own hash pass
collided := build(FullMatcher("ab.com"), DomainMatcher("ba.com"))
mphMultipliers = saved
if collided.mph.mul == mphMultipliers[0] {
t.Fatal("collided matcher uses the first multiplier")
}
matchers := []*MphValueMatcher{
build(DomainMatcher("example.com"), FullMatcher("full.org"), DomainMatcher(".dot.io")),
collided,
build(regex, SubstrMatcher("keyword")),
build(),
build(DomainMatcher("com"), DomainMatcher("a.b.c.d.e.f.g.h.i.j.k.l.m.n.o.p.q.r.s")),
}
var s MphValueMatcherCombiner
for i, m := range matchers {
s.Add(m, uint32(10+i))
}
inputs := []string{
"", ".", "..", "com", "example.com", "www.example.com", "xexample.com", "example.com.", "full.org", "x.full.org",
"dot.io", "x.dot.io", ".dot.io", "ab.com", "x.ab.com", "ba.com", "x.ba.com", "a12.net", "a12.net.x", "my-keyword.org",
"a.b.c.d.e.f.g.h.i.j.k.l.m.n.o.p.q.r.s", "0.a.b.c.d.e.f.g.h.i.j.k.l.m.n.o.p.q.r.s", "b.c.d.e.f.g.h.i.j.k.l.m.n.o.p.q.r.s",
"x.y.z.1.2.3.4.5.6.7.8.9.10.11.12.13.14.15.16.17.ab.com", "x.y.z.1.2.3.4.5.6.7.8.9.10.11.12.13.14.15.16.17.org",
}
for _, input := range inputs {
var want []uint32
for i, m := range matchers {
if m.MatchAny(input) {
want = append(want, uint32(10+i))
}
}
if got := s.Match(input); !slices.Equal(got, want) {
t.Errorf("Match(%q) = %v, want %v", input, got, want)
}
if got := s.MatchAny(input); got != (len(want) > 0) {
t.Errorf("MatchAny(%q) = %v", input, got)
}
}
if n := testing.AllocsPerRun(100, func() { s.MatchAny("www.a.b.c.example.org") }); n != 0 {
t.Errorf("MatchAny allocates %v times", n)
}
}
@@ -4,7 +4,6 @@ import (
"math/rand"
"reflect"
"slices"
"strings"
"testing"
"github.com/xtls/xray-core/common"
@@ -305,7 +304,7 @@ func TestMphMatcherGroupRandom(t *testing.T) {
domain["."+p] = append(domain["."+p], value)
}
}
common.Must(g.Build())
g.Build()
for _, input := range inputs {
keys := []string{input} // Whole input first, then "." suffixes from longest to shortest
for i := range len(input) {
@@ -317,10 +316,7 @@ func TestMphMatcherGroupRandom(t *testing.T) {
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)) {
if m := g.Match(input); !slices.Equal(m, want) {
t.Fatalf("seed %d: Match(%q) = %v, want %v", seed, input, m, want)
}
if m := g.MatchAny(input); m != (len(want) > 0) {
@@ -342,79 +338,3 @@ func TestMphMatcherGroupAppend(t *testing.T) {
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)
}
+1 -237
View File
@@ -2,12 +2,10 @@ package strmatcher
import (
"errors"
"math/bits"
"regexp"
"regexp/syntax"
"slices"
"strings"
"unicode"
"unicode/utf8"
"golang.org/x/net/idna"
@@ -77,9 +75,7 @@ func (m SubstrMatcher) Match(s string) bool {
// RegexMatcher is an implementation of Matcher.
type RegexMatcher struct {
pattern *regexp.Regexp
literals []string // every match contains all of them, longest first
tail []byteSet // tail[i] holds the bytes a matching input can have i bytes before its end
rest *byteSet // the bytes it can have further before, nil if any
literals []string // every match contains all of them, longest first
}
func newRegexMatcher(pattern string) (Matcher, error) {
@@ -91,239 +87,10 @@ func newRegexMatcher(pattern string) (Matcher, error) {
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 {
@@ -359,9 +126,6 @@ func (m *RegexMatcher) String() string {
}
func (m *RegexMatcher) Match(s string) bool {
if !m.mayMatch(s) {
return false
}
for _, l := range m.literals {
if !strings.Contains(s, l) {
return false
@@ -1,16 +1,9 @@
package strmatcher
import (
"hash/fnv"
"math/rand/v2"
"regexp"
"regexp/syntax"
"slices"
"strconv"
"strings"
"testing"
"unicode"
"unicode/utf8"
)
var regexLiteralCases = []struct {
@@ -44,147 +37,6 @@ func TestRegexRequiredLiterals(t *testing.T) {
}
}
var regexTailCases = []struct {
pattern string
guard bool
match []string // inputs the pattern matches
reject []string // inputs the tail guard alone rejects
}{
{`^[a-z]([a-z0-9-]{0,61}[a-z0-9])?$`, true, []string{"a", "localhost", "x-1"}, []string{"www.example.com", "localhost.", "LOCALHOST", "a b"}},
{`(^|\.)[a-z][1-9][0-9][a-z]\.com$`, true, []string{"a12b.com", "x.q10z.com"}, []string{"google.com", "a12b.co", "a12b.com.", "ab12.com"}},
{`^hses[1-7]?\.akamaized\.net$`, true, []string{"hses.akamaized.net", "hses3.akamaized.net"}, []string{"xhses.akamaized.net", "www.hses.akamaized.net"}},
{`(?i)k\.net$`, true, []string{"k.net", "K.NET", "\u212a.net"}, []string{"x.net", "k.nex"}},
{`[^.]+\.cn$`, true, []string{"a.cn", "\xff.cn", "\u4e2d.cn"}, []string{"a.cnn", "a.c"}},
{`\x{FFFD}$`, true, []string{"\xff", "a\xc3", "\uFFFD"}, []string{"a", "\xff."}},
{`^.\.cn$`, true, []string{"a.cn", "\u4E2D.cn", "\xff.cn"}, []string{"ab.cn"}},
{`^$`, true, []string{""}, []string{"a"}},
{`(^|\.)youyuapi\..+$`, false, []string{"youyuapi.com"}, nil},
{`abc`, false, []string{"abc", "xabcx"}, nil},
{`^ab`, false, []string{"ab", "abc"}, nil},
{`a$|b`, false, []string{"a", "bx"}, nil},
{`(?m)a$`, false, []string{"a", "a\nb"}, nil},
{strings.Repeat(`(?:abcdefgh(?:a`, 20) + strings.Repeat(`)*)*`, 20) + `\.com$`, false, []string{".com", "abcdefgha.com"}, nil}, // over tailBudget
}
func TestRegexTailGuard(t *testing.T) {
for _, test := range regexTailCases {
m, err := newRegexMatcher(test.pattern)
if err != nil {
t.Fatal(err)
}
rm := m.(*RegexMatcher)
if guard := rm.tail != nil || rm.rest != nil; guard != test.guard {
t.Errorf("%s: guard %v, want %v", test.pattern, guard, test.guard)
}
for _, s := range test.match {
if !rm.pattern.MatchString(s) || !rm.Match(s) {
t.Errorf("%s: %q does not match", test.pattern, s)
}
}
for _, s := range test.reject {
if rm.pattern.MatchString(s) || rm.mayMatch(s) {
t.Errorf("%s: %q passes the guard", test.pattern, s)
}
}
}
}
// TestRegexTailGuardFlatAlternation checks that a long but non-recursive pattern keeps its
// guard. Only nested repeats are charged against tailBudget, so a flat alternation of many
// names, however large, is walked once and guarded; its guard is checked against regexp.
func TestRegexTailGuardFlatAlternation(t *testing.T) {
var sb strings.Builder
sb.WriteString("(?:")
for i := 0; i < 20000; i++ {
if i > 0 {
sb.WriteByte('|')
}
sb.WriteString("name")
sb.WriteString(strconv.Itoa(i))
}
sb.WriteString(`)\.example\.com$`)
m, err := newRegexMatcher(sb.String())
if err != nil {
t.Fatal(err)
}
rm := m.(*RegexMatcher)
if rm.tail == nil && rm.rest == nil {
t.Fatal("flat alternation of 20000 names lost its guard")
}
for _, s := range []string{"name0.example.com", "name19999.example.com", "x.name12345.example.com"} {
if !rm.pattern.MatchString(s) || !rm.Match(s) {
t.Errorf("%q should match", s)
}
}
for _, s := range []string{"name0.example.org", "name0.example.com.", "name0.example.con", "google.com"} {
if rm.pattern.MatchString(s) {
t.Fatalf("test bug: %q matches the pattern", s)
}
if rm.mayMatch(s) {
t.Errorf("%q should be rejected by the guard", s)
}
}
}
// sampleMatch appends a string that re matches, assertions aside, unless it runs out of
// budget, which it spends one per call so that nested repeats stay cheap.
func sampleMatch(sb *strings.Builder, re *syntax.Regexp, rnd *rand.Rand, budget *int) {
if *budget <= 0 {
return
}
*budget--
switch re.Op {
case syntax.OpLiteral:
for _, r := range re.Rune {
if re.Flags&syntax.FoldCase != 0 {
for n := rnd.IntN(4); n > 0; n-- {
r = unicode.SimpleFold(r)
}
}
sampleRune(sb, r, rnd)
}
case syntax.OpCharClass:
if len(re.Rune) > 0 {
i := rnd.IntN(len(re.Rune)/2) * 2
sampleRune(sb, re.Rune[i]+rnd.Int32N(min(re.Rune[i+1]-re.Rune[i]+1, 300)), rnd)
}
case syntax.OpAnyChar, syntax.OpAnyCharNotNL:
sampleRune(sb, []rune{'a', '.', '\n', 0xe9, 0x212a, utf8.RuneError}[rnd.IntN(6)], rnd)
case syntax.OpCapture:
sampleMatch(sb, re.Sub[0], rnd, budget)
case syntax.OpConcat:
for _, sub := range re.Sub {
sampleMatch(sb, sub, rnd, budget)
}
case syntax.OpAlternate:
sampleMatch(sb, re.Sub[rnd.IntN(len(re.Sub))], rnd, budget)
case syntax.OpQuest, syntax.OpStar, syntax.OpPlus, syntax.OpRepeat:
lo, hi := 0, 3
switch re.Op {
case syntax.OpQuest:
hi = 1
case syntax.OpPlus:
lo = 1
case syntax.OpRepeat:
lo, hi = re.Min, re.Min+3
if re.Max >= 0 {
hi = min(hi, re.Max)
}
}
for n := lo + rnd.IntN(hi-lo+1); n > 0; n-- {
sampleMatch(sb, re.Sub[0], rnd, budget)
}
}
}
func sampleRune(sb *strings.Builder, r rune, rnd *rand.Rand) {
if r == utf8.RuneError && rnd.IntN(2) == 0 {
sb.WriteByte(0x80 | byte(rnd.IntN(0x80))) // regexp matches an invalid byte as U+FFFD
return
}
sb.WriteRune(r)
}
func FuzzRegexMatcher(f *testing.F) {
inputs := []string{
"", "x", "yy", "abd", "ccc", "ABC", "abCDef", "abcdef", "abababcc", "a.b", "a\xffb", "a\uFFFDb",
@@ -195,39 +47,14 @@ func FuzzRegexMatcher(f *testing.F) {
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:])
}
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)
}
})
}
+12 -67
View File
@@ -46,9 +46,7 @@ func (g *MphValueMatcher) Add(matcher Matcher, value uint32) {
func (g *MphValueMatcher) Build() error {
if g.mph != nil {
runtime.GC() // peak mem
if err := g.mph.Build(); err != nil {
return err
}
g.mph.Build()
}
runtime.GC() // peak mem
if g.ac != nil {
@@ -60,17 +58,23 @@ func (g *MphValueMatcher) Build() error {
// Match implements ValueMatcher.Match.
func (g *MphValueMatcher) Match(input string) []uint32 {
var result []uint32
result := make([][]uint32, 0, 5)
if g.mph != nil {
result = g.mph.Match(input) // a new slice, returned without another copy
if matches := g.mph.Match(input); len(matches) > 0 {
result = append(result, matches)
}
}
if g.ac != nil {
result = append(result, g.ac.Match(input)...)
if matches := g.ac.Match(input); len(matches) > 0 {
result = append(result, matches)
}
}
if g.regex != nil {
result = append(result, g.regex.Match(input)...)
if matches := g.regex.Match(input); len(matches) > 0 {
result = append(result, matches)
}
}
return result
return CompositeMatches(result)
}
// MatchAny implements ValueMatcher.MatchAny.
@@ -83,62 +87,3 @@ func (g *MphValueMatcher) MatchAny(input string) bool {
}
return g.regex != nil && g.regex.MatchAny(input)
}
func (g *MphValueMatcher) matchAnyHashed(input string, parents []mphSuffix, h, mul uint64) bool {
if g.mph != nil && g.mph.matchAnyHashed(input, parents, h, mul) {
return true
}
if g.ac != nil && g.ac.MatchAny(input) {
return true
}
return g.regex != nil && g.regex.MatchAny(input)
}
// MphValueMatcherCombiner combines several built MphValueMatchers, each bound to one value, and matches an input
// against them as their MatchAny would, hashing the input once for all of them.
type MphValueMatcherCombiner struct {
matchers []*MphValueMatcher
values []uint32
}
// Add adds a built matcher that stands for value.
func (s *MphValueMatcherCombiner) Add(m *MphValueMatcher, value uint32) {
s.matchers = append(s.matchers, m)
s.values = append(s.values, value)
}
// Match returns the values of the matchers that match input, in Add order.
func (s *MphValueMatcherCombiner) Match(input string) []uint32 {
if len(s.matchers) == 0 {
return nil
}
var stack [16]mphSuffix
mul := mphMultipliers[0]
parents, h := mphSuffixes(stack[:0], mul, input)
var result []uint32
for i, m := range s.matchers {
if m.matchAnyHashed(input, parents, h, mul) {
result = append(result, s.values[i])
}
}
return result
}
// MatchAny returns true as soon as one matcher matches input.
func (s *MphValueMatcherCombiner) MatchAny(input string) bool {
switch len(s.matchers) {
case 0:
return false
case 1:
return s.matchers[0].MatchAny(input) // nothing to share, and it stops at the first matching suffix
}
var stack [16]mphSuffix
mul := mphMultipliers[0]
parents, h := mphSuffixes(stack[:0], mul, input)
for _, m := range s.matchers {
if m.matchAnyHashed(input, parents, h, mul) {
return true
}
}
return false
}
-61
View File
@@ -1,61 +0,0 @@
package log
import (
"path/filepath"
"strings"
lua "github.com/yuin/gopher-lua"
)
// RegisterLua makes xray.log available to require in an LState.
func RegisterLua(L *lua.LState) {
L.PreloadModule("xray.log", func(L *lua.LState) int {
module := L.CreateTable(0, 4)
var source, prefix string // cache
for name, severity := range map[string]Severity{
"Debug": Severity_Debug,
"Info": Severity_Info,
"Warning": Severity_Warning,
"Error": Severity_Error,
} {
module.RawSetString(name, L.NewFunction(func(L *lua.LState) int {
if GetSeverity() < severity {
return 0
}
var content strings.Builder
// Prefix with the calling script's filename.
if caller, ok := L.GetStack(1); ok {
if _, err := L.GetInfo("S", caller, lua.LNil); err == nil && caller.Source != "" {
if caller.Source != source {
source = caller.Source
prefix = filepath.Base(strings.TrimPrefix(source, "@")) + ": "
}
content.WriteString(prefix)
}
}
for i := 1; i <= L.GetTop(); i++ {
content.WriteString(luaLogString(L, L.Get(i)))
}
Record(&GeneralMessage{
Severity: severity,
Content: content.String(),
})
return 0
}))
}
L.Push(module)
return 1
})
}
func luaLogString(L *lua.LState, value lua.LValue) string {
if ud, ok := value.(*lua.LUserData); ok {
if err, ok := ud.Value.(error); ok {
return err.Error()
}
}
if _, ok := L.GetMetaField(value, "__tostring").(*lua.LFunction); ok {
return L.ToStringMeta(value).String()
}
return value.String()
}
-213
View File
@@ -1,213 +0,0 @@
package log
import (
"errors"
"fmt"
"os"
"path/filepath"
"testing"
lua "github.com/yuin/gopher-lua"
)
type luaLogHandler struct {
messages []Message
}
func (h *luaLogHandler) Handle(msg Message) {
h.messages = append(h.messages, msg)
}
func TestLuaLog(t *testing.T) {
previous := logHandler.Load()
t.Cleanup(func() { logHandler.Store(previous) })
handler := &luaLogHandler{}
RegisterHandler(handler)
L := lua.NewState()
defer L.Close()
RegisterLua(L)
nativeError := L.NewUserData()
nativeError.Value = fmt.Errorf("lookup failed: %w", errors.New("upstream timeout"))
L.SetGlobal("nativeError", nativeError)
path := filepath.Join(t.TempDir(), "logging.lua")
if err := os.WriteFile(path, []byte(`
local log = require("xray.log")
assert(log == require("xray.log"))
log.Debug("query: ", "example.com")
log.Info("count=", 42, ", enabled=", true, ", value=", nil)
log.Warning(setmetatable({}, {
__tostring = function() return "fallback" end
}))
assert(select("#", log.Error("failed")) == 0)
log.Error("DNS failed: ", nativeError)
log.Warning(nativeError)
local ok, err = pcall(function() error("Lua failure", 0) end)
assert(not ok)
log.Error(err)
local calls = 0
local custom = setmetatable({}, {
__tostring = function() calls = calls + 1; return "custom" end
})
log.Info(custom, custom)
assert(calls == 2)
log.Info("a", "b", "c", "d", "e", "f", "g", "h", "i", "j", "k", "l")
log.Info()
function logHook()
log.Info("hook")
end
`), 0o600); err != nil {
t.Fatal(err)
}
if err := L.DoFile(path); err != nil {
t.Fatal(err)
}
if err := L.DoString(`
logHook()
require("xray.log").Info("anonymous")
`); err != nil {
t.Fatal(err)
}
other := filepath.Join(t.TempDir(), "other.lua")
if err := os.WriteFile(other, []byte(`
local log = require("xray.log")
log.Info("other")
logHook()
log.Info("other again")
`), 0o600); err != nil {
t.Fatal(err)
}
if err := L.DoFile(other); err != nil {
t.Fatal(err)
}
want := []struct {
severity Severity
message string
}{
{Severity_Debug, "[Debug] logging.lua: query: example.com"},
{Severity_Info, "[Info] logging.lua: count=42, enabled=true, value=nil"},
{Severity_Warning, "[Warning] logging.lua: fallback"},
{Severity_Error, "[Error] logging.lua: failed"},
{Severity_Error, "[Error] logging.lua: DNS failed: lookup failed: upstream timeout"},
{Severity_Warning, "[Warning] logging.lua: lookup failed: upstream timeout"},
{Severity_Error, "[Error] logging.lua: Lua failure"},
{Severity_Info, "[Info] logging.lua: customcustom"},
{Severity_Info, "[Info] logging.lua: abcdefghijkl"},
{Severity_Info, "[Info] logging.lua: "},
{Severity_Info, "[Info] logging.lua: hook"},
{Severity_Info, "[Info] <string>: anonymous"},
{Severity_Info, "[Info] other.lua: other"},
{Severity_Info, "[Info] logging.lua: hook"},
{Severity_Info, "[Info] other.lua: other again"},
}
if len(handler.messages) != len(want) {
t.Fatalf("logged %d messages, want %d", len(handler.messages), len(want))
}
for i, expected := range want {
msg, ok := handler.messages[i].(*GeneralMessage)
if !ok {
t.Fatalf("message %d has type %T, want *GeneralMessage", i, handler.messages[i])
}
if msg.Severity != expected.severity || msg.String() != expected.message {
t.Errorf("message %d = %q with severity %v, want %q with severity %v", i, msg.String(), msg.Severity, expected.message, expected.severity)
}
}
}
type luaSeverityLogHandler struct {
luaLogHandler
level Severity
}
func (h *luaSeverityLogHandler) Severity() Severity { return h.level }
func TestLuaLogSeverity(t *testing.T) {
previous := logHandler.Load()
t.Cleanup(func() { logHandler.Store(previous) })
L := lua.NewState()
defer L.Close()
RegisterLua(L)
for _, level := range []Severity{Severity_Unknown, Severity_Error, Severity_Warning, Severity_Info, Severity_Debug, Severity_Warning} {
t.Run(level.String(), func(t *testing.T) {
handler := &luaSeverityLogHandler{level: level}
RegisterHandler(handler)
want := []Severity{}
for _, severity := range []Severity{Severity_Error, Severity_Warning, Severity_Info, Severity_Debug} {
if severity <= level {
want = append(want, severity)
}
}
if err := L.DoString(fmt.Sprintf(`
local log = require("xray.log")
local calls = 0
local value = setmetatable({}, {
__tostring = function() calls = calls + 1; return "message" end
})
for _, write in ipairs({log.Error, log.Warning, log.Info, log.Debug}) do
assert(select("#", write(value)) == 0)
end
assert(calls == %d)
`, len(want))); err != nil {
t.Fatal(err)
}
if len(handler.messages) != len(want) {
t.Fatalf("logged %d messages, want %d", len(handler.messages), len(want))
}
for i, severity := range want {
msg := handler.messages[i].(*GeneralMessage)
if msg.Severity != severity || msg.Content != "<string>: message" {
t.Errorf("message %d = %v, want severity %v and content %q", i, msg, severity, "<string>: message")
}
}
})
}
}
type luaDiscardLogHandler struct{ level Severity }
func (luaDiscardLogHandler) Handle(Message) {}
func (h luaDiscardLogHandler) Severity() Severity { return h.level }
func BenchmarkLuaLog(b *testing.B) {
benchmarkLuaLog(b, Severity_Debug)
}
func BenchmarkLuaLogFiltered(b *testing.B) {
benchmarkLuaLog(b, Severity_Warning)
}
func benchmarkLuaLog(b *testing.B, level Severity) {
previous := logHandler.Load()
b.Cleanup(func() { logHandler.Store(previous) })
RegisterHandler(luaDiscardLogHandler{level: level})
L := lua.NewState()
defer L.Close()
RegisterLua(L)
if err := L.DoString(`custom = setmetatable({}, {__tostring = function() return "custom" end})`); err != nil {
b.Fatal(err)
}
for _, benchmark := range []struct {
name, arguments string
}{
{"strings", `"query: ", "example.com"`},
{"mixed", `"count=", 42, ", enabled=", true, ", value=", nil`},
{"many_arguments", `"a", "b", "c", "d", "e", "f", "g", "h", "i", "j", "k", "l"`},
{"tostring", "custom"},
} {
b.Run(benchmark.name, func(b *testing.B) {
if err := L.DoString(fmt.Sprintf(`local log = require("xray.log")
function benchmarkLog() log.Info(%s) end`, benchmark.arguments)); err != nil {
b.Fatal(err)
}
fn := L.GetGlobal("benchmarkLog")
b.ReportAllocs()
b.ResetTimer()
for i := 0; i < b.N; i++ {
if err := L.CallByParam(lua.P{Fn: fn, NRet: 0, Protect: true}); err != nil {
b.Fatal(err)
}
}
})
}
}
-3
View File
@@ -1,3 +0,0 @@
// Package lua provides shared GopherLua programs, state management, and value
// conversion and validation helpers for Xray scripts.
package lua
-65
View File
@@ -1,65 +0,0 @@
package lua
import (
glua "github.com/yuin/gopher-lua"
luar "layeh.com/gopher-luar"
)
// NewSlicePusher captures luar's slice metatable during state initialization.
// The returned function wraps slices without reflection or metatable lookup,
// and pushes nil for nil slices. Use it with this state or its coroutines.
func NewSlicePusher[T any](L *glua.LState) func(*glua.LState, []T) {
metatable := luar.New(L, []T{}).(*glua.LUserData).Metatable
return func(L *glua.LState, values []T) {
if values == nil {
L.Push(glua.LNil)
return
}
userdata := L.NewUserData()
userdata.Value = values
userdata.Metatable = metatable
L.Push(userdata)
}
}
// DirectMethod handles a Lua call without luar's reflected method invocation.
// It returns the result count and whether it handled the arguments. On false,
// it must leave the stack unchanged for the original luar wrapper.
type DirectMethod func(L *glua.LState) (nresults int, handled bool)
// PushWithDirectMethods pushes a luar userdata with typed Go method bindings.
// Handled calls bypass luar's argument conversion and reflect.Call; method lookup
// uses the methods table directly instead of luar's reflected __index handler.
// value must expose methods only. Bindings and their closures are installed once
// per Go type per LState, outside the method-call hot path.
func PushWithDirectMethods(L *glua.LState, value any, directMethods map[string]DirectMethod) {
userdata := luar.New(L, value).(*glua.LUserData)
metatable := userdata.Metatable.(*glua.LTable)
methods := metatable.RawGetString("methods").(*glua.LTable)
if metatable.RawGetString("__index") != methods {
for name, direct := range directMethods {
original := methods.RawGetString(name)
fn := L.NewFunction(func(L *glua.LState) int {
if nresults, handled := direct(L); handled {
return nresults
}
return callLuarMethod(L, original)
})
// Keep luar's method aliases on the same direct binding.
for key, method := methods.Next(glua.LNil); key != glua.LNil; key, method = methods.Next(key) {
if method == original {
methods.RawSet(key, fn)
}
}
}
metatable.RawSetString("__index", methods)
}
L.Push(userdata)
}
func callLuarMethod(L *glua.LState, method glua.LValue) int {
nargs := L.GetTop()
L.Insert(method, 1)
L.Call(nargs, glua.MultRet)
return L.GetTop()
}
-84
View File
@@ -1,84 +0,0 @@
package lua
import (
"net"
"testing"
glua "github.com/yuin/gopher-lua"
luar "layeh.com/gopher-luar"
)
func TestSlicePusher(t *testing.T) {
L := glua.NewState()
defer L.Close()
push := NewSlicePusher[int](L)
values := []int{3, 5}
L.SetGlobal("getValues", L.NewFunction(func(L *glua.LState) int {
push(L, values)
return 1
}))
if err := L.DoString(`
local values = getValues()
assert(#values == 2 and values[1] == 3 and values[2] == 5)
values[2] = 7
local co = coroutine.create(function()
local values = getValues()
assert(#values == 2 and values[1] == 3 and values[2] == 7)
return true
end)
local ok, result = coroutine.resume(co)
assert(ok and result == true)
`); err != nil {
t.Fatal(err)
}
if values[1] != 7 {
t.Fatal("slice storage was copied")
}
push(L, nil)
if L.Get(-1) != glua.LNil {
t.Fatal("nil slice must push Lua nil")
}
L.Pop(1)
push(L, []int{})
L.SetGlobal("empty", L.Get(-1))
L.Pop(1)
if err := L.DoString(`assert(type(empty) == "userdata" and #empty == 0)`); err != nil {
t.Fatal(err)
}
}
func TestSlicePusherMetatablePerState(t *testing.T) {
first := glua.NewState()
defer first.Close()
second := glua.NewState()
defer second.Close()
NewSlicePusher[int](first)(first, []int{1})
NewSlicePusher[int](second)(second, []int{1})
if first.Get(-1).(*glua.LUserData).Metatable == second.Get(-1).(*glua.LUserData).Metatable {
t.Fatal("independent states share a slice metatable")
}
}
func BenchmarkSlicePusher(b *testing.B) {
L := glua.NewState()
defer L.Close()
ips := []net.IP{net.ParseIP("127.0.0.1")}
pushIPs := NewSlicePusher[net.IP](L)
for _, benchmark := range []struct {
name string
push func(*glua.LState, []net.IP)
}{
{"bare", func(L *glua.LState, ips []net.IP) { PushUserData(L, ips) }},
{"luar", func(L *glua.LState, ips []net.IP) { L.Push(luar.New(L, ips)) }},
{"cached", pushIPs},
} {
b.Run(benchmark.name, func(b *testing.B) {
b.ReportAllocs()
b.ResetTimer()
for i := 0; i < b.N; i++ {
benchmark.push(L, ips)
L.Pop(1)
}
})
}
}
-151
View File
@@ -1,151 +0,0 @@
package lua
import (
"context"
"errors"
"sync"
"time"
glua "github.com/yuin/gopher-lua"
)
const maxIdleStates = 16
// Pool lends each state to one caller at a time. It grows on contention and
// keeps up to maxIdleStates idle states until Close. Acquire/Release callers
// decide reusability; WithState uses its callback's error.
type Pool struct {
ctx context.Context
cancel context.CancelFunc
timeout time.Duration
factory LStateFactory
idle []*glua.LState
top int
mu sync.Mutex
active sync.WaitGroup
closed bool
}
// NewPool tests the factory by creating one state during initialization.
func NewPool(ctx context.Context, timeout time.Duration, factory LStateFactory) (*Pool, error) {
if timeout <= 0 {
return nil, errors.New("Lua pool timeout must be positive")
}
poolCtx, cancel := context.WithCancel(ctx)
state, err := factory(poolCtx)
if err != nil {
cancel()
return nil, err
}
return &Pool{ctx: poolCtx, cancel: cancel, timeout: timeout, factory: factory, idle: []*glua.LState{state}, top: state.GetTop()}, nil
}
// Acquire returns an initialized exclusive state, growing the pool if necessary.
// ctx is passed to the factory for state creation; nil uses the pool context.
func (p *Pool) Acquire(ctx context.Context) (*glua.LState, error) {
p.mu.Lock()
if p.closed {
p.mu.Unlock()
return nil, errors.New("Lua pool is closed")
}
if err := p.ctx.Err(); err != nil {
p.mu.Unlock()
return nil, err
}
if ctx == nil {
ctx = p.ctx
} else if err := ctx.Err(); err != nil {
p.mu.Unlock()
return nil, err
}
p.active.Add(1)
n := len(p.idle)
if n != 0 {
state := p.idle[n-1]
p.idle[n-1] = nil
p.idle = p.idle[:n-1]
p.mu.Unlock()
return state, nil
}
p.mu.Unlock()
// TODO: Limit the total number of states. When the limit is reached, wait
// for a Release instead of creating another state; allow the wait to be
// cancelled by the caller or by Close.
state, err := p.factory(ctx)
if err != nil {
p.active.Done()
return nil, err
}
return state, nil
}
// WithState runs work on an exclusive state and releases it afterward.
// Nil ctx and zero timeout use pool defaults. The timeout starts after acquisition.
func (p *Pool) WithState(ctx context.Context, timeout time.Duration, work func(*glua.LState) error) error {
state, err := p.Acquire(ctx)
if err != nil {
return err
}
if ctx == nil {
ctx = p.ctx
}
if timeout == 0 {
timeout = p.timeout
}
ctx, cancel := context.WithTimeout(ctx, timeout)
state.SetContext(ctx)
reusable := false
defer func() {
cancel()
p.Release(state, reusable)
}()
err = work(state)
reusable = err == nil
return err
}
// Release resets a state for reuse or closes it.
func (p *Pool) Release(state *glua.LState, reusable bool) {
if reusable {
state.RemoveContext()
state.SetTop(p.top)
p.mu.Lock()
if !p.closed && p.ctx.Err() == nil && len(p.idle) < maxIdleStates {
p.idle = append(p.idle, state)
} else {
reusable = false
}
p.mu.Unlock()
}
if !reusable {
state.Close()
}
p.active.Done()
}
// Close cancels the pool context, closes idle states, and waits for borrowed states.
func (p *Pool) Close() {
p.mu.Lock()
if !p.closed {
p.closed = true
p.cancel()
for _, state := range p.idle {
state.Close()
}
p.idle = nil
}
p.mu.Unlock()
p.active.Wait()
}
-466
View File
@@ -1,466 +0,0 @@
package lua
import (
"context"
"errors"
"testing"
"time"
glua "github.com/yuin/gopher-lua"
)
func newTestPool(t testing.TB, ctx context.Context, timeout time.Duration, factory LStateFactory) *Pool {
t.Helper()
pool, err := NewPool(ctx, timeout, factory)
if err != nil {
t.Fatal(err)
}
t.Cleanup(pool.Close)
return pool
}
func assertPoolCloseBlocked(t *testing.T, done <-chan struct{}) {
t.Helper()
select {
case <-done:
t.Fatal("Close returned while work was still active")
case <-time.After(20 * time.Millisecond):
}
}
func TestPoolTimeoutValidation(t *testing.T) {
for _, tc := range []struct {
name string
timeout time.Duration
wantErr bool
}{
{"zero", 0, true},
{"negative", -time.Nanosecond, true},
{"positive", time.Nanosecond, false},
} {
t.Run(tc.name, func(t *testing.T) {
called := false
pool, err := NewPool(context.Background(), tc.timeout, func(context.Context) (*glua.LState, error) {
called = true
return glua.NewState(), nil
})
if pool != nil {
t.Cleanup(pool.Close)
}
if (err != nil) != tc.wantErr {
t.Fatalf("NewPool error = %v, want error %t", err, tc.wantErr)
}
if tc.wantErr && (pool != nil || called) {
t.Fatal("invalid timeout created a pool or called the factory")
}
})
}
}
func TestPoolFactoryFailure(t *testing.T) {
failure := errors.New("factory failed")
_, err := NewPool(context.Background(), time.Second, func(context.Context) (*glua.LState, error) {
return nil, failure
})
if !errors.Is(err, failure) {
t.Fatalf("NewPool error = %v, want original factory error", err)
}
calls := 0
pool := newTestPool(t, context.Background(), time.Second, func(context.Context) (*glua.LState, error) {
calls++
if calls == 1 {
return glua.NewState(), nil
}
return nil, failure
})
state, err := pool.Acquire(nil)
if err != nil {
t.Fatal(err)
}
defer pool.Release(state, true)
err = pool.WithState(nil, 0, func(*glua.LState) error {
t.Error("work ran after factory failure")
return nil
})
if !errors.Is(err, failure) {
t.Fatalf("WithState error = %v, want original factory error", err)
}
}
func TestPoolReusesStatesAndLimitsIdle(t *testing.T) {
created := 0
pool := newTestPool(t, context.Background(), time.Second, func(context.Context) (*glua.LState, error) {
created++
return glua.NewState(), nil
})
var borrowed []*glua.LState
defer func() {
for _, state := range borrowed {
pool.Release(state, false)
}
}()
for range maxIdleStates + 3 {
state, err := pool.Acquire(nil)
if err != nil {
t.Fatal(err)
}
borrowed = append(borrowed, state)
state.SetContext(context.Background())
}
states := borrowed
for _, state := range states {
pool.Release(state, true)
}
borrowed = nil
open := 0
for _, state := range states {
if !state.IsClosed() {
if state.Context() != nil {
t.Fatal("Release left a context on a reusable state")
}
open++
}
}
if open != maxIdleStates {
t.Fatalf("retained %d states, want %d", open, maxIdleStates)
}
if err := pool.WithState(nil, 0, func(*glua.LState) error { return nil }); err != nil {
t.Fatal(err)
}
if created != len(states) {
t.Fatalf("created %d states, want %d", created, len(states))
}
pool.Close()
for _, state := range states {
if !state.IsClosed() {
t.Fatal("Close left an idle state open")
}
}
}
func TestPoolWithStateOptions(t *testing.T) {
key := struct{}{}
parent := context.WithValue(context.Background(), key, "pool")
caller := context.WithValue(context.Background(), key, "caller")
pool := newTestPool(t, parent, time.Second, func(context.Context) (*glua.LState, error) {
return glua.NewState(), nil
})
for _, tc := range []struct {
name string
ctx context.Context
timeout time.Duration
wantValue string
wantTimeout time.Duration
}{
{"defaults", nil, 0, "pool", time.Second},
{"context", caller, 0, "caller", time.Second},
{"timeout", nil, 2 * time.Second, "pool", 2 * time.Second},
{"both", caller, 2 * time.Second, "caller", 2 * time.Second},
} {
t.Run(tc.name, func(t *testing.T) {
started := time.Now()
err := pool.WithState(tc.ctx, tc.timeout, func(L *glua.LState) error {
ctx := L.Context()
if ctx.Value(key) != tc.wantValue {
t.Errorf("context value = %v, want %q", ctx.Value(key), tc.wantValue)
}
deadline, ok := ctx.Deadline()
if !ok || deadline.Before(started.Add(tc.wantTimeout)) || deadline.After(time.Now().Add(tc.wantTimeout)) {
t.Errorf("deadline = %v, want timeout %v", deadline, tc.wantTimeout)
}
return nil
})
if err != nil {
t.Fatal(err)
}
})
}
}
func TestPoolFactoryContext(t *testing.T) {
caller, cancel := context.WithTimeout(context.Background(), time.Minute)
defer cancel()
for _, tc := range []struct {
name string
ctx context.Context
}{
{"default", nil},
{"caller", caller},
} {
t.Run(tc.name, func(t *testing.T) {
var contexts []context.Context
pool := newTestPool(t, context.Background(), time.Second, func(ctx context.Context) (*glua.LState, error) {
contexts = append(contexts, ctx)
return glua.NewState(), nil
})
state, err := pool.Acquire(nil)
if err != nil {
t.Fatal(err)
}
defer pool.Release(state, true)
if err := pool.WithState(tc.ctx, 2*time.Second, func(*glua.LState) error { return nil }); err != nil {
t.Fatal(err)
}
want := tc.ctx
if want == nil {
want = pool.ctx
}
if len(contexts) != 2 || contexts[0] != pool.ctx || contexts[1] != want {
t.Fatal("factory did not receive the initialization and acquisition contexts unchanged")
}
})
}
}
func TestPoolWithStateLifecycle(t *testing.T) {
failure := errors.New("work failed")
for _, tc := range []struct {
name string
work func(*glua.LState, context.CancelFunc) error
reusable bool
wantPanic bool
wantErr error
}{
{"success", func(*glua.LState, context.CancelFunc) error { return nil }, true, false, nil},
{"canceled success", func(_ *glua.LState, cancel context.CancelFunc) error {
cancel()
return nil
}, true, false, nil},
{"error", func(*glua.LState, context.CancelFunc) error { return failure }, false, false, failure},
{"timeout", func(L *glua.LState, _ context.CancelFunc) error { return L.DoString("while true do end") }, false, false, nil},
{"panic", func(*glua.LState, context.CancelFunc) error { panic(failure) }, false, true, nil},
} {
t.Run(tc.name, func(t *testing.T) {
pool := newTestPool(t, context.Background(), 10*time.Millisecond, func(context.Context) (*glua.LState, error) {
state := glua.NewState()
state.Push(glua.LTrue)
return state, nil
})
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
var state *glua.LState
var workCtx context.Context
var recovered any
err := func() (err error) {
defer func() { recovered = recover() }()
return pool.WithState(ctx, 0, func(L *glua.LState) error {
state, workCtx = L, L.Context()
L.Push(glua.LFalse)
return tc.work(L, cancel)
})
}()
if tc.wantPanic {
if recovered != failure {
t.Fatalf("panic = %v, want original panic", recovered)
}
} else {
if recovered != nil || (err == nil) != tc.reusable {
t.Fatalf("WithState error = %v, panic = %v", err, recovered)
}
if tc.wantErr != nil && !errors.Is(err, tc.wantErr) {
t.Fatalf("WithState error = %v, want %v", err, tc.wantErr)
}
}
if workCtx.Err() == nil {
t.Fatal("WithState did not cancel the execution context")
}
if closed := state.IsClosed(); closed == tc.reusable {
t.Fatalf("state closed = %t, want %t", closed, !tc.reusable)
}
if tc.reusable && (state.Context() != nil || state.GetTop() != 1 || state.Get(1) != glua.LTrue) {
t.Fatal("WithState did not reset the state for reuse")
}
if err := pool.WithState(nil, 0, func(L *glua.LState) error {
if (L == state) != tc.reusable {
t.Error("unexpected state reuse")
}
return nil
}); err != nil {
t.Fatal(err)
}
})
}
}
func TestPoolClose(t *testing.T) {
pool := newTestPool(t, context.Background(), time.Second, func(context.Context) (*glua.LState, error) {
return glua.NewState(), nil
})
finishCtx, finish := context.WithCancel(context.Background())
t.Cleanup(finish)
started, done := make(chan *glua.LState, 1), make(chan error, 1)
var workCtx context.Context
go func() {
done <- pool.WithState(nil, 0, func(L *glua.LState) error {
workCtx = L.Context()
started <- L
<-finishCtx.Done()
return nil
})
}()
var state *glua.LState
select {
case state = <-started:
case <-time.After(time.Second):
t.Fatal("WithState did not start")
}
closed := make(chan struct{})
go func() {
pool.Close()
close(closed)
}()
select {
case <-workCtx.Done():
case <-time.After(time.Second):
t.Fatal("Close did not cancel work using the pool context")
}
if !errors.Is(workCtx.Err(), context.Canceled) {
t.Fatalf("work context error = %v, want context.Canceled", workCtx.Err())
}
assertPoolCloseBlocked(t, closed)
finish()
select {
case err := <-done:
if err != nil {
t.Fatalf("successful work returned an error: %v", err)
}
case <-time.After(time.Second):
t.Fatal("WithState did not finish")
}
select {
case <-closed:
case <-time.After(time.Second):
t.Fatal("Close did not finish after WithState")
}
if !state.IsClosed() {
t.Fatal("Release returned a state to a closed pool")
}
if state, err := pool.Acquire(nil); state != nil || err == nil || errors.Is(err, context.Canceled) {
t.Fatalf("Acquire after Close = %v, %v; want closed pool error", state, err)
}
pool.Close()
}
func TestPoolCloseWaitsForFactory(t *testing.T) {
finishCtx, finish := context.WithCancel(context.Background())
started, canceled := make(chan struct{}), make(chan struct{})
first := true
pool := newTestPool(t, context.Background(), time.Second, func(ctx context.Context) (*glua.LState, error) {
if first {
first = false
return glua.NewState(), nil
}
close(started)
<-ctx.Done()
close(canceled)
<-finishCtx.Done()
return nil, ctx.Err()
})
t.Cleanup(finish)
state, err := pool.Acquire(nil)
if err != nil {
t.Fatal(err)
}
pool.Release(state, false)
acquireDone := make(chan error, 1)
go func() {
_, err := pool.Acquire(nil)
acquireDone <- err
}()
select {
case <-started:
case <-time.After(time.Second):
t.Fatal("state creation did not start")
}
closed := make(chan struct{})
go func() {
pool.Close()
close(closed)
}()
select {
case <-canceled:
case <-time.After(time.Second):
t.Fatal("Close did not cancel state creation")
}
assertPoolCloseBlocked(t, closed)
finish()
select {
case err := <-acquireDone:
if !errors.Is(err, context.Canceled) {
t.Fatalf("Acquire error = %v, want context.Canceled", err)
}
case <-time.After(time.Second):
t.Fatal("state creation did not finish")
}
select {
case <-closed:
case <-time.After(time.Second):
t.Fatal("Close did not finish after state creation")
}
}
func TestPoolCloseWaitsForCallerContext(t *testing.T) {
pool := newTestPool(t, context.Background(), time.Minute, func(context.Context) (*glua.LState, error) {
return glua.NewState(), nil
})
ctx, cancel := context.WithCancel(context.Background())
t.Cleanup(cancel)
started, done := make(chan context.Context, 1), make(chan error, 1)
go func() {
done <- pool.WithState(ctx, 0, func(L *glua.LState) error {
started <- L.Context()
<-L.Context().Done()
return L.Context().Err()
})
}()
var workCtx context.Context
select {
case workCtx = <-started:
case <-time.After(time.Second):
t.Fatal("WithState did not start")
}
closed := make(chan struct{})
go func() {
pool.Close()
close(closed)
}()
select {
case <-pool.ctx.Done():
case <-time.After(time.Second):
t.Fatal("Close did not cancel the pool context")
}
assertPoolCloseBlocked(t, closed)
if workCtx.Err() != nil || ctx.Err() != nil {
t.Fatal("Close canceled the caller's execution context")
}
cancel()
select {
case err := <-done:
if !errors.Is(err, context.Canceled) {
t.Fatalf("WithState error = %v, want context.Canceled", err)
}
case <-time.After(time.Second):
t.Fatal("WithState did not stop after caller cancellation")
}
select {
case <-closed:
case <-time.After(time.Second):
t.Fatal("Close did not finish after WithState")
}
}
func BenchmarkPoolAcquireRelease(b *testing.B) {
pool := newTestPool(b, context.Background(), time.Second, func(context.Context) (*glua.LState, error) {
return glua.NewState(), nil
})
b.ReportAllocs()
b.ResetTimer()
for i := 0; i < b.N; i++ {
state, err := pool.Acquire(nil)
if err != nil {
b.Fatal(err)
}
pool.Release(state, true)
}
}
-77
View File
@@ -1,77 +0,0 @@
package lua
import (
"bufio"
"context"
"os"
"time"
glua "github.com/yuin/gopher-lua"
"github.com/yuin/gopher-lua/parse"
)
// Program holds immutable bytecode that can be run by independent LStates.
type Program struct {
proto *glua.FunctionProto
}
// LStateFactory returns a fully initialized state or nil and an error.
// Implementations must close partial states on failure; callers own successful states.
type LStateFactory func(context.Context) (*glua.LState, error)
// CompileFile reads and compiles a Lua file once.
func CompileFile(path string) (*Program, error) {
f, err := os.Open(path)
if err != nil {
return nil, err
}
defer f.Close()
chunk, err := parse.Parse(bufio.NewReader(f), path)
if err != nil {
return nil, err
}
proto, err := glua.Compile(chunk, path)
if err != nil {
return nil, err
}
return &Program{proto: proto}, nil
}
// NewState creates a state, runs register, executes the program under ctx, and
// runs validate. It removes the initialization context before returning a state
// owned by the caller.
func (p *Program) NewState(ctx context.Context, register func(*glua.LState), validate func(*glua.LState) error) (*glua.LState, error) {
L := glua.NewState()
valid := false
defer func() {
if !valid {
L.Close()
}
}()
L.SetContext(ctx)
defer L.RemoveContext()
if register != nil {
register(L)
}
L.Push(L.NewFunctionFromProto(p.proto))
// Execute the Lua script's top level.
if err := L.PCall(0, 0, nil); err != nil {
return nil, err
}
if validate != nil {
if err := validate(L); err != nil {
return nil, err
}
}
valid = true
return L, nil
}
// NewStateFactory returns a factory that gives each state an initialization timeout.
func (p *Program) NewStateFactory(initTimeout time.Duration, register func(*glua.LState), validate func(*glua.LState) error) LStateFactory {
return func(ctx context.Context) (*glua.LState, error) {
initCtx, cancel := context.WithTimeout(ctx, initTimeout)
defer cancel()
return p.NewState(initCtx, register, validate)
}
}
-76
View File
@@ -1,76 +0,0 @@
package lua
import (
"context"
"errors"
"os"
"path/filepath"
"testing"
glua "github.com/yuin/gopher-lua"
)
func TestProgramStatesAreIndependent(t *testing.T) {
path := filepath.Join(t.TempDir(), "state.lua")
if err := os.WriteFile(path, []byte("value = (value or 0) + 1"), 0o600); err != nil {
t.Fatal(err)
}
program, err := CompileFile(path)
if err != nil {
t.Fatal(err)
}
first, err := program.NewState(context.Background(), nil, nil)
if err != nil {
t.Fatal(err)
}
defer first.Close()
first.SetGlobal("value", glua.LNumber(42))
second, err := program.NewState(context.Background(), nil, nil)
if err != nil {
t.Fatal(err)
}
defer second.Close()
if got := second.GetGlobal("value"); got != glua.LNumber(1) {
t.Fatalf("second state value = %v, want 1", got)
}
}
func TestProgramInitializationObservesCancellation(t *testing.T) {
path := filepath.Join(t.TempDir(), "loop.lua")
if err := os.WriteFile(path, []byte("while true do end"), 0o600); err != nil {
t.Fatal(err)
}
program, err := CompileFile(path)
if err != nil {
t.Fatal(err)
}
ctx, cancel := context.WithCancel(context.Background())
cancel()
state, err := program.NewState(ctx, nil, nil)
if err == nil || state != nil {
if state != nil {
state.Close()
}
t.Fatalf("NewState with canceled context = %v, %v; want nil state and error", state, err)
}
}
func TestNewStateClosesFailedValidation(t *testing.T) {
path := filepath.Join(t.TempDir(), "state.lua")
if err := os.WriteFile(path, []byte("value = 1"), 0o600); err != nil {
t.Fatal(err)
}
program, err := CompileFile(path)
if err != nil {
t.Fatal(err)
}
wantErr := errors.New("invalid script")
var checked *glua.LState
L, err := program.NewState(context.Background(), nil, func(L *glua.LState) error {
checked = L
return wantErr
})
if L != nil || !errors.Is(err, wantErr) || checked == nil || !checked.IsClosed() {
t.Fatalf("state = %v, error = %v, checked state closed = %t", L, err, checked != nil && checked.IsClosed())
}
}
-95
View File
@@ -1,95 +0,0 @@
package lua
import (
"math"
"github.com/xtls/xray-core/common/errors"
glua "github.com/yuin/gopher-lua"
)
type number interface {
~int | ~int8 | ~int16 | ~int32 | ~int64 |
~uint | ~uint8 | ~uint16 | ~uint32 | ~uint64 | ~uintptr |
~float32 | ~float64
}
// PushNumber converts a Go number to a Lua number and pushes it.
func PushNumber[T number](L *glua.LState, value T) {
L.Push(glua.LNumber(value))
}
// PushString converts a Go string to a Lua string and pushes it.
func PushString(L *glua.LState, value string) {
L.Push(glua.LString(value))
}
// PushNil pushes Lua nil.
func PushNil(L *glua.LState) {
L.Push(glua.LNil)
}
// PushUserData pushes a native Go value without copying it.
func PushUserData(L *glua.LState, value any) {
ud := L.NewUserData()
ud.Value = value
L.Push(ud)
}
// PushError pushes nil or the original Go error as userdata.
func PushError(L *glua.LState, err error) {
if err == nil {
L.Push(glua.LNil)
return
}
PushUserData(L, err)
}
// ReadUserData reads a native Go value of type T without copying it.
// Other Lua values or userdata containing a different type return invalidMessage.
func ReadUserData[T any](value glua.LValue, invalidMessage string) (T, error) {
if ud, ok := value.(*glua.LUserData); ok {
if result, ok := ud.Value.(T); ok {
return result, nil
}
}
var zero T
return zero, errors.New(invalidMessage)
}
// ReadError accepts nil, a native Go error, or a Lua string.
// Native errors retain their identity; other values return invalidMessage.
func ReadError(value glua.LValue, invalidMessage string) error {
if value == glua.LNil {
return nil
}
if ud, ok := value.(*glua.LUserData); ok {
if err, ok := ud.Value.(error); ok {
return err
}
}
if message, ok := value.(glua.LString); ok {
return errors.New(string(message))
}
return errors.New(invalidMessage)
}
// ReadUint32 accepts only integral Lua numbers in the uint32 range.
func ReadUint32(value glua.LValue, invalidMessage string) (uint32, error) {
number, ok := value.(glua.LNumber)
if !ok || number < 0 || number > math.MaxUint32 || math.Trunc(float64(number)) != float64(number) {
return 0, errors.New(invalidMessage)
}
return uint32(number), nil
}
// ReadOptionalString accepts a Lua string or nil, which becomes an empty string.
// It does not coerce other values to strings.
func ReadOptionalString(value glua.LValue, invalidMessage string) (string, error) {
if value == glua.LNil {
return "", nil
}
if result, ok := value.(glua.LString); ok {
return string(result), nil
}
return "", errors.New(invalidMessage)
}
-121
View File
@@ -1,121 +0,0 @@
package lua
import (
"errors"
"math"
"strings"
"testing"
glua "github.com/yuin/gopher-lua"
)
func TestReadUint32(t *testing.T) {
for _, tc := range []struct {
name string
value glua.LValue
want uint32
wantErr bool
}{
{name: "zero", value: glua.LNumber(0)},
{name: "integer", value: glua.LNumber(45), want: 45},
{name: "maximum", value: glua.LNumber(math.MaxUint32), want: math.MaxUint32},
{name: "fraction", value: glua.LNumber(1.5), wantErr: true},
{name: "negative", value: glua.LNumber(-1), wantErr: true},
{name: "overflow", value: glua.LNumber(math.MaxUint32 + 1), wantErr: true},
{name: "NaN", value: glua.LNumber(math.NaN()), wantErr: true},
{name: "positive infinity", value: glua.LNumber(math.Inf(1)), wantErr: true},
{name: "negative infinity", value: glua.LNumber(math.Inf(-1)), wantErr: true},
{name: "nil", value: glua.LNil, wantErr: true},
{name: "numeric string", value: glua.LString("45"), wantErr: true},
{name: "boolean", value: glua.LTrue, wantErr: true},
} {
t.Run(tc.name, func(t *testing.T) {
got, err := ReadUint32(tc.value, "invalid number")
if got != tc.want || (err != nil) != tc.wantErr {
t.Fatalf("ReadUint32() = %d, %v; want %d, error %t", got, err, tc.want, tc.wantErr)
}
if err != nil && !strings.Contains(err.Error(), "invalid number") {
t.Fatalf("error = %v, want invalid number", err)
}
})
}
}
func TestReadOptionalString(t *testing.T) {
for _, tc := range []struct {
name string
value glua.LValue
want string
wantErr bool
}{
{name: "nil", value: glua.LNil},
{name: "empty", value: glua.LString("")},
{name: "string", value: glua.LString("out"), want: "out"},
{name: "number", value: glua.LNumber(1), wantErr: true},
{name: "boolean", value: glua.LFalse, wantErr: true},
} {
t.Run(tc.name, func(t *testing.T) {
got, err := ReadOptionalString(tc.value, "invalid string")
if got != tc.want || (err != nil) != tc.wantErr {
t.Fatalf("ReadOptionalString() = %q, %v; want %q, error %t", got, err, tc.want, tc.wantErr)
}
if err != nil && !strings.Contains(err.Error(), "invalid string") {
t.Fatalf("error = %v, want invalid string", err)
}
})
}
}
func TestUserDataRoundTrip(t *testing.T) {
L := glua.NewState()
defer L.Close()
want := []int{1, 2}
PushUserData(L, want)
if L.GetTop() != 1 {
t.Fatalf("stack top = %d, want 1", L.GetTop())
}
got, err := ReadUserData[[]int](L.Get(-1), "invalid userdata")
if err != nil || len(got) != len(want) || &got[0] != &want[0] {
t.Fatalf("userdata = %v, %v; want original slice", got, err)
}
PushUserData(L, []int(nil))
if got, err := ReadUserData[[]int](L.Get(-1), "invalid userdata"); err != nil || got != nil {
t.Fatalf("nil slice userdata = %v, %v", got, err)
}
for _, value := range []glua.LValue{glua.LNil, glua.LString("1"), L.NewTable(), L.Get(1)} {
if got, err := ReadUserData[int](value, "invalid userdata"); got != 0 || err == nil || !strings.Contains(err.Error(), "invalid userdata") {
t.Fatalf("ReadUserData(%v) = %d, %v; want invalid userdata", value, got, err)
}
}
}
func TestErrorRoundTrip(t *testing.T) {
L := glua.NewState()
defer L.Close()
want := errors.New("upstream failed")
for _, err := range []error{nil, want} {
PushError(L, err)
if L.GetTop() != 1 {
t.Fatalf("stack top = %d, want 1", L.GetTop())
}
if err == nil && L.Get(-1) != glua.LNil {
t.Fatalf("nil error pushed as %v", L.Get(-1))
}
if got := ReadError(L.Get(-1), "invalid error"); got != err {
t.Fatalf("ReadError() = %v, want original error %v", got, err)
}
L.Pop(1)
}
for _, message := range []string{"script failed", ""} {
if err := ReadError(glua.LString(message), "invalid error"); err == nil || !strings.Contains(err.Error(), message) {
t.Fatalf("string error = %v, want %q", err, message)
}
}
wrong := L.NewUserData()
wrong.Value = "not a native error"
for _, value := range []glua.LValue{glua.LTrue, glua.LNumber(1), L.NewTable(), wrong, L.NewUserData()} {
if err := ReadError(value, "invalid error"); err == nil || !strings.Contains(err.Error(), "invalid error") {
t.Fatalf("ReadError(%v) = %v, want invalid error", value, err)
}
}
}
+3 -1
View File
@@ -314,10 +314,12 @@ func (m *ClientWorker) Dispatch(ctx context.Context, link *transport.Link) bool
}
sm := m.sessionManager
s := sm.Allocate(&m.strategy, link.Reader, link.Writer)
s := sm.Allocate(&m.strategy)
if s == nil {
return false
}
s.input = link.Reader
s.output = link.Writer
go fetchInput(ctx, s, m.link.Writer)
if _, ok := link.Reader.(*pipe.Reader); !ok {
select {
-47
View File
@@ -7,11 +7,9 @@ import (
"github.com/golang/mock/gomock"
"github.com/xtls/xray-core/common"
"github.com/xtls/xray-core/common/buf"
"github.com/xtls/xray-core/common/errors"
"github.com/xtls/xray-core/common/mux"
"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/testing/mocks"
"github.com/xtls/xray-core/transport"
@@ -116,48 +114,3 @@ func TestClientWorkerClose(t *testing.T) {
common.Must(w2.Close())
}
func TestClientWorkerUDPSource(t *testing.T) {
downR, downW := pipe.New(pipe.WithoutSizeLimit())
upR, upW := pipe.New(pipe.WithoutSizeLimit())
worker, err := mux.NewClientWorker(transport.Link{Reader: downR, Writer: upW}, mux.ClientStrategy{})
common.Must(err)
inR, inW := pipe.New(pipe.WithoutSizeLimit())
outR, outW := pipe.New(pipe.WithoutSizeLimit())
ctx := session.ContextWithOutbounds(context.Background(), []*session.Outbound{{
Target: net.UDPDestination(net.ParseAddress("8.8.8.8"), 53),
}})
if !worker.Dispatch(ctx, &transport.Link{Reader: inR, Writer: outW}) {
t.Fatal("failed to dispatch")
}
b := buf.New()
b.WriteString("query")
common.Must(inW.WriteMultiBuffer(buf.MultiBuffer{b}))
mb, err := upR.ReadMultiBuffer() // New frame, the session is UDP from now on
common.Must(err)
buf.ReleaseMulti(mb)
srcs := []net.Destination{
net.UDPDestination(net.ParseAddress("1.1.1.1"), 1111),
net.UDPDestination(net.DomainAddress("example.com"), 2222),
net.UDPDestination(net.ParseAddress("3.3.3.3"), 3333),
}
w := mux.NewResponseWriter(1, downW, protocol.TransferTypePacket)
var got buf.MultiBuffer
for i := range srcs {
b := buf.New()
b.WriteString("reply")
b.UDP = &srcs[i]
common.Must(w.WriteMultiBuffer(buf.MultiBuffer{b}))
// keep earlier replies around while the next frame is parsed
mb, err := outR.ReadMultiBuffer()
common.Must(err)
got = append(got, mb...)
}
for i, b := range got {
if b.UDP == nil || *b.UDP != srcs[i] {
t.Errorf("reply %d: source = %v, want %v", i, b.UDP, srcs[i])
}
}
}
+4 -5
View File
@@ -14,16 +14,15 @@ import (
type PacketReader struct {
reader io.Reader
eof bool
dest net.Destination
dest *net.Destination
}
// NewPacketReader creates a new PacketReader.
// dest is copied because the caller reuses it for the next frame.
func NewPacketReader(reader io.Reader, dest *net.Destination) *PacketReader {
return &PacketReader{
reader: reader,
eof: false,
dest: *dest,
dest: dest,
}
}
@@ -48,8 +47,8 @@ func (r *PacketReader) ReadMultiBuffer() (buf.MultiBuffer, error) {
return nil, err
}
r.eof = true
if r.dest.Network == net.Network_UDP {
b.UDP = &r.dest // only one packet is read, so b owns r.dest
if r.dest != nil && r.dest.Network == net.Network_UDP {
b.UDP = r.dest
}
return buf.MultiBuffer{b}, nil
}
+1 -3
View File
@@ -51,7 +51,7 @@ func (m *SessionManager) Count() int {
return int(m.count)
}
func (m *SessionManager) Allocate(Strategy *ClientStrategy, input buf.Reader, output buf.Writer) *Session {
func (m *SessionManager) Allocate(Strategy *ClientStrategy) *Session {
m.Lock()
defer m.Unlock()
@@ -64,8 +64,6 @@ func (m *SessionManager) Allocate(Strategy *ClientStrategy, input buf.Reader, ou
m.count++
s := &Session{
input: input,
output: output,
ID: m.count,
parent: m,
done: done.New(),
+3 -3
View File
@@ -9,7 +9,7 @@ import (
func TestSessionManagerAdd(t *testing.T) {
m := NewSessionManager()
s := m.Allocate(&ClientStrategy{}, nil, nil)
s := m.Allocate(&ClientStrategy{})
if s.ID != 1 {
t.Error("id: ", s.ID)
}
@@ -17,7 +17,7 @@ func TestSessionManagerAdd(t *testing.T) {
t.Error("size: ", m.Size())
}
s = m.Allocate(&ClientStrategy{}, nil, nil)
s = m.Allocate(&ClientStrategy{})
if s.ID != 2 {
t.Error("id: ", s.ID)
}
@@ -39,7 +39,7 @@ func TestSessionManagerAdd(t *testing.T) {
func TestSessionManagerClose(t *testing.T) {
m := NewSessionManager()
s := m.Allocate(&ClientStrategy{}, nil, nil)
s := m.Allocate(&ClientStrategy{})
if m.CloseIfNoSessionAndIdle(m.Size(), m.Count()) {
t.Error("able to close")
-20
View File
@@ -1,20 +0,0 @@
package net
// PacketConnWrapper wraps a PacketConn into a Conn with a fixed destination address.
type PacketConnWrapper struct {
PacketConn
Dest Addr
}
func (c *PacketConnWrapper) Read(p []byte) (int, error) {
n, _, err := c.PacketConn.ReadFrom(p)
return n, err
}
func (c *PacketConnWrapper) Write(p []byte) (int, error) {
return c.PacketConn.WriteTo(p, c.Dest)
}
func (c *PacketConnWrapper) RemoteAddr() Addr {
return c.Dest
}
-48
View File
@@ -1,8 +1,6 @@
package platform // import "github.com/xtls/xray-core/common/platform"
import (
"errors"
"fmt"
"os"
"path/filepath"
"strconv"
@@ -92,49 +90,3 @@ func GetConfDirPath() string {
configPath := NewEnvFlag(ConfdirLocation).GetValue(func() string { return "" })
return configPath
}
// ResolveLuaFile finds a local Lua script and returns its absolute path.
// Relative paths: XRAY_LOCATION_CONFDIR > XRAY_LOCATION_CONFIG > working dir > executable dir.
func ResolveLuaFile(path string) (string, error) {
if path == "" {
return "", errors.New("Lua file path is empty")
}
paths := []string{path}
if !filepath.IsAbs(path) {
paths = nil
for _, dir := range []string{
GetConfDirPath(),
NewEnvFlag(ConfigLocation).GetValue(func() string { return "" }),
".",
getExecutableDir(),
} {
if dir != "" {
paths = append(paths, filepath.Join(dir, path))
}
}
}
return resolveFile(paths)
}
func resolveFile(paths []string) (string, error) {
var tried []string
for _, path := range paths {
path, err := filepath.Abs(path)
if err != nil {
return "", fmt.Errorf("failed to resolve file path: %w", err)
}
tried = append(tried, path)
info, err := os.Stat(path)
if errors.Is(err, os.ErrNotExist) {
continue
}
if err != nil {
return "", fmt.Errorf("failed to inspect file %q: %w", path, err)
}
if !info.Mode().IsRegular() {
return "", fmt.Errorf("file is not a regular file: %s", path)
}
return path, nil
}
return "", fmt.Errorf("file not found; tried %q: %w", tried, os.ErrNotExist)
}
-51
View File
@@ -1,7 +1,6 @@
package platform_test
import (
"errors"
"os"
"path/filepath"
"runtime"
@@ -65,53 +64,3 @@ func TestGetAssetLocation(t *testing.T) {
}
}
}
func TestResolveLuaFile(t *testing.T) {
workingDir := t.TempDir()
t.Chdir(workingDir)
executable, err := os.Executable()
common.Must(err)
file, err := os.CreateTemp(filepath.Dir(executable), "lua-*.lua")
common.Must(err)
common.Must(file.Close())
defer os.Remove(file.Name())
name := filepath.Base(file.Name())
paths := []string{
filepath.Join(t.TempDir(), name),
filepath.Join(t.TempDir(), name),
filepath.Join(workingDir, name),
file.Name(),
}
t.Setenv(ConfdirLocation, filepath.Dir(paths[0]))
t.Setenv(ConfigLocation, filepath.Dir(paths[1]))
for _, path := range paths[:3] {
common.Must(os.WriteFile(path, nil, 0o600))
}
if got, err := ResolveLuaFile(paths[2]); err != nil || got != paths[2] {
t.Fatalf("absolute path = %q, %v; want %q", got, err, paths[2])
}
for i, want := range paths {
if i == 2 {
t.Setenv(ConfdirLocation, "")
t.Setenv(ConfigLocation, "")
}
if got, err := ResolveLuaFile(name); err != nil || got != want {
t.Fatalf("resolved path = %q, %v; want %q", got, err, want)
}
common.Must(os.Remove(want))
}
if _, err := ResolveLuaFile(name); !errors.Is(err, os.ErrNotExist) {
t.Fatalf("missing file error = %v", err)
}
t.Setenv(ConfdirLocation, filepath.Dir(paths[0]))
t.Setenv(ConfigLocation, filepath.Dir(paths[1]))
common.Must(os.Mkdir(paths[0], 0o700))
common.Must(os.WriteFile(paths[1], nil, 0o600))
for _, path := range []string{"", name, filepath.Join(t.TempDir(), name)} {
if _, err := ResolveLuaFile(path); err == nil {
t.Fatalf("accepted invalid path %q", path)
}
}
}
-4
View File
@@ -146,10 +146,6 @@ func SniffQUIC(b []byte) (*SniffHeader, error) {
}
restPayload := b[hdrLen+int(packetLen):]
// cachedReader can concatenate zero-padded UDP datagrams.
for len(restPayload) > 0 && restPayload[0] == 0 {
restPayload = restPayload[1:]
}
if !isQUICInitial { // Skip this packet if it's not initial packet
b = restPayload
continue
File diff suppressed because one or more lines are too long
+1 -1
View File
@@ -62,7 +62,7 @@ func (s *Service) Cleanup() error {
}
for name, subs := range s.subs {
newSub := make([]*Subscriber, 0, len(subs))
newSub := make([]*Subscriber, 0, len(s.subs))
for _, sub := range subs {
if !sub.IsClosed() {
newSub = append(newSub, sub)
+2 -2
View File
@@ -19,8 +19,8 @@ import (
var (
Version_x byte = 26
Version_y byte = 10
Version_z byte = 10
Version_y byte = 9
Version_z byte = 9
)
var (
+1 -3
View File
@@ -92,8 +92,6 @@ type Instance struct {
// Instance state
func (server *Instance) IsRunning() bool {
server.statusLock.Lock()
defer server.statusLock.Unlock()
return server.running
}
@@ -322,7 +320,7 @@ func (s *Instance) RequireFeatures(callback interface{}, optional bool) error {
// AddFeature registers a feature into current Instance.
func (s *Instance) AddFeature(feature features.Feature) error {
if s.IsRunning() {
if s.running {
if err := feature.Start(); err != nil {
errors.LogInfoInner(s.ctx, err, "failed to start feature")
}
-3
View File
@@ -21,7 +21,6 @@ require (
github.com/stretchr/testify v1.12.1
github.com/vishvananda/netlink v1.3.1
github.com/xtls/reality v0.0.0-20260908062103-8cdf7bf9c7f0
github.com/yuin/gopher-lua v1.1.2
go4.org/netipx v0.0.0-20231129151722-fdeea329fbba
golang.org/x/crypto v0.57.0
golang.org/x/exp v0.0.0-20240506185415-9bf2ced13842
@@ -35,7 +34,6 @@ require (
google.golang.org/protobuf v1.36.12
gvisor.dev/gvisor v0.0.0-20260122175437-89a5d21be8f0
h12.io/socks v1.0.3
layeh.com/gopher-luar v1.0.11
lukechampine.com/blake3 v1.4.1
mvdan.cc/gofumpt v0.12.0
)
@@ -57,7 +55,6 @@ require (
github.com/vishvananda/netns v0.0.5 // indirect
github.com/wlynxg/anet v0.0.5 // indirect
go.yaml.in/yaml/v3 v3.0.5 // indirect
golang.org/x/crypto/x509roots/fallback v0.0.0-20261005185213-c3db4df58582 // indirect
golang.org/x/text v0.42.0 // indirect
golang.org/x/time v0.14.0 // indirect
golang.org/x/tools v0.49.0 // indirect
-11
View File
@@ -2,9 +2,6 @@ github.com/andybalholm/brotli v1.0.6 h1:Yf9fFpf49Zrxb9NlQaluyE92/+X7UVHlhMNJN2sx
github.com/andybalholm/brotli v1.0.6/go.mod h1:fO7iG3H7G2nSZ7m0zPUDn85XEX2GTukHGRSepvi9Eig=
github.com/apernet/quic-go v0.61.1-0.20260806010916-184d081eef3e h1:5mgtR5gwIgBKMiGI1QdXldZZ+SNor06Nbu1wCBulQBg=
github.com/apernet/quic-go v0.61.1-0.20260806010916-184d081eef3e/go.mod h1:x7qxEvX6MCVtDuBKHj3E+88+BtrbEMuAL5qGUKItjW8=
github.com/chzyer/logex v1.1.10/go.mod h1:+Ywpsq7O8HXn0nuIou7OrIPyXbp3wmkHB+jjWRnGsAI=
github.com/chzyer/readline v0.0.0-20180603132655-2972be24d48e/go.mod h1:nSuG5e5PlCu98SY8svDHJxuZscDgtXS6KTTbou5AhLI=
github.com/chzyer/test v0.0.0-20180213035817-a1ea475d72b1/go.mod h1:Q3SI9o4m/ZMnBNeIyt5eFwwo7qiLfzFZmjNmxjkiQlU=
github.com/cloudflare/circl v1.6.5 h1:O64F26HEqNhznd/hrC5KZXVKYuKM2rx4deZDTc4ihQA=
github.com/cloudflare/circl v1.6.5/go.mod h1:h5LNyxAc5nTue9DS5jT+48en2PSDYt3zdGnz5OstK6c=
github.com/ghodss/yaml v1.0.1-0.20220118164431-d8423dcdf344 h1:Arcl6UOIS/kgO2nW3A65HN+7CMjSDP/gofXL4CZt1V4=
@@ -84,9 +81,6 @@ github.com/wlynxg/anet v0.0.5/go.mod h1:eay5PRQr7fIVAMbTbchTnO9gG65Hg/uYGdc7mguH
github.com/xtls/reality v0.0.0-20260908062103-8cdf7bf9c7f0 h1:rb+fKQFhz+5I2PPuQsNYxI5mUU840XWYtRF0ZBjvkws=
github.com/xtls/reality v0.0.0-20260908062103-8cdf7bf9c7f0/go.mod h1:DsJblcWDGt76+FVqBVwbwRhxyyNJsGV48gJLch0OOWI=
github.com/yuin/goldmark v1.4.1/go.mod h1:mwnBkeHKe2W/ZEtQ+71ViKU8L12m81fl3OWwC1Zlc8k=
github.com/yuin/gopher-lua v0.0.0-20190206043414-8bfc7677f583/go.mod h1:gqRgreBUhTSL0GeU64rtZ3Uq3wtjOa/TB2YfrtkCbVQ=
github.com/yuin/gopher-lua v1.1.2 h1:yF/FjE3hD65tBbt0VXLE13HWS9h34fdzJmrWRXwobGA=
github.com/yuin/gopher-lua v1.1.2/go.mod h1:7aRmXIWl37SqRf0koeyylBEzJ+aPt8A+mmkQ4f1ntR8=
go.uber.org/mock v0.5.2 h1:LbtPTcP8A5k9WPXj54PPPbjcI4Y6lhyOZXn+VS7wNko=
go.uber.org/mock v0.5.2/go.mod h1:wLlUxC2vVTPTaE3UD51E0BGOAElKrILxhVSDYQLld5o=
go.yaml.in/yaml/v3 v3.0.5 h1:N6y/pJk8buWs9NY5ERU2HSMfm+IuD/OtfdAnq6kESPw=
@@ -97,8 +91,6 @@ golang.org/x/crypto v0.0.0-20190308221718-c2843e01d9a2/go.mod h1:djNgcEr1/C05ACk
golang.org/x/crypto v0.0.0-20191011191535-87dc89f01550/go.mod h1:yigFU9vqHzYiE8UmvKecakEJjdnWj3jj499lnFckfCI=
golang.org/x/crypto v0.57.0 h1:3ZVCjf8Ggz7zneR/EHRVx68Ctf+2pmIMP2UFhh9cC6M=
golang.org/x/crypto v0.57.0/go.mod h1:Fdz0i5U6CoizGwLda9DttjSk6qlZo25zYNtR+ycvuZA=
golang.org/x/crypto/x509roots/fallback v0.0.0-20261005185213-c3db4df58582 h1:wjDBrGbfLuifgrVLFEWUBJYAfh5Q1wkMc5t0FY0tCbs=
golang.org/x/crypto/x509roots/fallback v0.0.0-20261005185213-c3db4df58582/go.mod h1:HPze8vhfG6fO06AM+VSvxRm4E3+5Yk375mgrJ5M2z1E=
golang.org/x/exp v0.0.0-20240506185415-9bf2ced13842 h1:vr/HnozRka3pE4EsMEg1lgkXJkTFJCVUX+S/ZT6wYzM=
golang.org/x/exp v0.0.0-20240506185415-9bf2ced13842/go.mod h1:XtvwrStGgqGPLc4cjQfWqZHG1YFdYs6swckp8vpsjnc=
golang.org/x/lint v0.0.0-20200302205851-738671d3881b/go.mod h1:3xt1FjdF8hUf6vQPIChWIBhFzV8gjjsPE/fR3IyQdNY=
@@ -113,7 +105,6 @@ golang.org/x/sync v0.0.0-20190423024810-112230192c58/go.mod h1:RxMgew5VJxzue5/jJ
golang.org/x/sync v0.0.0-20210220032951-036812b2e83c/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
golang.org/x/sync v0.23.0 h1:KameEIfc1IkluZyXWLn39Wd4tURc6GbCiISGiZm2bQk=
golang.org/x/sync v0.23.0/go.mod h1:sUUOizhqBxiL6pEWpqNLUiaJn1ShEbZ6BBqskPbjZm0=
golang.org/x/sys v0.0.0-20190204203706-41f3e6584952/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY=
golang.org/x/sys v0.0.0-20190215142949-d0b11bdaac8a/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY=
golang.org/x/sys v0.0.0-20190412213103-97732733099d/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs=
golang.org/x/sys v0.0.0-20201119102817-f84b799fce68/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs=
@@ -164,8 +155,6 @@ gvisor.dev/gvisor v0.0.0-20260122175437-89a5d21be8f0 h1:Lk6hARj5UPY47dBep70OD/TI
gvisor.dev/gvisor v0.0.0-20260122175437-89a5d21be8f0/go.mod h1:QkHjoMIBaYtpVufgwv3keYAbln78mBoCuShZrPrer1Q=
h12.io/socks v1.0.3 h1:Ka3qaQewws4j4/eDQnOdpr4wXsC//dXtWvftlIcCQUo=
h12.io/socks v1.0.3/go.mod h1:AIhxy1jOId/XCz9BO+EIgNL2rQiPTBNnOfnVnQ+3Eck=
layeh.com/gopher-luar v1.0.11 h1:8zJudpKI6HWkoh9eyyNFaTM79PY6CAPcIr6X/KTiliw=
layeh.com/gopher-luar v1.0.11/go.mod h1:TPnIVCZ2RJBndm7ohXyaqfhzjlZ+OA2SZR/YwL8tECk=
lukechampine.com/blake3 v1.4.1 h1:I3Smz7gso8w4/TunLKec6K2fn+kyKtDxr/xcQEN84Wg=
lukechampine.com/blake3 v1.4.1/go.mod h1:QFosUxmjB8mnrWFSNwKmvxHpfY72bmD2tQ0kBMM3kwo=
mvdan.cc/gofumpt v0.12.0 h1:1Lbudkz2kpM9Cjz2pL4M19u7q+GaEhCTNf7N9mfpcho=
-14
View File
@@ -14,11 +14,9 @@ import (
"github.com/xtls/xray-core/common/errors"
"github.com/xtls/xray-core/common/geodata"
"github.com/xtls/xray-core/common/net"
"github.com/xtls/xray-core/common/platform"
)
type NameServerConfig struct {
ID string `json:"id"`
Address *Address `json:"address"`
ClientIP *Address `json:"clientIp"`
Port uint16 `json:"port"`
@@ -45,7 +43,6 @@ func (c *NameServerConfig) UnmarshalJSON(data []byte) error {
}
var advanced struct {
ID string `json:"id"`
Address *Address `json:"address"`
ClientIP *Address `json:"clientIp"`
Port uint16 `json:"port"`
@@ -63,7 +60,6 @@ func (c *NameServerConfig) UnmarshalJSON(data []byte) error {
UnexpectedIPs StringList `json:"unexpectedIPs"`
}
if err := json.Unmarshal(data, &advanced); err == nil {
c.ID = advanced.ID
c.Address = advanced.Address
c.ClientIP = advanced.ClientIP
c.Port = advanced.Port
@@ -138,7 +134,6 @@ func (c *NameServerConfig) Build() (*dns.NameServer, error) {
}
return &dns.NameServer{
Id: c.ID,
Address: &net.Endpoint{
Network: net.Network_UDP,
Address: c.Address.Build(),
@@ -164,7 +159,6 @@ func (c *NameServerConfig) Build() (*dns.NameServer, error) {
// DNSConfig is a JSON serializable object for dns.Config
type DNSConfig struct {
Servers []*NameServerConfig `json:"servers"`
Script string `json:"script"`
Hosts *HostsWrapper `json:"hosts"`
ClientIP *Address `json:"clientIp"`
Tag string `json:"tag"`
@@ -284,14 +278,6 @@ func (c *DNSConfig) Build() (*dns.Config, error) {
QueryStrategy: resolveQueryStrategy(c.QueryStrategy),
}
if c.Script != "" {
path, err := platform.ResolveLuaFile(c.Script)
if err != nil {
return nil, errors.New("failed to resolve DNS script: ", c.Script).Base(err)
}
config.Script = path
}
if c.ClientIP != nil {
if !c.ClientIP.Family().IsIP() {
return nil, errors.New("not an IP address:", c.ClientIP.String())
-50
View File
@@ -2,8 +2,6 @@ package conf_test
import (
"encoding/json"
"os"
"path/filepath"
"testing"
"github.com/google/go-cmp/cmp"
@@ -124,51 +122,3 @@ func TestDNSConfigParsing(t *testing.T) {
}
}
}
func TestDNSScriptConfig(t *testing.T) {
dir := t.TempDir()
t.Setenv("xray.location.confdir", dir)
path := filepath.Join(dir, "lookup.lua")
if err := os.WriteFile(path, []byte("function HandleDNSQuery(domain, ipv4, ipv6, fake) end"), 0o600); err != nil {
t.Fatal(err)
}
for _, tc := range []struct {
name string
script string
wantError bool
}{
{"relative", "lookup.lua", false},
{"absolute", path, false},
{"missing", "missing.lua", true},
{"directory", dir, true},
} {
t.Run(tc.name, func(t *testing.T) {
built, err := (&DNSConfig{Script: tc.script}).Build()
if tc.wantError {
if err == nil {
t.Fatal("Build accepted an invalid script path")
}
return
}
if err != nil {
t.Fatal(err)
}
if built.Script != path {
t.Fatalf("script path = %q, want %q", built.Script, path)
}
})
}
var parsed DNSConfig
if err := json.Unmarshal([]byte(`{"servers":[{"id":"primary","address":"1.1.1.1"}]}`), &parsed); err != nil {
t.Fatal(err)
}
built, err := parsed.Build()
if err != nil {
t.Fatal(err)
}
if len(built.NameServer) != 1 || built.NameServer[0].Id != "primary" {
t.Fatalf("nameserver IDs = %v, want primary", built.NameServer)
}
}
-122
View File
@@ -1,14 +1,7 @@
package conf_test
import (
"crypto/ecdsa"
"crypto/ed25519"
"crypto/elliptic"
"crypto/rand"
"crypto/x509"
"encoding/base64"
"encoding/json"
"encoding/pem"
"testing"
"github.com/xtls/xray-core/common/protocol"
@@ -74,121 +67,6 @@ func TestMasqueConfig(t *testing.T) {
}
}
func TestMasqueWarpConfig(t *testing.T) {
creator := func() Buildable {
return new(MasqueConfig)
}
key, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader)
if err != nil {
t.Fatal(err)
}
pkcs8, err := x509.MarshalPKCS8PrivateKey(key)
if err != nil {
t.Fatal(err)
}
sec1, err := x509.MarshalECPrivateKey(key)
if err != nil {
t.Fatal(err)
}
quote := func(s string) string {
b, _ := json.Marshal(s)
return string(b)
}
server, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader)
if err != nil {
t.Fatal(err)
}
publicKey, err := x509.MarshalPKIXPublicKey(&server.PublicKey)
if err != nil {
t.Fatal(err)
}
publicPEM := string(pem.EncodeToMemory(&pem.Block{Type: "PUBLIC KEY", Bytes: publicKey}))
warpInput := func(key string, extra string) string {
return `{` + extra + `"warp": {"privateKey": ` + quote(key) + `, "publicKey": ` + quote(publicPEM) + `, "address": ["172.16.0.2", "2606:4700:110:8a36::2/128"]}}`
}
address := []string{"172.16.0.2/32", "2606:4700:110:8a36::2/128"}
for _, input := range []string{
string(pem.EncodeToMemory(&pem.Block{Type: "PRIVATE KEY", Bytes: pkcs8})),
string(pem.EncodeToMemory(&pem.Block{Type: "EC PRIVATE KEY", Bytes: sec1})),
base64.StdEncoding.EncodeToString(pkcs8),
base64.StdEncoding.EncodeToString(sec1),
} {
runMultiTestCase(t, []TestCase{
{
Input: warpInput(input, ""),
Parser: loadJSON(creator),
Output: &masque.Config{
Host: "cloudflareaccess.com",
Path: "/",
Warp: &masque.Warp{PrivateKey: pkcs8, PublicKey: publicKey, Address: address},
},
},
})
}
runMultiTestCase(t, []TestCase{
{
Input: warpInput(base64.StdEncoding.EncodeToString(sec1), `"host": "example.com", "path": "/warp", `),
Parser: loadJSON(creator),
Output: &masque.Config{
Host: "example.com",
Path: "/warp",
Warp: &masque.Warp{PrivateKey: pkcs8, PublicKey: publicKey, Address: address},
},
},
})
p384, err := ecdsa.GenerateKey(elliptic.P384(), rand.Reader)
if err != nil {
t.Fatal(err)
}
p384DER, err := x509.MarshalPKCS8PrivateKey(p384)
if err != nil {
t.Fatal(err)
}
ed, err := x509.MarshalPKCS8PrivateKey(ed25519.NewKeyFromSeed(make([]byte, ed25519.SeedSize)))
if err != nil {
t.Fatal(err)
}
withAddress := func(address string) string {
return `{"warp": {"privateKey": ` + quote(base64.StdEncoding.EncodeToString(pkcs8)) + `, "publicKey": ` + quote(publicPEM) + `, "address": ` + address + `}}`
}
withPublicKey := func(key string) string {
return `{"warp": {"privateKey": ` + quote(base64.StdEncoding.EncodeToString(pkcs8)) + `, "publicKey": ` + quote(key) + `, "address": ["172.16.0.2"]}}`
}
runMultiTestCase(t, []TestCase{
{
Input: withPublicKey(base64.StdEncoding.EncodeToString(publicKey)),
Parser: loadJSON(creator),
Output: &masque.Config{
Host: "cloudflareaccess.com",
Path: "/",
Warp: &masque.Warp{PrivateKey: pkcs8, PublicKey: publicKey, Address: []string{"172.16.0.2/32"}},
},
},
})
for _, input := range []string{
`{"warp": {}}`,
withAddress(`[]`),
withPublicKey(""),
withPublicKey("not a key"),
withPublicKey(base64.StdEncoding.EncodeToString([]byte("not a key"))),
withPublicKey(base64.StdEncoding.EncodeToString(pkcs8)),
withAddress(`["172.16.0"]`),
withAddress(`["172.16.0.2", "172.16.0.3"]`),
withAddress(`["2606:4700::1", "2606:4700::2/128"]`),
warpInput("not a key", ""),
warpInput(base64.StdEncoding.EncodeToString([]byte("not a key")), ""),
warpInput(base64.StdEncoding.EncodeToString(p384DER), ""),
warpInput(base64.StdEncoding.EncodeToString(ed), ""),
warpInput(base64.StdEncoding.EncodeToString(pkcs8), `"user": "u", "pass": "p", `),
} {
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)
-11
View File
@@ -7,7 +7,6 @@ import (
"github.com/xtls/xray-core/app/router"
"github.com/xtls/xray-core/common/errors"
"github.com/xtls/xray-core/common/geodata"
"github.com/xtls/xray-core/common/platform"
"github.com/xtls/xray-core/common/serial"
"google.golang.org/protobuf/proto"
@@ -73,7 +72,6 @@ type RouterConfig struct {
RuleList []json.RawMessage `json:"rules"`
DomainStrategy *string `json:"domainStrategy"`
Balancers []*BalancingRule `json:"balancers"`
Script string `json:"script"`
}
func (c *RouterConfig) getDomainStrategy() router.Config_DomainStrategy {
@@ -94,15 +92,6 @@ func (c *RouterConfig) getDomainStrategy() router.Config_DomainStrategy {
func (c *RouterConfig) Build() (*router.Config, error) {
config := new(router.Config)
if c.Script != "" {
path, err := platform.ResolveLuaFile(c.Script)
if err != nil {
return nil, errors.New("failed to resolve routing script").Base(err)
}
config.Script = path
}
config.DomainStrategy = c.getDomainStrategy()
var rawRuleList []json.RawMessage
-38
View File
@@ -2,8 +2,6 @@ package conf_test
import (
"encoding/json"
"os"
"path/filepath"
"testing"
"time"
_ "unsafe"
@@ -238,39 +236,3 @@ func TestRouterConfig(t *testing.T) {
},
})
}
func TestRouterScriptConfig(t *testing.T) {
dir := t.TempDir()
t.Setenv("xray.location.confdir", dir)
path := filepath.Join(dir, "route.lua")
if err := os.WriteFile(path, []byte("function HandleRoute() end"), 0o600); err != nil {
t.Fatal(err)
}
for _, tc := range []struct {
name string
script string
wantError bool
}{
{"relative", "route.lua", false},
{"absolute", path, false},
{"missing", "missing.lua", true},
{"directory", dir, true},
} {
t.Run(tc.name, func(t *testing.T) {
built, err := (&RouterConfig{Script: tc.script}).Build()
if tc.wantError {
if err == nil {
t.Fatal("Build accepted invalid script path")
}
return
}
if err != nil {
t.Fatal(err)
}
if built.Script != path {
t.Fatalf("script path = %q, want %q", built.Script, path)
}
})
}
}
+25 -168
View File
@@ -1,7 +1,6 @@
package conf
import (
"context"
"crypto/x509"
"encoding/base64"
"encoding/hex"
@@ -82,7 +81,7 @@ var (
"noise": func() interface{} { return new(NoiseMask) },
"salamander": func() interface{} { return new(Salamander) },
"sudoku": func() interface{} { return new(Sudoku) },
"xdns": func() interface{} { return new(XDNS) },
"xdns": func() interface{} { return new(Xdns) },
"xicmp": func() interface{} { return new(Xicmp) },
"realm": func() interface{} { return new(Realm) },
"udphop": func() interface{} { return new(UDPHop) },
@@ -309,27 +308,14 @@ type NoiseMask struct {
}
func (c *NoiseMask) Build() (proto.Message, error) {
noiseSlice := make([]*noise.Item, 0, len(c.Noise))
for _, item := range c.Noise {
if len(item.Packet) > 0 && item.Rand.To > 0 {
return nil, errors.New("len(item.Packet) > 0 && item.Rand.To > 0")
}
if strings.ToLower(item.Type) == "exp" {
var exp string
if err := json.Unmarshal(item.Packet, &exp); err != nil {
return nil, errors.New(`"packet" of noise "type": "exp" must be a string`).Base(err)
}
segments, err := parseNoiseExp(exp)
if err != nil {
return nil, err
}
noiseSlice = append(noiseSlice, &noise.Item{
Segments: segments,
DelayMin: int64(item.Delay.From),
DelayMax: int64(item.Delay.To),
})
continue
}
}
noiseSlice := make([]*noise.Item, 0, len(c.Noise))
for _, item := range c.Noise {
if item.RandRange == nil {
item.RandRange = &Int32Range{From: 0, To: 255}
}
@@ -358,88 +344,6 @@ func (c *NoiseMask) Build() (proto.Message, error) {
}, nil
}
var noiseExpPattern = regexp.MustCompile(`<\s*([a-z]+)(?:\s+([^>]*?))?\s*>`)
func parseNoiseExp(exp string) ([]*noise.Segment, error) {
var segments []*noise.Segment
matches := noiseExpPattern.FindAllStringSubmatchIndex(exp, -1)
last := 0
for _, m := range matches {
if strings.TrimSpace(exp[last:m[0]]) != "" {
return nil, errors.New("invalid noise exp near ", exp[last:m[0]])
}
last = m[1]
key := exp[m[2]:m[3]]
arg := ""
if m[4] >= 0 {
arg = exp[m[4]:m[5]]
}
segment, err := buildNoiseSegment(key, arg)
if err != nil {
return nil, err
}
segments = append(segments, segment)
}
if strings.TrimSpace(exp[last:]) != "" {
return nil, errors.New("invalid noise exp near ", exp[last:])
}
if len(segments) == 0 {
return nil, errors.New("empty noise exp: ", exp)
}
return segments, nil
}
func buildNoiseSegment(key, arg string) (*noise.Segment, error) {
sizeSegment := func(kind noise.Segment_Kind) (*noise.Segment, error) {
if arg == "" {
return nil, errors.New("<", key, "> in noise exp needs a size")
}
lo, hi, err := ParseRangeString(arg)
if err != nil {
return nil, err
}
if lo < 0 || hi < lo || hi > 65535 {
return nil, errors.New("invalid size in noise exp: ", arg)
}
return &noise.Segment{Kind: kind, MinSize: int64(lo), MaxSize: int64(hi)}, nil
}
switch key {
case "b":
hexStr := strings.TrimPrefix(strings.TrimPrefix(strings.Join(strings.Fields(arg), ""), "0x"), "0X")
if len(hexStr) == 0 {
return nil, errors.New("empty bytes in noise exp")
}
raw, err := hex.DecodeString(hexStr)
if err != nil {
return nil, errors.New("invalid hex in noise exp: ", arg).Base(err)
}
return &noise.Segment{Kind: noise.Segment_BYTES, Bytes: raw}, nil
case "r":
return sizeSegment(noise.Segment_RANDOM)
case "rc":
return sizeSegment(noise.Segment_RANDOM_ASCII)
case "rd":
return sizeSegment(noise.Segment_RANDOM_DIGIT)
case "t":
if arg != "" {
return nil, errors.New("<t> in noise exp takes no argument")
}
return &noise.Segment{Kind: noise.Segment_TIMESTAMP}, nil
case "c":
if arg != "" {
return nil, errors.New("<c> in noise exp takes no argument")
}
return &noise.Segment{Kind: noise.Segment_COUNTER}, nil
case "n":
if arg != "" {
return nil, errors.New("<n> in noise exp takes no argument")
}
return &noise.Segment{Kind: noise.Segment_NONCE}, nil
default:
return nil, errors.New("unknown <", key, "> in noise exp")
}
}
type UDPItem struct {
Rand int32 `json:"rand"`
RandRange *Int32Range `json:"randRange"`
@@ -790,79 +694,32 @@ func (c *Sudoku) Build() (proto.Message, error) {
}, nil
}
type XDNSDomain struct {
Names []string `json:"names"`
LenLimit int32 `json:"lenLimit"`
LabelLimit int32 `json:"labelLimit"`
Types []int32 `json:"types"`
Edns0 int32 `json:"edns0"`
type Xdns struct {
Domain json.RawMessage `json:"domain"`
Domains []string `json:"domains"`
Resolvers []string `json:"resolvers"`
}
type XDNSResolver struct {
Addrs []string `json:"addrs"`
}
func (c *Xdns) Build() (proto.Message, error) {
if c.Domain != nil {
return nil, errors.PrintRemovedFeatureError("domain", "domains(server) & resolvers(client)")
}
type XDNS struct {
Domains []XDNSDomain `json:"domains"`
Resolvers []XDNSResolver `json:"resolvers"`
ExtraPoll int32 `json:"extraPoll"`
}
if len(c.Domains) == 0 && len(c.Resolvers) == 0 {
return nil, errors.New("empty domains & empty resolvers")
}
func (c *XDNS) Build() (proto.Message, error) {
var domains []*xdns.DomainProto
var resolvers []*xdns.ResolverProto
for i := range c.Domains {
if c.Domains[i].LenLimit == 0 {
c.Domains[i].LenLimit = 255
}
if c.Domains[i].LabelLimit == 0 {
c.Domains[i].LabelLimit = 63
}
for j := range c.Domains[i].Names {
domain, err := xdns.NewDomain(c.Domains[i].Names[j], int(c.Domains[i].LenLimit), int(c.Domains[i].LabelLimit), []uint16{1, 5, 16, 28}, uint16(c.Domains[i].Edns0))
if err != nil {
return nil, err
}
errors.LogInfo(context.Background(), domain.Show())
domains = append(domains, &xdns.DomainProto{
Name: c.Domains[i].Names[j],
LenLimit: c.Domains[i].LenLimit,
LabelLimit: c.Domains[i].LabelLimit,
Types: c.Domains[i].Types,
Edns0: c.Domains[i].Edns0,
})
for _, r := range c.Resolvers {
if !strings.Contains(r, "+udp://") {
return nil, errors.New("invalid resolver ", r)
}
}
for i := range c.Resolvers {
for j := range c.Resolvers[i].Addrs {
var u *url.URL
var e error
if !strings.Contains(c.Resolvers[i].Addrs[j], "://") {
u, e = url.Parse("udp://" + c.Resolvers[i].Addrs[j])
} else {
u, e = url.Parse(c.Resolvers[i].Addrs[j])
}
if e != nil {
return nil, e
}
switch u.Scheme {
case "tcp", "udp":
default:
return nil, errors.New("invalid protocol")
}
var host, port string
host = u.Hostname()
port = u.Port()
if port == "" {
port = "53"
}
resolvers = append(resolvers, &xdns.ResolverProto{Type: u.Scheme, Addr: net.JoinHostPort(host, port)})
}
}
if c.ExtraPoll < 0 || c.ExtraPoll > 3 {
return nil, errors.New("c.ExtraPoll < 0 || c.ExtraPoll > 3")
}
return &xdns.Config{Domains: domains, Resolvers: resolvers, ExtraPoll: c.ExtraPoll}, nil
return &xdns.Config{
Domains: c.Domains,
Resolvers: c.Resolvers,
}, nil
}
type XMC struct {
@@ -1,136 +0,0 @@
package conf
import (
"encoding/json"
"testing"
"github.com/xtls/xray-core/transport/internet/finalmask/noise"
)
func expPacket(exp string) json.RawMessage {
b, _ := json.Marshal(exp)
return b
}
func buildNoiseExp(exp string) (*noise.Config, error) {
msg, err := (&NoiseMask{Noise: []NoiseItem{{Type: "exp", Packet: expPacket(exp)}}}).Build()
if err != nil {
return nil, err
}
return msg.(*noise.Config), nil
}
func TestNoiseExp(t *testing.T) {
cfg, err := buildNoiseExp("<b 0d0a0d0a><t><r 24><rc 20-40><rd 8><c><n>")
if err != nil {
t.Fatal(err)
}
segments := cfg.Items[0].Segments
if len(segments) != 7 {
t.Fatalf("got %d segments, want 7", len(segments))
}
want := []struct {
kind noise.Segment_Kind
bytes []byte
min, max int64
}{
{noise.Segment_BYTES, []byte{0x0d, 0x0a, 0x0d, 0x0a}, 0, 0},
{noise.Segment_TIMESTAMP, nil, 0, 0},
{noise.Segment_RANDOM, nil, 24, 24},
{noise.Segment_RANDOM_ASCII, nil, 20, 40},
{noise.Segment_RANDOM_DIGIT, nil, 8, 8},
{noise.Segment_COUNTER, nil, 0, 0},
{noise.Segment_NONCE, nil, 0, 0},
}
for i, w := range want {
s := segments[i]
if s.Kind != w.kind || s.MinSize != w.min || s.MaxSize != w.max || string(s.Bytes) != string(w.bytes) {
t.Errorf("segment %d = %+v, want %+v", i, s, w)
}
}
}
func TestNoiseExpStripsHexPrefix(t *testing.T) {
cfg, err := buildNoiseExp("<b 0x16030100>")
if err != nil {
t.Fatal(err)
}
if got := cfg.Items[0].Segments[0].Bytes; string(got) != string([]byte{0x16, 0x03, 0x01, 0x00}) {
t.Errorf("got %x", got)
}
}
func TestNoiseExpWhitespace(t *testing.T) {
if _, err := buildNoiseExp(" <b 00> <t> "); err != nil {
t.Errorf("surrounding whitespace should be allowed: %v", err)
}
cfg, err := buildNoiseExp("<b 0d 0a 0d 0a>")
if err != nil {
t.Fatal(err)
}
if got := cfg.Items[0].Segments[0].Bytes; string(got) != "\r\n\r\n" {
t.Errorf("got %x", got)
}
}
func TestNoiseExpRejects(t *testing.T) {
for _, exp := range []string{
"<x 1>",
"<b>",
"<b zz>",
"<b 0d0>",
"<r>",
"<r -1>",
"<r 40-20>",
"<r 70000>",
"<t 5>",
"<n 5>",
"garbage<t>",
"<t> tail",
"<t><b>",
} {
if _, err := buildNoiseExp(exp); err == nil {
t.Errorf("expected an error for %q", exp)
}
}
}
func TestNoiseExpConflicts(t *testing.T) {
if _, err := (&NoiseMask{Noise: []NoiseItem{{Type: "exp", Packet: expPacket("<t>"), Rand: Int32Range{From: 10, To: 20}}}}).Build(); err == nil {
t.Error("exp with rand should be rejected")
}
for _, packet := range []string{``, `[1, 2]`, `5`} {
if _, err := (&NoiseMask{Noise: []NoiseItem{{Type: "exp", Packet: json.RawMessage(packet)}}}).Build(); err == nil {
t.Errorf("expected an error for packet %q", packet)
}
}
}
func TestNoiseExpFromJSON(t *testing.T) {
var mask NoiseMask
if err := json.Unmarshal([]byte(`{"noise": [
{"type": "exp", "packet": "<b 504f5354><rd 10-20>", "delay": "1-3"},
{"type": "EXP", "packet": "<t>"},
{"type": "str", "packet": "<t>"},
{"rand": "10-20"}
]}`), &mask); err != nil {
t.Fatal(err)
}
msg, err := mask.Build()
if err != nil {
t.Fatal(err)
}
items := msg.(*noise.Config).Items
if len(items[0].Segments) != 2 || items[0].DelayMin != 1 || items[0].DelayMax != 3 {
t.Errorf("item 0 = %+v", items[0])
}
if len(items[1].Segments) != 1 || items[1].Segments[0].Kind != noise.Segment_TIMESTAMP {
t.Errorf("item 1 = %+v", items[1])
}
if len(items[2].Segments) != 0 || string(items[2].Packet) != "<t>" {
t.Errorf("item 2 = %+v", items[2])
}
if len(items[3].Segments) != 0 || items[3].RandMin != 10 || items[3].RandMax != 20 {
t.Errorf("item 3 = %+v", items[3])
}
}
+4 -113
View File
@@ -1,15 +1,10 @@
package conf
import (
"crypto/ecdsa"
"crypto/elliptic"
"crypto/x509"
"encoding/base64"
"encoding/json"
"encoding/pem"
"maps"
"math/big"
"net/netip"
"net/url"
"sort"
"strconv"
@@ -795,49 +790,16 @@ func (c *HysteriaConfig) Build() (proto.Message, error) {
return config, nil
}
type MasqueWarpConfig struct {
PrivateKey string `json:"privateKey"`
PublicKey string `json:"publicKey"`
Address []string `json:"address"`
}
type MasqueConfig struct {
Host string `json:"host"`
Path string `json:"path"`
User string `json:"user"`
Pass string `json:"pass"`
Headers map[string]string `json:"headers"`
Warp *MasqueWarpConfig `json:"warp"`
}
func (c *MasqueConfig) Build() (proto.Message, error) {
var warp *masque.Warp
host := c.Host
path := c.Path
if c.Warp != nil {
if c.User != "" || c.Pass != "" {
return nil, errors.New(`"user" and "pass" can't be used with "warp"`)
}
key, err := parseWarpPrivateKey(c.Warp.PrivateKey)
if err != nil {
return nil, errors.New(`invalid "privateKey" in "warp"`).Base(err)
}
publicKey, err := parseWarpPublicKey(c.Warp.PublicKey)
if err != nil {
return nil, errors.New(`invalid "publicKey" in "warp"`).Base(err)
}
address, err := parseWarpAddress(c.Warp.Address)
if err != nil {
return nil, err
}
warp = &masque.Warp{PrivateKey: key, PublicKey: publicKey, Address: address}
if host == "" {
host = masque.WarpHost
}
if path == "" {
path = masque.WarpPath
}
}
if path == "" {
path = masque.DefaultPath
}
@@ -849,9 +811,9 @@ func (c *MasqueConfig) Build() (proto.Message, error) {
if !strings.HasPrefix(path, "/") || strings.ContainsAny(path, "{}") {
return nil, errors.New(`invalid "path": `, path, `, only the variables {target} and {ipproto} are supported`)
}
if host != "" {
if u, err := url.Parse("https://" + host); err != nil || u.Host != host {
return nil, errors.New(`invalid "host": `, host)
if c.Host != "" {
if u, err := url.Parse("https://" + c.Host); err != nil || u.Host != c.Host {
return nil, errors.New(`invalid "host": `, c.Host)
}
}
for k, v := range c.Headers {
@@ -879,83 +841,12 @@ func (c *MasqueConfig) Build() (proto.Message, error) {
headers["Authorization"] = "Basic " + base64.StdEncoding.EncodeToString([]byte(c.User+":"+c.Pass))
}
return &masque.Config{
Host: host,
Host: c.Host,
Path: path,
Headers: headers,
Warp: warp,
}, nil
}
func parseWarpAddress(list []string) ([]string, error) {
if len(list) == 0 {
return nil, errors.New(`"address" in "warp" is not set`)
}
var v4, v6 bool
address := make([]string, 0, len(list))
for _, s := range list {
prefix, err := netip.ParsePrefix(s)
if err != nil {
addr, err := netip.ParseAddr(s)
if err != nil {
return nil, errors.New(`invalid "address" in "warp": `, s)
}
prefix = netip.PrefixFrom(addr, addr.BitLen())
}
if prefix.Addr().Is4() && v4 || prefix.Addr().Is6() && v6 {
return nil, errors.New(`"address" in "warp" takes at most one IPv4 and one IPv6 address`)
}
v4 = v4 || prefix.Addr().Is4()
v6 = v6 || prefix.Addr().Is6()
address = append(address, prefix.String())
}
return address, nil
}
func decodeWarpKey(s string) ([]byte, error) {
s = strings.TrimSpace(s)
if s == "" {
return nil, errors.New("empty key")
}
if block, _ := pem.Decode([]byte(s)); block != nil {
return block.Bytes, nil
}
der, err := base64.StdEncoding.DecodeString(s)
if err != nil {
return nil, errors.New("neither PEM nor base64").Base(err)
}
return der, nil
}
func parseWarpPublicKey(s string) ([]byte, error) {
der, err := decodeWarpKey(s)
if err != nil {
return nil, err
}
if _, err := x509.ParsePKIXPublicKey(der); err != nil {
return nil, errors.New("not a PKIX public key").Base(err)
}
return der, nil
}
func parseWarpPrivateKey(s string) ([]byte, error) {
der, err := decodeWarpKey(s)
if err != nil {
return nil, err
}
var key any
key, err = x509.ParsePKCS8PrivateKey(der)
if err != nil {
if key, err = x509.ParseECPrivateKey(der); err != nil {
return nil, errors.New("neither a PKCS #8 nor a SEC 1 private key")
}
}
ecKey, ok := key.(*ecdsa.PrivateKey)
if !ok || ecKey.Curve != elliptic.P256() {
return nil, errors.New("not an ECDSA P-256 key")
}
return x509.MarshalPKCS8PrivateKey(ecKey)
}
func readFileOrString(f string, s []string) ([]byte, error) {
if len(f) > 0 {
return filesystem.ReadCert(f)
-2
View File
@@ -316,7 +316,6 @@ type TLSConfig struct {
ECHServerKeys string `json:"echServerKeys"`
ECHConfigList string `json:"echConfigList"`
ECHSocketSettings *SocketConfig `json:"echSockopt"`
UseSystemCA bool `json:"useSystemCA"`
}
// Build implements Buildable.
@@ -404,7 +403,6 @@ func (c *TLSConfig) Build() (proto.Message, error) {
}
config.EchSocketSettings = ss
}
config.UseSystemCa = c.UseSystemCA
return config, nil
}
@@ -12,7 +12,6 @@ import (
"github.com/xtls/xray-core/core"
"github.com/xtls/xray-core/infra/conf"
"github.com/xtls/xray-core/infra/conf/serial"
"github.com/xtls/xray-core/proxy/hysteria"
"github.com/xtls/xray-core/proxy/masque"
"github.com/xtls/xray-core/proxy/shadowsocks"
"github.com/xtls/xray-core/proxy/shadowsocks_2022"
@@ -92,8 +91,6 @@ func extractInboundUsers(inb *core.InboundHandlerConfig) []*protocol.User {
return ty.Users
case *masque.ServerConfig:
return ty.Users
case *hysteria.ServerConfig:
return ty.Users
default:
fmt.Println("unsupported inbound type")
}
+4 -4
View File
@@ -467,7 +467,7 @@ func NewPacketReader(conn net.Conn, h *Handler, defaultRule *FinalRule, UDPOverr
if statConn != nil {
counter = statConn.ReadCounter
}
if c, ok := iConn.(*net.PacketConnWrapper); ok {
if c, ok := iConn.(*internet.PacketConnWrapper); ok {
isOverridden := false
if UDPOverride.Address != nil || UDPOverride.Port != 0 {
isOverridden = true
@@ -487,7 +487,7 @@ func NewPacketReader(conn net.Conn, h *Handler, defaultRule *FinalRule, UDPOverr
}
type PacketReader struct {
*net.PacketConnWrapper
*internet.PacketConnWrapper
stats.Counter
Handler *Handler
DefaultRule *FinalRule
@@ -542,7 +542,7 @@ func NewPacketWriter(conn net.Conn, h *Handler, defaultRule *FinalRule, UDPOverr
if statConn != nil {
counter = statConn.WriteCounter
}
if c, ok := iConn.(*net.PacketConnWrapper); ok {
if c, ok := iConn.(*internet.PacketConnWrapper); ok {
// If DialDest is a domain, it will be resolved in dialer
// check this behavior and add it to map
resolvedUDPAddr := utils.NewTypedSyncMap[string, net.Address]()
@@ -563,7 +563,7 @@ func NewPacketWriter(conn net.Conn, h *Handler, defaultRule *FinalRule, UDPOverr
}
type PacketWriter struct {
*net.PacketConnWrapper
*internet.PacketConnWrapper
stats.Counter
*Handler
DefaultRule *FinalRule
+1 -1
View File
@@ -151,7 +151,7 @@ func (c *Client) Process(ctx context.Context, link *transport.Link, dialer inter
}
defer conn.Close()
uc := &wireguard.UDPConnClient{
PacketConn: conn.(*net.PacketConnWrapper).PacketConn,
PacketConn: conn.(*internet.PacketConnWrapper).PacketConn,
Dest: conn.RemoteAddr().(*net.UDPAddr),
}
reader = uc
-15
View File
@@ -277,7 +277,6 @@ func (w *VisionReader) ReadMultiBuffer() (buf.MultiBuffer, error) {
w.ob.CanSpliceCopy = 1
}
}
SuppressOuterCloseNotify(w.conn)
readerConn, readCounter, _ := UnwrapRawConn(w.conn)
w.directReadCounter = readCounter
w.Reader = buf.NewReader(readerConn)
@@ -341,7 +340,6 @@ func (w *VisionWriter) WriteMultiBuffer(mb buf.MultiBuffer) error {
// w.ob.CanSpliceCopy = 1
// }
}
SuppressOuterCloseNotify(w.conn)
rawConn, _, writerCounter := UnwrapRawConn(w.conn)
w.Writer = buf.NewWriter(rawConn)
w.directWriteCounter = writerCounter
@@ -671,19 +669,6 @@ func XtlsFilterTls(buffer buf.MultiBuffer, trafficState *TrafficState, ctx conte
}
}
type CloseNotifySuppressor interface {
SuppressCloseNotify()
}
// Close our local TLS conn instance might send a incorrect close_notify alert
// if the XTLS direct copy mode is enabled and cause TLS BAD_RECORD_MAC on users' browser
// Close the underlying connection directly to avoid this issue.
func SuppressOuterCloseNotify(conn net.Conn) {
if suppressor, ok := stat.TryUnwrapStatsConn(conn).(CloseNotifySuppressor); ok {
suppressor.SuppressCloseNotify()
}
}
// UnwrapRawConn support unwrap encryption, stats, mask wrappers, tls, utls, reality, proxyproto, uds-wrapper conn and get raw tcp/uds conn from it
func UnwrapRawConn(conn net.Conn) (net.Conn, stats.Counter, stats.Counter) {
var readCounter, writerCounter stats.Counter
+118 -38
View File
@@ -2,6 +2,7 @@ package shadowsocks_2022
import (
"context"
"io"
"time"
"github.com/xtls/xray-core/common"
@@ -12,6 +13,9 @@ import (
"github.com/xtls/xray-core/common/net"
"github.com/xtls/xray-core/common/protocol"
"github.com/xtls/xray-core/common/session"
"github.com/xtls/xray-core/common/signal"
"github.com/xtls/xray-core/common/task"
"github.com/xtls/xray-core/common/utils"
"github.com/xtls/xray-core/core"
"github.com/xtls/xray-core/features/policy"
"github.com/xtls/xray-core/features/routing"
@@ -97,29 +101,35 @@ func (i *Inbound) processTCP(ctx context.Context, conn net.Conn, dispatcher rout
return errors.New("unable to set read deadline").Base(err)
}
// 1. Single read call for Salt + Fixed-length header chunk per SIP022 §3.1.4
headerLen := i.method.KeySaltLength + RequestHeaderFixedChunkLength + AEADTagSize
headerBuf := make([]byte, headerLen)
n, err := conn.Read(headerBuf)
if err != nil || n < headerLen {
ResetTCPConn(conn)
return errors.New("failed to read complete handshake header")
}
var salt [32]byte
copy(salt[:i.method.KeySaltLength], headerBuf[:i.method.KeySaltLength])
saltSlice := salt[:i.method.KeySaltLength]
fixedChunk := headerBuf[i.method.KeySaltLength:]
reader, reqHeader, err := InitServerStream(conn, i.method, i.psk, saltSlice, salt, fixedChunk, i.saltFilter)
if err != nil {
ResetTCPConn(conn)
if _, err := io.ReadFull(conn, saltSlice); err != nil {
return err
}
if !i.saltFilter.Check(salt) {
return ErrSaltNotUnique
}
sessionKey := DeriveSessionSubKey(i.psk, saltSlice, i.method.KeySaltLength)
aead, err := i.method.NewAEAD(sessionKey)
if err != nil {
return err
}
reader := NewStreamReader(conn, aead)
reqHeader, err := ReadClientRequestHeader(conn, reader)
if err != nil {
return err
}
conn.SetReadDeadline(time.Time{})
dest := reqHeader.Destination
writer := NewServerStreamWriter(conn, i.method, i.psk, saltSlice)
writer, err := WriteTCPResponse(conn, i.method, i.psk, saltSlice, nil)
if err != nil {
return err
}
ctx = log.ContextWithAccessMessage(ctx, &log.AccessMessage{
From: conn.RemoteAddr(),
@@ -136,17 +146,42 @@ func (i *Inbound) processTCP(ctx context.Context, conn net.Conn, dispatcher rout
}
if len(reqHeader.EarlyData) > 0 {
mb := buf.MergeBytes(nil, reqHeader.EarlyData)
if err := link.Writer.WriteMultiBuffer(mb); err != nil {
earlyBuf := buf.New()
earlyBuf.Write(reqHeader.EarlyData)
if err := link.Writer.WriteMultiBuffer(buf.MultiBuffer{earlyBuf}); err != nil {
return err
}
}
return TransportTCP(ctx, i.policyManager.ForLevel(uint32(i.user.Level)), reader, writer, link)
sessionPolicy = i.policyManager.ForLevel(uint32(i.user.Level))
ctx, cancel := context.WithCancel(ctx)
timer := signal.CancelAfterInactivity(ctx, cancel, sessionPolicy.Timeouts.ConnectionIdle)
ctx = policy.ContextWithBufferPolicy(ctx, sessionPolicy.Buffer)
requestDone := func() error {
defer timer.SetTimeout(sessionPolicy.Timeouts.DownlinkOnly)
return buf.Copy(reader, link.Writer, buf.UpdateActivity(timer))
}
responseDone := func() error {
defer timer.SetTimeout(sessionPolicy.Timeouts.UplinkOnly)
return buf.Copy(link.Reader, writer, buf.UpdateActivity(timer))
}
responseDoneAndCloseWriter := task.OnSuccess(responseDone, task.Close(link.Writer))
return task.Run(ctx, requestDone, responseDoneAndCloseWriter)
}
func (i *Inbound) processUDP(ctx context.Context, conn stat.Connection, dispatcher routing.Dispatcher) error {
reader := buf.NewPacketReader(conn)
udpConns := utils.NewTypedSyncMap[uint64, *udpConnEntry]()
defer func() {
udpConns.Range(func(key uint64, entry *udpConnEntry) bool {
entry.timer.SetTimeout(0)
return true
})
}()
reader := buf.NewReader(conn)
for {
mb, err := reader.ReadMultiBuffer()
if err != nil {
@@ -156,30 +191,75 @@ func (i *Inbound) processUDP(ctx context.Context, conn stat.Connection, dispatch
for _, b := range mb {
decoded, err := i.udpCodec.DecodePacket(b.Bytes())
b.Release()
if err != nil || decoded.HeaderType != HeaderTypeClient {
continue
}
sessionItem := i.udpCodec.GetSession(decoded.SessionID)
if sessionItem.User == nil {
sessionItem.Lock()
if sessionItem.User == nil {
sessionItem.User = i.user
}
sessionItem.Unlock()
}
link, err := sessionItem.EnsureLink(ctx, conn, decoded.Destination, dispatcher, i.policyManager, func(dest net.Destination, payload []byte) ([]byte, error) {
return i.udpCodec.EncodeServerPacket(decoded.SessionID, dest, payload)
})
if err != nil {
b.Release()
continue
}
entry, ok := udpConns.Load(decoded.SessionID)
if !ok {
sessCtx, cancel := context.WithCancel(ctx)
sessCtx = log.ContextWithAccessMessage(sessCtx, &log.AccessMessage{
From: conn.RemoteAddr(),
To: decoded.Destination,
Status: log.AccessAccepted,
Email: i.user.Email,
})
link, err := dispatcher.Dispatch(sessCtx, decoded.Destination)
if err != nil {
cancel()
b.Release()
continue
}
newEntry := &udpConnEntry{
link: link,
cancel: cancel,
}
sessionPolicy := i.policyManager.ForLevel(uint32(i.user.Level))
newEntry.timer = signal.CancelAfterInactivity(sessCtx, func() {
udpConns.Delete(decoded.SessionID)
common.Interrupt(link.Reader)
common.Interrupt(link.Writer)
cancel()
}, sessionPolicy.Timeouts.ConnectionIdle)
actual, loaded := udpConns.LoadOrStore(decoded.SessionID, newEntry)
if loaded {
// Another goroutine/packet beat us to storing, terminate our redundant link
newEntry.timer.SetTimeout(0)
entry = actual
} else {
entry = newEntry
go func(sessID uint64, dest net.Destination, cEntry *udpConnEntry) {
defer func() {
cEntry.timer.SetTimeout(0)
}()
for {
resMb, err := cEntry.link.Reader.ReadMultiBuffer()
if err != nil {
return
}
cEntry.timer.Update()
for _, rb := range resMb {
encPacket, err := i.udpCodec.EncodeServerPacket(sessID, dest, rb.Bytes())
rb.Release()
if err != nil {
continue
}
_, _ = conn.Write(encPacket)
}
}
}(decoded.SessionID, decoded.Destination, entry)
}
}
entry.timer.Update()
payloadBuf := buf.New()
payloadBuf.Write(decoded.Payload)
payloadBuf.UDP = &decoded.Destination
_ = link.Writer.WriteMultiBuffer(buf.MultiBuffer{payloadBuf})
b.Release()
_ = entry.link.Writer.WriteMultiBuffer(buf.MultiBuffer{payloadBuf})
}
}
}
+201 -48
View File
@@ -4,6 +4,7 @@ import (
"context"
"crypto/cipher"
"encoding/binary"
"io"
"strconv"
"strings"
"sync"
@@ -18,6 +19,8 @@ import (
"github.com/xtls/xray-core/common/net"
"github.com/xtls/xray-core/common/protocol"
"github.com/xtls/xray-core/common/session"
"github.com/xtls/xray-core/common/signal"
"github.com/xtls/xray-core/common/task"
"github.com/xtls/xray-core/common/utils"
"github.com/xtls/xray-core/common/uuid"
"github.com/xtls/xray-core/core"
@@ -204,46 +207,64 @@ func (i *MultiUserInbound) processTCP(ctx context.Context, conn net.Conn, dispat
return errors.New("unable to set read deadline").Base(err)
}
// 1. Single read call for Salt + EIH + Fixed-length header chunk per SIP022 §3.1.4
headerLen := i.method.KeySaltLength + AESBlockSize + RequestHeaderFixedChunkLength + AEADTagSize
headerBuf := make([]byte, headerLen)
n, err := conn.Read(headerBuf)
if err != nil || n < headerLen {
ResetTCPConn(conn)
return errors.New("failed to read complete handshake header")
}
// 1. Read Request Salt (16 or 32 bytes)
var salt [32]byte
copy(salt[:i.method.KeySaltLength], headerBuf[:i.method.KeySaltLength])
saltSlice := salt[:i.method.KeySaltLength]
eih := headerBuf[i.method.KeySaltLength : i.method.KeySaltLength+AESBlockSize]
fixedChunk := headerBuf[i.method.KeySaltLength+AESBlockSize:]
decryptedHash, err := DecryptEIH(i.method, i.masterPSK, saltSlice, eih)
if err != nil {
ResetTCPConn(conn)
if _, err := io.ReadFull(conn, saltSlice); err != nil {
return err
}
if !i.saltFilter.Check(salt) {
return ErrSaltNotUnique
}
// 2. Read Extended Identity Header (16 bytes)
var eih [AESBlockSize]byte
if _, err := io.ReadFull(conn, eih[:]); err != nil {
return err
}
// Decrypt EIH with IdentitySubKey derived from masterPSK and salt
identitySubkey := DeriveIdentitySubKey(i.masterPSK, saltSlice, i.method.KeySaltLength)
block, err := i.method.NewBlock(identitySubkey)
if err != nil {
return err
}
var decryptedHash [AESBlockSize]byte
block.Decrypt(decryptedHash[:], eih[:])
// Lookup user
user, ok := i.usersByHash.Load(decryptedHash)
if !ok {
ResetTCPConn(conn)
if !ok || user == nil {
return ErrInvalidRequest
}
userPSK := user.Account.(*MemoryAccount).Key
reader, reqHeader, err := InitServerStream(conn, i.method, userPSK, saltSlice, salt, fixedChunk, i.saltFilter)
// 3. Derive Session Subkey using matched user's PSK
sessionKey := DeriveSessionSubKey(userPSK, saltSlice, i.method.KeySaltLength)
aead, err := i.method.NewAEAD(sessionKey)
if err != nil {
ResetTCPConn(conn)
return err
}
reader := NewStreamReader(conn, aead)
// 4 & 5. Read Client Request Header
reqHeader, err := ReadClientRequestHeader(conn, reader)
if err != nil {
return err
}
conn.SetReadDeadline(time.Time{})
dest := reqHeader.Destination
writer := NewServerStreamWriter(conn, i.method, userPSK, saltSlice)
// 6. Send Server Response Handshake
writer, err := WriteTCPResponse(conn, i.method, userPSK, saltSlice, nil)
if err != nil {
return err
}
// Dispatch Connection to Xray routing with matched User
// 7. Dispatch Connection to Xray routing with matched User
inbound := session.InboundFromContext(ctx)
inbound.User = user
@@ -262,17 +283,42 @@ func (i *MultiUserInbound) processTCP(ctx context.Context, conn net.Conn, dispat
}
if len(reqHeader.EarlyData) > 0 {
mb := buf.MergeBytes(nil, reqHeader.EarlyData)
if err := link.Writer.WriteMultiBuffer(mb); err != nil {
earlyBuf := buf.New()
earlyBuf.Write(reqHeader.EarlyData)
if err := link.Writer.WriteMultiBuffer(buf.MultiBuffer{earlyBuf}); err != nil {
return err
}
}
return TransportTCP(ctx, i.policyManager.ForLevel(user.Level), reader, writer, link)
sessionPolicy = i.policyManager.ForLevel(user.Level)
ctx, cancel := context.WithCancel(ctx)
timer := signal.CancelAfterInactivity(ctx, cancel, sessionPolicy.Timeouts.ConnectionIdle)
ctx = policy.ContextWithBufferPolicy(ctx, sessionPolicy.Buffer)
requestDone := func() error {
defer timer.SetTimeout(sessionPolicy.Timeouts.DownlinkOnly)
return buf.Copy(reader, link.Writer, buf.UpdateActivity(timer))
}
responseDone := func() error {
defer timer.SetTimeout(sessionPolicy.Timeouts.UplinkOnly)
return buf.Copy(link.Reader, writer, buf.UpdateActivity(timer))
}
responseDoneAndCloseWriter := task.OnSuccess(responseDone, task.Close(link.Writer))
return task.Run(ctx, requestDone, responseDoneAndCloseWriter)
}
func (i *MultiUserInbound) processUDP(ctx context.Context, conn stat.Connection, dispatcher routing.Dispatcher) error {
reader := buf.NewPacketReader(conn)
udpConns := utils.NewTypedSyncMap[uint64, *udpConnEntry]()
defer func() {
udpConns.Range(func(key uint64, entry *udpConnEntry) bool {
entry.timer.SetTimeout(0)
return true
})
}()
reader := buf.NewReader(conn)
for {
mb, err := reader.ReadMultiBuffer()
if err != nil {
@@ -296,61 +342,168 @@ func (i *MultiUserInbound) processUDP(ctx context.Context, conn stat.Connection,
sessionID := binary.BigEndian.Uint64(rawHeader[:8])
packetID := binary.BigEndian.Uint64(rawHeader[8:16])
// Replay protection & session lookup
sessionItem := i.udpSessions.GetOrCreate(sessionID)
if !sessionItem.CheckPacketID(packetID) {
sessionItem.Lock()
if !sessionItem.Window.Check(packetID) {
sessionItem.Unlock()
b.Release()
continue
}
var userPSK []byte
var currentUser *protocol.MemoryUser
sessionItem.Lock()
currentUser = sessionItem.User
userPSK = sessionItem.UserPSK
sessionItem.Unlock()
if currentUser == nil {
if sessionItem.User != nil {
currentUser = sessionItem.User
userPSK = sessionItem.UserPSK
sessionItem.Unlock()
} else {
sessionItem.Unlock()
// Decrypt EIH
decryptedHash := DecryptUDPEIH(i.udpMasterCipher, rawHeader[:], packetBytes[16:32])
identitySubkey := DeriveIdentitySubKey(i.masterPSK, rawHeader[:8], i.method.KeySaltLength)
idBlock, err := i.method.NewBlock(identitySubkey)
if err != nil {
b.Release()
continue
}
var decryptedHash [16]byte
idBlock.Decrypt(decryptedHash[:], packetBytes[16:32])
user, ok := i.usersByHash.Load(decryptedHash)
if !ok {
if !ok || user == nil {
b.Release()
continue
}
currentUser = user
userPSK = user.Account.(*MemoryAccount).Key
sessionItem.Lock()
sessionItem.User = user
sessionItem.UserPSK = userPSK
sessionItem.Unlock()
}
decoded, err := sessionItem.DecryptAESPayload(i.method, userPSK, sessionID, packetID, rawHeader[:], packetBytes[32:])
// Decrypt Body (with AEAD caching per session)
bodyAead := sessionItem.GetRemoteCipher()
if bodyAead == nil {
bodyKey := DeriveSessionSubKey(userPSK, rawHeader[:8], i.method.KeySaltLength)
var err error
bodyAead, err = i.method.NewAEAD(bodyKey)
if err != nil {
b.Release()
continue
}
sessionItem.SetRemoteCipher(bodyAead)
}
bodyNonce := rawHeader[4:16]
bodyCipher := packetBytes[32:]
bodyPlain, err := bodyAead.Open(nil, bodyNonce, bodyCipher, nil)
b.Release()
if err != nil {
if err != nil || len(bodyPlain) < 1+8+2 {
continue
}
sessionItem.Lock()
if sessionItem.User == nil {
sessionItem.User = currentUser
sessionItem.UserPSK = userPSK
}
sessionItem.Window.Add(packetID)
sessionItem.Unlock()
link, err := sessionItem.EnsureLink(ctx, conn, decoded.Destination, dispatcher, i.policyManager, func(replyDest net.Destination, payload []byte) ([]byte, error) {
return i.encodeServerUDPPacket(sessionID, userPSK, replyDest, payload)
})
if bodyPlain[0] != HeaderTypeClient {
continue
}
epoch := binary.BigEndian.Uint64(bodyPlain[1:9])
diff := time.Now().Unix() - int64(epoch)
if diff < -30 || diff > 30 {
continue
}
paddingLen := int(binary.BigEndian.Uint16(bodyPlain[9:11]))
offset := 11 + paddingLen
if len(bodyPlain) < offset {
continue
}
dest, addrLen, err := parseAddressPort(bodyPlain[offset:])
if err != nil {
continue
}
payload := bodyPlain[offset+addrLen:]
entry, ok := udpConns.Load(sessionID)
if !ok {
sessCtx, cancel := context.WithCancel(ctx)
inbound := session.InboundFromContext(sessCtx)
inbound.User = currentUser
sessCtx = log.ContextWithAccessMessage(sessCtx, &log.AccessMessage{
From: conn.RemoteAddr(),
To: dest,
Status: log.AccessAccepted,
Email: currentUser.Email,
})
link, err := dispatcher.Dispatch(sessCtx, dest)
if err != nil {
cancel()
continue
}
newEntry := &udpConnEntry{
link: link,
cancel: cancel,
}
sessionPolicy := i.policyManager.ForLevel(currentUser.Level)
newEntry.timer = signal.CancelAfterInactivity(sessCtx, func() {
udpConns.Delete(sessionID)
common.Interrupt(link.Reader)
common.Interrupt(link.Writer)
cancel()
}, sessionPolicy.Timeouts.ConnectionIdle)
actual, loaded := udpConns.LoadOrStore(sessionID, newEntry)
if loaded {
newEntry.timer.SetTimeout(0)
entry = actual
} else {
entry = newEntry
go func(sessID uint64, uPSK []byte, d net.Destination, cEntry *udpConnEntry) {
defer func() {
cEntry.timer.SetTimeout(0)
}()
for {
resMb, err := cEntry.link.Reader.ReadMultiBuffer()
if err != nil {
return
}
cEntry.timer.Update()
for _, rb := range resMb {
encPacket, err := i.encodeServerUDPPacket(sessID, uPSK, d, rb.Bytes())
rb.Release()
if err != nil {
continue
}
_, _ = conn.Write(encPacket)
}
}
}(sessionID, userPSK, dest, entry)
}
}
entry.timer.Update()
pBuf := buf.New()
pBuf.Write(decoded.Payload)
pBuf.UDP = &decoded.Destination
_ = link.Writer.WriteMultiBuffer(buf.MultiBuffer{pBuf})
pBuf.Write(payload)
_ = entry.link.Writer.WriteMultiBuffer(buf.MultiBuffer{pBuf})
}
}
}
func (i *MultiUserInbound) encodeServerUDPPacket(clientSessionID uint64, userPSK []byte, dest net.Destination, payload []byte) ([]byte, error) {
return i.udpSessions.EncodeServerPacket(i.method, userPSK, clientSessionID, dest, payload)
sessionItem := i.udpSessions.GetOrCreate(clientSessionID)
if err := sessionItem.EnsureServerState(i.method, i.udpMasterCipher, nil, userPSK); err != nil {
return nil, err
}
return sessionItem.EncodeServerPacket(i.method, clientSessionID, dest, payload)
}
+124 -59
View File
@@ -4,6 +4,7 @@ import (
"context"
"crypto/cipher"
"encoding/binary"
"io"
"strconv"
"time"
@@ -14,6 +15,9 @@ import (
"github.com/xtls/xray-core/common/net"
"github.com/xtls/xray-core/common/protocol"
"github.com/xtls/xray-core/common/session"
"github.com/xtls/xray-core/common/signal"
"github.com/xtls/xray-core/common/task"
"github.com/xtls/xray-core/common/utils"
"github.com/xtls/xray-core/common/uuid"
"github.com/xtls/xray-core/core"
"github.com/xtls/xray-core/features/policy"
@@ -31,17 +35,18 @@ type relayDest struct {
destination net.Destination
email string
level uint32
key []byte
blockCipher cipher.Block
}
type RelayInbound struct {
networks []net.Network
method *CipherMethod
relayPSK []byte
relayBlock cipher.Block
destinations map[[AESBlockSize]byte]*relayDest
udpSessions *UDPSessionManager
policyManager policy.Manager
networks []net.Network
method *CipherMethod
relayPSK []byte
relayBlock cipher.Block
destinations map[[AESBlockSize]byte]*relayDest
rawDestinations []*RelayDestination
policyManager policy.Manager
}
func NewRelayServer(ctx context.Context, config *RelayServerConfig) (*RelayInbound, error) {
@@ -73,13 +78,13 @@ func NewRelayServer(ctx context.Context, config *RelayServerConfig) (*RelayInbou
v := core.MustFromContext(ctx)
i := &RelayInbound{
networks: networks,
method: method,
relayPSK: relayPSK,
relayBlock: relayBlock,
destinations: make(map[[AESBlockSize]byte]*relayDest),
udpSessions: NewUDPSessionManager(500 * time.Second),
policyManager: v.GetFeature(policy.ManagerType()).(policy.Manager),
networks: networks,
method: method,
relayPSK: relayPSK,
relayBlock: relayBlock,
destinations: make(map[[AESBlockSize]byte]*relayDest),
rawDestinations: config.Destinations,
policyManager: v.GetFeature(policy.ManagerType()).(policy.Manager),
}
for idx, d := range config.Destinations {
@@ -103,6 +108,7 @@ func NewRelayServer(ctx context.Context, config *RelayServerConfig) (*RelayInbou
destination: net.TCPDestination(d.Address.AsAddress(), net.Port(d.Port)),
email: d.Email,
level: uint32(d.Level),
key: destKey,
blockCipher: destBlock,
}
}
@@ -133,36 +139,28 @@ func (i *RelayInbound) processTCP(ctx context.Context, conn net.Conn, dispatcher
return errors.New("unable to set read deadline").Base(err)
}
// Read initial handshake in a single read call per SIP022 §3.1.3 & §3.1.4
// Read Salt + Outer EIH
needed := i.method.KeySaltLength + AESBlockSize
requestHeader := buf.New()
n, err := requestHeader.ReadFrom(conn)
if err != nil {
requestHeader.Release()
ResetTCPConn(conn)
var headerBuf [48]byte
headerSlice := headerBuf[:needed]
if _, err := io.ReadFull(conn, headerSlice); err != nil {
return err
}
if int(n) < needed {
requestHeader.Release()
ResetTCPConn(conn)
return ErrInvalidRequest
}
headerSlice := requestHeader.Bytes()
salt := headerSlice[:i.method.KeySaltLength]
eih := headerSlice[i.method.KeySaltLength:needed]
eih := headerSlice[i.method.KeySaltLength:]
decryptedHash, err := DecryptEIH(i.method, i.relayPSK, salt, eih)
identitySubkey := DeriveIdentitySubKey(i.relayPSK, salt, i.method.KeySaltLength)
block, err := i.method.NewBlock(identitySubkey)
if err != nil {
requestHeader.Release()
ResetTCPConn(conn)
return err
}
var decryptedHash [AESBlockSize]byte
block.Decrypt(decryptedHash[:], eih)
targetDest, ok := i.destinations[decryptedHash]
if !ok {
requestHeader.Release()
ResetTCPConn(conn)
return ErrInvalidRequest
}
conn.SetReadDeadline(time.Time{})
@@ -184,26 +182,45 @@ func (i *RelayInbound) processTCP(ctx context.Context, conn net.Conn, dispatcher
link, err := dispatcher.Dispatch(ctx, targetDest.destination)
if err != nil {
requestHeader.Release()
return err
}
// Unwrap outer EIH: send client salt and remaining handshake bytes to next hop
// in a single write call, satisfying downstream server's single-read handshake expectation (SIP022 §3.1.3).
var saltCopy [32]byte
copy(saltCopy[:i.method.KeySaltLength], salt)
copy(requestHeader.Bytes()[AESBlockSize:AESBlockSize+i.method.KeySaltLength], saltCopy[:i.method.KeySaltLength])
requestHeader.Advance(AESBlockSize)
if err := link.Writer.WriteMultiBuffer(buf.MultiBuffer{requestHeader}); err != nil {
// Unwrap outer EIH: send client salt to next hop, stripping this hop's EIH
saltBuf := buf.New()
saltBuf.Write(salt)
if err := link.Writer.WriteMultiBuffer(buf.MultiBuffer{saltBuf}); err != nil {
return err
}
return TransportTCP(ctx, i.policyManager.ForLevel(targetDest.level), buf.NewReader(conn), buf.NewWriter(conn), link)
sessionPolicy = i.policyManager.ForLevel(targetDest.level)
ctx, cancel := context.WithCancel(ctx)
timer := signal.CancelAfterInactivity(ctx, cancel, sessionPolicy.Timeouts.ConnectionIdle)
ctx = policy.ContextWithBufferPolicy(ctx, sessionPolicy.Buffer)
requestDone := func() error {
defer timer.SetTimeout(sessionPolicy.Timeouts.DownlinkOnly)
return buf.Copy(buf.NewReader(conn), link.Writer, buf.UpdateActivity(timer))
}
responseDone := func() error {
defer timer.SetTimeout(sessionPolicy.Timeouts.UplinkOnly)
return buf.Copy(link.Reader, buf.NewWriter(conn), buf.UpdateActivity(timer))
}
responseDoneAndCloseWriter := task.OnSuccess(responseDone, task.Close(link.Writer))
return task.Run(ctx, requestDone, responseDoneAndCloseWriter)
}
func (i *RelayInbound) processUDP(ctx context.Context, conn stat.Connection, dispatcher routing.Dispatcher) error {
reader := buf.NewPacketReader(conn)
udpConns := utils.NewTypedSyncMap[uint64, *udpConnEntry]()
defer func() {
udpConns.Range(func(key uint64, entry *udpConnEntry) bool {
entry.timer.SetTimeout(0)
return true
})
}()
reader := buf.NewReader(conn)
for {
mb, err := reader.ReadMultiBuffer()
if err != nil {
@@ -221,7 +238,11 @@ func (i *RelayInbound) processUDP(ctx context.Context, conn stat.Connection, dis
var packetHeader [AESBlockSize]byte
i.relayBlock.Decrypt(packetHeader[:], data[:AESBlockSize])
eiHeader := DecryptUDPEIH(i.relayBlock, packetHeader[:], data[AESBlockSize:2*AESBlockSize])
var eiHeader [AESBlockSize]byte
i.relayBlock.Decrypt(eiHeader[:], data[AESBlockSize:2*AESBlockSize])
for idx := 0; idx < AESBlockSize; idx++ {
eiHeader[idx] ^= packetHeader[idx]
}
targetDest, ok := i.destinations[eiHeader]
if !ok {
@@ -242,24 +263,68 @@ func (i *RelayInbound) processUDP(ctx context.Context, conn stat.Connection, dis
dest := targetDest.destination
dest.Network = net.Network_UDP
sessionItem := i.udpSessions.GetOrCreate(sessionID)
if sessionItem.User == nil {
sessionItem.Lock()
if sessionItem.User == nil {
sessionItem.User = &protocol.MemoryUser{
Email: targetDest.email,
Level: targetDest.level,
}
entry, ok := udpConns.Load(sessionID)
if !ok {
sessCtx, cancel := context.WithCancel(ctx)
inbound := session.InboundFromContext(sessCtx)
inbound.User = &protocol.MemoryUser{
Email: targetDest.email,
Level: targetDest.level,
}
sessCtx = log.ContextWithAccessMessage(sessCtx, &log.AccessMessage{
From: conn.RemoteAddr(),
To: dest,
Status: log.AccessAccepted,
Email: targetDest.email,
})
link, err := dispatcher.Dispatch(sessCtx, dest)
if err != nil {
cancel()
b.Release()
continue
}
newEntry := &udpConnEntry{
link: link,
cancel: cancel,
}
sessionPolicy := i.policyManager.ForLevel(targetDest.level)
newEntry.timer = signal.CancelAfterInactivity(sessCtx, func() {
udpConns.Delete(sessionID)
common.Interrupt(link.Reader)
common.Interrupt(link.Writer)
cancel()
}, sessionPolicy.Timeouts.ConnectionIdle)
actual, loaded := udpConns.LoadOrStore(sessionID, newEntry)
if loaded {
newEntry.timer.SetTimeout(0)
entry = actual
} else {
entry = newEntry
go func(cEntry *udpConnEntry) {
defer func() {
cEntry.timer.SetTimeout(0)
}()
for {
resMb, err := cEntry.link.Reader.ReadMultiBuffer()
if err != nil {
return
}
cEntry.timer.Update()
for _, rb := range resMb {
_, _ = conn.Write(rb.Bytes())
rb.Release()
}
}
}(entry)
}
sessionItem.Unlock()
}
link, err := sessionItem.EnsureLink(ctx, conn, dest, dispatcher, i.policyManager, nil)
if err != nil {
b.Release()
continue
}
_ = link.Writer.WriteMultiBuffer(buf.MultiBuffer{b})
entry.timer.Update()
_ = entry.link.Writer.WriteMultiBuffer(buf.MultiBuffer{b})
}
}
}
-11
View File
@@ -61,14 +61,3 @@ func DeriveUserPSKHash(userPSK []byte) [AESBlockSize]byte {
copy(out[:], h[:AESBlockSize])
return out
}
func DecryptEIH(method *CipherMethod, key, salt, eih []byte) ([AESBlockSize]byte, error) {
identitySubkey := DeriveIdentitySubKey(key, salt, method.KeySaltLength)
block, err := method.NewBlock(identitySubkey)
if err != nil {
return [AESBlockSize]byte{}, err
}
var decryptedHash [AESBlockSize]byte
block.Decrypt(decryptedHash[:], eih)
return decryptedHash, nil
}
+13 -33
View File
@@ -4,6 +4,7 @@ import (
"context"
"crypto/rand"
"io"
"time"
"github.com/xtls/xray-core/common"
"github.com/xtls/xray-core/common/buf"
@@ -45,12 +46,8 @@ func NewClient(ctx context.Context, config *ClientConfig) (*Outbound, error) {
return nil, errors.New("invalid key: ", config.Key).Base(err)
}
if method.IsChaCha && len(pskList) > 1 {
return nil, errors.New("multi-key is not supported for chacha20-poly1305")
}
finalPSK := pskList[len(pskList)-1]
udpCodec, err := NewUDPPacketCodec(method, pskList)
udpCodec, err := NewUDPPacketCodec(method, finalPSK)
if err != nil {
return nil, errors.New("failed to create udp packet codec").Base(err)
}
@@ -129,30 +126,18 @@ func (o *Outbound) Process(ctx context.Context, link *transport.Link, dialer int
requestDone := func() error {
defer timer.SetTimeout(sessionPolicy.Timeouts.DownlinkOnly)
var initialPayload []byte
var firstBuf *buf.Buffer
var remainingMB buf.MultiBuffer
if timeoutReader, ok := link.Reader.(buf.TimeoutReader); ok {
if mb, err := timeoutReader.ReadMultiBufferTimeout(0); err == nil && !mb.IsEmpty() {
remainingMB, firstBuf = buf.SplitFirst(mb)
initialPayload = firstBuf.Bytes()
}
}
bodyWriter, err := WriteTCPRequest(conn, o.method, o.pskList, destination, clientSaltSlice, initialPayload)
if firstBuf != nil {
firstBuf.Release()
}
bufferedWriter := buf.NewBufferedWriter(buf.NewWriter(conn))
bodyWriter, err := WriteTCPRequest(bufferedWriter, o.method, o.pskList, destination, clientSaltSlice, nil)
if err != nil {
buf.ReleaseMulti(remainingMB)
return errors.New("failed to write request").Base(err)
}
if !remainingMB.IsEmpty() {
if err := bodyWriter.WriteMultiBuffer(remainingMB); err != nil {
return err
}
if err = buf.CopyOnceTimeout(link.Reader, bodyWriter, time.Millisecond*100); err != nil && err != buf.ErrNotTimeoutReader && err != buf.ErrReadTimeout {
return errors.New("failed to write A request payload").Base(err)
}
if err := bufferedWriter.SetBuffered(false); err != nil {
return err
}
return buf.Copy(link.Reader, bodyWriter, buf.UpdateActivity(timer))
@@ -178,18 +163,13 @@ func (o *Outbound) Process(ctx context.Context, link *transport.Link, dialer int
}
if network == net.Network_UDP {
session, err := o.udpCodec.NewClientSession()
if err != nil {
return errors.New("failed to create client udp session").Base(err)
}
requestDone := func() error {
defer timer.SetTimeout(sessionPolicy.Timeouts.DownlinkOnly)
writer := &UDPWriter{
Writer: conn,
Destination: destination,
Session: session,
Codec: o.udpCodec,
}
if err := buf.Copy(link.Reader, writer, buf.UpdateActivity(timer)); err != nil {
@@ -202,8 +182,8 @@ func (o *Outbound) Process(ctx context.Context, link *transport.Link, dialer int
defer timer.SetTimeout(sessionPolicy.Timeouts.UplinkOnly)
reader := &UDPReader{
Reader: conn,
Session: session,
Reader: conn,
Codec: o.udpCodec,
}
if err := buf.Copy(reader, link.Writer, buf.UpdateActivity(timer)); err != nil {
+187 -441
View File
@@ -16,13 +16,14 @@ import (
)
type UDPCodec struct {
method *CipherMethod
pskList [][]byte
psk []byte
blockCipher cipher.Block
blockCiphers []cipher.Block
chachaCipher cipher.AEAD
sessions *UDPSessionManager
method *CipherMethod
psk []byte
blockCipher cipher.Block
chachaCipher cipher.AEAD
clientBodyCipher cipher.AEAD
clientSessionID uint64
nextPacketID atomic.Uint64
sessions *UDPSessionManager
}
type (
@@ -47,23 +48,22 @@ func newUDPCodec(method *CipherMethod, psk []byte) (*UDPCodec, error) {
return c, nil
}
func NewUDPPacketCodec(method *CipherMethod, pskList [][]byte) (*UDPCodec, error) {
if method.IsChaCha && len(pskList) > 1 {
return nil, errors.New("multi-key is not supported for chacha20-poly1305")
}
finalPSK := pskList[len(pskList)-1]
c, err := newUDPCodec(method, finalPSK)
func NewUDPPacketCodec(method *CipherMethod, psk []byte) (*UDPCodec, error) {
c, err := newUDPCodec(method, psk)
if err != nil {
return nil, err
}
c.pskList = pskList
if len(pskList) > 1 {
c.blockCiphers = make([]cipher.Block, len(pskList))
for i, psk := range pskList {
c.blockCiphers[i], err = method.NewBlock(psk)
if err != nil {
return nil, err
}
var sessID [8]byte
if _, err := io.ReadFull(rand.Reader, sessID[:]); err != nil {
return nil, err
}
c.clientSessionID = binary.BigEndian.Uint64(sessID[:])
if !method.IsChaCha {
clientBodyKey := DeriveSessionSubKey(psk, sessID[:], method.KeySaltLength)
c.clientBodyCipher, err = method.NewAEAD(clientBodyKey)
if err != nil {
return nil, err
}
}
return c, nil
@@ -78,37 +78,108 @@ func NewUDPServerCodec(method *CipherMethod, psk []byte, sessionTimeout time.Dur
return c, nil
}
func (c *UDPCodec) Sessions() *UDPSessionManager {
return c.sessions
}
func (c *UDPCodec) EncodeClientPacket(dest net.Destination, payload []byte) (*buf.Buffer, error) {
packetID := c.nextPacketID.Add(1)
sessID := c.clientSessionID
func (c *UDPCodec) GetSession(sessionID uint64) *ServerUDPSession {
if c.sessions == nil {
return nil
// Padding determination (e.g. DNS port 53 disguise)
var paddingLen int
if dest.Port == 53 && len(payload) < MaxPaddingLength {
paddingLen = mrand.IntN(MaxPaddingLength-len(payload)) + 1
}
return c.sessions.GetOrCreate(sessionID)
addrPortLen := AddrPortLength(dest)
if c.method.IsChaCha {
// ChaCha20 mode: 24-byte nonce + plaintext header (27B) + padding + dest + payload + AEAD tag (16B)
totalLen := PacketNonceSize + 27 + paddingLen + addrPortLen + len(payload) + AEADTagSize
if totalLen > buf.Size {
return nil, ErrPacketTooLarge
}
outBuf := buf.New()
var nonce [PacketNonceSize]byte
if _, err := io.ReadFull(rand.Reader, nonce[:]); err != nil {
outBuf.Release()
return nil, err
}
outBuf.Write(nonce[:])
var hdr [16 + 1 + 8 + 2]byte
binary.BigEndian.PutUint64(hdr[0:8], sessID)
binary.BigEndian.PutUint64(hdr[8:16], packetID)
hdr[16] = HeaderTypeClient
binary.BigEndian.PutUint64(hdr[17:25], uint64(time.Now().Unix()))
binary.BigEndian.PutUint16(hdr[25:27], uint16(paddingLen))
outBuf.Write(hdr[:])
if paddingLen > 0 {
outBuf.Write(zeroPadding[:paddingLen])
}
if err := WriteAddressPort(outBuf, dest); err != nil {
outBuf.Release()
return nil, err
}
outBuf.Write(payload)
plainBytes := outBuf.Bytes()[PacketNonceSize:]
outBuf.Extend(int32(c.chachaCipher.Overhead()))
c.chachaCipher.Seal(plainBytes[:0], nonce[:], plainBytes, nil)
return outBuf, nil
}
// AES mode:
// 16B Encrypted Header + (11B header + padding + dest + payload + 16B AEAD tag)
totalLen := 16 + 11 + paddingLen + addrPortLen + len(payload) + AEADTagSize
if totalLen > buf.Size {
return nil, ErrPacketTooLarge
}
outBuf := buf.New()
var rawHeader [16]byte
binary.BigEndian.PutUint64(rawHeader[:8], sessID)
binary.BigEndian.PutUint64(rawHeader[8:16], packetID)
var encryptedHeader [16]byte
c.blockCipher.Encrypt(encryptedHeader[:], rawHeader[:])
outBuf.Write(encryptedHeader[:])
bodyAead := c.clientBodyCipher
var hdr [1 + 8 + 2]byte
hdr[0] = HeaderTypeClient
binary.BigEndian.PutUint64(hdr[1:9], uint64(time.Now().Unix()))
binary.BigEndian.PutUint16(hdr[9:11], uint16(paddingLen))
outBuf.Write(hdr[:])
if paddingLen > 0 {
outBuf.Write(zeroPadding[:paddingLen])
}
if err := WriteAddressPort(outBuf, dest); err != nil {
outBuf.Release()
return nil, err
}
outBuf.Write(payload)
plainBytes := outBuf.Bytes()[16:]
bodyNonce := rawHeader[4:16]
outBuf.Extend(int32(bodyAead.Overhead()))
bodyAead.Seal(plainBytes[:0], bodyNonce, plainBytes, nil)
return outBuf, nil
}
type DecodedUDPPacket struct {
SessionID uint64
PacketID uint64
HeaderType byte
Timestamp uint64
ClientSessionID uint64
Destination net.Destination
Payload []byte
SessionID uint64
PacketID uint64
HeaderType byte
Timestamp uint64
Destination net.Destination
Payload []byte
}
func DecryptUDPEIH(block cipher.Block, rawHeader, eih []byte) [AESBlockSize]byte {
var decryptedHash [AESBlockSize]byte
block.Decrypt(decryptedHash[:], eih)
for k := 0; k < AESBlockSize; k++ {
decryptedHash[k] ^= rawHeader[k]
}
return decryptedHash
}
func ParseAddressPort(data []byte) (net.Destination, int, error) {
func parseAddressPort(data []byte) (net.Destination, int, error) {
if len(data) < 1 {
return net.Destination{}, 0, ErrPacketTooShort
}
@@ -149,9 +220,6 @@ func parsePlainUDPPacket(sessionID, packetID uint64, bodyPlain []byte) (DecodedU
}
headerType := bodyPlain[0]
if headerType != HeaderTypeClient && headerType != HeaderTypeServer {
return DecodedUDPPacket{}, ErrBadHeaderType
}
epoch := binary.BigEndian.Uint64(bodyPlain[1:9])
diff := int(math.Abs(float64(time.Now().Unix() - int64(epoch))))
if diff > 30 {
@@ -159,13 +227,11 @@ func parsePlainUDPPacket(sessionID, packetID uint64, bodyPlain []byte) (DecodedU
}
offset := 9
var clientSessionID uint64
if headerType == HeaderTypeServer {
if len(bodyPlain) < offset+8+2 {
return DecodedUDPPacket{}, ErrPacketTooShort
}
clientSessionID = binary.BigEndian.Uint64(bodyPlain[offset : offset+8])
offset += 8
offset += 8 // skip clientSessionID
}
paddingLen := int(binary.BigEndian.Uint16(bodyPlain[offset : offset+2]))
@@ -176,20 +242,19 @@ func parsePlainUDPPacket(sessionID, packetID uint64, bodyPlain []byte) (DecodedU
}
offset += paddingLen
dest, addrLen, err := ParseAddressPort(bodyPlain[offset:])
dest, addrLen, err := parseAddressPort(bodyPlain[offset:])
if err != nil {
return DecodedUDPPacket{}, err
}
payload := bodyPlain[offset+addrLen:]
return DecodedUDPPacket{
SessionID: sessionID,
PacketID: packetID,
HeaderType: headerType,
Timestamp: epoch,
ClientSessionID: clientSessionID,
Destination: dest,
Payload: payload,
SessionID: sessionID,
PacketID: packetID,
HeaderType: headerType,
Timestamp: epoch,
Destination: dest,
Payload: payload,
}, nil
}
@@ -204,7 +269,7 @@ func (c *UDPCodec) DecodePacket(data []byte) (DecodedUDPPacket, error) {
}
nonce := data[:PacketNonceSize]
ciphertext := data[PacketNonceSize:]
plain, err := c.chachaCipher.Open(nil, nonce, ciphertext, nil)
plain, err := c.chachaCipher.Open(ciphertext[:0], nonce, ciphertext, nil)
if err != nil {
return DecodedUDPPacket{}, errors.New("failed to decrypt chacha udp packet").Base(err)
}
@@ -215,22 +280,17 @@ func (c *UDPCodec) DecodePacket(data []byte) (DecodedUDPPacket, error) {
sessionID := binary.BigEndian.Uint64(plain[:8])
packetID := binary.BigEndian.Uint64(plain[8:16])
sessionItem := c.sessions.GetOrCreate(sessionID)
if !sessionItem.CheckPacketID(packetID) {
return DecodedUDPPacket{}, ErrPacketIdNotUnique
if c.sessions != nil {
sessionItem := c.sessions.GetOrCreate(sessionID)
sessionItem.Lock()
if !sessionItem.Window.CheckAndAdd(packetID) {
sessionItem.Unlock()
return DecodedUDPPacket{}, ErrPacketIdNotUnique
}
sessionItem.Unlock()
}
decoded, err := parsePlainUDPPacket(sessionID, packetID, plain[16:])
if err != nil {
return DecodedUDPPacket{}, err
}
if decoded.HeaderType != HeaderTypeClient {
return DecodedUDPPacket{}, ErrBadHeaderType
}
sessionItem.AddPacketID(packetID)
return decoded, nil
return parsePlainUDPPacket(sessionID, packetID, plain[16:])
}
// AES mode
@@ -239,52 +299,54 @@ func (c *UDPCodec) DecodePacket(data []byte) (DecodedUDPPacket, error) {
sessionID := binary.BigEndian.Uint64(rawHeader[:8])
packetID := binary.BigEndian.Uint64(rawHeader[8:16])
sessionItem := c.sessions.GetOrCreate(sessionID)
if !sessionItem.CheckPacketID(packetID) {
return DecodedUDPPacket{}, ErrPacketIdNotUnique
}
var bodyAead cipher.AEAD
var sessionItem *ServerUDPSession
return sessionItem.DecryptAESPayload(c.method, c.psk, sessionID, packetID, rawHeader[:], data[16:])
}
if c.sessions != nil {
sessionItem = c.sessions.GetOrCreate(sessionID)
sessionItem.Lock()
if !sessionItem.Window.Check(packetID) {
sessionItem.Unlock()
return DecodedUDPPacket{}, ErrPacketIdNotUnique
}
sessionItem.Unlock()
func (s *ServerUDPSession) DecryptAESPayload(method *CipherMethod, psk []byte, sessionID, packetID uint64, rawHeader, bodyCipher []byte) (DecodedUDPPacket, error) {
bodyAead := s.clientBodyCipher
isNewCipher := false
if bodyAead == nil {
bodyKey := DeriveSessionSubKey(psk, rawHeader[:8], method.KeySaltLength)
bodyAead = sessionItem.GetRemoteCipher()
if bodyAead == nil {
bodyKey := DeriveSessionSubKey(c.psk, rawHeader[:8], c.method.KeySaltLength)
var err error
bodyAead, err = c.method.NewAEAD(bodyKey)
if err != nil {
return DecodedUDPPacket{}, err
}
sessionItem.SetRemoteCipher(bodyAead)
}
} else {
bodyKey := DeriveSessionSubKey(c.psk, rawHeader[:8], c.method.KeySaltLength)
var err error
bodyAead, err = method.NewAEAD(bodyKey)
bodyAead, err = c.method.NewAEAD(bodyKey)
if err != nil {
return DecodedUDPPacket{}, err
}
isNewCipher = true
}
bodyNonce := rawHeader[4:16]
bodyPlain, err := bodyAead.Open(nil, bodyNonce, bodyCipher, nil)
bodyCipher := data[16:]
bodyPlain, err := bodyAead.Open(bodyCipher[:0], bodyNonce, bodyCipher, nil)
if err != nil {
return DecodedUDPPacket{}, errors.New("failed to decrypt aes udp body").Base(err)
}
decoded, err := parsePlainUDPPacket(sessionID, packetID, bodyPlain)
if err != nil {
return DecodedUDPPacket{}, err
if sessionItem != nil {
sessionItem.Lock()
sessionItem.Window.Add(packetID)
sessionItem.Unlock()
}
if decoded.HeaderType != HeaderTypeClient {
return DecodedUDPPacket{}, ErrBadHeaderType
}
s.AddPacketID(packetID)
if isNewCipher {
s.clientBodyCipher = bodyAead
}
return decoded, nil
return parsePlainUDPPacket(sessionID, packetID, bodyPlain)
}
func (s *ServerUDPSession) EnsureServerState(method *CipherMethod, psk []byte) error {
func (s *ServerUDPSession) EnsureServerState(method *CipherMethod, headerBlock cipher.Block, chachaCipher cipher.AEAD, psk []byte) error {
s.Lock()
defer s.Unlock()
if s.ServerSessionID != 0 {
@@ -301,29 +363,23 @@ func (s *ServerUDPSession) EnsureServerState(method *CipherMethod, psk []byte) e
}
}
if method.IsChaCha {
var err error
s.serverChaCha, err = method.NewUDPCipher(psk)
return err
}
var err error
s.serverHeaderBlock, err = method.NewBlock(psk)
if err != nil {
s.ServerSessionID = 0
return err
}
bodyKey := DeriveSessionSubKey(psk, sidBuf[:], method.KeySaltLength)
s.serverBodyCipher, err = method.NewAEAD(bodyKey)
if err != nil {
s.ServerSessionID = 0
return err
s.ServerChaCha = chachaCipher
} else {
s.ServerBlockCipher = headerBlock
bodyKey := DeriveSessionSubKey(psk, sidBuf[:], method.KeySaltLength)
bodyAead, err := method.NewAEAD(bodyKey)
if err != nil {
s.ServerSessionID = 0
return err
}
s.ServerCipher = bodyAead
}
return nil
}
func (s *ServerUDPSession) EncodeServerPacket(method *CipherMethod, clientSessionID uint64, dest net.Destination, payload []byte) ([]byte, error) {
serverSessionID := s.ServerSessionID
serverPacketID := s.ServerPacketID.Add(1) - 1
serverPacketID := s.ServerPacketID.Add(1)
if method.IsChaCha {
var nonce [PacketNonceSize]byte
@@ -348,7 +404,7 @@ func (s *ServerUDPSession) EncodeServerPacket(method *CipherMethod, clientSessio
}
plainBuf.Write(payload)
sealed := s.serverChaCha.Seal(nil, nonce[:], plainBuf.Bytes(), nil)
sealed := s.ServerChaCha.Seal(nil, nonce[:], plainBuf.Bytes(), nil)
res := make([]byte, PacketNonceSize+len(sealed))
copy(res[:PacketNonceSize], nonce[:])
copy(res[PacketNonceSize:], sealed)
@@ -361,7 +417,7 @@ func (s *ServerUDPSession) EncodeServerPacket(method *CipherMethod, clientSessio
binary.BigEndian.PutUint64(rawHeader[8:16], serverPacketID)
var encryptedHeader [16]byte
s.serverHeaderBlock.Encrypt(encryptedHeader[:], rawHeader[:])
s.ServerBlockCipher.Encrypt(encryptedHeader[:], rawHeader[:])
bodyBuf := buf.New()
defer bodyBuf.Release()
@@ -379,7 +435,7 @@ func (s *ServerUDPSession) EncodeServerPacket(method *CipherMethod, clientSessio
bodyBuf.Write(payload)
bodyNonce := rawHeader[4:16]
sealedBody := s.serverBodyCipher.Seal(nil, bodyNonce, bodyBuf.Bytes(), nil)
sealedBody := s.ServerCipher.Seal(nil, bodyNonce, bodyBuf.Bytes(), nil)
res := make([]byte, 16+len(sealedBody))
copy(res[:16], encryptedHeader[:])
@@ -388,327 +444,17 @@ func (s *ServerUDPSession) EncodeServerPacket(method *CipherMethod, clientSessio
}
func (c *UDPCodec) EncodeServerPacket(clientSessionID uint64, dest net.Destination, payload []byte) ([]byte, error) {
return c.sessions.EncodeServerPacket(c.method, c.psk, clientSessionID, dest, payload)
}
type serverSessionState struct {
sessionID uint64
window *SlidingWindow
cipher cipher.AEAD
lastSeen atomic.Int64
}
func (st *serverSessionState) check(packetID uint64) bool {
if st.window == nil {
st.window = new(SlidingWindow)
}
return st.window.Check(packetID)
}
func (st *serverSessionState) add(packetID uint64) {
if st.window == nil {
st.window = new(SlidingWindow)
}
st.window.Add(packetID)
}
type ClientUDPSession struct {
codec *UDPCodec
clientSessionID uint64
nextPacketID atomic.Uint64
clientBodyCipher cipher.AEAD
current atomic.Pointer[serverSessionState]
old atomic.Pointer[serverSessionState]
}
func (c *UDPCodec) NewClientSession() (*ClientUDPSession, error) {
var sessID [8]byte
if _, err := io.ReadFull(rand.Reader, sessID[:]); err != nil {
sessionItem := c.sessions.GetOrCreate(clientSessionID)
if err := sessionItem.EnsureServerState(c.method, c.blockCipher, c.chachaCipher, c.psk); err != nil {
return nil, err
}
clientSessionID := binary.BigEndian.Uint64(sessID[:])
var clientBodyCipher cipher.AEAD
var err error
if !c.method.IsChaCha {
finalPSK := c.psk
clientBodyKey := DeriveSessionSubKey(finalPSK, sessID[:], c.method.KeySaltLength)
clientBodyCipher, err = c.method.NewAEAD(clientBodyKey)
if err != nil {
return nil, err
}
}
return &ClientUDPSession{
codec: c,
clientSessionID: clientSessionID,
clientBodyCipher: clientBodyCipher,
}, nil
}
func (s *ClientUDPSession) getServerSession(sessionID uint64, now int64) (*serverSessionState, error) {
cur := s.current.Load()
if cur != nil && cur.sessionID == sessionID {
return cur, nil
}
old := s.old.Load()
if old != nil && old.sessionID == sessionID {
if now-old.lastSeen.Load() > 60 {
s.old.CompareAndSwap(old, nil)
return nil, errors.New("old server session expired")
}
return old, nil
}
// New server session:
// Spec §3.2.4: reject newer server sessions when the last packet received from the old session is less than 1 minute old.
if old != nil && now-old.lastSeen.Load() < 60 {
return nil, errors.New("newer server session rejected: old session is less than 1 minute old")
}
var bodyAead cipher.AEAD
if !s.codec.method.IsChaCha {
var sessBytes [8]byte
binary.BigEndian.PutUint64(sessBytes[:], sessionID)
bodyKey := DeriveSessionSubKey(s.codec.psk, sessBytes[:], s.codec.method.KeySaltLength)
var err error
bodyAead, err = s.codec.method.NewAEAD(bodyKey)
if err != nil {
return nil, err
}
}
newState := &serverSessionState{
sessionID: sessionID,
cipher: bodyAead,
}
newState.lastSeen.Store(now)
if cur == nil {
s.current.CompareAndSwap(nil, newState)
return s.current.Load(), nil
}
s.old.Store(cur)
s.current.Store(newState)
return newState, nil
}
func (s *ClientUDPSession) ClientSessionID() uint64 {
return s.clientSessionID
}
func (s *ClientUDPSession) EncodePacket(dest net.Destination, payload []byte) (*buf.Buffer, error) {
packetID := s.nextPacketID.Add(1) - 1
sessID := s.clientSessionID
var paddingLen int
if dest.Port == 53 && len(payload) < MaxPaddingLength {
paddingLen = mrand.IntN(MaxPaddingLength) + 1
}
addrPortLen := AddrPortLength(dest)
if s.codec.method.IsChaCha {
totalLen := PacketNonceSize + 27 + paddingLen + addrPortLen + len(payload) + AEADTagSize
if totalLen > buf.Size {
return nil, ErrPacketTooLarge
}
outBuf := buf.New()
var nonce [PacketNonceSize]byte
if _, err := io.ReadFull(rand.Reader, nonce[:]); err != nil {
outBuf.Release()
return nil, err
}
outBuf.Write(nonce[:])
var hdr [16 + 1 + 8 + 2]byte
binary.BigEndian.PutUint64(hdr[0:8], sessID)
binary.BigEndian.PutUint64(hdr[8:16], packetID)
hdr[16] = HeaderTypeClient
binary.BigEndian.PutUint64(hdr[17:25], uint64(time.Now().Unix()))
binary.BigEndian.PutUint16(hdr[25:27], uint16(paddingLen))
outBuf.Write(hdr[:])
if paddingLen > 0 {
outBuf.Write(zeroPadding[:paddingLen])
}
if err := WriteAddressPort(outBuf, dest); err != nil {
outBuf.Release()
return nil, err
}
outBuf.Write(payload)
plainBytes := outBuf.Bytes()[PacketNonceSize:]
outBuf.Extend(int32(s.codec.chachaCipher.Overhead()))
s.codec.chachaCipher.Seal(plainBytes[:0], nonce[:], plainBytes, nil)
return outBuf, nil
}
// AES mode
var sessBytes [8]byte
binary.BigEndian.PutUint64(sessBytes[:], sessID)
var rawHeader [16]byte
copy(rawHeader[:8], sessBytes[:])
binary.BigEndian.PutUint64(rawHeader[8:16], packetID)
eihCount := 0
if len(s.codec.pskList) > 1 {
eihCount = len(s.codec.pskList) - 1
}
totalLen := 16 + eihCount*16 + 11 + paddingLen + addrPortLen + len(payload) + AEADTagSize
if totalLen > buf.Size {
return nil, ErrPacketTooLarge
}
outBuf := buf.New()
if len(s.codec.pskList) > 1 {
var encryptedHeader [16]byte
s.codec.blockCiphers[0].Encrypt(encryptedHeader[:], rawHeader[:])
outBuf.Write(encryptedHeader[:])
for i := 0; i < len(s.codec.pskList)-1; i++ {
nextPSK := s.codec.pskList[i+1]
pskHash := DeriveUserPSKHash(nextPSK)
var eihPlain [16]byte
for k := 0; k < 16; k++ {
eihPlain[k] = pskHash[k] ^ rawHeader[k]
}
var encryptedEIH [16]byte
s.codec.blockCiphers[i].Encrypt(encryptedEIH[:], eihPlain[:])
outBuf.Write(encryptedEIH[:])
}
} else {
var encryptedHeader [16]byte
s.codec.blockCipher.Encrypt(encryptedHeader[:], rawHeader[:])
outBuf.Write(encryptedHeader[:])
}
bodyAead := s.clientBodyCipher
var hdr [1 + 8 + 2]byte
hdr[0] = HeaderTypeClient
binary.BigEndian.PutUint64(hdr[1:9], uint64(time.Now().Unix()))
binary.BigEndian.PutUint16(hdr[9:11], uint16(paddingLen))
outBuf.Write(hdr[:])
if paddingLen > 0 {
outBuf.Write(zeroPadding[:paddingLen])
}
if err := WriteAddressPort(outBuf, dest); err != nil {
outBuf.Release()
return nil, err
}
outBuf.Write(payload)
headerOffset := 16 + eihCount*16
plainBytes := outBuf.Bytes()[headerOffset:]
bodyNonce := rawHeader[4:16]
outBuf.Extend(int32(bodyAead.Overhead()))
bodyAead.Seal(plainBytes[:0], bodyNonce, plainBytes, nil)
return outBuf, nil
}
func (s *ClientUDPSession) DecodePacket(data []byte) (DecodedUDPPacket, error) {
if len(data) < PacketMinimalHeaderSize {
return DecodedUDPPacket{}, ErrPacketTooShort
}
if s.codec.method.IsChaCha {
if len(data) < PacketNonceSize+AEADTagSize {
return DecodedUDPPacket{}, ErrPacketTooShort
}
nonce := data[:PacketNonceSize]
ciphertext := data[PacketNonceSize:]
plain, err := s.codec.chachaCipher.Open(nil, nonce, ciphertext, nil)
if err != nil {
return DecodedUDPPacket{}, errors.New("failed to decrypt chacha udp packet").Base(err)
}
if len(plain) < 16+1+8+2 {
return DecodedUDPPacket{}, ErrPacketTooShort
}
sessionID := binary.BigEndian.Uint64(plain[:8])
packetID := binary.BigEndian.Uint64(plain[8:16])
now := time.Now().Unix()
st, err := s.getServerSession(sessionID, now)
if err != nil {
return DecodedUDPPacket{}, err
}
if !st.check(packetID) {
return DecodedUDPPacket{}, ErrPacketIdNotUnique
}
decoded, err := parsePlainUDPPacket(sessionID, packetID, plain[16:])
if err != nil {
return DecodedUDPPacket{}, err
}
if decoded.HeaderType != HeaderTypeServer {
return DecodedUDPPacket{}, ErrBadHeaderType
}
if decoded.ClientSessionID != s.clientSessionID {
return DecodedUDPPacket{}, errors.New("client session ID mismatch")
}
st.add(packetID)
st.lastSeen.Store(now)
return decoded, nil
}
// AES mode
var rawHeader [16]byte
s.codec.blockCipher.Decrypt(rawHeader[:], data[:16])
sessionID := binary.BigEndian.Uint64(rawHeader[:8])
packetID := binary.BigEndian.Uint64(rawHeader[8:16])
now := time.Now().Unix()
st, err := s.getServerSession(sessionID, now)
if err != nil {
return DecodedUDPPacket{}, err
}
if !st.check(packetID) {
return DecodedUDPPacket{}, ErrPacketIdNotUnique
}
bodyAead := st.cipher
bodyNonce := rawHeader[4:16]
bodyCipher := data[16:]
bodyPlain, err := bodyAead.Open(nil, bodyNonce, bodyCipher, nil)
if err != nil {
return DecodedUDPPacket{}, errors.New("failed to decrypt aes udp body").Base(err)
}
decoded, err := parsePlainUDPPacket(sessionID, packetID, bodyPlain)
if err != nil {
return DecodedUDPPacket{}, err
}
if decoded.HeaderType != HeaderTypeServer {
return DecodedUDPPacket{}, ErrBadHeaderType
}
if decoded.ClientSessionID != s.clientSessionID {
return DecodedUDPPacket{}, errors.New("client session ID mismatch")
}
st.add(packetID)
st.lastSeen.Store(now)
return decoded, nil
return sessionItem.EncodeServerPacket(c.method, clientSessionID, dest, payload)
}
type UDPWriter struct {
Writer io.Writer
Destination net.Destination
Session *ClientUDPSession
Codec *UDPPacketCodec
}
func (w *UDPWriter) WriteMultiBuffer(mb buf.MultiBuffer) error {
@@ -722,7 +468,7 @@ func (w *UDPWriter) WriteMultiBuffer(mb buf.MultiBuffer) error {
if b.UDP != nil {
dest = *b.UDP
}
pktBuf, err := w.Session.EncodePacket(dest, b.Bytes())
pktBuf, err := w.Codec.EncodeClientPacket(dest, b.Bytes())
b.Release()
if err != nil {
buf.ReleaseMulti(mb)
@@ -739,8 +485,8 @@ func (w *UDPWriter) WriteMultiBuffer(mb buf.MultiBuffer) error {
}
type UDPReader struct {
Reader io.Reader
Session *ClientUDPSession
Reader io.Reader
Codec *UDPPacketCodec
}
func (r *UDPReader) ReadMultiBuffer() (buf.MultiBuffer, error) {
@@ -752,7 +498,7 @@ func (r *UDPReader) ReadMultiBuffer() (buf.MultiBuffer, error) {
return nil, err
}
decoded, err := r.Session.DecodePacket(buffer.Bytes())
decoded, err := r.Codec.DecodePacket(buffer.Bytes())
if err != nil {
buffer.Release()
continue
-105
View File
@@ -2,11 +2,9 @@ package shadowsocks_2022_test
import (
"context"
"crypto/rand"
"encoding/base64"
"encoding/binary"
"errors"
"io"
gonet "net"
"sync"
"sync/atomic"
@@ -271,106 +269,3 @@ func (c *dummyStatConn) WriteMultiBuffer(mb buf.MultiBuffer) error {
}
return nil
}
func TestRelayTCPHandshakeForwarding(t *testing.T) {
methods := []string{MethodAES128GCM, MethodAES256GCM}
for _, methodName := range methods {
t.Run(methodName, func(t *testing.T) {
method, err := GetCipherMethod(methodName)
common.Must(err)
relayKey := make([]byte, method.KeySaltLength)
destKey := make([]byte, method.KeySaltLength)
_, _ = io.ReadFull(rand.Reader, relayKey)
_, _ = io.ReadFull(rand.Reader, destKey)
targetPort := uint32(54321)
relayConfig := &RelayServerConfig{
Method: methodName,
Key: base64.StdEncoding.EncodeToString(relayKey),
Destinations: []*RelayDestination{
{
Key: base64.StdEncoding.EncodeToString(destKey),
Address: net.NewIPOrDomain(net.LocalHostIP),
Port: targetPort,
Email: "test@xray.com",
},
},
}
testCtx := newTestContext()
inbound, err := NewRelayServer(testCtx, relayConfig)
common.Must(err)
targetDest := net.TCPDestination(net.LocalHostIP, net.Port(targetPort))
downstreamR, downstreamW := gonet.Pipe()
defer downstreamR.Close()
defer downstreamW.Close()
disp := &dummyDispatcher{
onDispatch: func(ctx context.Context, dest net.Destination) (*transport.Link, error) {
inLink := &transport.Link{
Reader: buf.NewReader(downstreamR),
Writer: &customWriter{
write: func(mb buf.MultiBuffer) error {
defer buf.ReleaseMulti(mb)
for _, b := range mb {
if _, err := downstreamW.Write(b.Bytes()); err != nil {
return err
}
}
return nil
},
},
}
return inLink, nil
},
}
clientConn, relayConn := gonet.Pipe()
defer clientConn.Close()
defer relayConn.Close()
go func() {
_ = inbound.Process(testCtx, net.Network_TCP, &dummyStatConn{Conn: relayConn}, disp)
}()
clientSalt := make([]byte, method.KeySaltLength)
_, _ = io.ReadFull(rand.Reader, clientSalt)
pskList := [][]byte{relayKey, destKey}
go func() {
_, err := WriteTCPRequest(clientConn, method, pskList, targetDest, clientSalt, []byte("relay payload"))
if err != nil {
t.Errorf("WriteTCPRequest failed: %v", err)
}
}()
// Downstream server must be able to read Salt + Fixed chunk in a single Read call!
headerLen := method.KeySaltLength + RequestHeaderFixedChunkLength + AEADTagSize
headerBuf := make([]byte, headerLen)
n, err := downstreamR.Read(headerBuf)
if err != nil {
t.Fatalf("downstream failed to read handshake: %v", err)
}
if n < headerLen {
t.Fatalf("downstream expected single read >= %d bytes, got %d", headerLen, n)
}
// Verify downstream can decode the fixed chunk and subsequent payload
sessionKey := DeriveSessionSubKey(destKey, headerBuf[:method.KeySaltLength], method.KeySaltLength)
aead, err := method.NewAEAD(sessionKey)
common.Must(err)
reader := NewStreamReader(downstreamR, aead)
reqHeader, err := ReadClientRequestHeaderWithFixed(reader, headerBuf[method.KeySaltLength:])
if err != nil {
t.Fatalf("downstream failed to parse client request header: %v", err)
}
if string(reqHeader.EarlyData) != "relay payload" {
t.Fatalf("payload mismatch: expected 'relay payload', got '%s'", string(reqHeader.EarlyData))
}
})
}
}
+16 -41
View File
@@ -6,11 +6,8 @@ import (
"sync/atomic"
"time"
"github.com/xtls/xray-core/common/net"
"github.com/xtls/xray-core/common/protocol"
"github.com/xtls/xray-core/common/signal"
"github.com/xtls/xray-core/common/utils"
"github.com/xtls/xray-core/transport"
)
const (
@@ -77,42 +74,30 @@ func (f *SlidingWindow) CheckAndAdd(counter uint64) bool {
type ServerUDPSession struct {
sync.Mutex
SessionID uint64
Window *SlidingWindow
User *protocol.MemoryUser
UserPSK []byte
LastActive atomic.Int64 // Unix timestamp in seconds
clientBodyCipher cipher.AEAD
SessionID uint64
RemoteCipher atomic.Pointer[cipher.AEAD]
Window SlidingWindow
User *protocol.MemoryUser
UserPSK []byte
LastActive atomic.Int64 // Unix timestamp in seconds
ServerSessionID uint64
ServerPacketID atomic.Uint64
serverBodyCipher cipher.AEAD
serverHeaderBlock cipher.Block
serverChaCha cipher.AEAD
manager *UDPSessionManager
link atomic.Pointer[transport.Link]
timer *signal.ActivityTimer
currentConn atomic.Value // stores stat.Connection
ServerCipher cipher.AEAD
ServerBlockCipher cipher.Block
ServerChaCha cipher.AEAD
}
func (s *ServerUDPSession) CheckPacketID(packetID uint64) bool {
s.Lock()
defer s.Unlock()
if s.Window == nil {
s.Window = new(SlidingWindow)
func (s *ServerUDPSession) GetRemoteCipher() cipher.AEAD {
ptr := s.RemoteCipher.Load()
if ptr == nil {
return nil
}
return s.Window.Check(packetID)
return *ptr
}
func (s *ServerUDPSession) AddPacketID(packetID uint64) {
s.Lock()
defer s.Unlock()
if s.Window == nil {
s.Window = new(SlidingWindow)
}
s.Window.Add(packetID)
func (s *ServerUDPSession) SetRemoteCipher(c cipher.AEAD) {
s.RemoteCipher.Store(&c)
}
type UDPSessionManager struct {
@@ -137,7 +122,6 @@ func (m *UDPSessionManager) GetOrCreate(sessionID uint64) *ServerUDPSession {
s := &ServerUDPSession{
SessionID: sessionID,
manager: m,
}
s.LastActive.Store(now)
@@ -164,7 +148,6 @@ func (m *UDPSessionManager) cleanup(now int64) {
m.sessions.Range(func(k uint64, v *ServerUDPSession) bool {
if now-v.LastActive.Load() > timeoutSec {
m.sessions.Delete(k)
v.Close()
}
return true
})
@@ -173,11 +156,3 @@ func (m *UDPSessionManager) cleanup(now int64) {
func (m *UDPSessionManager) Delete(sessionID uint64) {
m.sessions.Delete(sessionID)
}
func (m *UDPSessionManager) EncodeServerPacket(method *CipherMethod, psk []byte, clientSessionID uint64, dest net.Destination, payload []byte) ([]byte, error) {
sessionItem := m.GetOrCreate(clientSessionID)
if err := sessionItem.EnsureServerState(method, psk); err != nil {
return nil, err
}
return sessionItem.EncodeServerPacket(method, clientSessionID, dest, payload)
}
+6 -149
View File
@@ -2,161 +2,18 @@ package shadowsocks_2022
import (
"context"
"sync"
"github.com/xtls/xray-core/common"
"github.com/xtls/xray-core/common/buf"
"github.com/xtls/xray-core/common/errors"
"github.com/xtls/xray-core/common/log"
"github.com/xtls/xray-core/common/net"
"github.com/xtls/xray-core/common/session"
"github.com/xtls/xray-core/common/signal"
"github.com/xtls/xray-core/features/policy"
"github.com/xtls/xray-core/features/routing"
"github.com/xtls/xray-core/proxy"
"github.com/xtls/xray-core/transport"
"github.com/xtls/xray-core/transport/internet/stat"
)
func (s *ServerUDPSession) UpdateConn(conn stat.Connection) {
if s.currentConn.Load() == nil {
s.currentConn.Store(conn)
}
if s.timer != nil {
s.timer.Update()
}
}
func (s *ServerUDPSession) WriteToClient(b []byte) error {
connVal := s.currentConn.Load()
if connVal == nil {
return errors.New("client connection closed")
}
conn, ok := connVal.(stat.Connection)
if !ok || conn == nil {
return errors.New("client connection closed")
}
_, err := conn.Write(b)
return err
}
func (s *ServerUDPSession) Close() {
if s.timer != nil {
s.timer.SetTimeout(0)
}
if link := s.link.Load(); link != nil {
common.Interrupt(link.Reader)
common.Interrupt(link.Writer)
}
}
func (s *ServerUDPSession) EnsureLink(
ctx context.Context,
conn stat.Connection,
dest net.Destination,
dispatcher routing.Dispatcher,
policyManager policy.Manager,
responseEncoder func(dest net.Destination, payload []byte) ([]byte, error),
) (*transport.Link, error) {
s.UpdateConn(conn)
if link := s.link.Load(); link != nil {
return link, nil
}
s.Lock()
defer s.Unlock()
if link := s.link.Load(); link != nil {
return link, nil
}
sessCtx, cancel := context.WithCancel(ctx)
inbound := session.InboundFromContext(sessCtx)
if inbound != nil && s.User != nil {
inbound.User = s.User
}
var email string
var level uint32
if s.User != nil {
email = s.User.Email
level = s.User.Level
}
sessCtx = log.ContextWithAccessMessage(sessCtx, &log.AccessMessage{
From: conn.RemoteAddr(),
To: dest,
Status: log.AccessAccepted,
Email: email,
})
link, err := dispatcher.Dispatch(sessCtx, dest)
if err != nil {
cancel()
return nil, err
}
s.link.Store(link)
sessionPolicy := policyManager.ForLevel(level)
s.timer = signal.CancelAfterInactivity(sessCtx, func() {
if s.manager != nil {
s.manager.Delete(s.SessionID)
}
s.Close()
cancel()
}, sessionPolicy.Timeouts.ConnectionIdle)
go handleUDPResponse(s, link, dest, responseEncoder)
return link, nil
}
// ResetTCPConn sets SO_LINGER to 0 per SIP022 §3.1.4 to consistently send RST on close
// when handshake or header validation fails.
func ResetTCPConn(conn net.Conn) {
rawConn, _, _ := proxy.UnwrapRawConn(conn)
if tcpConn, ok := rawConn.(*net.TCPConn); ok {
_ = tcpConn.SetLinger(0)
}
}
func handleUDPResponse(s *ServerUDPSession, link *transport.Link, fallbackDest net.Destination, encode func(dest net.Destination, payload []byte) ([]byte, error)) {
defer func() {
if s.timer != nil {
s.timer.SetTimeout(0)
}
}()
for {
resMb, err := link.Reader.ReadMultiBuffer()
if err != nil {
return
}
if s.timer != nil {
s.timer.Update()
}
for i, rb := range resMb {
b := rb.Bytes()
if encode != nil {
replyDest := fallbackDest
if rb.UDP != nil {
replyDest = *rb.UDP
}
encPacket, err := encode(replyDest, b)
rb.Release()
if err != nil {
continue
}
if err := s.WriteToClient(encPacket); err != nil {
buf.ReleaseMulti(resMb[i+1:])
return
}
} else {
err := s.WriteToClient(b)
rb.Release()
if err != nil {
buf.ReleaseMulti(resMb[i+1:])
return
}
}
}
}
type udpConnEntry struct {
sync.Mutex
link *transport.Link
timer *signal.ActivityTimer
cancel context.CancelFunc
}
const (
+36 -172
View File
@@ -182,48 +182,57 @@ func TestTCPStream(t *testing.T) {
common.Must(err)
IncreaseNonce(reader.Nonce())
dest, addrLen, err := ParseAddressPort(plainVar)
vBuf := buf.New()
vBuf.Write(plainVar)
receivedDest, err = ReadAddressPort(vBuf)
common.Must(err)
receivedDest = net.TCPDestination(dest.Address, dest.Port)
plainVar = plainVar[addrLen:]
padLen := int(binary.BigEndian.Uint16(plainVar[:2]))
receivedPayload = plainVar[2+padLen:]
// Server sends response stream with receivedPayload as first payload
writer := NewServerStreamWriter(serverConn, method, rawKey, salt)
pBuf := buf.New()
pBuf.Write(receivedPayload)
_ = writer.WriteMultiBuffer(buf.MultiBuffer{pBuf})
// Skip padding
var padBytes [2]byte
_, _ = vBuf.Read(padBytes[:])
padLen := int(padBytes[0])<<8 | int(padBytes[1])
vBuf.Advance(int32(padLen))
// Read and echo additional stream data
receivedPayload = make([]byte, vBuf.Len())
copy(receivedPayload, vBuf.Bytes())
vBuf.Release()
// Server sends response handshake
serverSalt := make([]byte, method.KeySaltLength)
_, _ = rand.Read(serverSalt)
respKey := DeriveSessionSubKey(rawKey, serverSalt, method.KeySaltLength)
respAead, err := method.NewAEAD(respKey)
writer := NewStreamWriter(serverConn, respAead)
_, _ = serverConn.Write(serverSalt)
fixedResp := make([]byte, 1+8+method.KeySaltLength+2)
fixedResp[0] = HeaderTypeServer
binary.BigEndian.PutUint64(fixedResp[1:9], uint64(time.Now().Unix()))
copy(fixedResp[9:9+method.KeySaltLength], salt)
binary.BigEndian.PutUint16(fixedResp[9+method.KeySaltLength:11+method.KeySaltLength], 0)
fixedChunk := respAead.Seal(nil, writer.Nonce(), fixedResp, nil)
IncreaseNonce(writer.Nonce())
_, _ = serverConn.Write(fixedChunk)
// Echo stream data
mb, err := reader.ReadMultiBuffer()
common.Must(err)
_ = writer.WriteMultiBuffer(mb)
_ = writer.Close()
}()
// Client goroutine
go func() {
defer wg.Done()
clientSalt := make([]byte, method.KeySaltLength)
common.Must2(io.ReadFull(rand.Reader, clientSalt))
writer, err := WriteTCPRequest(clientConn, method, [][]byte{rawKey}, dest, clientSalt, testPayload)
clientSalt, writer, err := ClientHandshake(clientConn, method, [][]byte{rawKey}, dest, testPayload)
common.Must(err)
reader, err := ReadTCPResponse(clientConn, method, rawKey, clientSalt)
reader, _, err := ClientVerifyServerResponse(clientConn, method, rawKey, clientSalt)
common.Must(err)
// The first ReadMultiBuffer drains initialPayload from reader cache
mbInit, err := reader.ReadMultiBuffer()
common.Must(err)
if !bytes.Equal(mbInit[0].Bytes(), testPayload) {
t.Errorf("drained initial payload mismatch: got %s, want %s", mbInit[0].Bytes(), testPayload)
}
buf.ReleaseMulti(mbInit)
// Send additional stream data
streamData := []byte("stream chunk test")
_ = writer.WriteMultiBuffer(buf.MultiBuffer{buf.FromBytes(streamData)})
_ = writer.WriteChunk(streamData)
mb, err := reader.ReadMultiBuffer()
common.Must(err)
@@ -263,14 +272,12 @@ func TestUDPCodec(t *testing.T) {
psk := make([]byte, method.KeySaltLength)
_, _ = rand.Read(psk)
clientCodec, err := NewUDPPacketCodec(method, [][]byte{psk})
clientCodec, err := NewUDPPacketCodec(method, psk)
common.Must(err)
serverCodec, err := NewUDPServerCodec(method, psk, time.Minute)
common.Must(err)
session, err := clientCodec.NewClientSession()
common.Must(err)
pktBuf, err := session.EncodePacket(dest, payload)
pktBuf, err := clientCodec.EncodeClientPacket(dest, payload)
common.Must(err)
defer pktBuf.Release()
@@ -353,146 +360,3 @@ func TestMultiUserManager(t *testing.T) {
t.Fatal("user1 should have been removed")
}
}
func TestLargeStreamTransfer(t *testing.T) {
method, err := GetCipherMethod(MethodAES128GCM)
common.Must(err)
sessionKey := make([]byte, 16)
_, _ = rand.Read(sessionKey)
clientAead, err := method.NewAEAD(sessionKey)
common.Must(err)
serverAead, err := method.NewAEAD(sessionKey)
common.Must(err)
r, w := io.Pipe()
defer r.Close()
defer w.Close()
writer := NewStreamWriter(w, clientAead)
reader := NewStreamReader(r, serverAead)
const totalSize = 100 * 1024 // 100 KB
data := make([]byte, totalSize)
_, _ = rand.Read(data)
errCh := make(chan error, 1)
go func() {
// Write using Write (which splits by MaxPacketSize = 65535)
_, werr := writer.Write(data)
if werr != nil {
errCh <- werr
return
}
_ = w.Close()
errCh <- nil
}()
var received []byte
for {
mb, rerr := reader.ReadMultiBuffer()
if !mb.IsEmpty() {
for _, b := range mb {
received = append(received, b.Bytes()...)
}
buf.ReleaseMulti(mb)
}
if rerr != nil {
if rerr == io.EOF {
break
}
t.Fatalf("ReadMultiBuffer error: %v", rerr)
}
}
if werr := <-errCh; werr != nil {
t.Fatalf("writer error: %v", werr)
}
if len(received) != totalSize {
t.Fatalf("received size mismatch: got %d, want %d", len(received), totalSize)
}
if !bytes.Equal(received, data) {
t.Fatal("received data does not match sent data")
}
}
func TestClientUDPSessionMultiDestination(t *testing.T) {
for _, methodName := range []string{MethodAES128GCM, MethodAES256GCM, MethodChaCha20Poly1305} {
t.Run(methodName, func(t *testing.T) {
method, err := GetCipherMethod(methodName)
common.Must(err)
rawKey := make([]byte, method.KeySaltLength)
_, _ = rand.Read(rawKey)
clientCodec, err := NewUDPPacketCodec(method, [][]byte{rawKey})
common.Must(err)
serverCodec, err := NewUDPServerCodec(method, rawKey, time.Minute)
common.Must(err)
session, err := clientCodec.NewClientSession()
common.Must(err)
dest1 := net.UDPDestination(net.LocalHostIP, net.Port(53))
dest2 := net.UDPDestination(net.IPAddress([]byte{127, 0, 0, 2}), net.Port(53))
payload1 := []byte("query-google-dns")
payload2 := []byte("query-cloudflare-dns")
// Client sends to dest1 and dest2 using SAME session
pkt1, err := session.EncodePacket(dest1, payload1)
common.Must(err)
defer pkt1.Release()
pkt2, err := session.EncodePacket(dest2, payload2)
common.Must(err)
defer pkt2.Release()
// Server decodes both
dec1, err := serverCodec.DecodePacket(pkt1.Bytes())
common.Must(err)
dec2, err := serverCodec.DecodePacket(pkt2.Bytes())
common.Must(err)
if dec1.SessionID != session.ClientSessionID() || dec2.SessionID != session.ClientSessionID() {
t.Fatalf("both packets must share client session ID %d, got %d and %d", session.ClientSessionID(), dec1.SessionID, dec2.SessionID)
}
if dec1.Destination.String() != dest1.String() {
t.Fatalf("expected dest1 %s, got %s", dest1, dec1.Destination)
}
if dec2.Destination.String() != dest2.String() {
t.Fatalf("expected dest2 %s, got %s", dest2, dec2.Destination)
}
if !bytes.Equal(dec1.Payload, payload1) || !bytes.Equal(dec2.Payload, payload2) {
t.Fatal("payload mismatch")
}
// Server replies to dest1 and dest2
respPayload1 := []byte("reply-google-dns")
respPayload2 := []byte("reply-cloudflare-dns")
respPkt1, err := serverCodec.EncodeServerPacket(dec1.SessionID, dest1, respPayload1)
common.Must(err)
respPkt2, err := serverCodec.EncodeServerPacket(dec2.SessionID, dest2, respPayload2)
common.Must(err)
// Client decodes replies
clientDec1, err := session.DecodePacket(respPkt1)
common.Must(err)
if clientDec1.Destination.String() != dest1.String() {
t.Fatalf("expected client dec1 dest %s, got %s", dest1, clientDec1.Destination)
}
if !bytes.Equal(clientDec1.Payload, respPayload1) {
t.Fatal("reply payload 1 mismatch")
}
clientDec2, err := session.DecodePacket(respPkt2)
common.Must(err)
if clientDec2.Destination.String() != dest2.String() {
t.Fatalf("expected client dec2 dest %s, got %s", dest2, clientDec2.Destination)
}
if !bytes.Equal(clientDec2.Payload, respPayload2) {
t.Fatal("reply payload 2 mismatch")
}
})
}
}
+115 -235
View File
@@ -1,25 +1,18 @@
package shadowsocks_2022
import (
"context"
"crypto/cipher"
"crypto/rand"
"encoding/binary"
"io"
"math"
mrand "math/rand/v2"
"sync"
"time"
"github.com/xtls/xray-core/common/antireplay"
"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/common/protocol"
"github.com/xtls/xray-core/common/signal"
"github.com/xtls/xray-core/common/task"
"github.com/xtls/xray-core/features/policy"
"github.com/xtls/xray-core/transport"
)
var addrParser = protocol.NewAddressParser(
@@ -45,6 +38,15 @@ func WriteAddressPort(w io.Writer, dest net.Destination) error {
return addrParser.WriteAddressPort(w, dest.Address, dest.Port)
}
// ReadAddressPort reads a destination address and port in SOCKS5 format
func ReadAddressPort(r io.Reader) (net.Destination, error) {
addr, port, err := addrParser.ReadAddressPort(nil, r)
if err != nil {
return net.Destination{}, err
}
return net.TCPDestination(addr, port), nil
}
// AddrPortLength returns the serialized length of a destination in SOCKS5 format
func AddrPortLength(dest net.Destination) int {
switch dest.Address.Family() {
@@ -117,16 +119,8 @@ func (w *StreamWriter) Write(p []byte) (int, error) {
func (w *StreamWriter) WriteMultiBuffer(mb buf.MultiBuffer) error {
defer buf.ReleaseMulti(mb)
for _, b := range mb {
p := b.Bytes()
for len(p) > 0 {
chunkSize := len(p)
if chunkSize > MaxPacketSize {
chunkSize = MaxPacketSize
}
if err := w.WriteChunk(p[:chunkSize]); err != nil {
return err
}
p = p[chunkSize:]
if err := w.WriteChunk(b.Bytes()); err != nil {
return err
}
}
return nil
@@ -174,7 +168,7 @@ func (r *StreamReader) Read(p []byte) (int, error) {
IncreaseNonce(r.nonce[:])
payloadLen := int(binary.BigEndian.Uint16(decryptedLen))
if payloadLen == 0 || payloadLen > MaxPacketSize {
if payloadLen == 0 {
return 0, ErrInvalidRequest
}
@@ -200,10 +194,11 @@ func (r *StreamReader) Read(p []byte) (int, error) {
func (r *StreamReader) ReadMultiBuffer() (buf.MultiBuffer, error) {
if r.cached > 0 {
mb := buf.MergeBytes(nil, r.buffer[r.offset:r.offset+r.cached])
b := buf.New()
b.Write(r.buffer[r.offset : r.offset+r.cached])
r.cached = 0
r.offset = 0
return mb, nil
return buf.MultiBuffer{b}, nil
}
if _, err := io.ReadFull(r.reader, r.lenBuf[:]); err != nil {
@@ -217,7 +212,7 @@ func (r *StreamReader) ReadMultiBuffer() (buf.MultiBuffer, error) {
IncreaseNonce(r.nonce[:])
payloadLen := int(binary.BigEndian.Uint16(decryptedLen))
if payloadLen == 0 || payloadLen > MaxPacketSize {
if payloadLen == 0 {
return nil, ErrInvalidRequest
}
@@ -232,8 +227,9 @@ func (r *StreamReader) ReadMultiBuffer() (buf.MultiBuffer, error) {
}
IncreaseNonce(r.nonce[:])
mb := buf.MergeBytes(nil, decryptedPayload)
return mb, nil
b := buf.New()
b.Write(decryptedPayload)
return buf.MultiBuffer{b}, nil
}
type ClientRequestHeader struct {
@@ -241,8 +237,13 @@ type ClientRequestHeader struct {
EarlyData []byte
}
func ReadClientRequestHeaderWithFixed(reader *StreamReader, fixedChunk []byte) (*ClientRequestHeader, error) {
plainFixed, err := reader.cipher.Open(fixedChunk[:0], reader.Nonce(), fixedChunk, nil)
func ReadClientRequestHeader(conn io.Reader, reader *StreamReader) (*ClientRequestHeader, error) {
var fixedBuf [RequestHeaderFixedChunkLength + AEADTagSize]byte
if _, err := io.ReadFull(conn, fixedBuf[:]); err != nil {
return nil, err
}
plainFixed, err := reader.cipher.Open(fixedBuf[:0], reader.Nonce(), fixedBuf[:], nil)
if err != nil {
return nil, errors.New("failed to decrypt client request header").Base(err)
}
@@ -271,7 +272,7 @@ func ReadClientRequestHeaderWithFixed(reader *StreamReader, fixedChunk []byte) (
} else {
varChunkCipher = make([]byte, needed)
}
if _, err := io.ReadFull(reader.reader, varChunkCipher); err != nil {
if _, err := io.ReadFull(conn, varChunkCipher); err != nil {
return nil, err
}
@@ -281,34 +282,31 @@ func ReadClientRequestHeaderWithFixed(reader *StreamReader, fixedChunk []byte) (
}
IncreaseNonce(reader.Nonce())
dest, addrLen, err := ParseAddressPort(plainVar)
b := buf.New()
b.Write(plainVar)
defer b.Release()
dest, err := ReadAddressPort(b)
if err != nil {
return nil, err
}
dest.Network = net.Network_TCP
offset := addrLen
if len(plainVar) < offset+2 {
return nil, ErrPacketTooShort
var padLenBytes [2]byte
if _, err := b.Read(padLenBytes[:]); err != nil {
return nil, err
}
paddingLen := int(binary.BigEndian.Uint16(plainVar[offset : offset+2]))
offset += 2
if len(plainVar) < offset+paddingLen {
paddingLen := int(binary.BigEndian.Uint16(padLenBytes[:]))
if int(b.Len()) < paddingLen {
return nil, ErrNoPadding
}
offset += paddingLen
var earlyData []byte
var payloadLen int
if len(plainVar) > offset {
earlyData = plainVar[offset:]
payloadLen = len(earlyData)
if paddingLen > 0 {
b.Advance(int32(paddingLen))
}
// SIP022 §3.1.4: Servers MUST reject the request if the variable-length header chunk does not contain payload and the padding length is 0.
if paddingLen == 0 && payloadLen == 0 {
return nil, errors.New("request without payload and padding is not allowed")
var earlyData []byte
if b.Len() > 0 {
earlyData = make([]byte, b.Len())
copy(earlyData, b.Bytes())
}
return &ClientRequestHeader{
@@ -317,6 +315,34 @@ func ReadClientRequestHeaderWithFixed(reader *StreamReader, fixedChunk []byte) (
}, nil
}
// ClientHandshake writes the full client request header to w
func ClientHandshake(w io.Writer, method *CipherMethod, pskList [][]byte, dest net.Destination, payload []byte) ([]byte, *StreamWriter, error) {
salt := make([]byte, method.KeySaltLength)
if _, err := io.ReadFull(rand.Reader, salt); err != nil {
return nil, nil, err
}
writer, err := WriteTCPRequest(w, method, pskList, dest, salt, payload)
if err != nil {
return nil, nil, err
}
return salt, writer.(*StreamWriter), nil
}
// ClientVerifyServerResponse reads and verifies the server's handshake response
func ClientVerifyServerResponse(r io.Reader, method *CipherMethod, psk []byte, clientSalt []byte) (*StreamReader, []byte, error) {
reader, err := ReadTCPResponse(r, method, psk, clientSalt)
if err != nil {
return nil, nil, err
}
sr := reader.(*StreamReader)
var initialPayload []byte
if sr.cached > 0 {
initialPayload = make([]byte, sr.cached)
copy(initialPayload, sr.buffer[sr.offset:sr.offset+sr.cached])
}
return sr, initialPayload, nil
}
// WriteTCPRequest writes the Shadowsocks 2022 request header into w and returns a body writer.
func WriteTCPRequest(w io.Writer, method *CipherMethod, pskList [][]byte, dest net.Destination, clientSalt []byte, payload []byte) (buf.Writer, error) {
finalPSK := pskList[len(pskList)-1]
@@ -328,16 +354,7 @@ func WriteTCPRequest(w io.Writer, method *CipherMethod, pskList [][]byte, dest n
writer := NewStreamWriter(w, aead)
payloadLen := len(payload)
var paddingLen int
if payloadLen < MaxPaddingLength {
paddingLen = mrand.IntN(MaxPaddingLength) + 1
}
addrPortLen := AddrPortLength(dest)
varHeaderLen := addrPortLen + 2 + paddingLen + payloadLen
totalHandshakeLen := int32(method.KeySaltLength + len(pskList)*AESBlockSize + RequestHeaderFixedChunkLength + AEADTagSize + varHeaderLen + AEADTagSize)
handshakeBuf := buf.NewWithSize(totalHandshakeLen)
handshakeBuf := buf.New()
defer handshakeBuf.Release()
handshakeBuf.Write(clientSalt)
@@ -355,6 +372,14 @@ func WriteTCPRequest(w io.Writer, method *CipherMethod, pskList [][]byte, dest n
handshakeBuf.Write(encryptedEIH[:])
}
payloadLen := len(payload)
var paddingLen int
if payloadLen < MaxPaddingLength {
paddingLen = mrand.IntN(MaxPaddingLength-payloadLen) + 1
}
addrPortLen := AddrPortLength(dest)
varHeaderLen := addrPortLen + 2 + paddingLen + payloadLen
var fixedHeaderPlaintext [RequestHeaderFixedChunkLength]byte
fixedHeaderPlaintext[0] = HeaderTypeClient
binary.BigEndian.PutUint64(fixedHeaderPlaintext[1:9], uint64(time.Now().Unix()))
@@ -364,7 +389,7 @@ func WriteTCPRequest(w io.Writer, method *CipherMethod, pskList [][]byte, dest n
IncreaseNonce(writer.nonce[:])
handshakeBuf.Write(fixedChunk)
varHeaderBuf := buf.NewWithSize(int32(varHeaderLen))
varHeaderBuf := buf.New()
defer varHeaderBuf.Release()
if err := WriteAddressPort(varHeaderBuf, dest); err != nil {
@@ -396,21 +421,12 @@ func WriteTCPRequest(w io.Writer, method *CipherMethod, pskList [][]byte, dest n
// ReadTCPResponse reads and verifies the server's handshake response and returns a reader for the stream.
func ReadTCPResponse(r io.Reader, method *CipherMethod, psk []byte, clientSalt []byte) (buf.Reader, error) {
fixedPlainLen := 1 + 8 + method.KeySaltLength + 2
chunkCipherLen := fixedPlainLen + AEADTagSize
headerLen := method.KeySaltLength + chunkCipherLen
// Single read call for Salt + Fixed-length response header chunk per SIP022 §3.1.4
var headerBuf [128]byte
headerSlice := headerBuf[:headerLen]
n, err := r.Read(headerSlice)
if err != nil || n < headerLen {
return nil, errors.New("failed to read complete server response header")
var serverSalt [32]byte
serverSaltSlice := serverSalt[:method.KeySaltLength]
if _, err := io.ReadFull(r, serverSaltSlice); err != nil {
return nil, err
}
serverSaltSlice := headerSlice[:method.KeySaltLength]
chunkSlice := headerSlice[method.KeySaltLength:headerLen]
sessionKey := DeriveSessionSubKey(psk, serverSaltSlice, method.KeySaltLength)
aead, err := method.NewAEAD(sessionKey)
if err != nil {
@@ -419,6 +435,14 @@ func ReadTCPResponse(r io.Reader, method *CipherMethod, psk []byte, clientSalt [
reader := NewStreamReader(r, aead)
fixedPlainLen := 1 + 8 + method.KeySaltLength + 2
chunkCipherLen := fixedPlainLen + AEADTagSize
var chunkBuf [64]byte
chunkSlice := chunkBuf[:chunkCipherLen]
if _, err := io.ReadFull(r, chunkSlice); err != nil {
return nil, err
}
decryptedFixed, err := reader.cipher.Open(chunkSlice[:0], reader.nonce[:], chunkSlice, nil)
if err != nil {
return nil, errors.New("failed to decrypt server response header").Base(err)
@@ -460,190 +484,46 @@ func ReadTCPResponse(r io.Reader, method *CipherMethod, psk []byte, clientSalt [
return reader, nil
}
// ServerStreamWriter lazily sends the response header along with the first payload chunk per SIP022 §3.1.2 & §3.1.4.
type ServerStreamWriter struct {
mu sync.Mutex
w io.Writer
method *CipherMethod
psk []byte
clientSalt []byte
streamWriter *StreamWriter
}
func NewServerStreamWriter(w io.Writer, method *CipherMethod, psk []byte, clientSalt []byte) *ServerStreamWriter {
return &ServerStreamWriter{
w: w,
method: method,
psk: psk,
clientSalt: clientSalt,
}
}
func (s *ServerStreamWriter) sendHeaderWithFirstPayload(payload []byte) (*StreamWriter, error) {
// WriteTCPResponse writes the server handshake response and returns a body writer for server stream.
func WriteTCPResponse(w io.Writer, method *CipherMethod, psk []byte, clientSalt []byte, initialPayload []byte) (buf.Writer, error) {
var serverSalt [32]byte
serverSaltSlice := serverSalt[:s.method.KeySaltLength]
serverSaltSlice := serverSalt[:method.KeySaltLength]
if _, err := io.ReadFull(rand.Reader, serverSaltSlice); err != nil {
return nil, err
}
respKey := DeriveSessionSubKey(s.psk, serverSaltSlice, s.method.KeySaltLength)
respAead, err := s.method.NewAEAD(respKey)
respKey := DeriveSessionSubKey(psk, serverSaltSlice, method.KeySaltLength)
respAead, err := method.NewAEAD(respKey)
if err != nil {
return nil, err
}
sw := NewStreamWriter(s.w, respAead)
writer := NewStreamWriter(w, respAead)
totalHeaderLen := int32(s.method.KeySaltLength + 1 + 8 + s.method.KeySaltLength + 2 + AEADTagSize + len(payload) + AEADTagSize)
outBuf := buf.NewWithSize(totalHeaderLen)
defer outBuf.Release()
respBuf := buf.New()
defer respBuf.Release()
outBuf.Write(serverSaltSlice)
respBuf.Write(serverSaltSlice)
var fixedRespPlain [1 + 8 + 32 + 2]byte
fixedRespSlice := fixedRespPlain[:1+8+s.method.KeySaltLength+2]
fixedRespSlice := fixedRespPlain[:1+8+method.KeySaltLength+2]
fixedRespSlice[0] = HeaderTypeServer
binary.BigEndian.PutUint64(fixedRespSlice[1:9], uint64(time.Now().Unix()))
copy(fixedRespSlice[9:9+s.method.KeySaltLength], s.clientSalt)
binary.BigEndian.PutUint16(fixedRespSlice[9+s.method.KeySaltLength:11+s.method.KeySaltLength], uint16(len(payload)))
copy(fixedRespSlice[9:9+method.KeySaltLength], clientSalt)
binary.BigEndian.PutUint16(fixedRespSlice[9+method.KeySaltLength:11+method.KeySaltLength], uint16(len(initialPayload)))
fixedRespChunk := sw.cipher.Seal(nil, sw.nonce[:], fixedRespSlice, nil)
IncreaseNonce(sw.nonce[:])
outBuf.Write(fixedRespChunk)
fixedRespChunk := writer.cipher.Seal(nil, writer.nonce[:], fixedRespSlice, nil)
IncreaseNonce(writer.nonce[:])
respBuf.Write(fixedRespChunk)
if len(payload) > 0 {
payloadChunk := sw.cipher.Seal(nil, sw.nonce[:], payload, nil)
IncreaseNonce(sw.nonce[:])
outBuf.Write(payloadChunk)
if len(initialPayload) > 0 {
initialChunk := writer.cipher.Seal(nil, writer.nonce[:], initialPayload, nil)
IncreaseNonce(writer.nonce[:])
respBuf.Write(initialChunk)
}
if _, err := s.w.Write(outBuf.Bytes()); err != nil {
if _, err := w.Write(respBuf.Bytes()); err != nil {
return nil, err
}
return sw, nil
}
func (s *ServerStreamWriter) WriteMultiBuffer(mb buf.MultiBuffer) error {
if mb.IsEmpty() {
return nil
}
if s.streamWriter == nil {
s.mu.Lock()
if s.streamWriter == nil {
firstBuf := mb[0]
firstBytes := firstBuf.Bytes()
chunkSize := len(firstBytes)
if chunkSize > MaxPacketSize {
chunkSize = MaxPacketSize
}
firstPayload := firstBytes[:chunkSize]
sw, err := s.sendHeaderWithFirstPayload(firstPayload)
if err != nil {
s.mu.Unlock()
buf.ReleaseMulti(mb)
return err
}
s.streamWriter = sw
firstBuf.Advance(int32(chunkSize))
if firstBuf.IsEmpty() {
firstBuf.Release()
mb = mb[1:]
}
}
s.mu.Unlock()
if len(mb) == 0 {
return nil
}
}
return s.streamWriter.WriteMultiBuffer(mb)
}
func (s *ServerStreamWriter) Write(p []byte) (int, error) {
n := len(p)
if s.streamWriter == nil {
s.mu.Lock()
if s.streamWriter == nil {
chunkSize := len(p)
if chunkSize > MaxPacketSize {
chunkSize = MaxPacketSize
}
firstPayload := p[:chunkSize]
sw, err := s.sendHeaderWithFirstPayload(firstPayload)
if err != nil {
s.mu.Unlock()
return 0, err
}
s.streamWriter = sw
p = p[chunkSize:]
}
s.mu.Unlock()
if len(p) == 0 {
return n, nil
}
}
_, err := s.streamWriter.Write(p)
return n, err
}
func (s *ServerStreamWriter) Close() error {
if s.streamWriter == nil {
s.mu.Lock()
defer s.mu.Unlock()
if s.streamWriter == nil {
sw, err := s.sendHeaderWithFirstPayload(nil)
if err != nil {
return err
}
s.streamWriter = sw
}
}
return nil
}
// InitServerStream decrypts the client request header, verifies the timestamp and replay filter,
// and returns a StreamReader for subsequent stream chunks.
func InitServerStream(conn net.Conn, method *CipherMethod, psk, saltSlice []byte, salt [32]byte, fixedChunk []byte, saltFilter *antireplay.ReplayFilter[[32]byte]) (*StreamReader, *ClientRequestHeader, error) {
sessionKey := DeriveSessionSubKey(psk, saltSlice, method.KeySaltLength)
aead, err := method.NewAEAD(sessionKey)
if err != nil {
return nil, nil, err
}
reader := NewStreamReader(conn, aead)
reqHeader, err := ReadClientRequestHeaderWithFixed(reader, fixedChunk)
if err != nil {
return nil, nil, err
}
_ = conn.SetReadDeadline(time.Time{})
if !saltFilter.Check(salt) {
return nil, nil, ErrSaltNotUnique
}
return reader, reqHeader, nil
}
func TransportTCP(ctx context.Context, sessionPolicy policy.Session, reader buf.Reader, writer buf.Writer, link *transport.Link) error {
ctx, cancel := context.WithCancel(ctx)
timer := signal.CancelAfterInactivity(ctx, cancel, sessionPolicy.Timeouts.ConnectionIdle)
ctx = policy.ContextWithBufferPolicy(ctx, sessionPolicy.Buffer)
requestDone := func() error {
defer timer.SetTimeout(sessionPolicy.Timeouts.DownlinkOnly)
return buf.Copy(reader, link.Writer, buf.UpdateActivity(timer))
}
responseDone := func() error {
defer timer.SetTimeout(sessionPolicy.Timeouts.UplinkOnly)
if c, ok := writer.(io.Closer); ok {
defer c.Close()
}
return buf.Copy(link.Reader, writer, buf.UpdateActivity(timer))
}
responseDoneAndCloseWriter := task.OnSuccess(responseDone, task.Close(link.Writer))
return task.Run(ctx, requestDone, responseDoneAndCloseWriter)
return writer, nil
}
+1 -9
View File
@@ -170,17 +170,9 @@ func (s *Server) processTCP(ctx context.Context, conn stat.Connection, dispatche
return errors.New("UDP associate with listen port failed")
}
tempUDPConn.SetTimeout(plcy.Timeouts.ConnectionIdle)
var udpConn stat.Connection = tempUDPConn
if counters, ok := conn.(*stat.CounterConnection); ok {
udpConn = &stat.CounterConnection{
Connection: tempUDPConn,
ReadCounter: counters.ReadCounter,
WriteCounter: counters.WriteCounter,
}
}
errCh := make(chan error, 1)
go func() {
errCh <- s.handleUDPPayload(ctx, udpConn, dispatcher)
errCh <- s.handleUDPPayload(ctx, tempUDPConn, dispatcher)
}()
// Associated TCP keeps the UDP alive
// Close UDP if TCP connection is closed
-2
View File
@@ -213,8 +213,6 @@ If the filters cannot be added, Xray does not start. They are removed when Xray
`autoSystemWfpBlockLeak` (Windows only) is empty by default, as the filters break some setups: with `"dns"`, a local DNS resolver other programs use (e.g. on `127.0.0.1:53`), the DNS of another VPN on its own interface, virtual machines whose NAT resolves names on the host, or signing in to a captive portal; with `"misconfigtun"`, IPv4 or IPv6 on the local network while no route of that version leads to the TUN. Without the filters, DNS may leak as described above. To keep an IP version out of the TUN on purpose while still blocking DNS leaks, use only `["dns"]`.
`autoOutboundsInterface` (the default with `autoSystemRoutingTable`) keeps Xray's own connections out of the TUN by binding them to another interface, which Windows only honors while that interface has weak host send and forwarding off for the IP versions routed to the TUN. Otherwise, Windows sends them into the TUN, from that interface's address, and they stall. While the TUN runs, Xray therefore turns weak host send off on that interface, and on again when it stops or another interface takes over. Forwarding cannot be turned off this way, as Mobile Hotspot and Internet Connection Sharing need it, so a warning is logged while it is on. Having the hotspot share the TUN instead of that interface (Settings, Mobile hotspot, Share my internet connection from) moves forwarding to the TUN, where it does no harm, and sends the hotspot's devices through Xray as well.
You can give the adapter ip address manually, you can live Windows to give it autogenerated ip address (which take few seconds), it doesn't matter, the traffic going _through_ the interface will be forwarded into the app for proxying. \
Minimal configuration that will work for local machine is routing passing the traffic on-link through the interface.
You will need the interface id for that, unfortunately it is going to change with every Xray start due to implementation ambiguity between Xray and wintun driver.
+1 -1
View File
@@ -101,7 +101,7 @@ func (t *stackGVisor) Start() error {
// Use custom UDP packet handler, instead of strict gVisor forwarder, for FullCone NAT support
udpForwarder := newUdpConnectionHandler(t.handler.HandleConnection, t.writeRawUDPPacket)
ipStack.SetTransportProtocolHandler(udp.ProtocolNumber, func(id stack.TransportEndpointID, pkt *stack.PacketBuffer) bool {
data := pkt.Data().AsRange().ToSlice()
data := pkt.Clone().Data().AsRange().ToSlice()
// if len(data) == 0 {
// return false
// }
+6 -10
View File
@@ -155,20 +155,12 @@ func NewTun(options *Config) (Tun, error) {
fdStr := platform.NewEnvFlag(platform.TunFdKey).GetValue(func() string { return "" })
if fdStr != "" {
// iOS: use provided fd from NetworkExtension
providedFd, err := strconv.Atoi(fdStr)
if err != nil {
return nil, err
}
// duplicate NetworkExtension fd so Xray can close its own handle
// without closing the original.
fd, err := unix.FcntlInt(uintptr(providedFd), unix.F_DUPFD_CLOEXEC, 0)
fd, err := strconv.Atoi(fdStr)
if err != nil {
return nil, err
}
if err = unix.SetNonblock(fd, true); err != nil {
_ = unix.Close(fd)
return nil, err
}
@@ -240,7 +232,11 @@ func (t *DarwinTun) Close() error {
t.waitKq.close()
}
routeErr := t.unsetSystemRoutes()
return xerrors.Combine(routeErr, t.tunFile.Close())
if t.ownsFd {
return xerrors.Combine(routeErr, t.tunFile.Close())
}
// iOS: don't close the fd, it's owned by NetworkExtension
return routeErr
}
func (t *DarwinTun) monitorRouteChanges() {
-19
View File
@@ -46,7 +46,6 @@ type WindowsTun struct {
luid winipcfg.LUID
cbr winipcfg.ChangeCallback
cbi winipcfg.ChangeCallback
guard outboundGuard
wfp windows.Handle
resolver *savedResolver
skipStop chan struct{}
@@ -179,11 +178,6 @@ startOver:
}
ipif, err := t.luid.IPInterface(family)
if err != nil {
// With IPv6 disabled system-wide (DisabledComponents), the adapter has no
// IPv6 interface at all. Skip the family unless the config asks for it.
if err == windows.ERROR_NOT_FOUND && family == windows.AF_INET6 && !address6 && !route6 {
continue
}
return err
}
ipif.RouterDiscoveryBehavior = winipcfg.RouterDiscoveryDisabled
@@ -298,21 +292,10 @@ startOver:
}
if updater != nil {
// Xray's own connections have to stay out of the IP versions routed
// to the TUN, which needs Windows to honor the binding to updater's
// interface.
if route4 {
t.guard.families = append(t.guard.families, windows.AF_INET)
}
if route6 {
t.guard.families = append(t.guard.families, windows.AF_INET6)
}
t.guard.check()
// Only a registered callback goes into the fields: a nil pointer in
// them would not compare equal to nil in Close.
cbr, err := winipcfg.RegisterRouteChangeCallback(func(notificationType winipcfg.MibNotificationType, route *winipcfg.MibIPforwardRow2) {
updater.Update()
t.guard.check()
})
if err != nil {
return err
@@ -320,7 +303,6 @@ startOver:
t.cbr = cbr
cbi, err := winipcfg.RegisterInterfaceChangeCallback(func(notificationType winipcfg.MibNotificationType, iface *winipcfg.MibIPInterfaceRow) {
updater.Update()
t.guard.check()
})
if err != nil {
return err
@@ -344,7 +326,6 @@ func (t *WindowsTun) Close() error {
if t.cbi != nil {
t.cbi.Unregister()
}
t.guard.restore()
if t.luid != 0 {
t.luid.FlushRoutes(windows.AF_INET)
t.luid.FlushIPAddresses(windows.AF_INET)
-120
View File
@@ -1,120 +0,0 @@
//go:build windows
package tun
import (
"context"
"slices"
"strings"
"sync"
"github.com/xtls/xray-core/common/errors"
"golang.org/x/sys/windows"
"golang.zx2c4.com/wireguard/windows/tunnel/winipcfg"
)
// outboundGuard keeps Windows to the binding of autoOutboundsInterface, which
// keeps Xray's own connections out of the TUN. With weak host send or
// forwarding on for an IP version on the bound interface, Windows sends them
// where the routes lead, into the TUN, from that interface's address, and
// drops what comes back to that address through the TUN, so they stall.
//
// For the IP versions routed to the TUN, weak host send is turned off on the
// bound interface while the TUN runs, and turned on again when the TUN stops
// or another interface takes over. Forwarding is what Mobile Hotspot and
// Internet Connection Sharing need, so it is only reported.
type outboundGuard struct {
sync.Mutex
families []winipcfg.AddressFamily
luid winipcfg.LUID // of the interface last checked
name string // of that interface
turnedOff []winipcfg.AddressFamily // where weak host send was turned off on it
forwarding bool // whether forwarding was on there
stopped bool
}
// check turns weak host send off on the bound interface, and warns when
// forwarding comes on there, but not again while it stays on.
func (g *outboundGuard) check() {
g.Lock()
defer g.Unlock()
if g.stopped {
return
}
var luid winipcfg.LUID
var name string
if iface := updater.Get(); iface != nil {
luid, _ = winipcfg.LUIDFromIndex(uint32(iface.Index))
name = iface.Name
}
if luid != g.luid {
g.restoreLocked()
g.luid, g.name = luid, name
g.forwarding = false // to warn about the new interface as well
}
if luid == 0 {
return
}
var forwarding []string
for _, family := range g.families {
row, err := luid.IPInterface(family)
if err != nil {
continue // the interface lacks that IP version
}
if row.ForwardingEnabled {
forwarding = append(forwarding, familyName(family))
}
if !row.WeakHostSend {
continue
}
if err := setWeakHostSend(row, false); err != nil {
errors.LogWarningInner(context.Background(), err, "[tun] unable to turn weak host send off for ", familyName(family), " on ", name)
continue
}
if !slices.Contains(g.turnedOff, family) {
g.turnedOff = append(g.turnedOff, family)
errors.LogInfo(context.Background(), "[tun] weak host send turned off for ", familyName(family), " on ", name, " while the TUN runs, as Windows would ignore autoOutboundsInterface")
}
}
wasOn := g.forwarding
g.forwarding = len(forwarding) > 0
if g.forwarding && !wasOn {
errors.LogWarning(context.Background(), "[tun] forwarding is on for ", strings.Join(forwarding, " and "), " on ", name, " (Mobile Hotspot and Internet Connection Sharing turn it on), so Windows ignores autoOutboundsInterface there, and Xray's own connections go into the TUN and stall: turn the hotspot off, or have it share the TUN instead of ", name)
}
}
// restore turns weak host send on again where check turned it off, for good.
func (g *outboundGuard) restore() {
g.Lock()
defer g.Unlock()
g.restoreLocked()
g.stopped = true
}
func (g *outboundGuard) restoreLocked() {
for _, family := range g.turnedOff {
row, err := g.luid.IPInterface(family)
if err == nil {
err = setWeakHostSend(row, true)
}
if err != nil {
errors.LogWarningInner(context.Background(), err, "[tun] unable to turn weak host send on again for ", familyName(family), " on ", g.name)
}
}
g.turnedOff = nil
}
func setWeakHostSend(row *winipcfg.MibIPInterfaceRow, on bool) error {
row.WeakHostSend = on
if row.Family == windows.AF_INET {
row.SitePrefixLength = 0 // as SetIpInterfaceEntry requires for IPv4
}
return row.Set()
}
func familyName(family winipcfg.AddressFamily) string {
if family == windows.AF_INET {
return "IPv4"
}
return "IPv6"
}
+2 -12
View File
@@ -52,12 +52,9 @@ func (b *bind) Open(port uint16) (fns []conn.ReceiveFunc, actualPort uint16, err
case <-ch:
default:
errors.LogErrorInner(context.Background(), err, "unexpected closed")
b.mu.Lock()
downFunc := b.downFunc
b.mu.Unlock()
if downFunc != nil {
if b.downFunc != nil {
go func() {
common.Must(downFunc())
common.Must(b.downFunc())
}()
}
}
@@ -79,13 +76,6 @@ func (b *bind) Open(port uint16) (fns []conn.ReceiveFunc, actualPort uint16, err
}, uint16(c.LocalAddr().(*net.UDPAddr).Port), nil
}
// setDownFunc sets downFunc after the device is created, since the device may already be using the bind.
func (b *bind) setDownFunc(f func() error) {
b.mu.Lock()
defer b.mu.Unlock()
b.downFunc = f
}
func (b *bind) Close() error {
b.mu.Lock()
defer b.mu.Unlock()
+9 -11
View File
@@ -27,6 +27,7 @@ import (
"github.com/xtls/xray-core/features/stats"
"github.com/xtls/xray-core/transport"
"github.com/xtls/xray-core/transport/internet"
"github.com/xtls/xray-core/transport/internet/finalmask"
"golang.zx2c4.com/wireguard/device"
)
@@ -199,7 +200,7 @@ func (h *Handler) Process(ctx context.Context, link *transport.Link, dialer inte
}
defer conn.Close()
c := &UDPConnClient{
PacketConn: conn.(*net.PacketConnWrapper).PacketConn,
PacketConn: conn.(*internet.PacketConnWrapper).PacketConn,
Dest: conn.RemoteAddr().(*net.UDPAddr),
}
reader = c
@@ -263,14 +264,14 @@ func (h *Handler) init(ctx context.Context) error {
if err != nil {
return nil, errors.New("failed to dial to dest").Base(err)
}
pktConn = conn.(*net.PacketConnWrapper).PacketConn
pktConn = conn.(*finalmask.PacketConnWrapper).PacketConn
} else {
conn, err := internet.DialSystem(ctx, dest, h.streamSettings.SocketSettings)
if err != nil {
return nil, errors.New("failed to dial to dest").Base(err)
}
switch c := conn.(type) {
case *net.PacketConnWrapper:
case *internet.PacketConnWrapper:
pktConn = c.PacketConn
case *cnc.Connection:
pktConn = &internet.FakePacketConn{Conn: c}
@@ -287,13 +288,7 @@ func (h *Handler) init(ctx context.Context) error {
}
return pktConn, nil
}
// device.NewDevice may use the bind right away (Up -> BindUpdate -> Open),
// so everything it reads must be set before creating the device.
bind := &bind{
resolveFunc: resolveFunc,
listenFunc: listenFunc,
reserved: h.conf.Reserved,
}
bind := &bind{}
logger := &device.Logger{
Verbosef: func(format string, args ...any) {
log.Record(&log.GeneralMessage{
@@ -309,7 +304,10 @@ func (h *Handler) init(ctx context.Context) error {
},
}
dev := device.NewDevice(h.tun, bind, logger)
bind.setDownFunc(dev.Down)
bind.resolveFunc = resolveFunc
bind.listenFunc = listenFunc
bind.downFunc = dev.Down
bind.reserved = h.conf.Reserved
var cfg strings.Builder
cfg.WriteString("private_key=" + h.conf.SecretKey + "\n")
for _, peer := range h.conf.Peers {
-91
View File
@@ -1,91 +0,0 @@
package wireguard
import (
"context"
"github.com/xtls/xray-core/common/errors"
tunicmp "github.com/xtls/xray-core/proxy/tun/icmp"
"gvisor.dev/gvisor/pkg/buffer"
"gvisor.dev/gvisor/pkg/tcpip"
"gvisor.dev/gvisor/pkg/tcpip/header"
"gvisor.dev/gvisor/pkg/tcpip/stack"
"gvisor.dev/gvisor/pkg/tcpip/transport/icmp"
)
// CreateICMPEchoResponder answers ICMP echo requests from peers locally, the way
// the TUN inbound does: ICMP is not proxied, but ping and connectivity checks
// through the tunnel get a reply instead of timing out.
//
// In promiscuous mode gVisor skips its own IPv4 echo reply for addresses that are
// not assigned to the NIC and leaves it to a custom handler; IPv6 is registered
// too so both families behave the same.
func CreateICMPEchoResponder(gstack *stack.Stack) {
gstack.SetTransportProtocolHandler(icmp.ProtocolNumber4, func(id stack.TransportEndpointID, pkt *stack.PacketBuffer) bool {
return handleICMPEcho(gstack, header.IPv4ProtocolNumber, id, pkt)
})
gstack.SetTransportProtocolHandler(icmp.ProtocolNumber6, func(id stack.TransportEndpointID, pkt *stack.PacketBuffer) bool {
return handleICMPEcho(gstack, header.IPv6ProtocolNumber, id, pkt)
})
}
func handleICMPEcho(gstack *stack.Stack, netProto tcpip.NetworkProtocolNumber, id stack.TransportEndpointID, pkt *stack.PacketBuffer) bool {
srcIP := id.RemoteAddress
dstIP := id.LocalAddress
if srcIP.Len() == 0 || dstIP.Len() == 0 {
return true
}
headerBytes := pkt.TransportHeader().Slice()
payloadBytes := pkt.Data().AsRange().ToSlice()
message := make([]byte, len(headerBytes)+len(payloadBytes))
copy(message, headerBytes)
copy(message[len(headerBytes):], payloadBytes)
if _, _, ok := tunicmp.ParseEchoRequest(netProto, message); !ok {
return true
}
reply, err := tunicmp.BuildLocalEchoReply(netProto, message, dstIP, srcIP)
if err != nil {
errors.LogInfoInner(context.Background(), err, "failed to build local icmp echo reply")
return true
}
if err := writeRawICMPPacket(gstack, netProto, reply, dstIP, srcIP); err != nil {
errors.LogInfoInner(context.Background(), err, "failed to write local icmp echo reply")
}
return true
}
func writeRawICMPPacket(gstack *stack.Stack, netProto tcpip.NetworkProtocolNumber, message []byte, srcIP, dstIP tcpip.Address) error {
pkt := stack.NewPacketBuffer(stack.PacketBufferOptions{
ReserveHeaderBytes: header.IPv6MinimumSize,
Payload: buffer.MakeWithData(message),
})
defer pkt.DecRef()
if netProto == header.IPv4ProtocolNumber {
ipHdr := header.IPv4(pkt.NetworkHeader().Push(header.IPv4MinimumSize))
ipHdr.Encode(&header.IPv4Fields{
TotalLength: uint16(header.IPv4MinimumSize + len(message)),
TTL: 64,
Protocol: uint8(header.ICMPv4ProtocolNumber),
SrcAddr: srcIP,
DstAddr: dstIP,
})
ipHdr.SetChecksum(^ipHdr.CalculateChecksum())
} else {
ipHdr := header.IPv6(pkt.NetworkHeader().Push(header.IPv6MinimumSize))
ipHdr.Encode(&header.IPv6Fields{
PayloadLength: uint16(len(message)),
TransportProtocol: header.ICMPv6ProtocolNumber,
HopLimit: 64,
SrcAddr: srcIP,
DstAddr: dstIP,
})
}
if err := gstack.WriteRawPacket(1, netProto, buffer.MakeWithView(pkt.ToView())); err != nil {
return errors.New("failed to write raw icmp packet back to stack ", err)
}
return nil
}
-176
View File
@@ -1,176 +0,0 @@
package wireguard
import (
"bytes"
"net/netip"
"testing"
"time"
"github.com/xtls/xray-core/common/net"
"gvisor.dev/gvisor/pkg/tcpip"
"gvisor.dev/gvisor/pkg/tcpip/checksum"
"gvisor.dev/gvisor/pkg/tcpip/header"
)
func newICMPTestStack(t *testing.T) *netTun {
t.Helper()
dev, _, gstack, err := CreateNetTUN([]netip.Addr{
netip.MustParseAddr("10.66.0.1"),
netip.MustParseAddr("fd00::1"),
}, nil, 1420, false)
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { dev.Close() })
CreateForwarder(gstack, func(conn net.Conn, dest net.Destination) { conn.Close() })
CreateICMPEchoResponder(gstack)
return dev.(*netTun)
}
// startReader must run before the request is written: the stack may answer
// synchronously inside Write, and netTun hands packets over an unbuffered channel.
func startReader(dev *netTun) <-chan []byte {
got := make(chan []byte, 1)
go func() {
buf := make([]byte, 2048)
sizes := make([]int, 1)
if _, err := dev.Read([][]byte{buf}, sizes, 0); err == nil {
got <- buf[:sizes[0]]
}
}()
return got
}
func awaitPacket(t *testing.T, got <-chan []byte) []byte {
t.Helper()
select {
case p := <-got:
return p
case <-time.After(2 * time.Second):
t.Fatal("no echo reply from the stack")
return nil
}
}
func TestICMPv4EchoReply(t *testing.T) {
dev := newICMPTestStack(t)
src := tcpip.AddrFrom4([4]byte{10, 66, 0, 2})
dst := tcpip.AddrFrom4([4]byte{1, 1, 1, 1})
payload := []byte("xray wireguard ping")
icmpMsg := make([]byte, header.ICMPv4MinimumSize+len(payload))
req := header.ICMPv4(icmpMsg)
req.SetType(header.ICMPv4Echo)
req.SetIdent(0x1234)
req.SetSequence(7)
copy(req.Payload(), payload)
req.SetChecksum(header.ICMPv4Checksum(req[:header.ICMPv4MinimumSize], checksum.Checksum(payload, 0)))
pkt := make([]byte, header.IPv4MinimumSize+len(icmpMsg))
ip := header.IPv4(pkt)
ip.Encode(&header.IPv4Fields{
TotalLength: uint16(len(pkt)),
TTL: 64,
Protocol: uint8(header.ICMPv4ProtocolNumber),
SrcAddr: src,
DstAddr: dst,
})
ip.SetChecksum(^ip.CalculateChecksum())
copy(pkt[header.IPv4MinimumSize:], icmpMsg)
got := startReader(dev)
if _, err := dev.Write([][]byte{pkt}, 0); err != nil {
t.Fatal(err)
}
reply := header.IPv4(awaitPacket(t, got))
if !reply.IsValid(len(reply)) {
t.Fatal("invalid ipv4 reply")
}
if reply.SourceAddress() != dst || reply.DestinationAddress() != src {
t.Fatalf("reply addresses %v -> %v, want %v -> %v", reply.SourceAddress(), reply.DestinationAddress(), dst, src)
}
if reply.TransportProtocol() != header.ICMPv4ProtocolNumber {
t.Fatalf("reply protocol %v, want icmpv4", reply.TransportProtocol())
}
echo := header.ICMPv4(reply.Payload())
if echo.Type() != header.ICMPv4EchoReply {
t.Fatalf("reply type %v, want echo reply", echo.Type())
}
if echo.Ident() != 0x1234 || echo.Sequence() != 7 {
t.Fatalf("reply ident/seq %#x/%d, want 0x1234/7", echo.Ident(), echo.Sequence())
}
if !bytes.Equal(echo.Payload(), payload) {
t.Fatalf("reply payload %q, want %q", echo.Payload(), payload)
}
if checksum.Checksum(echo, 0) != 0xffff {
t.Fatal("bad icmpv4 checksum")
}
}
func TestICMPv6EchoReply(t *testing.T) {
dev := newICMPTestStack(t)
src := tcpip.AddrFrom16([16]byte{0xfd, 15: 2})
dst := tcpip.AddrFrom16([16]byte{0x26, 0x06, 0x47, 0x00, 0x47, 0x00, 15: 0x11})
payload := []byte("xray wireguard ping6")
icmpMsg := make([]byte, header.ICMPv6MinimumSize+len(payload))
req := header.ICMPv6(icmpMsg)
req.SetType(header.ICMPv6EchoRequest)
req.SetIdent(0x4321)
req.SetSequence(9)
copy(req.Payload(), payload)
req.SetChecksum(header.ICMPv6Checksum(header.ICMPv6ChecksumParams{
Header: req[:header.ICMPv6MinimumSize],
Src: src,
Dst: dst,
PayloadCsum: checksum.Checksum(payload, 0),
PayloadLen: len(payload),
}))
pkt := make([]byte, header.IPv6MinimumSize+len(icmpMsg))
ip := header.IPv6(pkt)
ip.Encode(&header.IPv6Fields{
PayloadLength: uint16(len(icmpMsg)),
TransportProtocol: header.ICMPv6ProtocolNumber,
HopLimit: 64,
SrcAddr: src,
DstAddr: dst,
})
copy(pkt[header.IPv6MinimumSize:], icmpMsg)
got := startReader(dev)
if _, err := dev.Write([][]byte{pkt}, 0); err != nil {
t.Fatal(err)
}
reply := header.IPv6(awaitPacket(t, got))
if !reply.IsValid(len(reply)) {
t.Fatal("invalid ipv6 reply")
}
if reply.SourceAddress() != dst || reply.DestinationAddress() != src {
t.Fatalf("reply addresses %v -> %v, want %v -> %v", reply.SourceAddress(), reply.DestinationAddress(), dst, src)
}
echo := header.ICMPv6(reply.Payload())
if echo.Type() != header.ICMPv6EchoReply {
t.Fatalf("reply type %v, want echo reply", echo.Type())
}
if echo.Ident() != 0x4321 || echo.Sequence() != 9 {
t.Fatalf("reply ident/seq %#x/%d, want 0x4321/9", echo.Ident(), echo.Sequence())
}
if !bytes.Equal(echo.Payload(), payload) {
t.Fatalf("reply payload %q, want %q", echo.Payload(), payload)
}
zeroed := header.ICMPv6(append([]byte(nil), echo[:header.ICMPv6MinimumSize]...))
zeroed.SetChecksum(0)
want := header.ICMPv6Checksum(header.ICMPv6ChecksumParams{
Header: zeroed,
Src: dst,
Dst: src,
PayloadCsum: checksum.Checksum(echo.Payload(), 0),
PayloadLen: len(echo.Payload()),
})
if echo.Checksum() != want {
t.Fatalf("icmpv6 checksum %#x, want %#x", echo.Checksum(), want)
}
}
+2 -2
View File
@@ -21,7 +21,7 @@ import (
"syscall"
"time"
xnet "github.com/xtls/xray-core/common/net"
"github.com/xtls/xray-core/transport/internet"
"golang.zx2c4.com/wireguard/tun"
"golang.org/x/net/dns/dnsmessage"
@@ -220,7 +220,7 @@ func (tun *netTun) DialUDPAddrPort(laddr, raddr netip.AddrPort) (net.Conn, error
if err != nil {
return nil, err
}
return &xnet.PacketConnWrapper{
return &internet.PacketConnWrapper{
PacketConn: conn,
Dest: net.UDPAddrFromAddrPort(raddr),
}, nil
-80
View File
@@ -1,80 +0,0 @@
package wireguard
import (
"runtime"
"testing"
)
const benchBatch = 64
// Raw cost of queueing and draining a small burst, as one flow's reader does.
func BenchmarkQueueBurstChan(b *testing.B) {
ch := make(chan *packet, udpQueueLimit)
p := &packet{}
b.ReportAllocs()
for i := 0; i < b.N; i++ {
for j := 0; j < benchBatch; j++ {
ch <- p
}
for j := 0; j < benchBatch; j++ {
<-ch
}
}
}
func BenchmarkQueueBurstPacketQueue(b *testing.B) {
q := newPacketQueue(udpQueueLimit)
p := &packet{}
b.ReportAllocs()
for i := 0; i < b.N; i++ {
for j := 0; j < benchBatch; j++ {
q.push(p)
}
for j := 0; j < benchBatch; j++ {
q.pop()
}
}
}
// Producer and consumer on different goroutines; the producer yields when the
// queue is full instead of spinning, like a blocking channel send would.
func BenchmarkQueueStreamChan(b *testing.B) {
ch := make(chan *packet, udpQueueLimit)
p := &packet{}
done := make(chan struct{})
go func() {
for range ch {
}
close(done)
}()
b.ReportAllocs()
b.ResetTimer()
for i := 0; i < b.N; i++ {
ch <- p
}
close(ch)
<-done
}
func BenchmarkQueueStreamPacketQueue(b *testing.B) {
q := newPacketQueue(udpQueueLimit)
p := &packet{}
done := make(chan struct{})
go func() {
for {
if _, ok := q.pop(); !ok {
break
}
}
close(done)
}()
b.ReportAllocs()
b.ResetTimer()
for i := 0; i < b.N; i++ {
for !q.push(p) {
runtime.Gosched()
}
}
q.close()
<-done
}
+3 -6
View File
@@ -113,7 +113,7 @@ func NewServer(ctx context.Context, conf *DeviceConfig) (*Server, error) {
users.Store(user.Account.(*MemoryAccount).Pub, user)
}
s := &Server{
return &Server{
conf: conf,
ctx: core.ToBackgroundDetachedContext(ctx),
policyManager: p,
@@ -131,11 +131,7 @@ func NewServer(ctx context.Context, conf *DeviceConfig) (*Server, error) {
pub: pub,
users: users,
}
// Install the stack's protocol handlers before the device can deliver packets to it (Start -> dev.Up).
CreateForwarder(stack, s.HandleConnection)
CreateICMPEchoResponder(stack)
return s, nil
}, nil
}
func (s *Server) AddUser(ctx context.Context, user *protocol.MemoryUser) error {
@@ -324,6 +320,7 @@ func (s *Server) Start() error {
return err
}
s.dev = dev
CreateForwarder(s.stack, s.HandleConnection)
return nil
}
+18 -73
View File
@@ -85,7 +85,7 @@ func CreateForwarder(gstack *stack.Stack, handler func(conn net.Conn, dest net.D
}
gstack.SetTransportProtocolHandler(udp.ProtocolNumber, func(id stack.TransportEndpointID, pkt *stack.PacketBuffer) bool {
data := pkt.Data().AsRange().ToSlice()
data := pkt.Clone().Data().AsRange().ToSlice()
// if len(data) == 0 {
// return false
// }
@@ -112,7 +112,12 @@ func (m *udpManager) feed(src net.Destination, dst net.Destination, data []byte)
m.mutex.RLock()
uc, ok := m.m[src.NetAddr()]
if ok {
if !uc.queue.push(&packet{p: data, dest: &dst}) {
select {
case uc.queue <- &packet{
p: data,
dest: &dst,
}:
default:
errors.LogDebug(context.Background(), "drop udp with size ", len(data), " to ", dst.NetAddr(), " original ", uc.dst.NetAddr(), " > queue full")
}
m.mutex.RUnlock()
@@ -126,7 +131,7 @@ func (m *udpManager) feed(src net.Destination, dst net.Destination, data []byte)
uc, ok = m.m[src.NetAddr()]
if !ok {
uc = &udpConn{
queue: newPacketQueue(udpQueueLimit),
queue: make(chan *packet, 1024),
src: src,
dst: dst,
}
@@ -140,7 +145,12 @@ func (m *udpManager) feed(src net.Destination, dst net.Destination, data []byte)
go m.handler(uc, dst)
}
if !uc.queue.push(&packet{p: data, dest: &dst}) {
select {
case uc.queue <- &packet{
p: data,
dest: &dst,
}:
default:
errors.LogDebug(context.Background(), "drop udp with size ", len(data), " to ", dst.NetAddr(), " original ", uc.dst.NetAddr(), " > queue full 2")
}
}
@@ -148,7 +158,7 @@ func (m *udpManager) feed(src net.Destination, dst net.Destination, data []byte)
func (m *udpManager) close(uc *udpConn) {
if !uc.closed {
uc.closed = true
uc.queue.close()
close(uc.queue)
delete(m.m, uc.src.NetAddr())
}
}
@@ -222,7 +232,7 @@ type packet struct {
}
type udpConn struct {
queue *packetQueue
queue chan *packet
src net.Destination
dst net.Destination
writeFunc func(payload []byte, src net.Destination, dst net.Destination) error
@@ -232,7 +242,7 @@ type udpConn struct {
func (c *udpConn) ReadMultiBuffer() (buf.MultiBuffer, error) {
for {
q, ok := c.queue.pop()
q, ok := <-c.queue
if !ok {
return nil, io.EOF
}
@@ -251,7 +261,7 @@ func (c *udpConn) ReadMultiBuffer() (buf.MultiBuffer, error) {
}
func (c *udpConn) Read(p []byte) (int, error) {
q, ok := c.queue.pop()
q, ok := <-c.queue
if !ok {
return 0, io.EOF
}
@@ -314,68 +324,3 @@ func (c *udpConn) SetReadDeadline(t time.Time) error {
func (c *udpConn) SetWriteDeadline(t time.Time) error {
return nil
}
// udpQueueLimit bounds the packets waiting for one UDP flow; more are dropped.
const udpQueueLimit = 1024
// packetQueue holds the packets waiting for one UDP flow. Unlike a buffered
// channel of the same bound it only allocates for packets actually queued, so
// the many idle flows kept until the idle timeout cost next to nothing.
type packetQueue struct {
mu sync.Mutex
items []*packet
limit int
notify chan struct{}
closed bool
}
func newPacketQueue(limit int) *packetQueue {
return &packetQueue{limit: limit, notify: make(chan struct{}, 1)}
}
// push queues p and reports whether it was accepted.
func (q *packetQueue) push(p *packet) bool {
q.mu.Lock()
defer q.mu.Unlock()
if q.closed || len(q.items) >= q.limit {
return false
}
q.items = append(q.items, p)
select {
case q.notify <- struct{}{}:
default:
}
return true
}
// pop blocks until a packet is queued or the queue is closed and drained.
func (q *packetQueue) pop() (*packet, bool) {
for {
q.mu.Lock()
if len(q.items) > 0 {
p := q.items[0]
q.items[0] = nil
q.items = q.items[1:]
if len(q.items) == 0 {
q.items = nil
}
q.mu.Unlock()
return p, true
}
if q.closed {
q.mu.Unlock()
return nil, false
}
q.mu.Unlock()
<-q.notify
}
}
func (q *packetQueue) close() {
q.mu.Lock()
defer q.mu.Unlock()
if !q.closed {
q.closed = true
close(q.notify)
}
}
+23 -18
View File
@@ -9,21 +9,32 @@ import (
"net"
"net/netip"
"os"
"sync/atomic"
"sync"
"syscall"
"golang.org/x/sys/unix"
"github.com/vishvananda/netlink"
"github.com/xtls/xray-core/common/errors"
xnet "github.com/xtls/xray-core/common/net"
"github.com/xtls/xray-core/transport/internet"
"golang.zx2c4.com/wireguard/tun"
)
var tableIndex atomic.Uint32
var (
tableIndex int = 10230
mu sync.Mutex
)
func init() {
tableIndex.Store(10230)
func allocateIPv6TableIndex() int {
mu.Lock()
defer mu.Unlock()
if tableIndex > 10230 {
errors.LogInfo(context.Background(), "allocate new ipv6 table index: ", tableIndex)
}
currentIndex := tableIndex
tableIndex++
return currentIndex
}
type kernelTun struct {
@@ -100,23 +111,17 @@ func createKernelTun(localAddresses, dnsServers []netip.Addr, mtu int) (tdev tun
}
}
var ipv6TableIndex int
ipv6TableIndex := allocateIPv6TableIndex()
if v6 != nil {
r := &netlink.Route{}
r := &netlink.Route{Table: ipv6TableIndex}
for {
ipv6TableIndex = int(tableIndex.Add(1)) - 1
r.Table = ipv6TableIndex
routeList, fErr := netlink.RouteListFiltered(netlink.FAMILY_V6, r, netlink.RT_FILTER_TABLE)
if fErr != nil {
return nil, nil, errors.New("failed to pre check routes for table: ", ipv6TableIndex).Base(fErr)
}
if len(routeList) == 0 {
errors.LogInfo(context.Background(), "allocate new ipv6 table index: ", ipv6TableIndex)
if len(routeList) == 0 || fErr != nil {
break
}
// to prevent infinite loop
if ipv6TableIndex > 65535 {
return nil, nil, errors.New("failed to find available ipv6 table index")
ipv6TableIndex--
if ipv6TableIndex < 0 {
return nil, nil, fmt.Errorf("failed to find available ipv6 table index")
}
}
}
@@ -258,7 +263,7 @@ func (tun *kernelTun) DialUDPAddrPort(laddr, raddr netip.AddrPort) (net.Conn, er
if err != nil {
return nil, err
}
return &xnet.PacketConnWrapper{
return &internet.PacketConnWrapper{
PacketConn: conn,
Dest: net.UDPAddrFromAddrPort(raddr),
}, nil
-91
View File
@@ -1,91 +0,0 @@
package wireguard
import (
"testing"
"time"
"github.com/xtls/xray-core/common/net"
)
// BenchmarkUDPManagerNewSession measures what one new UDP flow costs the
// inbound while it stays open: QUIC and DNS open many short flows, and each
// one lives until the connection idle timeout.
func BenchmarkUDPManagerNewSession(b *testing.B) {
m := &udpManager{
handler: func(conn net.Conn, dest net.Destination) {},
m: make(map[string]*udpConn),
}
dst := net.UDPDestination(net.ParseAddress("1.1.1.1"), 443)
payload := make([]byte, 1200)
b.ReportAllocs()
b.ResetTimer()
for i := 0; i < b.N; i++ {
src := net.UDPDestination(net.IPAddress([]byte{10, byte(i >> 16), byte(i >> 8), byte(i)}), net.Port(1024+i%60000))
m.feed(src, dst, payload)
}
}
func TestPacketQueueOrderAndClose(t *testing.T) {
q := newPacketQueue(udpQueueLimit)
for i := 0; i < 3; i++ {
if !q.push(&packet{p: []byte{byte(i)}}) {
t.Fatalf("push %d rejected", i)
}
}
for i := 0; i < 3; i++ {
p, ok := q.pop()
if !ok || p.p[0] != byte(i) {
t.Fatalf("pop %d: got %v, %v", i, p, ok)
}
}
q.close()
if _, ok := q.pop(); ok {
t.Fatal("pop after close returned a packet")
}
if q.push(&packet{}) {
t.Fatal("push after close accepted")
}
}
func TestPacketQueueLimit(t *testing.T) {
q := newPacketQueue(udpQueueLimit)
for i := 0; i < udpQueueLimit; i++ {
if !q.push(&packet{}) {
t.Fatalf("push %d rejected below the limit", i)
}
}
if q.push(&packet{}) {
t.Fatal("push above the limit accepted")
}
}
func TestPacketQueueCloseUnblocksReader(t *testing.T) {
q := newPacketQueue(udpQueueLimit)
done := make(chan bool)
go func() {
_, ok := q.pop()
done <- ok
}()
q.close()
select {
case ok := <-done:
if ok {
t.Fatal("blocked pop returned a packet after close")
}
case <-time.After(time.Second):
t.Fatal("close did not wake the reader")
}
}
func TestPacketQueueDropsDrainedStorage(t *testing.T) {
q := newPacketQueue(udpQueueLimit)
for i := 0; i < 100; i++ {
q.push(&packet{})
}
for i := 0; i < 100; i++ {
q.pop()
}
if q.items != nil {
t.Fatalf("drained queue still holds %d slots", cap(q.items))
}
}
-6
View File
@@ -115,9 +115,6 @@ func TestDokodemoTCP(t *testing.T) {
defer CloseServer(server)
break
}
if server != nil {
CloseServer(server)
}
retry++
if retry > 5 {
t.Fatal("All attempts failed to start client")
@@ -212,9 +209,6 @@ func TestDokodemoUDP(t *testing.T) {
defer CloseServer(server)
break
}
if server != nil {
CloseServer(server)
}
retry++
if retry > 5 {
t.Fatal("All attempts failed to start client")

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