Compare commits

...
27 Commits
Author SHA1 Message Date
Meo597 5d1d8200d9 dns: unify the Lua API for local and configured DNS
Expose localdns.Client in xray.dns.Servers with the ID "localhost",
so Lua scripts can use the same server API for both clients.
2026-10-01 21:58:09 +08:00
Meo597 2440f53cdd feat(router): add Lua scripting for routing script 2026-10-01 18:25:53 +08:00
Meo597 1c52c65872 geodata: rename Lua matcher constructors to BuildDomainMatcher and BuildIPMatcher 2026-09-28 22:52:45 +08:00
Meo597 5724db08f4 lua: standardize hooks and host APIs on PascalCase 2026-09-28 22:52:14 +08:00
Meo597 3d3306503d geodata: flatten Lua matcher rule arguments 2026-09-28 22:23:04 +08:00
Meo597 5e1bb92b98 dns: reduce Lua allocations with flat arguments and returns and native IP slice userdata 2026-09-28 22:12:47 +08:00
Meo597 459301d42e refine Lua script path fallback 2026-09-27 23:15:51 +08:00
Meo597 3982028a9c lua: lower camel case 2026-09-26 17:58:25 +08:00
Meo597 70b8e9a61d log caller 2026-09-26 15:42:57 +08:00
Meo597 219f758060 add log module 2026-09-26 14:40:15 +08:00
Meo597 9628003594 Reduce memory usage 2026-09-26 06:20:34 +08:00
Meo597 72d9ab50b9 add tests 2026-09-26 06:16:54 +08:00
Meo597 235843c5d2 feat(dns): add Lua scripting for DNS queries 2026-09-26 05:37:23 +08:00
Hossin Asaadi 60e2a0c502 WireGuard proxy: Release packet views after use (#6801)
https://github.com/XTLS/Xray-core/pull/6801#issuecomment-5807228464
2026-09-24 05:46:26 +00:00
Esko Mobius 7d3e44fee2 Proxy: Add MASQUE outbound & transport (IETF CONNECT-IP, RFC 9484) (#6807)
Closes https://github.com/XTLS/Xray-core/issues/5495#issuecomment-3710683679
2026-09-24 02:13:37 +00:00
dependabot[bot] 9927942aaa Bump google.golang.org/grpc from 1.83.2 to 1.84.0 (#6793)
Bumps [google.golang.org/grpc](https://github.com/grpc/grpc-go) from 1.83.2 to 1.84.0.
- [Release notes](https://github.com/grpc/grpc-go/releases)
- [Commits](https://github.com/grpc/grpc-go/compare/v1.83.2...v1.84.0)

---
updated-dependencies:
- dependency-name: google.golang.org/grpc
  dependency-version: 1.84.0
  dependency-type: direct:production
  update-type: version-update:semver-minor
...

Signed-off-by: dependabot[bot] <support@github.com>
Co-authored-by: dependabot[bot] <49699333+dependabot[bot]@users.noreply.github.com>
2026-09-24 02:04:39 +00:00
LjhAUMEM a308ded2e6 WireGuard outbound: Fix endpoint IP address (#6804)
Fixes https://github.com/XTLS/Xray-core/issues/6803
2026-09-24 01:59:28 +00:00
dependabot[bot] 7741e9e77e Bump golang.zx2c4.com/wireguard/windows from 1.0.1 to 1.1.1 (#6809)
Bumps golang.zx2c4.com/wireguard/windows from 1.0.1 to 1.1.1.

---
updated-dependencies:
- dependency-name: golang.zx2c4.com/wireguard/windows
  dependency-version: 1.1.1
  dependency-type: direct:production
  update-type: version-update:semver-minor
...

Signed-off-by: dependabot[bot] <support@github.com>
Co-authored-by: dependabot[bot] <49699333+dependabot[bot]@users.noreply.github.com>
2026-09-24 01:57:16 +00:00
Жора ЗмейкинandLjhAUMEM d562d8947d Hysteria outbound: Fix UDP DATAGRAM truncation with ChromeParrot (#6788)
https://github.com/XTLS/Xray-core/pull/6788#issuecomment-5751428127

---------

Co-authored-by: LjhAUMEM <llnu14702@gmail.com>
2026-09-20 22:51:03 +00:00
Hossin Asaadi dbb1ea30ba WireGuard proxy: Prevent panic when closing netTun with a write in flight (#6778)
https://github.com/XTLS/Xray-core/pull/6778#pullrequestreview-5251957308
2026-09-20 20:07:54 +00:00
LjhAUMEM efc9e6da62 WireGuard outbound: Remove remoteDNS' "local" mode and domainStrategy (#6771)
https://github.com/XTLS/Xray-core/issues/6567#issuecomment-5660953757
https://github.com/XTLS/Xray-core/pull/6771#issuecomment-5751826982
https://github.com/XTLS/Xray-core/pull/6771#issuecomment-5751831637
2026-09-20 19:07:51 +00:00
LjhAUMEM 8267cf953a Transport: Refactor to be based on Finalmask's dialer & listener (#6754)
https://github.com/XTLS/Xray-core/pull/6327#issuecomment-5645958010
https://github.com/XTLS/Xray-core/pull/6754#issuecomment-5720818254
https://github.com/XTLS/Xray-core/pull/6754#issuecomment-5751585906
2026-09-20 18:12:50 +00:00
dependabot[bot] 24e6f6d551 Bump golang.org/x/net from 0.58.0 to 0.59.0 (#6756)
Bumps [golang.org/x/net](https://github.com/golang/net) from 0.58.0 to 0.59.0.
- [Commits](https://github.com/golang/net/compare/v0.58.0...v0.59.0)

---
updated-dependencies:
- dependency-name: golang.org/x/net
  dependency-version: 0.59.0
  dependency-type: direct:production
  update-type: version-update:semver-minor
...

Signed-off-by: dependabot[bot] <support@github.com>
Co-authored-by: dependabot[bot] <49699333+dependabot[bot]@users.noreply.github.com>
2026-09-20 17:46:51 +00:00
RisaroandRPRX dcdfc57ccd XDRIVE transport: Add the Google Drive and "template" backend (#6748)
https://github.com/XTLS/Xray-core/pull/5645#issuecomment-3849778103
https://github.com/XTLS/Xray-core/pull/5645#issuecomment-3851839033
https://github.com/XTLS/Xray-core/pull/6745#issuecomment-5627294204
https://github.com/XTLS/Xray-core/pull/6748#issuecomment-5660444946
https://github.com/XTLS/Xray-core/pull/6748#issuecomment-5719209642

---------

Co-authored-by: RPRX <63339210+RPRX@users.noreply.github.com>
2026-09-19 09:03:11 +00:00
RPRXandRisaro 3461c511aa XDRIVE transport: The universal remote-storage/any-service-based proxy, ignoring IP whitelists (#5645)
https://github.com/XTLS/Xray-core/pull/5414#issuecomment-3796734827
https://github.com/XTLS/Xray-core/pull/5581#issuecomment-3797134147
https://github.com/XTLS/Xray-core/pull/5645#issuecomment-3899873945
https://github.com/XTLS/Xray-core/pull/6745#issuecomment-5627420177
https://github.com/XTLS/Xray-core/pull/6748#issuecomment-5740443122

---------

Co-authored-by: Risaro <62798663+Risaro@users.noreply.github.com>
2026-09-19 08:23:38 +00:00
SVLAVRand风扇滑翔翼 c412e77a9b TUN inbound: Preserve UDP packet destinations with traffic stats (#6747)
https://github.com/XTLS/Xray-core/pull/6747#issuecomment-5647964551

---------

Co-authored-by: 风扇滑翔翼 <Fangliding.fshxy@outlook.com>
2026-09-12 20:09:35 +00:00
Kiangand风扇滑翔翼 ccb69ea5e2 common/buf/readv_windows.go: Fix raw WSARecv blocking thread and preventing connection close (#6743)
https://github.com/XTLS/Xray-core/pull/6743#issuecomment-5647971870

---------

Co-authored-by: 风扇滑翔翼 <Fangliding.fshxy@outlook.com>
2026-09-12 19:51:59 +00:00
142 changed files with 14307 additions and 1279 deletions
+3
View File
@@ -470,6 +470,9 @@ func (d *DefaultDispatcher) routedDispatch(ctx context.Context, link *transport.
return // DO NOT CHANGE: the traffic shouldn't be processed by default outbound if the specified outbound tag doesn't exist (yet), e.g., VLESS Reverse Proxy
}
} else {
if err != common.ErrNoClue {
errors.LogErrorInner(ctx, err, "failed to pick route for ", destination)
}
errors.LogInfo(ctx, "default route for ", destination)
}
}
+25 -6
View File
@@ -93,6 +93,7 @@ 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
}
@@ -239,6 +240,13 @@ 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.
@@ -258,8 +266,10 @@ 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"`
unknownFields protoimpl.UnknownFields
sizeCache protoimpl.SizeCache
// 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
}
func (x *Config) Reset() {
@@ -369,6 +379,13 @@ 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"`
@@ -435,7 +452,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\"\xde\x05\n" +
"\x14app/dns/config.proto\x12\fxray.app.dns\x1a\x1ccommon/net/destination.proto\x1a\x1bcommon/geodata/geodat.proto\"\xee\x05\n" +
"\n" +
"NameServer\x123\n" +
"\aaddress\x18\x01 \x01(\v2\x19.xray.common.net.EndpointR\aaddress\x12\x1b\n" +
@@ -461,10 +478,11 @@ const file_app_dns_config_proto_rawDesc = "" +
"\n" +
"actUnprior\x18\x0e \x01(\bR\n" +
"actUnprior\x12\x1a\n" +
"\bpolicyID\x18\x11 \x01(\rR\bpolicyIDB\x0f\n" +
"\bpolicyID\x18\x11 \x01(\rR\bpolicyID\x12\x0e\n" +
"\x02id\x18\x12 \x01(\tR\x02idB\x0f\n" +
"\r_disableCacheB\r\n" +
"\v_serveStaleB\x12\n" +
"\x10_serveExpiredTTLJ\x04\b\x04\x10\x05\"\x82\x05\n" +
"\x10_serveExpiredTTLJ\x04\b\x04\x10\x05\"\x9a\x05\n" +
"\x06Config\x129\n" +
"\vname_server\x18\x05 \x03(\v2\x18.xray.app.dns.NameServerR\n" +
"nameServer\x12\x1b\n" +
@@ -480,7 +498,8 @@ 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\x1a}\n" +
"\x13enableParallelQuery\x18\x0e \x01(\bR\x13enableParallelQuery\x12\x16\n" +
"\x06script\x18\x0f \x01(\tR\x06script\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,6 +27,7 @@ message NameServer {
repeated xray.common.geodata.IPRule unexpected_ip = 13;
bool actUnprior = 14;
uint32 policyID = 17;
string id = 18;
}
enum QueryStrategy {
@@ -73,4 +74,7 @@ message Config {
bool disableFallbackIfMatch = 11;
bool enableParallelQuery = 14;
// Absolute path to the Lua DNS query script.
string script = 15;
}
+16
View File
@@ -31,6 +31,8 @@ type DNS struct {
domainMatcher geodata.DomainMatcher
matcherInfos []*DomainMatcherInfo
checkSystem bool
script *scriptEngine
scriptPath string
}
// DomainMatcherInfo contains information attached to index returned by Server.domainMatcher.
@@ -180,6 +182,7 @@ func New(ctx context.Context, config *Config) (*DNS, error) {
disableFallbackIfMatch: config.DisableFallbackIfMatch,
enableParallelQuery: config.EnableParallelQuery,
checkSystem: checkSystem,
scriptPath: config.Script,
}, nil
}
@@ -190,11 +193,21 @@ 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
}
@@ -257,6 +270,9 @@ 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 {
+201
View File
@@ -0,0 +1,201 @@
package dns
import (
"context"
"math"
"strings"
"github.com/xtls/xray-core/common/errors"
"github.com/xtls/xray-core/common/net"
featureDNS "github.com/xtls/xray-core/features/dns"
"github.com/xtls/xray-core/features/dns/localdns"
lua "github.com/yuin/gopher-lua"
)
// luaDNSServer adapts configured and local DNS to the same Lua API.
type luaDNSServer struct {
id string
name string
query func(context.Context, string, featureDNS.IPOption) ([]net.IP, uint32, error)
}
// RegisterLua makes xray.dns available to scripts backed by client.
func RegisterLua(L *lua.LState, client featureDNS.Client) {
var servers []luaDNSServer
switch client := client.(type) {
case *DNS:
servers = luaServers(client)
case *localdns.Client:
servers = []luaDNSServer{{
id: "localhost",
name: "localhost",
query: func(_ context.Context, domain string, option featureDNS.IPOption) ([]net.IP, uint32, error) {
return client.LookupIP(domain, option)
},
}}
}
registerLua(L, servers, client)
}
// RegisterLua makes xray.dns available to DNS scripts.
func (s *DNS) RegisterLua(L *lua.LState) {
registerLua(L, luaServers(s), nil)
}
func luaServers(s *DNS) []luaDNSServer {
servers := make([]luaDNSServer, len(s.clients))
for i, client := range s.clients {
servers[i] = luaDNSServer{id: client.id, name: client.Name(), query: client.QueryIP}
}
return servers
}
func registerLua(L *lua.LState, servers []luaDNSServer, client featureDNS.Client) {
L.PreloadModule("xray.dns", func(L *lua.LState) int {
serverList := L.NewTable()
for i, client := range servers {
server := L.NewTable()
server.RawSetString("ID", lua.LString(client.id))
server.RawSetString("Query", L.NewFunction(func(L *lua.LState) int {
domain, ok := L.Get(2).(lua.LString)
if !ok {
L.RaiseError("server:Query requires a domain")
return 0
}
option := featureDNS.IPOption{
IPv4Enable: L.CheckBool(3),
IPv6Enable: L.CheckBool(4),
FakeEnable: L.CheckBool(5),
}
ctx := L.Context()
if ctx == nil {
L.RaiseError("server:Query requires an active DNS query")
return 0
}
var ips []net.IP
var ttl uint32
var err error
if !option.FakeEnable && strings.EqualFold(client.name, "FakeDNS") {
err = featureDNS.ErrEmptyResponse
} else {
ips, ttl, err = client.query(ctx, string(domain), option)
}
addresses := L.NewUserData()
addresses.Value = ips
L.Push(addresses)
L.Push(lua.LNumber(ttl))
if err != nil {
ud := L.NewUserData()
ud.Value = err
L.Push(ud)
} else {
L.Push(lua.LNil)
}
return 3
}))
serverList.RawSetInt(i+1, server)
}
module := L.NewTable()
if servers != nil {
module.RawSetString("Servers", serverList)
}
if client != nil {
module.RawSetString("Query", newLuaClientQuery(L, client))
}
L.Push(module)
return 1
})
}
func newLuaClientQuery(L *lua.LState, client featureDNS.Client) *lua.LFunction {
return L.NewFunction(func(L *lua.LState) int {
domain, ok := L.Get(1).(lua.LString)
if !ok {
L.RaiseError("dns.Query requires a domain")
return 0
}
option := featureDNS.IPOption{
IPv4Enable: L.CheckBool(2),
IPv6Enable: L.CheckBool(3),
FakeEnable: L.CheckBool(4),
}
if L.Context() == nil {
L.RaiseError("dns.Query requires an active DNS query")
return 0
}
ips, ttl, err := client.LookupIP(string(domain), option)
addresses := L.NewUserData()
addresses.Value = ips
L.Push(addresses)
L.Push(lua.LNumber(ttl))
if err != nil {
ud := L.NewUserData()
ud.Value = err
L.Push(ud)
} else {
L.Push(lua.LNil)
}
return 3
})
}
// CallLuaHook invokes HandleDNSQuery in the supplied state.
// Returned slices and IP bytes may share storage with DNS caches or matcher inputs.
func (s *DNS) CallLuaHook(L *lua.LState, ctx context.Context, domain string, option featureDNS.IPOption) ([]net.IP, uint32, error) {
previous, top := L.Context(), L.GetTop()
L.SetContext(ctx)
defer func() {
L.SetTop(top)
if previous == nil {
L.RemoveContext()
} else {
L.SetContext(previous)
}
}()
fn := L.GetGlobal("HandleDNSQuery")
if fn.Type() != lua.LTFunction {
return nil, 0, errors.New("DNS script must define HandleDNSQuery(domain, ipv4, ipv6, fake)")
}
if err := L.CallByParam(lua.P{Fn: fn, NRet: 3, Protect: true},
lua.LString(strings.ToLower(domain)), lua.LBool(option.IPv4Enable),
lua.LBool(option.IPv6Enable), lua.LBool(option.FakeEnable)); err != nil {
return nil, 0, err
}
return readLuaDNSResult(L.Get(-3), L.Get(-2), L.Get(-1))
}
func readLuaDNSResult(addresses, ttlValue, errorValue lua.LValue) ([]net.IP, uint32, error) {
if errorValue != lua.LNil {
if ud, ok := errorValue.(*lua.LUserData); ok {
if err, ok := ud.Value.(error); ok {
return nil, 0, err
}
}
if s, ok := errorValue.(lua.LString); ok {
return nil, 0, errors.New(string(s))
}
return nil, 0, errors.New("DNS script error must be an error or string")
}
ttl, ok := ttlValue.(lua.LNumber)
if !ok || ttl < 0 || ttl > math.MaxUint32 || math.Trunc(float64(ttl)) != float64(ttl) {
return nil, 0, errors.New("DNS script returned invalid TTL")
}
if addresses == lua.LNil {
return nil, 0, featureDNS.ErrEmptyResponse
}
ud, ok := addresses.(*lua.LUserData)
if !ok {
return nil, 0, errors.New("DNS script IPs must be native IP slice userdata")
}
ips, ok := ud.Value.([]net.IP)
if !ok {
return nil, 0, errors.New("DNS script IPs must be native IP slice userdata")
}
if len(ips) == 0 {
return nil, 0, featureDNS.ErrEmptyResponse
}
return ips, uint32(ttl), nil
}
+296
View File
@@ -0,0 +1,296 @@
package dns
import (
"context"
go_errors "errors"
"math"
"strings"
"testing"
"time"
"github.com/xtls/xray-core/common/geodata"
"github.com/xtls/xray-core/common/net"
featureDNS "github.com/xtls/xray-core/features/dns"
"github.com/xtls/xray-core/features/dns/localdns"
lua "github.com/yuin/gopher-lua"
)
func TestReadLuaDNSResult(t *testing.T) {
L := lua.NewState()
defer L.Close()
want := []net.IP{net.ParseIP("8.8.8.8"), {127, 0, 0, 1}, net.ParseIP("::1")}
addresses := L.NewUserData()
addresses.Value = want
ips, ttl, err := readLuaDNSResult(addresses, lua.LNumber(45), lua.LNil)
if err != nil || ttl != 45 || len(ips) != len(want) {
t.Fatalf("readLuaDNSResult() = %v, TTL %d, %v", ips, ttl, err)
}
for i := range want {
if !ips[i].Equal(want[i]) {
t.Fatalf("IP %d = %v, want %v", i, ips[i], want[i])
}
}
}
func TestReadLuaDNSResultValidation(t *testing.T) {
L := lua.NewState()
defer L.Close()
for _, tc := range []struct {
name string
change func(*[3]lua.LValue)
want string
}{
{"fractional TTL", func(v *[3]lua.LValue) { v[1] = lua.LNumber(1.5) }, "invalid TTL"},
{"oversized TTL", func(v *[3]lua.LValue) { v[1] = lua.LNumber(4294967296) }, "invalid TTL"},
{"negative TTL", func(v *[3]lua.LValue) { v[1] = lua.LNumber(-1) }, "invalid TTL"},
{"NaN TTL", func(v *[3]lua.LValue) { v[1] = lua.LNumber(math.NaN()) }, "invalid TTL"},
{"missing TTL", func(v *[3]lua.LValue) { v[1] = lua.LNil }, "invalid TTL"},
{"string IPs", func(v *[3]lua.LValue) { v[0] = lua.LString("127.0.0.1") }, "native IP slice"},
{"wrong userdata", func(v *[3]lua.LValue) { v[0].(*lua.LUserData).Value = net.ParseIP("127.0.0.1") }, "native IP slice"},
{"script error", func(v *[3]lua.LValue) { v[2] = lua.LString("blocked by script") }, "blocked by script"},
{"invalid error", func(v *[3]lua.LValue) { v[2] = lua.LTrue }, "error or string"},
} {
t.Run(tc.name, func(t *testing.T) {
addresses := L.NewUserData()
addresses.Value = []net.IP{net.ParseIP("127.0.0.1")}
values := [3]lua.LValue{addresses, lua.LNumber(60), lua.LNil}
tc.change(&values)
_, _, err := readLuaDNSResult(values[0], values[1], values[2])
if err == nil || !strings.Contains(err.Error(), tc.want) {
t.Fatalf("readLuaDNSResult error = %v, want %q", err, tc.want)
}
})
}
addresses := L.NewUserData()
addresses.Value = []net.IP(nil)
for _, empty := range []lua.LValue{addresses, lua.LNil} {
if _, _, err := readLuaDNSResult(empty, lua.LNumber(0), lua.LNil); !go_errors.Is(err, featureDNS.ErrEmptyResponse) {
t.Fatalf("empty result error = %v, want ErrEmptyResponse", err)
}
}
wantErr := go_errors.New("upstream failed")
errorValue := L.NewUserData()
errorValue.Value = wantErr
if _, _, err := readLuaDNSResult(lua.LNil, lua.LNil, errorValue); err != wantErr {
t.Fatalf("upstream error = %v, want original error %v", err, wantErr)
}
}
func TestCallLuaHookCancellation(t *testing.T) {
L := lua.NewState()
defer L.Close()
if err := L.DoString(`function HandleDNSQuery(domain, ipv4, ipv6, fake) while true do end end`); err != nil {
t.Fatal(err)
}
ctx, cancel := context.WithTimeout(context.Background(), 50*time.Millisecond)
defer cancel()
_, _, err := (&DNS{}).CallLuaHook(L, ctx, "example.com", featureDNS.IPOption{IPv4Enable: true})
if err == nil {
t.Fatal("CallLuaHook did not stop after context cancellation")
}
if L.Context() != nil {
t.Fatal("CallLuaHook left the canceled context on the Lua state")
}
}
func TestCallLuaHookNormalizesDomain(t *testing.T) {
L := lua.NewState()
defer L.Close()
addresses := L.NewUserData()
addresses.Value = []net.IP{net.ParseIP("127.0.0.1")}
L.SetGlobal("ips", addresses)
if err := L.DoString(`
function HandleDNSQuery(domain, ipv4, ipv6, fake)
assert(domain == "example.com")
assert(ipv4 and not ipv6 and not fake)
return ips, 60, nil
end
`); err != nil {
t.Fatal(err)
}
s := &DNS{}
if _, _, err := s.CallLuaHook(L, context.Background(), "ExAmPlE.CoM", featureDNS.IPOption{IPv4Enable: true}); err != nil {
t.Fatal(err)
}
}
func TestCallLuaHookRestoresState(t *testing.T) {
for _, tc := range []struct {
name string
body string
wantErr bool
}{
{"success", `return ips, 60`, false},
{"error", `error("failed")`, true},
} {
t.Run(tc.name, func(t *testing.T) {
L := lua.NewState()
defer L.Close()
addresses := L.NewUserData()
addresses.Value = []net.IP{net.ParseIP("127.0.0.1")}
L.SetGlobal("ips", addresses)
if err := L.DoString("function HandleDNSQuery() " + tc.body + " end"); err != nil {
t.Fatal(err)
}
previous, cancel := context.WithCancel(context.Background())
defer cancel()
L.SetContext(previous)
L.Push(lua.LTrue)
_, _, err := (&DNS{}).CallLuaHook(L, context.Background(), "example.com", featureDNS.IPOption{IPv4Enable: true})
if (err != nil) != tc.wantErr {
t.Fatalf("hook error = %v, want error %t", err, tc.wantErr)
}
if L.Context() != previous || L.GetTop() != 1 || L.Get(1) != lua.LTrue {
t.Fatal("hook did not restore the previous context and stack")
}
})
}
}
func TestLuaDNSServerQuery(t *testing.T) {
L := lua.NewState()
defer L.Close()
geodata.RegisterLua(L)
option := featureDNS.IPOption{IPv4Enable: true}
ips := []net.IP{net.ParseIP("127.0.0.1"), net.ParseIP("8.8.8.8")}
server := &DNS{clients: []*Client{{server: &benchmarkLuaNameServer{ips: ips}, ipOption: &option, timeoutMs: time.Second}}}
server.RegisterLua(L)
if err := L.DoString(`
local server = require("xray.dns").Servers[1]
local matcher = require("xray.geodata").BuildIPMatcher("127.0.0.0/8")
function HandleDNSQuery(domain, ipv4, ipv6, fake)
local ips, ttl, err = server:Query(domain, ipv4, ipv6, fake)
assert(type(ips) == "userdata" and not err)
assert(matcher:AnyMatch(ips))
local matched = matcher:FilterIPs(ips)
return matched, ttl, err
end
`); err != nil {
t.Fatal(err)
}
got, ttl, err := server.CallLuaHook(L, context.Background(), "example.com", option)
if err != nil || ttl != 60 || len(got) != 1 || !got[0].Equal(ips[0]) {
t.Fatalf("server query = %v, TTL %d, %v", got, ttl, err)
}
}
type luaDNSClient struct {
featureDNS.Client
lookup func(string, featureDNS.IPOption) ([]net.IP, uint32, error)
}
func (c *luaDNSClient) LookupIP(domain string, option featureDNS.IPOption) ([]net.IP, uint32, error) {
return c.lookup(domain, option)
}
func TestLuaDNSClientQuery(t *testing.T) {
L := lua.NewState()
defer L.Close()
L.SetContext(context.Background())
defer L.RemoveContext()
geodata.RegisterLua(L)
want := []net.IP{{127, 0, 0, 1}}
client := &luaDNSClient{lookup: func(domain string, option featureDNS.IPOption) ([]net.IP, uint32, error) {
if domain != "MiXeD.Example." || !option.IPv4Enable || option.IPv6Enable || !option.FakeEnable {
t.Fatalf("dns.Query arguments = %q, %+v", domain, option)
}
return want, 42, nil
}}
RegisterLua(L, client)
if err := L.DoString(`
local dns = require("xray.dns")
local matcher = require("xray.geodata").BuildIPMatcher("127.0.0.1")
assert(dns.Servers == nil)
ips, ttl, err = dns.Query("MiXeD.Example.", true, false, true)
assert(not err and ttl == 42 and matcher:AnyMatch(ips))
`); err != nil {
t.Fatal(err)
}
got := L.GetGlobal("ips").(*lua.LUserData).Value.([]net.IP)
if &got[0] != &want[0] {
t.Fatal("dns.Query copied the IP slice")
}
}
func TestLuaDNSLocalClient(t *testing.T) {
L := lua.NewState()
defer L.Close()
L.SetContext(context.Background())
defer L.RemoveContext()
RegisterLua(L, localdns.New())
if err := L.DoString(`
local dns = require("xray.dns")
assert(dns.Servers[1].ID == "localhost")
serverIPs, _, serverErr = dns.Servers[1]:Query("127.0.0.1", true, false, false)
clientIPs, _, clientErr = dns.Query("127.0.0.1", true, false, false)
assert(not serverErr and not clientErr)
`); err != nil {
t.Fatal(err)
}
for _, name := range []string{"serverIPs", "clientIPs"} {
ips := L.GetGlobal(name).(*lua.LUserData).Value.([]net.IP)
if len(ips) != 1 || !ips[0].Equal(net.ParseIP("127.0.0.1")) {
t.Fatalf("%s = %v", name, ips)
}
}
}
type benchmarkLuaNameServer struct {
ips []net.IP
}
func (*benchmarkLuaNameServer) Name() string { return "benchmark" }
func (*benchmarkLuaNameServer) IsDisableCache() bool { return true }
func (s *benchmarkLuaNameServer) QueryIP(context.Context, string, featureDNS.IPOption) ([]net.IP, uint32, error) {
return s.ips, 60, nil
}
// BenchmarkLuaDNSHookCall isolates a preloaded Lua hook and its server:Query bridge.
// The direct case measures the same DNS client without Lua.
func BenchmarkLuaDNSHookCall(b *testing.B) {
option := featureDNS.IPOption{IPv4Enable: true}
ip := net.ParseIP("127.0.0.1")
upstream := &benchmarkLuaNameServer{ips: []net.IP{ip}}
client := &Client{server: upstream, ipOption: &option, timeoutMs: time.Second}
server := &DNS{clients: []*Client{client}}
L := lua.NewState()
defer L.Close()
server.RegisterLua(L)
if err := L.DoString(`
local server = require("xray.dns").Servers[1]
function HandleDNSQuery(domain, ipv4, ipv6, fake)
return server:Query(domain, ipv4, ipv6, fake)
end
`); err != nil {
b.Fatal(err)
}
ctx := context.Background()
for _, bench := range []struct {
name string
query func() ([]net.IP, uint32, error)
}{
{"direct", func() ([]net.IP, uint32, error) { return client.QueryIP(ctx, "example.com", option) }},
{"lua_hook", func() ([]net.IP, uint32, error) { return server.CallLuaHook(L, ctx, "example.com", option) }},
} {
b.Run(bench.name, func(b *testing.B) {
b.ReportAllocs()
b.ResetTimer()
var ips []net.IP
var ttl uint32
var err error
for i := 0; i < b.N; i++ {
ips, ttl, err = bench.query()
if err != nil {
b.Fatal(err)
}
}
b.StopTimer()
if ttl != 60 || len(ips) != 1 || !ips[0].Equal(ip) {
b.Fatalf("query() = %v, TTL %d; want %v, TTL 60", ips, ttl, ip)
}
})
}
}
+2 -1
View File
@@ -29,6 +29,7 @@ type Server interface {
// Client is the interface for DNS client.
type Client struct {
id string
server Server
skipFallback bool
expectedIPs geodata.IPMatcher
@@ -97,7 +98,7 @@ func NewClient(
ipOption dns.IPOption,
updateRules func(bool),
) (*Client, error) {
client := &Client{}
client := &Client{id: ns.Id}
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{server: NewLocalNameServer(), ipOption: &ipOption}
return &Client{id: "localhost", server: NewLocalNameServer(), ipOption: &ipOption}
}
+73
View File
@@ -0,0 +1,73 @@
package dns
import (
"context"
"time"
"github.com/xtls/xray-core/common/errors"
"github.com/xtls/xray-core/common/geodata"
"github.com/xtls/xray-core/common/log"
luamgr "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 = 10 * time.Second
type scriptEngine struct {
dns *DNS
pool *luamgr.Pool
}
func newScriptEngine(path string, server *DNS) (*scriptEngine, error) {
program, err := luamgr.CompileFile(path)
if err != nil {
return nil, err
}
e := &scriptEngine{dns: server}
e.pool, err = luamgr.NewPool(server.ctx, func(poolCtx context.Context) (*lua.LState, error) {
initCtx, cancel := context.WithTimeout(poolCtx, scriptExecutionTimeout)
defer cancel()
L, err := program.NewState(initCtx, func(L *lua.LState) {
geodata.RegisterLua(L)
log.RegisterLua(L)
server.RegisterLua(L)
})
if err != nil {
return nil, err
}
if L.GetGlobal("HandleDNSQuery").Type() != lua.LTFunction {
L.Close()
return nil, errors.New("DNS script must define HandleDNSQuery(domain, ipv4, ipv6, fake)")
}
return L, nil
})
if err != nil {
return nil, err
}
errors.LogInfo(server.ctx, "DNS script initialized from ", path)
return e, nil
}
func (e *scriptEngine) close() {
e.pool.Close()
}
func (e *scriptEngine) query(domain string, option dns.IPOption) ([]net.IP, uint32, error) {
L, err := e.pool.Acquire()
if err != nil {
return nil, 0, err
}
reusable := false
defer func() {
e.pool.Release(L, reusable)
}()
queryCtx, cancel := context.WithTimeout(e.pool.Context(), scriptExecutionTimeout)
defer cancel()
ips, ttl, err := e.dns.CallLuaHook(L, queryCtx, domain, option)
if err == nil {
reusable = true
}
return ips, ttl, err
}
+197
View File
@@ -0,0 +1,197 @@
package dns
import (
"context"
"os"
"path/filepath"
"strings"
"testing"
"time"
"github.com/xtls/xray-core/common/net"
featureDNS "github.com/xtls/xray-core/features/dns"
)
type geoIPScriptNameServer struct {
name string
answers map[string]net.IP
ttl uint32
calls int
}
func (s *geoIPScriptNameServer) Name() string { return s.name }
func (s *geoIPScriptNameServer) IsDisableCache() bool { return true }
func (s *geoIPScriptNameServer) QueryIP(ctx context.Context, domain string, _ featureDNS.IPOption) ([]net.IP, uint32, error) {
if err := ctx.Err(); err != nil {
return nil, 0, err
}
s.calls++
ip, ok := s.answers[domain]
if !ok {
return nil, 0, featureDNS.ErrEmptyResponse
}
return []net.IP{ip}, s.ttl, nil
}
func TestDNSScriptGeoIPFallback(t *testing.T) {
t.Setenv("xray.location.asset", filepath.Join("..", "..", "resources"))
script := `
local servers = require("xray.dns").Servers
local us_ips = require("xray.geodata").BuildIPMatcher("geoip:us")
local by_id = {}
for _, server in ipairs(servers) do
by_id[server.ID] = server
end
assert(by_id.primary and by_id.fallback, "primary and fallback DNS servers are required")
function HandleDNSQuery(domain, ipv4, ipv6, fake)
local ips, ttl, err = by_id.primary:Query(domain, ipv4, ipv6, fake)
if not err and us_ips:AnyMatch(ips) then
return ips, ttl, nil
end
return by_id.fallback:Query(domain, ipv4, ipv6, fake)
end
`
scriptPath := filepath.Join(t.TempDir(), "geoip_fallback.lua")
if err := os.WriteFile(scriptPath, []byte(script), 0o600); err != nil {
t.Fatal(err)
}
primary := &geoIPScriptNameServer{
name: "primary",
answers: map[string]net.IP{
"us.example": net.ParseIP("2001:4860:4860::8888"),
"other.example": net.ParseIP("127.0.0.1"),
},
ttl: 30,
}
fallback := &geoIPScriptNameServer{
name: "fallback",
answers: map[string]net.IP{"other.example": net.ParseIP("9.9.9.9")},
ttl: 60,
}
option := featureDNS.IPOption{IPv4Enable: true, IPv6Enable: true}
hosts, err := NewStaticHosts(nil)
if err != nil {
t.Fatal(err)
}
server := &DNS{
ctx: context.Background(),
hosts: hosts,
ipOption: &option,
scriptPath: scriptPath,
clients: []*Client{
{id: "primary", server: primary, ipOption: &option, timeoutMs: 2 * time.Second},
{id: "fallback", server: fallback, ipOption: &option, timeoutMs: 2 * time.Second},
},
}
if err := server.Start(); err != nil {
t.Fatal(err)
}
defer server.Close()
for _, tc := range []struct {
domain string
ip net.IP
ttl uint32
}{
{"Us.Example.", net.ParseIP("2001:4860:4860::8888"), 30},
{"other.example", net.ParseIP("9.9.9.9"), 60},
} {
ips, ttl, err := server.LookupIP(tc.domain, option)
if err != nil {
t.Fatalf("LookupIP(%q): %v", tc.domain, err)
}
if ttl != tc.ttl || len(ips) != 1 || !ips[0].Equal(tc.ip) {
t.Fatalf("LookupIP(%q) = %v, TTL %d; want %v, TTL %d", tc.domain, ips, ttl, tc.ip, tc.ttl)
}
}
if primary.calls != 2 || fallback.calls != 1 {
t.Fatalf("upstream calls: primary %d, fallback %d; want 2 and 1", primary.calls, fallback.calls)
}
}
func TestDNSScriptRejectsInvalidStartup(t *testing.T) {
for _, tc := range []struct {
name string
script string
}{
{"syntax", "function HandleDNSQuery("},
{"missing hook", "value = 1"},
{"top-level error", `error("setup failed")`},
} {
t.Run(tc.name, func(t *testing.T) {
path := filepath.Join(t.TempDir(), "script.lua")
if err := os.WriteFile(path, []byte(tc.script), 0o600); err != nil {
t.Fatal(err)
}
server := &DNS{ctx: context.Background(), scriptPath: path}
if err := server.Start(); err == nil {
t.Fatal("Start accepted an invalid DNS script")
}
if server.script != nil {
t.Fatal("Start retained a script engine after failure")
}
})
}
}
func TestDNSScriptHookErrorAndFakeDNSOption(t *testing.T) {
path := filepath.Join(t.TempDir(), "script.lua")
script := `
local server = require("xray.dns").Servers[1]
local log = require("xray.log")
log.Info("DNS script loaded")
function HandleDNSQuery(domain, ipv4, ipv6, fake)
log.Debug("DNS query: ", domain)
if domain == "bad.example" then error("script failure") end
local ips, ttl, err = server:Query(domain, ipv4, ipv6, fake)
if err then log.Error("DNS failed: ", err) end
return ips, ttl, err
end
`
if err := os.WriteFile(path, []byte(script), 0o600); err != nil {
t.Fatal(err)
}
option := featureDNS.IPOption{IPv4Enable: true}
hosts, err := NewStaticHosts(nil)
if err != nil {
t.Fatal(err)
}
upstream := &geoIPScriptNameServer{
name: "FakeDNS",
answers: map[string]net.IP{"good.example": net.ParseIP("198.18.0.1")},
ttl: 30,
}
server := &DNS{
ctx: context.Background(),
hosts: hosts,
ipOption: &option,
scriptPath: path,
clients: []*Client{{id: "fake", server: upstream, ipOption: &option, timeoutMs: time.Second}},
}
if err := server.Start(); err != nil {
t.Fatal(err)
}
defer server.Close()
if _, _, err := server.LookupIP("bad.example", option); err == nil || !strings.Contains(err.Error(), "script failure") {
t.Fatalf("hook failure = %v, want script failure", err)
}
if _, _, err := server.LookupIP("good.example", option); err != featureDNS.ErrEmptyResponse {
t.Fatalf("FakeDNS without FakeEnable = %v, want ErrEmptyResponse", err)
}
if upstream.calls != 0 {
t.Fatalf("FakeDNS was queried without FakeEnable: %d calls", upstream.calls)
}
withFake := featureDNS.IPOption{IPv4Enable: true, FakeEnable: true}
ips, ttl, err := server.LookupIP("good.example", withFake)
if err != nil || ttl != 30 || len(ips) != 1 || !ips[0].Equal(net.ParseIP("198.18.0.1")) {
t.Fatalf("FakeDNS with FakeEnable = %v, TTL %d, %v", ips, ttl, err)
}
if upstream.calls != 1 {
t.Fatalf("FakeDNS query count = %d, want 1", upstream.calls)
}
}
+14 -4
View File
@@ -587,8 +587,10 @@ type Config struct {
DomainStrategy Config_DomainStrategy `protobuf:"varint,1,opt,name=domain_strategy,json=domainStrategy,proto3,enum=xray.app.router.Config_DomainStrategy" json:"domain_strategy,omitempty"`
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"`
unknownFields protoimpl.UnknownFields
sizeCache protoimpl.SizeCache
// 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
}
func (x *Config) Reset() {
@@ -642,6 +644,13 @@ 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 = "" +
@@ -699,11 +708,12 @@ 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\"\x96\x02\n" +
"\ttolerance\x18\x06 \x01(\x02R\ttolerance\"\xae\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\"B\n" +
"\x0ebalancing_rule\x18\x03 \x03(\v2\x1e.xray.app.router.BalancingRuleR\rbalancingRule\x12\x16\n" +
"\x06script\x18\x04 \x01(\tR\x06script\"B\n" +
"\x0eDomainStrategy\x12\b\n" +
"\x04AsIs\x10\x00\x12\x10\n" +
"\fIpIfNonMatch\x10\x02\x12\x0e\n" +
+2
View File
@@ -110,4 +110,6 @@ 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;
}
+206
View File
@@ -0,0 +1,206 @@
package router
import (
"context"
"runtime"
"github.com/xtls/xray-core/common/errors"
"github.com/xtls/xray-core/common/net"
"github.com/xtls/xray-core/features/routing"
lua "github.com/yuin/gopher-lua"
)
const (
luaContextType = "xray.router.Context"
luaAttributesType = "xray.router.Attributes"
)
// RegisterLua makes xray.router available to routing scripts.
func (r *Router) RegisterLua(L *lua.LState) {
registerLuaContext(L)
L.PreloadModule("xray.router", func(L *lua.LState) int {
module := L.NewTable()
module.RawSetString("NetworkUnknown", lua.LNumber(net.Network_Unknown))
module.RawSetString("NetworkTCP", lua.LNumber(net.Network_TCP))
module.RawSetString("NetworkUDP", lua.LNumber(net.Network_UDP))
module.RawSetString("NetworkUNIX", lua.LNumber(net.Network_UNIX))
module.RawSetString("LocalOS", lua.LString(runtime.GOOS))
module.RawSetString("PickOutbound", L.NewFunction(func(L *lua.LState) int {
tag, ok := L.Get(2).(lua.LString)
if !ok {
L.ArgError(2, "balancer tag must be a string")
return 0
}
balancer, found := (*r.balancers.Load())[string(tag)]
if !found {
L.Push(lua.LNil)
pushLuaError(L, errors.New("balancer ", tag, " not found"))
return 2
}
outboundTag, err := balancer.PickOutbound()
L.Push(lua.LString(outboundTag))
pushLuaError(L, err)
return 2
}))
module.RawSetString("FindProcess", L.NewFunction(func(L *lua.LState) int {
pid, name, path, err := findProcess(checkLuaContext(L), net.FindProcess)
L.Push(lua.LNumber(pid))
L.Push(lua.LString(name))
L.Push(lua.LString(path))
pushLuaError(L, err)
return 4
}))
L.Push(module)
return 1
})
}
func registerLuaContext(L *lua.LState) {
attributes := L.NewTypeMetatable(luaAttributesType)
L.SetField(attributes, "__index", L.NewFunction(func(L *lua.LState) int {
values := L.CheckUserData(1).Value.(map[string]string)
key := L.CheckString(2)
if value, found := values[key]; found {
L.Push(lua.LString(value))
} else {
L.Push(lua.LNil)
}
return 1
}))
methods := L.NewTable()
L.SetFuncs(methods, map[string]lua.LGFunction{
"GetSourceIPs": func(L *lua.LState) int {
return pushLuaIPs(L, checkLuaContext(L).GetSourceIPs())
},
"GetTargetIPs": func(L *lua.LState) int {
return pushLuaIPs(L, checkLuaContext(L).GetTargetIPs())
},
"GetLocalIPs": func(L *lua.LState) int {
return pushLuaIPs(L, checkLuaContext(L).GetLocalIPs())
},
"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
}
func pushLuaIPs(L *lua.LState, ips []net.IP) int {
addresses := L.NewUserData()
addresses.Value = ips
L.Push(addresses)
return 1
}
func pushLuaError(L *lua.LState, err error) {
if err == nil {
L.Push(lua.LNil)
return
}
value := L.NewUserData()
value.Value = err
L.Push(value)
}
// CallLuaHook invokes HandleRoute in the supplied state.
func (r *Router) CallLuaHook(L *lua.LState, ctx context.Context, routeCtx routing.Context) (string, string, error) {
previous, top := L.Context(), L.GetTop()
L.SetContext(ctx)
defer func() {
L.SetTop(top)
if previous == nil {
L.RemoveContext()
} else {
L.SetContext(previous)
}
}()
fn := L.GetGlobal("HandleRoute")
if fn.Type() != lua.LTFunction {
return "", "", errors.New("routing script must define HandleRoute(...)")
}
value := L.NewUserData()
value.Value = routeCtx
L.SetMetatable(value, L.GetTypeMetatable(luaContextType))
if err := L.CallByParam(lua.P{Fn: fn, NRet: 3, Protect: true},
value, lua.LString(routeCtx.GetInboundTag()), lua.LNumber(routeCtx.GetSourcePort()),
lua.LNumber(routeCtx.GetTargetPort()), lua.LNumber(routeCtx.GetLocalPort()),
lua.LString(routeCtx.GetTargetDomain()), lua.LNumber(routeCtx.GetNetwork()),
lua.LString(routeCtx.GetProtocol()), lua.LString(routeCtx.GetUser()),
lua.LNumber(routeCtx.GetVlessRoute()), lua.LBool(routeCtx.GetSkipDNSResolve())); err != nil {
return "", "", err
}
return readLuaRouteResult(L.Get(-3), L.Get(-2), L.Get(-1))
}
func readLuaRouteResult(tagValue, ruleValue, errorValue lua.LValue) (string, string, error) {
if errorValue != lua.LNil {
if value, ok := errorValue.(*lua.LUserData); ok {
if err, ok := value.Value.(error); ok {
return "", "", err
}
}
if value, ok := errorValue.(lua.LString); ok {
return "", "", errors.New(string(value))
}
return "", "", errors.New("routing script error must be an error or string")
}
if tagValue == lua.LNil {
return "", "", nil
}
tag, ok := tagValue.(lua.LString)
if !ok {
return "", "", errors.New("routing script outboundTag must be a string or nil")
}
if tag == "" {
return "", "", nil
}
var ruleTag string
if ruleValue != lua.LNil {
value, ok := ruleValue.(lua.LString)
if !ok {
return "", "", errors.New("routing script ruleTag must be a string")
}
ruleTag = string(value)
}
return string(tag), ruleTag, nil
}
type processFinder func(string, string, uint16, string, uint16) (int, string, string, error)
func findProcess(ctx routing.Context, finder processFinder) (int, string, string, error) {
sources := ctx.GetSourceIPs()
if len(sources) == 0 {
return 0, "", "", errors.New("process lookup requires a source IP")
}
var network string
switch ctx.GetNetwork() {
case net.Network_TCP:
network = "tcp"
case net.Network_UDP:
network = "udp"
default:
return 0, "", "", errors.New("process lookup requires TCP or UDP")
}
targetIP, targetPort := "", uint16(0)
if targets := ctx.GetTargetIPs(); len(targets) > 0 {
targetIP, targetPort = targets[0].String(), uint16(ctx.GetTargetPort())
}
return finder(network, sources[0].String(), uint16(ctx.GetSourcePort()), targetIP, targetPort)
}
+292
View File
@@ -0,0 +1,292 @@
package router
import (
"context"
go_errors "errors"
"runtime"
"strings"
"testing"
"github.com/xtls/xray-core/common/geodata"
"github.com/xtls/xray-core/common/net"
"github.com/xtls/xray-core/common/protocol"
"github.com/xtls/xray-core/common/session"
"github.com/xtls/xray-core/features/routing"
routing_session "github.com/xtls/xray-core/features/routing/session"
lua "github.com/yuin/gopher-lua"
)
type luaRouteTestContext struct {
*routing_session.Context
sourceIPs, targetIPs, localIPs []net.IP
}
func (c *luaRouteTestContext) GetSourceIPs() []net.IP { return c.sourceIPs }
func (c *luaRouteTestContext) GetTargetIPs() []net.IP { return c.targetIPs }
func (c *luaRouteTestContext) GetLocalIPs() []net.IP { return c.localIPs }
func newLuaRouteTestContext() *luaRouteTestContext {
return &luaRouteTestContext{
Context: &routing_session.Context{
Inbound: &session.Inbound{
Tag: "in", VlessRoute: 4321,
Source: net.TCPDestination(net.LocalHostIP, 1234),
Local: net.TCPDestination(net.LocalHostIP, 5678),
User: &protocol.MemoryUser{Email: "user@example.com"},
},
Outbound: &session.Outbound{
Target: net.TCPDestination(net.LocalHostIP, 443),
RouteTarget: net.TCPDestination(net.DomainAddress("MiXeD.Example."), 443),
},
Content: &session.Content{
Protocol: "tls", Attributes: map[string]string{"key": "value"}, SkipDNSResolve: true,
},
},
sourceIPs: []net.IP{{127, 0, 0, 2}},
targetIPs: []net.IP{{127, 0, 0, 3}},
localIPs: []net.IP{{127, 0, 0, 1}},
}
}
func newLuaRouterState(t *testing.T, script string) (*Router, *lua.LState) {
t.Helper()
r := new(Router)
if err := r.Init(context.Background(), &Config{}, nil, nil, nil); err != nil {
t.Fatal(err)
}
L := lua.NewState()
t.Cleanup(L.Close)
r.RegisterLua(L)
geodata.RegisterLua(L)
if err := L.DoString(script); err != nil {
t.Fatal(err)
}
return r, L
}
func TestLuaRouteBinding(t *testing.T) {
r, L := newLuaRouterState(t, `
local router = require("xray.router")
local matcher = require("xray.geodata").BuildIPMatcher("127.0.0.0/8")
assert(router.NetworkUnknown == 0 and router.NetworkTCP == 2)
assert(router.NetworkUDP == 3 and router.NetworkUNIX == 4)
assert(router.BuildIPMatcher == nil and router.BuildDomainMatcher == nil)
function HandleRoute(ctx, inboundTag, sourcePort, targetPort, localPort,
targetDomain, network, protocol, user, vlessRoute, skipDNSResolve, ...)
assert(select("#", ...) == 0)
assert(inboundTag == "in" and sourcePort == 1234 and targetPort == 443 and localPort == 5678)
assert(targetDomain == "MiXeD.Example." and network == router.NetworkTCP)
assert(protocol == "tls" and user == "user@example.com" and vlessRoute == 4321 and skipDNSResolve)
assert(ctx.GetNetwork == nil and ctx.Context == nil)
savedContext = ctx
sourceIPs, targetIPs, localIPs = ctx:GetSourceIPs(), ctx:GetTargetIPs(), ctx:GetLocalIPs()
attributes = ctx:GetAttributes()
assert(matcher:AnyMatch(sourceIPs) and matcher:AnyMatch(targetIPs) and matcher:AnyMatch(localIPs))
assert(attributes.key == "value" and attributes.missing == nil)
assert(not pcall(function() attributes.key = "changed" end))
return "out", "rule"
end`)
ctx := newLuaRouteTestContext()
tag, rule, err := r.CallLuaHook(L, context.Background(), ctx)
if err != nil || tag != "out" || rule != "rule" {
t.Fatalf("hook = %q, %q, %v", tag, rule, err)
}
if L.GetGlobal("savedContext").(*lua.LUserData).Value != ctx {
t.Fatal("routing context was copied")
}
for _, tc := range []struct {
name string
want []net.IP
}{
{"sourceIPs", ctx.sourceIPs},
{"targetIPs", ctx.targetIPs},
{"localIPs", ctx.localIPs},
} {
got := L.GetGlobal(tc.name).(*lua.LUserData).Value.([]net.IP)
if &got[0] != &tc.want[0] {
t.Fatalf("%s storage was copied", tc.name)
}
}
ctx.Content.Attributes["key"] = "updated"
L.SetGlobal("expectedOS", lua.LString(runtime.GOOS))
if err := L.DoString(`
assert(attributes.key == "updated")
assert(require("xray.router").LocalOS == expectedOS)`); err != nil {
t.Fatal(err)
}
}
func TestLuaRouteResult(t *testing.T) {
nativeErr := go_errors.New("native failure")
for _, tc := range []struct {
name, body, tag, rule, wantErr string
native bool
}{
{name: "route", body: `return "out", "rule"`, tag: "out", rule: "rule"},
{name: "no match", body: `return nil`},
{name: "empty tag", body: `return ""`},
{name: "invalid tag", body: `return 1`, wantErr: "outboundTag"},
{name: "invalid rule", body: `return "out", false`, wantErr: "ruleTag"},
{name: "string error", body: `return nil, nil, "script failure"`, wantErr: "script failure"},
{name: "native error", body: `return nil, nil, nativeError`, native: true},
{name: "runtime error", body: `error("runtime failure")`, wantErr: "runtime failure"},
} {
t.Run(tc.name, func(t *testing.T) {
r, L := newLuaRouterState(t, "function HandleRoute() "+tc.body+" end")
value := L.NewUserData()
value.Value = nativeErr
L.SetGlobal("nativeError", value)
previous := context.WithValue(context.Background(), struct{}{}, true)
L.SetContext(previous)
L.Push(lua.LTrue)
tag, rule, err := r.CallLuaHook(L, context.Background(), &routing_session.Context{})
if tag != tc.tag || rule != tc.rule {
t.Fatalf("result = %q, %q, %v", tag, rule, err)
}
switch {
case tc.native:
if err != nativeErr {
t.Fatalf("error = %v, want original error", err)
}
case tc.wantErr != "":
if err == nil || !strings.Contains(err.Error(), tc.wantErr) {
t.Fatalf("error = %v, want %q", err, tc.wantErr)
}
case err != nil:
t.Fatal(err)
}
if L.Context() != previous || L.GetTop() != 1 || L.Get(1) != lua.LTrue {
t.Fatal("hook did not restore the previous context and stack")
}
})
}
}
func TestLuaRouteCancellation(t *testing.T) {
r, L := newLuaRouterState(t, `function HandleRoute() while true do end end`)
ctx, cancel := context.WithCancel(context.Background())
cancel()
if _, _, err := r.CallLuaHook(L, ctx, &routing_session.Context{}); err == nil {
t.Fatal("CallLuaHook did not stop after context cancellation")
}
if L.Context() != nil || L.GetTop() != 0 {
t.Fatal("CallLuaHook did not restore the Lua state")
}
}
func TestFindProcess(t *testing.T) {
for _, tc := range []struct {
name, network, target string
targetPort uint16
modify func(*luaRouteTestContext)
wantErr bool
}{
{name: "TCP", network: "tcp", target: "127.0.0.3", targetPort: 443},
{name: "UDP", network: "udp", target: "127.0.0.3", targetPort: 443, modify: func(c *luaRouteTestContext) {
c.Outbound.Target.Network = net.Network_UDP
}},
{name: "domain target", network: "tcp", modify: func(c *luaRouteTestContext) { c.targetIPs = nil }},
{name: "missing source", modify: func(c *luaRouteTestContext) { c.sourceIPs = nil }, wantErr: true},
{name: "unsupported network", modify: func(c *luaRouteTestContext) {
c.Outbound.Target.Network = net.Network_UNIX
}, wantErr: true},
} {
t.Run(tc.name, func(t *testing.T) {
ctx := newLuaRouteTestContext()
if tc.modify != nil {
tc.modify(ctx)
}
called := false
pid, name, path, err := findProcess(ctx, func(network, source string, sourcePort uint16, target string, targetPort uint16) (int, string, string, error) {
called = true
if network != tc.network || source != "127.0.0.2" || sourcePort != 1234 || target != tc.target || targetPort != tc.targetPort {
t.Fatalf("endpoints = %s %s:%d -> %s:%d", network, source, sourcePort, target, targetPort)
}
return 42, "process", "/path/process", nil
})
if tc.wantErr {
if err == nil || called {
t.Fatalf("findProcess = %d, %q, %q, %v", pid, name, path, err)
}
return
}
if err != nil || !called || pid != 42 || name != "process" || path != "/path/process" {
t.Fatalf("findProcess = %d, %q, %q, %v", pid, name, path, err)
}
})
}
}
// BenchmarkLuaRouteHookCall isolates a preloaded Lua hook and its routing context bridge.
// The direct case runs an equivalent native routing rule.
func BenchmarkLuaRouteHookCall(b *testing.B) {
r := new(Router)
if err := r.Init(context.Background(), &Config{Rule: []*RoutingRule{{
TargetTag: &RoutingRule_Tag{Tag: "out"},
RuleTag: "rule",
InboundTag: []string{"in"},
Networks: []net.Network{net.Network_TCP},
Ip: []*geodata.IPRule{{
Value: &geodata.IPRule_Custom{Custom: &geodata.CIDRRule{
Cidr: &geodata.CIDR{Ip: []byte{127, 0, 0, 0}, Prefix: 8},
}},
}},
}}}, nil, nil, nil); err != nil {
b.Fatal(err)
}
L := lua.NewState()
defer L.Close()
r.RegisterLua(L)
geodata.RegisterLua(L)
if err := L.DoString(`
local router = require("xray.router")
local matcher = require("xray.geodata").BuildIPMatcher("127.0.0.0/8")
function HandleRoute(ctx, inboundTag, sourcePort, targetPort, localPort,
targetDomain, network, protocol, user, vlessRoute, skipDNSResolve)
if inboundTag == "in" and network == router.NetworkTCP and matcher:AnyMatch(ctx:GetTargetIPs()) then
return "out", "rule"
end
end
`); err != nil {
b.Fatal(err)
}
ctx := context.Background()
routeCtx := newLuaRouteTestContext()
for _, benchmark := range []struct {
name string
route func() (string, string, error)
}{
{"direct", func() (string, string, error) {
route, err := r.PickRoute(routeCtx)
if err != nil {
return "", "", err
}
return route.GetOutboundTag(), route.GetRuleTag(), nil
}},
{"lua_hook", func() (string, string, error) {
return r.CallLuaHook(L, ctx, routeCtx)
}},
} {
b.Run(benchmark.name, func(b *testing.B) {
b.ReportAllocs()
b.ResetTimer()
var tag, rule string
var err error
for i := 0; i < b.N; i++ {
tag, rule, err = benchmark.route()
if err != nil {
b.Fatal(err)
}
}
b.StopTimer()
if tag != "out" || rule != "rule" {
b.Fatalf("route() = %q, %q; want out, rule", tag, rule)
}
})
}
}
var _ routing.Context = (*luaRouteTestContext)(nil)
+17
View File
@@ -20,6 +20,8 @@ import (
type Router struct {
domainStrategy Config_DomainStrategy
rules atomic.Pointer[[]*Rule]
scriptPath string
script *scriptEngine
balancers atomic.Pointer[map[string]*Balancer]
dns dns.Client
@@ -40,6 +42,7 @@ 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
@@ -52,6 +55,10 @@ 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 {
@@ -221,6 +228,13 @@ 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
}
@@ -235,6 +249,9 @@ 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())
+79
View File
@@ -0,0 +1,79 @@
package router
import (
"context"
"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"
luamgr "github.com/xtls/xray-core/common/lua"
"github.com/xtls/xray-core/features/routing"
lua "github.com/yuin/gopher-lua"
)
const scriptExecutionTimeout = 10 * time.Second
type scriptEngine struct {
router *Router
pool *luamgr.Pool
}
func newScriptEngine(path string, router *Router) (*scriptEngine, error) {
program, err := luamgr.CompileFile(path)
if err != nil {
return nil, err
}
e := &scriptEngine{router: router}
e.pool, err = luamgr.NewPool(router.ctx, func(poolCtx context.Context) (*lua.LState, error) {
initCtx, cancel := context.WithTimeout(poolCtx, scriptExecutionTimeout)
defer cancel()
L, err := program.NewState(initCtx, func(L *lua.LState) {
geodata.RegisterLua(L)
log.RegisterLua(L)
router.RegisterLua(L)
dns.RegisterLua(L, router.dns)
})
if err != nil {
return nil, err
}
if L.GetGlobal("HandleRoute").Type() != lua.LTFunction {
L.Close()
return nil, errors.New("routing script must define HandleRoute(...)")
}
return L, nil
})
if err != nil {
return nil, err
}
errors.LogInfo(router.ctx, "routing script initialized from ", path)
return e, nil
}
func (e *scriptEngine) close() {
e.pool.Close()
}
func (e *scriptEngine) pickRoute(ctx routing.Context) (routing.Route, error) {
L, err := e.pool.Acquire()
if err != nil {
return nil, err
}
reusable := false
defer func() {
e.pool.Release(L, reusable)
}()
callCtx, cancel := context.WithTimeout(e.pool.Context(), scriptExecutionTimeout)
defer cancel()
tag, ruleTag, err := e.router.CallLuaHook(L, callCtx, ctx)
if err != nil {
return nil, err
}
reusable = true
if tag == "" {
return nil, common.ErrNoClue
}
return &Route{Context: ctx, outboundTag: tag, ruleTag: ruleTag}, nil
}
+372
View File
@@ -0,0 +1,372 @@
package router
import (
"context"
stdnet "net"
"os"
"path/filepath"
"sync"
"sync/atomic"
"testing"
"time"
wireDNS "github.com/miekg/dns"
"github.com/xtls/xray-core/app/dispatcher"
appdns "github.com/xtls/xray-core/app/dns"
"github.com/xtls/xray-core/app/proxyman"
_ "github.com/xtls/xray-core/app/proxyman/outbound"
"github.com/xtls/xray-core/common"
"github.com/xtls/xray-core/common/net"
"github.com/xtls/xray-core/common/serial"
"github.com/xtls/xray-core/core"
featureDNS "github.com/xtls/xray-core/features/dns"
"github.com/xtls/xray-core/features/outbound"
"github.com/xtls/xray-core/features/routing"
routing_session "github.com/xtls/xray-core/features/routing/session"
"github.com/xtls/xray-core/proxy/blackhole"
"github.com/xtls/xray-core/proxy/freedom"
)
type luaRouteDNSClient struct {
featureDNS.Client
lookup func(string, featureDNS.IPOption) ([]net.IP, uint32, error)
}
func (d *luaRouteDNSClient) LookupIP(domain string, option featureDNS.IPOption) ([]net.IP, uint32, error) {
return d.lookup(domain, option)
}
type luaRouteOutboundManager struct{ outbound.Manager }
func (*luaRouteOutboundManager) Select(selectors []string) []string { return selectors }
func writeRouteScript(t *testing.T, script string) string {
t.Helper()
path := filepath.Join(t.TempDir(), "route.lua")
if err := os.WriteFile(path, []byte(script), 0o600); err != nil {
t.Fatal(err)
}
return path
}
func startLuaRouter(t *testing.T, script string, d featureDNS.Client, config *Config) *Router {
t.Helper()
if config == nil {
config = &Config{}
}
config.Script = writeRouteScript(t, script)
r := new(Router)
if err := r.Init(context.Background(), config, d, &luaRouteOutboundManager{}, nil); err != nil {
t.Fatal(err)
}
if err := r.Start(); err != nil {
t.Fatal(err)
}
t.Cleanup(func() {
if err := r.Close(); err != nil {
t.Error(err)
}
})
return r
}
func TestRouterScriptStartup(t *testing.T) {
for _, tc := range []struct{ name, script string }{
{"syntax error", "function HandleRoute("},
{"missing hook", "value = 1"},
{"initialization error", `error("setup failed")`},
} {
t.Run(tc.name, func(t *testing.T) {
r := new(Router)
if err := r.Init(context.Background(), &Config{Script: writeRouteScript(t, tc.script)}, nil, nil, nil); err != nil {
t.Fatal(err)
}
defer r.Close()
if err := r.Start(); err == nil {
t.Fatal("Start accepted an invalid routing script")
}
})
}
}
func TestRouterScriptRouting(t *testing.T) {
var dnsCalls atomic.Int32
d := &luaRouteDNSClient{lookup: func(string, featureDNS.IPOption) ([]net.IP, uint32, error) {
dnsCalls.Add(1)
return []net.IP{{1, 2, 3, 4}}, 60, nil
}}
r := startLuaRouter(t, `
function HandleRoute(ctx, inbound)
if inbound == "miss" then return nil end
return "lua-out", "lua-rule"
end`, d, &Config{
DomainStrategy: Config_IpOnDemand,
Rule: []*RoutingRule{{
TargetTag: &RoutingRule_Tag{Tag: "json-out"},
Networks: []net.Network{net.Network_TCP},
}},
})
ctx := newLuaRouteTestContext()
ctx.Content.SkipDNSResolve = false
route, err := r.PickRoute(ctx)
if err != nil || route.GetOutboundTag() != "lua-out" || route.GetRuleTag() != "lua-rule" || route.(*Route).Context != ctx {
t.Fatalf("route = %v, %v", route, err)
}
ctx.Inbound.Tag = "miss"
if route, err := r.PickRoute(ctx); route != nil || err != common.ErrNoClue {
t.Fatalf("miss = %v, %v", route, err)
}
if dnsCalls.Load() != 0 {
t.Fatal("script routing implicitly resolved DNS")
}
}
func TestRouterScriptModules(t *testing.T) {
ips := []net.IP{{127, 0, 0, 7}}
calls := 0
d := &luaRouteDNSClient{lookup: func(domain string, option featureDNS.IPOption) ([]net.IP, uint32, error) {
calls++
if domain != "MiXeD.Example." || !option.IPv4Enable || option.IPv6Enable || !option.FakeEnable {
t.Fatalf("dns.Query arguments = %q, %+v", domain, option)
}
return ips, 17, nil
}}
r := startLuaRouter(t, `
local dns = require("xray.dns")
local matcher = require("xray.geodata").BuildIPMatcher("127.0.0.0/8")
assert(dns.Servers == nil and type(dns.Query) == "function")
assert(type(require("xray.log").Info) == "function")
function HandleRoute(ctx, inbound, sourcePort, targetPort, localPort, domain)
local ips, ttl, err = dns.Query(domain, true, false, true)
assert(not err and ttl == 17)
assert(matcher:AnyMatch(ips) and matcher:AnyMatch(ctx:GetTargetIPs()))
return "out"
end`, d, nil)
if _, err := r.PickRoute(newLuaRouteTestContext()); err != nil {
t.Fatal(err)
}
if calls != 1 {
t.Fatalf("DNS calls = %d, want 1", calls)
}
}
func TestRouterScriptBalancerReload(t *testing.T) {
config := func(tag string) *Config {
return &Config{BalancingRule: []*BalancingRule{{
Tag: "balance", Strategy: "roundrobin", OutboundSelector: []string{tag},
}}}
}
r := startLuaRouter(t, `
local router = require("xray.router")
function HandleRoute()
local tag, err = router:PickOutbound("balance")
return tag, "balanced", err
end`, nil, config("old"))
pick := func(want string) {
t.Helper()
route, err := r.PickRoute(&routing_session.Context{})
if err != nil || route.GetOutboundTag() != want || route.GetRuleTag() != "balanced" {
t.Fatalf("route = %v, %v, want %q", route, err, want)
}
}
pick("old")
if err := r.SetOverrideTarget("balance", "override"); err != nil {
t.Fatal(err)
}
pick("override")
if err := r.SetOverrideTarget("balance", ""); err != nil {
t.Fatal(err)
}
if err := r.ReloadRules(config("new"), false); err != nil {
t.Fatal(err)
}
pick("new")
}
func TestRouterScriptConcurrentBalancerReload(t *testing.T) {
config := func(tag string) *Config {
return &Config{BalancingRule: []*BalancingRule{{
Tag: "balance", Strategy: "roundrobin", OutboundSelector: []string{tag},
}}}
}
r := startLuaRouter(t, `
local router = require("xray.router")
function HandleRoute()
local tag, err = router:PickOutbound("balance")
return tag, nil, err
end`, nil, config("a"))
var wg sync.WaitGroup
for range 4 {
wg.Go(func() {
for range 20 {
route, err := r.PickRoute(&routing_session.Context{})
if err != nil {
t.Errorf("PickRoute: %v", err)
return
}
if tag := route.GetOutboundTag(); tag != "a" && tag != "b" {
t.Errorf("unexpected tag %q", tag)
}
}
})
}
wg.Go(func() {
for range 20 {
for _, tag := range []string{"a", "b"} {
if err := r.ReloadRules(config(tag), false); err != nil {
t.Error(err)
return
}
}
}
})
wg.Wait()
}
func TestRouterScriptStateReuse(t *testing.T) {
r := startLuaRouter(t, `
local calls = 0
function HandleRoute(ctx, inbound)
calls = calls + 1
if inbound == "miss" then return nil end
if inbound == "fail" then error("failed") end
return tostring(calls)
end`, nil, nil)
ctx := newLuaRouteTestContext()
pick := func(want string) {
t.Helper()
route, err := r.PickRoute(ctx)
if err != nil || route.GetOutboundTag() != want {
t.Fatalf("route = %v, %v, want %q", route, err, want)
}
}
pick("1")
ctx.Inbound.Tag = "miss"
if _, err := r.PickRoute(ctx); err != common.ErrNoClue {
t.Fatalf("miss = %v", err)
}
ctx.Inbound.Tag = "in"
pick("3")
ctx.Inbound.Tag = "fail"
if _, err := r.PickRoute(ctx); err == nil {
t.Fatal("script error was ignored")
}
ctx.Inbound.Tag = "in"
pick("1")
}
func TestRouterScriptDNSDispatcherReentry(t *testing.T) {
conn, err := stdnet.ListenPacket("udp4", "127.0.0.1:0")
if err != nil {
t.Fatal(err)
}
port := conn.LocalAddr().(*stdnet.UDPAddr).Port
ready, stopped := make(chan struct{}), make(chan error, 1)
var queries atomic.Int32
server := &wireDNS.Server{
PacketConn: conn,
NotifyStartedFunc: func() {
close(ready)
},
Handler: wireDNS.HandlerFunc(func(w wireDNS.ResponseWriter, query *wireDNS.Msg) {
queries.Add(1)
response := new(wireDNS.Msg).SetReply(query)
for _, question := range query.Question {
if question.Name == "nested.example." && question.Qtype == wireDNS.TypeA {
response.Answer = append(response.Answer, &wireDNS.A{
Hdr: wireDNS.RR_Header{Name: question.Name, Rrtype: wireDNS.TypeA, Class: wireDNS.ClassINET, Ttl: 60},
A: stdnet.IP{127, 0, 0, 7},
})
}
}
if err := w.WriteMsg(response); err != nil {
t.Error(err)
}
}),
}
go func() { stopped <- server.ActivateAndServe() }()
defer func() {
server.Shutdown()
select {
case err := <-stopped:
if err != nil {
t.Error(err)
}
case <-time.After(3 * time.Second):
t.Error("DNS server did not stop")
}
}()
select {
case <-ready:
case err := <-stopped:
t.Fatalf("DNS server startup: %v", err)
case <-time.After(3 * time.Second):
t.Fatal("DNS server did not start")
}
dnsScript := writeRouteScript(t, `
local server = require("xray.dns").Servers[1]
function HandleDNSQuery(domain, ipv4, ipv6, fake)
return server:Query(domain, ipv4, ipv6, fake)
end`)
routerScript := writeRouteScript(t, `
local router = require("xray.router")
local dns = require("xray.dns")
local matcher = require("xray.geodata").BuildIPMatcher("127.0.0.7")
local active = false
function HandleRoute(ctx, inbound, sourcePort, targetPort, localPort, domain, network,
protocol, user, vlessRoute, skipDNSResolve)
assert(not active, "borrowed Router VM reentered")
if inbound == "dns" then
assert(network == router.NetworkUDP and skipDNSResolve == false)
return "direct", "dns-route"
end
active = true
local ips, ttl, err = dns.Query("nested.example", true, false, false)
assert(not err and matcher:AnyMatch(ips) and active)
active = false
return "direct", "outer-route"
end`)
instance, err := core.New(&core.Config{
App: []*serial.TypedMessage{
serial.ToTypedMessage(&appdns.Config{
Tag: "dns", Script: dnsScript, DisableCache: true,
NameServer: []*appdns.NameServer{{
Id: "upstream", TimeoutMs: 1000,
Address: &net.Endpoint{
Network: net.Network_UDP,
Address: &net.IPOrDomain{Address: &net.IPOrDomain_Ip{Ip: []byte{127, 0, 0, 1}}},
Port: uint32(port),
},
}},
}),
serial.ToTypedMessage(&Config{Script: routerScript}),
serial.ToTypedMessage(&dispatcher.Config{}),
serial.ToTypedMessage(&proxyman.OutboundConfig{}),
},
Outbound: []*core.OutboundHandlerConfig{
{Tag: "default", ProxySettings: serial.ToTypedMessage(&blackhole.Config{})},
{Tag: "direct", ProxySettings: serial.ToTypedMessage(&freedom.Config{
FinalRules: []*freedom.FinalRuleConfig{{Action: freedom.RuleAction_Allow}},
})},
},
})
if err != nil {
t.Fatal(err)
}
defer instance.Close()
if err := instance.Start(); err != nil {
t.Fatal(err)
}
r := instance.GetFeature(routing.RouterType()).(*Router)
route, err := r.PickRoute(newLuaRouteTestContext())
if err != nil || route.GetOutboundTag() != "direct" || route.GetRuleTag() != "outer-route" {
t.Fatalf("nested DNS routing = %v, %v", route, err)
}
if queries.Load() == 0 {
t.Fatal("DNS query did not pass through the dispatcher")
}
}
+11 -1
View File
@@ -5,7 +5,8 @@ import (
)
type windowsReader struct {
bufs []syscall.WSABuf
bufs []syscall.WSABuf
ready bool
}
func (r *windowsReader) Init(bs []*Buffer) {
@@ -15,6 +16,7 @@ func (r *windowsReader) Init(bs []*Buffer) {
for _, b := range bs {
r.bufs = append(r.bufs, syscall.WSABuf{Len: uint32(Size), Buf: &b.v[0]})
}
r.ready = false
}
func (r *windowsReader) Clear() {
@@ -25,6 +27,14 @@ func (r *windowsReader) Clear() {
}
func (r *windowsReader) Read(fd uintptr) int32 {
// On the first invocation, we return -1 to indicate "not ready"
// to make rawConn.Read wait for readability using the runtime's own mechanism
// because syscall.WSARecv() is a blocking call when used with nil OVERLAPPED
if !r.ready {
r.ready = true
return -1
}
var nBytes uint32
var flags uint32
err := syscall.WSARecv(syscall.Handle(fd), &r.bufs[0], uint32(len(r.bufs)), &nBytes, &flags, nil, nil)
+58
View File
@@ -0,0 +1,58 @@
package geodata
import (
lua "github.com/yuin/gopher-lua"
luar "layeh.com/gopher-luar"
)
// RegisterLua makes xray.geodata available to require in an LState.
func RegisterLua(L *lua.LState) {
L.PreloadModule("xray.geodata", func(L *lua.LState) int {
module := L.NewTable()
module.RawSetString("BuildDomainMatcher", L.NewFunction(func(L *lua.LState) int {
parsed, err := ParseDomainRules(luaRules(L), Domain_Domain)
if err != nil {
L.RaiseError("%v", err)
return 0
}
matcher, err := DomainReg.BuildDomainMatcher(parsed)
if err != nil {
L.RaiseError("%v", err)
return 0
}
L.Push(luar.New(L, matcher))
return 1
}))
module.RawSetString("BuildIPMatcher", L.NewFunction(func(L *lua.LState) int {
parsed, err := ParseIPRules(luaRules(L))
if err != nil {
L.RaiseError("%v", err)
return 0
}
matcher, err := IPReg.BuildIPMatcher(parsed)
if err != nil {
L.RaiseError("%v", err)
return 0
}
L.Push(luar.New(L, matcher))
return 1
}))
L.Push(module)
return 1
})
}
func luaRules(L *lua.LState) []string {
rules := make([]string, L.GetTop())
for i := range rules {
value, ok := L.Get(i + 1).(lua.LString)
if !ok {
L.RaiseError("geodata rules must be strings")
return nil
}
rules[i] = string(value)
}
return rules
}
+66
View File
@@ -0,0 +1,66 @@
package geodata
import (
"testing"
"github.com/xtls/xray-core/common/net"
lua "github.com/yuin/gopher-lua"
)
func TestLuaIPMatcher(t *testing.T) {
L := lua.NewState()
defer L.Close()
RegisterLua(L)
ip := L.NewUserData()
ip.Value = net.ParseIP("127.0.0.1")
L.SetGlobal("ip", ip)
ips := L.NewUserData()
ips.Value = []net.IP{ip.Value.(net.IP), net.ParseIP("8.8.8.8")}
L.SetGlobal("ips", ips)
if err := L.DoString(`
local matcher = require("xray.geodata").BuildIPMatcher("127.0.0.0/8", "::1")
assert(matcher:Match(ip))
assert(matcher:AnyMatch(ips))
assert(not matcher:Matches(ips))
local matched, unmatched = matcher:FilterIPs(ips)
assert(type(matched) == "userdata" and type(unmatched) == "userdata")
assert(#matched == 1 and #unmatched == 1)
`); err != nil {
t.Fatal(err)
}
}
func TestLuaDomainMatcher(t *testing.T) {
L := lua.NewState()
defer L.Close()
RegisterLua(L)
if err := L.DoString(`
local matcher = require("xray.geodata").BuildDomainMatcher("example.com", "full:other.com")
assert(matcher:MatchAny("example.com"))
assert(matcher:MatchAny("www.example.com"))
assert(matcher:MatchAny("other.com"))
assert(not matcher:MatchAny("www.other.com"))
assert(#(matcher:Match("www.example.com")) == 1)
`); err != nil {
t.Fatal(err)
}
}
func TestLuaMatchersRejectInvalidRules(t *testing.T) {
for _, tc := range []struct {
name string
script string
}{
{"IP rule", `require("xray.geodata").BuildIPMatcher("not-an-ip")`},
{"non-string domain rule", `require("xray.geodata").BuildDomainMatcher("example.com", true)`},
} {
t.Run(tc.name, func(t *testing.T) {
L := lua.NewState()
defer L.Close()
RegisterLua(L)
if err := L.DoString(tc.script); err == nil {
t.Fatal("invalid geodata rule was accepted")
}
})
}
}
+50
View File
@@ -0,0 +1,50 @@
package log
import (
"path/filepath"
"strings"
lua "github.com/yuin/gopher-lua"
)
// RegisterLua makes xray.log available to require in an LState.
func RegisterLua(L *lua.LState) {
L.PreloadModule("xray.log", func(L *lua.LState) int {
module := L.NewTable()
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 {
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 != "" {
content.WriteString(filepath.Base(strings.TrimPrefix(caller.Source, "@")))
content.WriteString(": ")
}
}
for i := 1; i <= L.GetTop(); i++ {
value := L.Get(i)
// Use Error() for Go errors in userdata.
if ud, ok := value.(*lua.LUserData); ok {
if err, ok := ud.Value.(error); ok {
content.WriteString(err.Error())
continue
}
}
content.WriteString(L.ToStringMeta(value).String())
}
Record(&GeneralMessage{
Severity: severity,
Content: content.String(),
})
return 0
}))
}
L.Push(module)
return 1
})
}
+92
View File
@@ -0,0 +1,92 @@
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) {
logHandler.RLock()
previous := logHandler.Handler
logHandler.RUnlock()
t.Cleanup(func() { RegisterHandler(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)
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)
}
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: hook"},
{Severity_Info, "[Info] <string>: anonymous"},
}
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)
}
}
}
+138
View File
@@ -0,0 +1,138 @@
package lua
import (
"context"
"errors"
"sync"
glua "github.com/yuin/gopher-lua"
)
const maxIdleStates = 16
// LStateFactory must initialize a state fully and observe ctx while doing so.
// The pool owns any non-nil state it returns, even when it also returns an error.
type LStateFactory func(ctx context.Context) (*glua.LState, error)
// Pool lends each state to one caller at a time. It grows on contention and
// keeps up to maxIdleStates idle states until Close. Callers decide whether a
// state is reusable.
type Pool struct {
ctx context.Context
cancel context.CancelFunc
factory LStateFactory
idle []*glua.LState
mu sync.Mutex
active sync.WaitGroup
closed bool
}
// NewPool initializes one state before returning, so top-level errors surface at startup.
func NewPool(ctx context.Context, factory LStateFactory) (*Pool, error) {
poolCtx, cancel := context.WithCancel(ctx)
// Create one state now to catch factory errors at startup.
state, err := factory(poolCtx)
if err != nil {
cancel()
if state != nil {
state.Close()
}
return nil, err
}
if state == nil {
cancel()
return nil, errors.New("Lua state factory returned nil")
}
if err := poolCtx.Err(); err != nil {
state.Close()
cancel()
return nil, err
}
return &Pool{ctx: poolCtx, cancel: cancel, factory: factory, idle: []*glua.LState{state}}, nil
}
// Context is cancelled by Close. Query contexts should derive from it.
func (p *Pool) Context() context.Context {
return p.ctx
}
// Acquire returns an initialized exclusive state, growing the pool if necessary.
func (p *Pool) Acquire() (*glua.LState, error) {
p.mu.Lock()
if p.closed || p.ctx.Err() != nil {
p.mu.Unlock()
return nil, p.ctx.Err()
}
p.active.Add(1)
n := len(p.idle)
if n != 0 {
state := p.idle[n-1]
p.idle = p.idle[:n-1]
p.mu.Unlock()
return state, nil
}
p.mu.Unlock()
// TODO: Limit the total number of states. When the limit is reached, wait
// for a Release instead of creating another state; allow the wait to be
// cancelled by the caller or by Close.
state, err := p.factory(p.ctx)
if err == nil && state == nil {
err = errors.New("Lua state factory returned nil")
}
if err != nil {
if state != nil {
state.Close()
}
p.active.Done()
return nil, err
}
if err := p.ctx.Err(); err != nil {
state.Close()
p.active.Done()
return nil, err
}
return state, nil
}
// Release returns a healthy state to the pool and closes a failed or cancelled one.
func (p *Pool) Release(state *glua.LState, reusable bool) {
if reusable {
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 active work, 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()
}
+197
View File
@@ -0,0 +1,197 @@
package lua
import (
"context"
"errors"
"testing"
"time"
glua "github.com/yuin/gopher-lua"
)
func TestPoolFactoryFailureClosesReturnedState(t *testing.T) {
failure := errors.New("factory failed")
state := glua.NewState()
_, err := NewPool(context.Background(), func(context.Context) (*glua.LState, error) {
return state, failure
})
if !errors.Is(err, failure) || !state.IsClosed() {
t.Fatalf("NewPool error = %v, state closed = %t", err, state.IsClosed())
}
var failedState *glua.LState
calls := 0
pool, err := NewPool(context.Background(), func(context.Context) (*glua.LState, error) {
calls++
if calls == 1 {
return glua.NewState(), nil
}
failedState = glua.NewState()
return failedState, failure
})
if err != nil {
t.Fatal(err)
}
defer pool.Close()
borrowed, err := pool.Acquire()
if err != nil {
t.Fatal(err)
}
defer pool.Release(borrowed, true)
_, err = pool.Acquire()
if !errors.Is(err, failure) || !failedState.IsClosed() {
t.Fatalf("Acquire error = %v, state closed = %t", err, failedState.IsClosed())
}
}
func TestPoolReusesStatesAndLimitsIdle(t *testing.T) {
created := 0
pool, err := NewPool(context.Background(), func(context.Context) (*glua.LState, error) {
created++
return glua.NewState(), nil
})
if err != nil {
t.Fatal(err)
}
defer pool.Close()
states := make([]*glua.LState, maxIdleStates+3)
for i := range states {
states[i], err = pool.Acquire()
if err != nil {
t.Fatal(err)
}
}
for _, state := range states {
pool.Release(state, true)
}
for i, state := range states {
if got, want := state.IsClosed(), i >= maxIdleStates; got != want {
t.Fatalf("state %d closed = %t, want %t", i, got, want)
}
}
borrowed, err := pool.Acquire()
if err != nil {
t.Fatal(err)
}
if created != len(states) {
t.Fatalf("Acquire created %d states, want %d", created, len(states))
}
pool.Release(borrowed, true)
}
func TestPoolCloseCancelsAndWaitsForBorrowedState(t *testing.T) {
pool, err := NewPool(context.Background(), func(context.Context) (*glua.LState, error) {
return glua.NewState(), nil
})
if err != nil {
t.Fatal(err)
}
state, err := pool.Acquire()
if err != nil {
t.Fatal(err)
}
done := make(chan struct{})
go func() {
pool.Close()
close(done)
}()
select {
case <-pool.Context().Done():
case <-time.After(time.Second):
pool.Release(state, false)
t.Fatal("Close did not cancel the pool context")
}
select {
case <-done:
pool.Release(state, false)
t.Fatal("Close returned while a state was borrowed")
default:
}
pool.Release(state, true)
select {
case <-done:
case <-time.After(time.Second):
t.Fatal("Close did not finish after Release")
}
if !state.IsClosed() {
t.Fatal("borrowed state was not closed")
}
if _, err := pool.Acquire(); !errors.Is(err, context.Canceled) {
t.Fatalf("Acquire after Close = %v, want context.Canceled", err)
}
pool.Close()
}
func TestPoolCloseCancelsStateCreation(t *testing.T) {
started := make(chan struct{})
first := true
pool, err := NewPool(context.Background(), func(ctx context.Context) (*glua.LState, error) {
if first {
first = false
return glua.NewState(), nil
}
close(started)
<-ctx.Done()
return nil, ctx.Err()
})
if err != nil {
t.Fatal(err)
}
borrowed, err := pool.Acquire()
if err != nil {
t.Fatal(err)
}
acquireDone := make(chan error, 1)
go func() {
_, err := pool.Acquire()
acquireDone <- err
}()
<-started
closeDone := make(chan struct{})
go func() {
pool.Close()
close(closeDone)
}()
select {
case err := <-acquireDone:
if !errors.Is(err, context.Canceled) {
t.Fatalf("Acquire during Close = %v, want context.Canceled", err)
}
case <-time.After(time.Second):
t.Fatal("state creation did not stop after Close")
}
select {
case <-closeDone:
pool.Release(borrowed, false)
t.Fatal("Close returned while the initial state was borrowed")
default:
}
pool.Release(borrowed, true)
select {
case <-closeDone:
case <-time.After(time.Second):
t.Fatal("Close did not finish after Release")
}
}
func BenchmarkPoolAcquireRelease(b *testing.B) {
pool, err := NewPool(context.Background(), func(context.Context) (*glua.LState, error) {
return glua.NewState(), nil
})
if err != nil {
b.Fatal(err)
}
defer pool.Close()
b.ReportAllocs()
b.ResetTimer()
for i := 0; i < b.N; i++ {
state, err := pool.Acquire()
if err != nil {
b.Fatal(err)
}
pool.Release(state, true)
}
b.StopTimer()
}
+56
View File
@@ -0,0 +1,56 @@
// Package lua provides shared GopherLua programs and state management for Xray scripts.
package lua
import (
"bufio"
"context"
"os"
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
}
// 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 VM, makes modules available, and executes the file top level.
// Module loaders run only when Lua calls require. Each state gets its own globals.
// The caller owns the returned state.
func (p *Program) NewState(ctx context.Context, register func(*glua.LState)) (*glua.LState, error) {
L := glua.NewState()
if register != nil {
register(L)
}
L.SetContext(ctx)
L.Push(L.NewFunctionFromProto(p.proto))
err := L.PCall(0, 0, nil)
L.RemoveContext()
if err == nil {
err = ctx.Err()
}
if err != nil {
L.Close()
return nil, err
}
return L, nil
}
+55
View File
@@ -0,0 +1,55 @@
package lua
import (
"context"
"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)
if err != nil {
t.Fatal(err)
}
defer first.Close()
first.SetGlobal("value", glua.LNumber(42))
second, err := program.NewState(context.Background(), 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)
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)
}
}
+48
View File
@@ -1,6 +1,8 @@
package platform // import "github.com/xtls/xray-core/common/platform"
import (
"errors"
"fmt"
"os"
"path/filepath"
"strconv"
@@ -90,3 +92,49 @@ 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,6 +1,7 @@
package platform_test
import (
"errors"
"os"
"path/filepath"
"runtime"
@@ -64,3 +65,53 @@ func TestGetAssetLocation(t *testing.T) {
}
}
}
func TestResolveLuaFile(t *testing.T) {
workingDir := t.TempDir()
t.Chdir(workingDir)
executable, err := os.Executable()
common.Must(err)
file, err := os.CreateTemp(filepath.Dir(executable), "lua-*.lua")
common.Must(err)
common.Must(file.Close())
defer os.Remove(file.Name())
name := filepath.Base(file.Name())
paths := []string{
filepath.Join(t.TempDir(), name),
filepath.Join(t.TempDir(), name),
filepath.Join(workingDir, name),
file.Name(),
}
t.Setenv(ConfdirLocation, filepath.Dir(paths[0]))
t.Setenv(ConfigLocation, filepath.Dir(paths[1]))
for _, path := range paths[:3] {
common.Must(os.WriteFile(path, nil, 0o600))
}
if got, err := ResolveLuaFile(paths[2]); err != nil || got != paths[2] {
t.Fatalf("absolute path = %q, %v; want %q", got, err, paths[2])
}
for i, want := range paths {
if i == 2 {
t.Setenv(ConfdirLocation, "")
t.Setenv(ConfigLocation, "")
}
if got, err := ResolveLuaFile(name); err != nil || got != want {
t.Fatalf("resolved path = %q, %v; want %q", got, err, want)
}
common.Must(os.Remove(want))
}
if _, err := ResolveLuaFile(name); !errors.Is(err, os.ErrNotExist) {
t.Fatalf("missing file error = %v", err)
}
t.Setenv(ConfdirLocation, filepath.Dir(paths[0]))
t.Setenv(ConfigLocation, filepath.Dir(paths[1]))
common.Must(os.Mkdir(paths[0], 0o700))
common.Must(os.WriteFile(paths[1], nil, 0o600))
for _, path := range []string{"", name, filepath.Join(t.TempDir(), name)} {
if _, err := ResolveLuaFile(path); err == nil {
t.Fatalf("accepted invalid path %q", path)
}
}
}
+10 -8
View File
@@ -23,19 +23,21 @@ 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.55.0
golang.org/x/crypto v0.57.0
golang.org/x/exp v0.0.0-20240506185415-9bf2ced13842
golang.org/x/net v0.58.0
golang.org/x/sync v0.22.0
golang.org/x/sys v0.47.0
golang.org/x/net v0.59.0
golang.org/x/sync v0.23.0
golang.org/x/sys v0.48.0
golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2
golang.zx2c4.com/wireguard v0.0.0-20250521234502-f333402bd9cb
golang.zx2c4.com/wireguard/windows v1.0.1
google.golang.org/grpc v1.83.2
golang.zx2c4.com/wireguard/windows v1.1.1
google.golang.org/grpc v1.84.0
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,9 +59,9 @@ 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/text v0.41.0 // 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
google.golang.org/genproto/googleapis/rpc v0.0.0-20260526163538-3dc84a4a5aaa // indirect
google.golang.org/genproto/googleapis/rpc v0.0.0-20260706201446-f0a921348800 // indirect
gopkg.in/yaml.v2 v2.4.0 // indirect
)
+25 -34
View File
@@ -2,16 +2,13 @@ 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/cespare/xxhash/v2 v2.3.0 h1:UL815xU9SqsFlibzuggzjXhog7bL6oX9BbNZnL2UFvs=
github.com/cespare/xxhash/v2 v2.3.0/go.mod h1:VGX0DQ3Q6kWi7AoAeZDth3/j3BFtOZR5XLFGgcrjCOs=
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=
github.com/ghodss/yaml v1.0.1-0.20220118164431-d8423dcdf344/go.mod h1:GIjDIg/heH5DOkXY3YJ/wNhfHsQHoXGjl8G8amsYQ1I=
github.com/go-logr/logr v1.4.3 h1:CjnDlHq8ikf6E492q6eKboGOC0T8CDaOvkHCIg8idEI=
github.com/go-logr/logr v1.4.3/go.mod h1:9T104GzyrTigFIr8wt5mBrctHMim0Nb2HLGrmQ40KvY=
github.com/go-logr/stdr v1.2.2 h1:hSWxHoqTgW2S2qGc0LTAI563KZ5YKYRhT3MFKZMbjag=
github.com/go-logr/stdr v1.2.2/go.mod h1:mMo/vtBO5dYbehREoey6XUKy/eSumjCCveDpRre4VKE=
github.com/go-quicktest/qt v1.102.0 h1:HSQxCeh5YZH3EL3W39ixjtyaEhcWSXQHtHnMBzSs474=
github.com/go-quicktest/qt v1.102.0/go.mod h1:p4lGIVX+8Wa6ZPNDvqcxq36XpUDLh42FLetFU7odllI=
github.com/golang/mock v1.7.0-rc.1 h1:YojYx61/OLFsiv6Rw1Z96LpldJIy31o+UHmwAUMJ6/U=
@@ -91,18 +88,9 @@ github.com/wlynxg/anet v0.0.5/go.mod h1:eay5PRQr7fIVAMbTbchTnO9gG65Hg/uYGdc7mguH
github.com/xtls/reality v0.0.0-20260908062103-8cdf7bf9c7f0 h1:rb+fKQFhz+5I2PPuQsNYxI5mUU840XWYtRF0ZBjvkws=
github.com/xtls/reality v0.0.0-20260908062103-8cdf7bf9c7f0/go.mod h1:DsJblcWDGt76+FVqBVwbwRhxyyNJsGV48gJLch0OOWI=
github.com/yuin/goldmark v1.4.1/go.mod h1:mwnBkeHKe2W/ZEtQ+71ViKU8L12m81fl3OWwC1Zlc8k=
go.opentelemetry.io/auto/sdk v1.2.1 h1:jXsnJ4Lmnqd11kwkBV2LgLoFMZKizbCi5fNZ/ipaZ64=
go.opentelemetry.io/auto/sdk v1.2.1/go.mod h1:KRTj+aOaElaLi+wW1kO/DZRXwkF4C5xPbEe3ZiIhN7Y=
go.opentelemetry.io/otel v1.44.0 h1:JjwHmHpA4iZ3wBxluu2fbbE7j4kqlE8jXyAyPXH7HqU=
go.opentelemetry.io/otel v1.44.0/go.mod h1:BMgjTHL9WPRlRjL2oZCBTL4whCGtXch2H4BhOPIAyYc=
go.opentelemetry.io/otel/metric v1.44.0 h1:1w0gILTcHdr3YI+ixLyjemwrVnsMURbTZFrSYCdDdmc=
go.opentelemetry.io/otel/metric v1.44.0/go.mod h1:8O7hanEPBNgEMmybD3s2VBKcgWOCsA6tzHBPODAiquo=
go.opentelemetry.io/otel/sdk v1.44.0 h1:nHYwb9lK+fJPU/dnT6s7W7Z8itMWyqrnVfbheVYrZ58=
go.opentelemetry.io/otel/sdk v1.44.0/go.mod h1:Osuydd3Se74nqjAKxid74N5eC+jfEqfTegHRnq58oK0=
go.opentelemetry.io/otel/sdk/metric v1.44.0 h1:3LlKgI+VjbVsjNRFZJZAJ30WjXC5VkNRks6si09iEfI=
go.opentelemetry.io/otel/sdk/metric v1.44.0/go.mod h1:5B5pMARnXxKhltooO4xUuCBorl65a4EpnTalObqOigA=
go.opentelemetry.io/otel/trace v1.44.0 h1:jxF5CsGYCe74MCRx2X4g7WsY/VBKRqqpNvXlX/6gtIk=
go.opentelemetry.io/otel/trace v1.44.0/go.mod h1:oLl1jrMQAVo6v3GAggN+1VH9VIz9iUSvW53sW1Q8PIE=
github.com/yuin/gopher-lua v0.0.0-20190206043414-8bfc7677f583/go.mod h1:gqRgreBUhTSL0GeU64rtZ3Uq3wtjOa/TB2YfrtkCbVQ=
github.com/yuin/gopher-lua v1.1.2 h1:yF/FjE3hD65tBbt0VXLE13HWS9h34fdzJmrWRXwobGA=
github.com/yuin/gopher-lua v1.1.2/go.mod h1:7aRmXIWl37SqRf0koeyylBEzJ+aPt8A+mmkQ4f1ntR8=
go.uber.org/mock v0.5.2 h1:LbtPTcP8A5k9WPXj54PPPbjcI4Y6lhyOZXn+VS7wNko=
go.uber.org/mock v0.5.2/go.mod h1:wLlUxC2vVTPTaE3UD51E0BGOAElKrILxhVSDYQLld5o=
go.yaml.in/yaml/v3 v3.0.5 h1:N6y/pJk8buWs9NY5ERU2HSMfm+IuD/OtfdAnq6kESPw=
@@ -111,8 +99,8 @@ go4.org/netipx v0.0.0-20231129151722-fdeea329fbba h1:0b9z3AuHCjxk0x/opv64kcgZLBs
go4.org/netipx v0.0.0-20231129151722-fdeea329fbba/go.mod h1:PLyyIXexvUFg3Owu6p/WfdlivPbZJsZdgWZlrGope/Y=
golang.org/x/crypto v0.0.0-20190308221718-c2843e01d9a2/go.mod h1:djNgcEr1/C05ACkg1iLfiJU5Ep61QUkGW8qpdssI0+w=
golang.org/x/crypto v0.0.0-20191011191535-87dc89f01550/go.mod h1:yigFU9vqHzYiE8UmvKecakEJjdnWj3jj499lnFckfCI=
golang.org/x/crypto v0.55.0 h1:+KWHjbgOaAQ66dh/YlkZKHlz9ZUlq61AFirAR9ntP8M=
golang.org/x/crypto v0.55.0/go.mod h1:uq0V9dE/fzQuJtbnL+2EhWOE63vo164FY8xqEnV9xis=
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/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=
@@ -121,12 +109,13 @@ golang.org/x/mod v0.5.1/go.mod h1:5OXOZSfqPIIbmVBIIKWRFfZjPR0E5r58TLhUjH0a2Ro=
golang.org/x/net v0.0.0-20190404232315-eb5bcb51f2a3/go.mod h1:t9HGtf8HONx5eT2rtn7q6eTqICYqUVnKs3thJo3Qplg=
golang.org/x/net v0.0.0-20190620200207-3b0461eec859/go.mod h1:z5CRVTTTmAJ677TzLLGU+0bjPO0LkuOLi4/5GtJWs/s=
golang.org/x/net v0.0.0-20211015210444-4f30a5c0130f/go.mod h1:9nx3DQGgdP8bBQD5qxJ1jj9UTztislL4KSBs9R2vV5Y=
golang.org/x/net v0.58.0 h1:ynWG7rqYi4ccpTEuPZ2QGWHktVEM9DMCj9yzDE0Q7To=
golang.org/x/net v0.58.0/go.mod h1:YwCddHnFlT7eLQqVprV19OnhLGtc5xOKgE0RyqgfWAU=
golang.org/x/net v0.59.0 h1:5zfYln+w5XCxwrnMMJPufRgNoXEaGxl0wo5GqPXyues=
golang.org/x/net v0.59.0/go.mod h1:2DA/G1UfVbCpQPeWTmMPGY7Cs2PkBkwu743bVX5PIVg=
golang.org/x/sync v0.0.0-20190423024810-112230192c58/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
golang.org/x/sync v0.0.0-20210220032951-036812b2e83c/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
golang.org/x/sync v0.22.0 h1:SZjpbeLmrCk4xhRSZFNZW5gFUeCeFgjekvI/+gfScek=
golang.org/x/sync v0.22.0/go.mod h1:9xrNwdLfx4jkKbNva9FpL6vEN7evnE43NNNJQ2LF3+0=
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=
@@ -134,14 +123,14 @@ golang.org/x/sys v0.0.0-20210423082822-04245dca01da/go.mod h1:h1NjWce9XRLGQEsW7w
golang.org/x/sys v0.0.0-20211019181941-9d821ace8654/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
golang.org/x/sys v0.2.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
golang.org/x/sys v0.10.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
golang.org/x/sys v0.47.0 h1:o7XGOvZQCADBQQ4Y7VNq2dRWQR7JmOUW8Kxx4ZsNgWs=
golang.org/x/sys v0.47.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw=
golang.org/x/sys v0.48.0 h1:bbX/i/6MgT9BVLM9RT1thmxL04yeTAhbEz4SyadbXoo=
golang.org/x/sys v0.48.0/go.mod h1:hNLxWAXmnKAxqDtdwIYC4bM9oQPEecfsnNMuSxOs3og=
golang.org/x/term v0.0.0-20201126162022-7de9c90e9dd1/go.mod h1:bj7SfCRtBDWHUb9snDiAeCFNEtKQo2Wmx5Cou7ajbmo=
golang.org/x/text v0.3.0/go.mod h1:NqM8EUOU14njkJ3fqMW+pc6Ldnwhi/IjpwHt7yyuwOQ=
golang.org/x/text v0.3.6/go.mod h1:5Zoc/QRtKVWzQhOtBMvqHzDpF6irO9z98xDceosuGiQ=
golang.org/x/text v0.3.7/go.mod h1:u+2+/6zg+i71rQMx5EYifcz6MCKuco9NR6JIITiCfzQ=
golang.org/x/text v0.41.0 h1:vz/seA0lnX87Othu2f/0L24RcgrXD9/YFTSuGjj3rH8=
golang.org/x/text v0.41.0/go.mod h1:jvf1O8ajNzZqhSrQBPbutR/EB83Cc0CFrezNQIwbb5M=
golang.org/x/text v0.42.0 h1:JbOZXgfeCPU9gacVtYliJqOhD+zhrEqK4LfdpmlUZqI=
golang.org/x/text v0.42.0/go.mod h1:ojzP1Z+2QtioaF8DTtO8K5q7JWVVYwZKenzujK0Zd0E=
golang.org/x/time v0.14.0 h1:MRx4UaLrDotUKUdCIqzPC48t1Y9hANFKIRpNx+Te8PI=
golang.org/x/time v0.14.0/go.mod h1:eL/Oa2bBBK0TkX57Fyni+NgnyQQN4LitPmob2Hjnqw4=
golang.org/x/tools v0.0.0-20180917221912-90fa682c2a6e/go.mod h1:n7NCudcB/nEzxVGmLbDWY5pfWTLqBcC2KZ6jyYvM4mQ=
@@ -157,14 +146,14 @@ golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2 h1:B82qJJgjvYKsXS9jeu
golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2/go.mod h1:deeaetjYA+DHMHg+sMSMI58GrEteJUUzzw7en6TJQcI=
golang.zx2c4.com/wireguard v0.0.0-20250521234502-f333402bd9cb h1:whnFRlWMcXI9d+ZbWg+4sHnLp52d5yiIPUxMBSt4X9A=
golang.zx2c4.com/wireguard v0.0.0-20250521234502-f333402bd9cb/go.mod h1:rpwXGsirqLqN2L0JDJQlwOboGHmptD5ZD6T2VmcqhTw=
golang.zx2c4.com/wireguard/windows v1.0.1 h1:eOxiDVbywPC+ZQqvdCK7x+ZwWXKbYv50TtH8ysFIbw8=
golang.zx2c4.com/wireguard/windows v1.0.1/go.mod h1:+fbT3FFdX4zzYDLwJh5+HPEcNN/3HyNdzhNSVsQM+zs=
golang.zx2c4.com/wireguard/windows v1.1.1 h1:8/H97U1v1PNDNcBsMZgU3KFuND9MQdTsU2NOwmCXArE=
golang.zx2c4.com/wireguard/windows v1.1.1/go.mod h1:+fbT3FFdX4zzYDLwJh5+HPEcNN/3HyNdzhNSVsQM+zs=
gonum.org/v1/gonum v0.17.0 h1:VbpOemQlsSMrYmn7T2OUvQ4dqxQXU+ouZFQsZOx50z4=
gonum.org/v1/gonum v0.17.0/go.mod h1:El3tOrEuMpv2UdMrbNlKEh9vd86bmQ6vqIcDwxEOc1E=
google.golang.org/genproto/googleapis/rpc v0.0.0-20260526163538-3dc84a4a5aaa h1:mZHHdPZl0dbGHCflZgAq/Q468DWVFcU2whhB2KAo8fk=
google.golang.org/genproto/googleapis/rpc v0.0.0-20260526163538-3dc84a4a5aaa/go.mod h1:4Hqkh8ycfw05ld/3BWL7rJOSfebL2Q+DVDeRgYgxUU8=
google.golang.org/grpc v1.83.2 h1:EManeRomTObA0BU7I8vXgg/78uE5MJ9M8B39EX2WscU=
google.golang.org/grpc v1.83.2/go.mod h1:YPI1hK3kDked6iHvgX3tR0y+nX/qpMFKhPgFsokw1S8=
google.golang.org/genproto/googleapis/rpc v0.0.0-20260706201446-f0a921348800 h1:qEHAMpSaUhtD0p3NbEEI83HwNGFxEwaSJ1G9PLnCBZE=
google.golang.org/genproto/googleapis/rpc v0.0.0-20260706201446-f0a921348800/go.mod h1:4Hqkh8ycfw05ld/3BWL7rJOSfebL2Q+DVDeRgYgxUU8=
google.golang.org/grpc v1.84.0 h1:soMyaPJ8pAak5PIQ0DGBUir0XRo2fRoMqhNWMLlLxO0=
google.golang.org/grpc v1.84.0/go.mod h1:ljCht0DrxQrXBDRTZp52Qxh3Ffk8CdYm2sj4O2QN2C0=
google.golang.org/protobuf v1.36.12 h1:pJOKDDOyeXErUroCihFAd5LQuwXBSpVnKGrj5o/fwxc=
google.golang.org/protobuf v1.36.12/go.mod h1:HTf+CrKn2C3g5S8VImy6tdcUvCska2kB7j23XfzDpco=
gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0=
@@ -177,6 +166,8 @@ 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,9 +14,11 @@ 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"`
@@ -43,6 +45,7 @@ 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"`
@@ -60,6 +63,7 @@ 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
@@ -134,6 +138,7 @@ func (c *NameServerConfig) Build() (*dns.NameServer, error) {
}
return &dns.NameServer{
Id: c.ID,
Address: &net.Endpoint{
Network: net.Network_UDP,
Address: c.Address.Build(),
@@ -159,6 +164,7 @@ 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"`
@@ -278,6 +284,14 @@ 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,6 +2,8 @@ package conf_test
import (
"encoding/json"
"os"
"path/filepath"
"testing"
"github.com/google/go-cmp/cmp"
@@ -122,3 +124,51 @@ func TestDNSConfigParsing(t *testing.T) {
}
}
}
func TestDNSScriptConfig(t *testing.T) {
dir := t.TempDir()
t.Setenv("xray.location.confdir", dir)
path := filepath.Join(dir, "lookup.lua")
if err := os.WriteFile(path, []byte("function HandleDNSQuery(domain, ipv4, ipv6, fake) end"), 0o600); err != nil {
t.Fatal(err)
}
for _, tc := range []struct {
name string
script string
wantError bool
}{
{"relative", "lookup.lua", false},
{"absolute", path, false},
{"missing", "missing.lua", true},
{"directory", dir, true},
} {
t.Run(tc.name, func(t *testing.T) {
built, err := (&DNSConfig{Script: tc.script}).Build()
if tc.wantError {
if err == nil {
t.Fatal("Build accepted an invalid script path")
}
return
}
if err != nil {
t.Fatal(err)
}
if built.Script != path {
t.Fatalf("script path = %q, want %q", built.Script, path)
}
})
}
var parsed DNSConfig
if err := json.Unmarshal([]byte(`{"servers":[{"id":"primary","address":"1.1.1.1"}]}`), &parsed); err != nil {
t.Fatal(err)
}
built, err := parsed.Build()
if err != nil {
t.Fatal(err)
}
if len(built.NameServer) != 1 || built.NameServer[0].Id != "primary" {
t.Fatalf("nameserver IDs = %v, want primary", built.NameServer)
}
}
+37
View File
@@ -0,0 +1,37 @@
package conf
import (
"net/netip"
"github.com/xtls/xray-core/common/errors"
"github.com/xtls/xray-core/common/protocol"
"github.com/xtls/xray-core/proxy/masque"
"google.golang.org/protobuf/proto"
)
type MasqueClientConfig struct {
Address *Address `json:"address"`
Port uint16 `json:"port"`
RemoteDNS []string `json:"remoteDNS"`
}
func (c *MasqueClientConfig) Build() (proto.Message, error) {
if c.Address == nil {
return nil, errors.New(`MASQUE: "address" is not set`)
}
if c.Port == 0 {
return nil, errors.New(`MASQUE: "port" is not set`)
}
for _, s := range c.RemoteDNS {
if _, err := netip.ParseAddr(s); err != nil {
return nil, errors.New(`MASQUE: invalid "remoteDNS" `, s).Base(err)
}
}
return &masque.ClientConfig{
Server: &protocol.ServerEndpoint{
Address: c.Address.Build(),
Port: uint32(c.Port),
},
RemoteDns: c.RemoteDNS,
}, nil
}
+85
View File
@@ -0,0 +1,85 @@
package conf_test
import (
"encoding/json"
"testing"
. "github.com/xtls/xray-core/infra/conf"
"github.com/xtls/xray-core/transport/internet/masque"
)
func TestMasqueConfig(t *testing.T) {
creator := func() Buildable {
return new(MasqueConfig)
}
runMultiTestCase(t, []TestCase{
{
Input: `{}`,
Parser: loadJSON(creator),
Output: &masque.Config{Path: "/.well-known/masque/ip/*/*/"},
},
{
Input: `{
"host": "example.com:8443",
"path": "/.well-known/masque/ip/{target}/{ipproto}/",
"headers": {"Authorization": "Basic dTpw"}
}`,
Parser: loadJSON(creator),
Output: &masque.Config{
Host: "example.com:8443",
Path: "/.well-known/masque/ip/*/*/",
Headers: map[string]string{"Authorization": "Basic dTpw"},
},
},
{
Input: `{"path": "/masque/ip{?target,ipproto}"}`,
Parser: loadJSON(creator),
Output: &masque.Config{Path: "/masque/ip?target=*&ipproto=*"},
},
})
for _, input := range []string{
`{"path": "/masque/{target}/{ipproto}/{dns}"}`,
`{"path": "masque"}`,
`{"host": "example.com/path"}`,
`{"headers": {"host": "example.com"}}`,
`{"headers": {"Capsule-Protocol": "?0"}}`,
`{"headers": {"X Token": "a"}}`,
`{"headers": {"X-Token": "a\r\nb"}}`,
} {
if _, err := loadJSON(creator)(input); err == nil {
t.Errorf("expected an error for %s", input)
}
}
}
func TestMasqueOutboundConfig(t *testing.T) {
build := func(s string) error {
c := new(OutboundDetourConfig)
if err := json.Unmarshal([]byte(s), c); err != nil {
return err
}
_, err := c.Build()
return err
}
if err := build(`{
"protocol": "masque",
"settings": {"address": "example.com", "port": 443},
"streamSettings": {"network": "masque", "security": "tls"},
"mux": {"enabled": false, "concurrency": -1}
}`); err != nil {
t.Error(err)
}
for _, input := range []string{
`{"protocol": "masque", "settings": {"address": "example.com"}, "streamSettings": {"network": "masque", "security": "tls"}}`,
`{"protocol": "masque", "settings": {"address": "example.com", "port": 443}, "streamSettings": {"network": "masque", "security": "tls"}, "mux": {"enabled": true}}`,
`{"protocol": "masque", "settings": {"address": "example.com", "port": 443}, "streamSettings": {"network": "masque", "security": "tls"}, "mux": {"enabled": true, "concurrency": -1}}`,
`{"protocol": "freedom", "streamSettings": {"network": "masque", "security": "tls"}}`,
} {
if err := build(input); err == nil {
t.Errorf("expected an error for %s", input)
}
}
}
+11
View File
@@ -7,6 +7,7 @@ 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"
@@ -72,6 +73,7 @@ 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 {
@@ -92,6 +94,15 @@ 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,6 +2,8 @@ package conf_test
import (
"encoding/json"
"os"
"path/filepath"
"testing"
"time"
_ "unsafe"
@@ -236,3 +238,39 @@ func TestRouterConfig(t *testing.T) {
},
})
}
func TestRouterScriptConfig(t *testing.T) {
dir := t.TempDir()
t.Setenv("xray.location.confdir", dir)
path := filepath.Join(dir, "route.lua")
if err := os.WriteFile(path, []byte("function HandleRoute() end"), 0o600); err != nil {
t.Fatal(err)
}
for _, tc := range []struct {
name string
script string
wantError bool
}{
{"relative", "route.lua", false},
{"absolute", path, false},
{"missing", "missing.lua", true},
{"directory", dir, true},
} {
t.Run(tc.name, func(t *testing.T) {
built, err := (&RouterConfig{Script: tc.script}).Build()
if tc.wantError {
if err == nil {
t.Fatal("Build accepted invalid script path")
}
return
}
if err != nil {
t.Fatal(err)
}
if built.Script != path {
t.Fatalf("script path = %q, want %q", built.Script, path)
}
})
}
}
+5 -16
View File
@@ -14,7 +14,6 @@ import (
googleuuid "github.com/google/uuid"
"github.com/xtls/xray-core/common/errors"
"github.com/xtls/xray-core/common/net"
"github.com/xtls/xray-core/transport/internet"
"github.com/xtls/xray-core/transport/internet/finalmask/fragment"
"github.com/xtls/xray-core/transport/internet/finalmask/header/custom"
"github.com/xtls/xray-core/transport/internet/finalmask/mkcp/aes128gcm"
@@ -909,22 +908,13 @@ func (c *Realm) Build() (proto.Message, error) {
}
type UDPHop struct {
Sockopt *SocketConfig `json:"sockopt"`
Mode string `json:"mode"`
Interval Int32Range `json:"interval"`
RemotePorts PortList `json:"remotePorts"`
RemoteIPs []string `json:"remoteIPs"`
Mode string `json:"mode"`
Interval Int32Range `json:"interval"`
RemoteIPs []string `json:"remoteIPs"`
RemotePorts PortList `json:"remotePorts"`
}
func (c *UDPHop) Build() (proto.Message, error) {
var sockopt *internet.SocketConfig
if c.Sockopt != nil {
var err error
sockopt, err = c.Sockopt.Build()
if err != nil {
return nil, err
}
}
var local, remote, remoteOnce bool
for _, mode := range strings.Split(c.Mode, ",") {
switch strings.ToLower(mode) {
@@ -953,14 +943,13 @@ func (c *UDPHop) Build() (proto.Message, error) {
return nil, errors.New("invalid ip ", ip)
}
return &udphop.Config{
Sockopt: sockopt,
Local: local,
Remote: remote,
RemoteOnce: remoteOnce,
IntervalMin: int64(c.Interval.From),
IntervalMax: int64(c.Interval.To),
RemotePorts: c.RemotePorts.Build().Ports(),
RemoteIPs: remoteIPs,
RemotePorts: c.RemotePorts.Build().Ports(),
}, nil
}
+26
View File
@@ -36,6 +36,10 @@ func (p TransportProtocol) Build() (string, error) {
return "", errors.PrintRemovedFeatureError("QUIC transport (without web service, etc.)", "XHTTP stream-one H3")
case "hysteria":
return "hysteria", nil
case "masque":
return "masque", nil
case "xdrive":
return "xdrive", nil
default:
return "", errors.New("Config: unknown transport protocol: ", p)
}
@@ -59,6 +63,8 @@ type StreamConfig struct {
WSSettings *WebSocketConfig `json:"wsSettings"`
HTTPUPGRADESettings *HttpUpgradeConfig `json:"httpupgradeSettings"`
HysteriaSettings *HysteriaConfig `json:"hysteriaSettings"`
MASQUESettings *MasqueConfig `json:"masqueSettings"`
XDRIVESettings *XDriveConfig `json:"xdriveSettings"`
SocketSettings *SocketConfig `json:"sockopt"`
}
@@ -192,6 +198,26 @@ func (c *StreamConfig) Build() (*internet.StreamConfig, error) {
Settings: serial.ToTypedMessage(hs),
})
}
if c.MASQUESettings != nil {
ms, err := c.MASQUESettings.Build()
if err != nil {
return nil, errors.New("Failed to build MASQUE config.").Base(err)
}
config.TransportSettings = append(config.TransportSettings, &internet.TransportConfig{
ProtocolName: "masque",
Settings: serial.ToTypedMessage(ms),
})
}
if c.XDRIVESettings != nil {
xs, err := c.XDRIVESettings.Build()
if err != nil {
return nil, errors.New("Failed to build XDRIVE config.").Base(err)
}
config.TransportSettings = append(config.TransportSettings, &internet.TransportConfig{
ProtocolName: "xdrive",
Settings: serial.ToTypedMessage(xs),
})
}
if c.SocketSettings != nil {
ss, err := c.SocketSettings.Build()
if err != nil {
+90
View File
@@ -20,9 +20,12 @@ import (
"github.com/xtls/xray-core/transport/internet/httpupgrade"
"github.com/xtls/xray-core/transport/internet/hysteria"
"github.com/xtls/xray-core/transport/internet/kcp"
"github.com/xtls/xray-core/transport/internet/masque"
"github.com/xtls/xray-core/transport/internet/splithttp"
"github.com/xtls/xray-core/transport/internet/tcp"
"github.com/xtls/xray-core/transport/internet/websocket"
"github.com/xtls/xray-core/transport/internet/xdrive"
"golang.org/x/net/http/httpguts"
"google.golang.org/protobuf/proto"
)
@@ -785,6 +788,46 @@ func (c *HysteriaConfig) Build() (proto.Message, error) {
return config, nil
}
type MasqueConfig struct {
Host string `json:"host"`
Path string `json:"path"`
Headers map[string]string `json:"headers"`
}
func (c *MasqueConfig) Build() (proto.Message, error) {
path := c.Path
if path == "" {
path = masque.DefaultPath
}
path = strings.NewReplacer(
"{target}", "*", "{ipproto}", "*",
"{?target,ipproto}", "?target=*&ipproto=*", "{?ipproto,target}", "?ipproto=*&target=*",
"{&target,ipproto}", "&target=*&ipproto=*", "{&ipproto,target}", "&ipproto=*&target=*",
).Replace(path)
if !strings.HasPrefix(path, "/") || strings.ContainsAny(path, "{}") {
return nil, errors.New(`invalid "path": `, path, `, only the variables {target} and {ipproto} are supported`)
}
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 {
if !httpguts.ValidHeaderFieldName(k) || !httpguts.ValidHeaderFieldValue(v) {
return nil, errors.New(`invalid header in "headers": `, strconv.Quote(k))
}
switch strings.ToLower(k) {
case "host", "capsule-protocol":
return nil, errors.New(`"headers" can't contain "`, k, `"`)
}
}
return &masque.Config{
Host: c.Host,
Path: path,
Headers: c.Headers,
}, nil
}
func readFileOrString(f string, s []string) ([]byte, error) {
if len(f) > 0 {
return filesystem.ReadCert(f)
@@ -794,3 +837,50 @@ func readFileOrString(f string, s []string) ([]byte, error) {
}
return nil, errors.New("both file and bytes are empty.")
}
type XDriveConfig struct {
RemoteFolder string `json:"remoteFolder"`
Service string `json:"service"`
Secrets []string `json:"secrets"`
SegmentBytes uint32 `json:"segmentBytes"`
FlushIntervalMs uint32 `json:"flushIntervalMs"`
PollIntervalMs uint32 `json:"pollIntervalMs"`
MaxPollIntervalMs uint32 `json:"maxPollIntervalMs"`
SessionTTLSeconds uint32 `json:"sessionTtlSeconds"`
Concurrency uint32 `json:"concurrency"`
EagerWindowMs uint32 `json:"eagerWindowMs"`
HoleTimeoutMs uint32 `json:"holeTimeoutMs"`
Template json.RawMessage `json:"template"`
}
// Build implements Buildable.
func (c *XDriveConfig) Build() (proto.Message, error) {
switch c.Service {
case "local":
case "Google Drive":
if len(c.Secrets) != 3 {
return nil, errors.New("Google Drive needs 3 secrets in order of ClientID, ClientSecret, RefreshToken")
}
case "template":
if len(c.Template) == 0 {
return nil, errors.New(`service "template" needs a "template" object`)
}
default:
return nil, errors.New("unsupported service")
}
config := &xdrive.Config{
RemoteFolder: c.RemoteFolder,
Service: c.Service,
Secrets: c.Secrets,
SegmentBytes: c.SegmentBytes,
FlushIntervalMs: c.FlushIntervalMs,
PollIntervalMs: c.PollIntervalMs,
MaxPollIntervalMs: c.MaxPollIntervalMs,
SessionTtlSeconds: c.SessionTTLSeconds,
Concurrency: c.Concurrency,
EagerWindowMs: c.EagerWindowMs,
HoleTimeoutMs: c.HoleTimeoutMs,
Template: string(c.Template),
}
return config, nil
}
+73
View File
@@ -291,3 +291,76 @@ func TestHeaderCustomUDPBuildRejectsExprWithoutArgs(t *testing.T) {
t.Fatalf("expected transform arg rejection, got %v", err)
}
}
func TestXDriveStreamConfig(t *testing.T) {
config := new(StreamConfig)
if err := json.Unmarshal([]byte(`{
"method": "xdrive",
"xdriveSettings": {
"remoteFolder": "/tmp/xdrive",
"service": "local"
}
}`), config); err != nil {
t.Fatalf("Unmarshal: %v", err)
}
built, err := config.Build()
if err != nil {
t.Fatalf("Build: %v", err)
}
if built.ProtocolName != "xdrive" {
t.Fatalf("ProtocolName is %q, want %q", built.ProtocolName, "xdrive")
}
if len(built.TransportSettings) != 1 || built.TransportSettings[0].ProtocolName != "xdrive" {
t.Fatalf("TransportSettings is %v, want a single xdrive entry", built.TransportSettings)
}
}
func TestXDriveRejectsUnknownService(t *testing.T) {
config := new(XDriveConfig)
if err := json.Unmarshal([]byte(`{"remoteFolder": "/tmp/xdrive", "service": "Dropbox"}`), config); err != nil {
t.Fatalf("Unmarshal: %v", err)
}
if _, err := config.Build(); err == nil {
t.Fatal("Build accepted an unsupported service")
}
}
func TestXDriveTemplateStreamConfig(t *testing.T) {
config := new(StreamConfig)
if err := json.Unmarshal([]byte(`{
"method": "xdrive",
"xdriveSettings": {
"remoteFolder": "folder",
"service": "template",
"secrets": ["user", "pass"],
"template": {
"flatten": true,
"auth": {"type": "basic", "username": "{secret0}", "password": "{secret1}"},
"put": {"method": "PUT", "url": "https://dav.example/{folder}/{name}"},
"get": {"method": "GET", "url": "https://dav.example/{folder}/{name}"},
"delete": {"method": "DELETE", "url": "https://dav.example/{folder}/{name}"},
"list": {"method": "PROPFIND", "url": "https://dav.example/{folder}/", "namesRegex": "<d:href>/folder/([^<]+)</d:href>"}
}
}
}`), config); err != nil {
t.Fatalf("Unmarshal: %v", err)
}
built, err := config.Build()
if err != nil {
t.Fatalf("Build: %v", err)
}
if built.ProtocolName != "xdrive" {
t.Fatalf("ProtocolName is %q, want xdrive", built.ProtocolName)
}
}
func TestXDriveTemplateNeedsTemplate(t *testing.T) {
config := new(XDriveConfig)
if err := json.Unmarshal([]byte(`{"remoteFolder": "f", "service": "template"}`), config); err != nil {
t.Fatalf("Unmarshal: %v", err)
}
if _, err := config.Build(); err == nil {
t.Fatal("Build accepted a template service without a template")
}
}
+7 -23
View File
@@ -59,14 +59,13 @@ func (c *WireGuardPeerConfig) Build() (*wireguard.PeerConfig, error) {
type WireGuardConfig struct {
IsClient bool `json:""`
NoKernelTun bool `json:"noKernelTun"`
SecretKey string `json:"secretKey"`
Address []string `json:"address"`
Peers []*WireGuardPeerConfig `json:"peers"`
MTU int32 `json:"mtu"`
Reserved []byte `json:"reserved"`
DomainStrategy string `json:"domainStrategy"`
DNS []string `json:"remoteDNS"`
NoKernelTun bool `json:"noKernelTun"`
SecretKey string `json:"secretKey"`
Address []string `json:"address"`
Peers []*WireGuardPeerConfig `json:"peers"`
MTU int32 `json:"mtu"`
Reserved []byte `json:"reserved"`
DNS []string `json:"remoteDNS"`
}
func (c *WireGuardConfig) Build() (proto.Message, error) {
@@ -125,21 +124,6 @@ func (c *WireGuardConfig) Build() (proto.Message, error) {
}
config.Reserved = c.Reserved
switch strings.ToLower(c.DomainStrategy) {
case "forceip", "":
config.DomainStrategy = wireguard.DeviceConfig_FORCE_IP
case "forceipv4":
config.DomainStrategy = wireguard.DeviceConfig_FORCE_IP4
case "forceipv6":
config.DomainStrategy = wireguard.DeviceConfig_FORCE_IP6
case "forceipv4v6":
config.DomainStrategy = wireguard.DeviceConfig_FORCE_IP46
case "forceipv6v4":
config.DomainStrategy = wireguard.DeviceConfig_FORCE_IP64
default:
return nil, errors.New("unsupported domain strategy: ", c.DomainStrategy)
}
config.IsClient = c.IsClient
config.NoKernelTun = c.NoKernelTun
config.DNS = c.DNS
+10
View File
@@ -16,6 +16,7 @@ import (
"github.com/xtls/xray-core/common/serial"
core "github.com/xtls/xray-core/core"
"github.com/xtls/xray-core/proxy/freedom"
"github.com/xtls/xray-core/proxy/masque"
"github.com/xtls/xray-core/transport/internet"
)
@@ -48,6 +49,7 @@ var (
"vmess": func() interface{} { return new(VMessOutboundConfig) },
"trojan": func() interface{} { return new(TrojanClientConfig) },
"hysteria": func() interface{} { return new(HysteriaClientConfig) },
"masque": func() interface{} { return new(MasqueClientConfig) },
"dns": func() interface{} { return new(DNSOutboundConfig) },
"wireguard": func() interface{} { return &WireGuardConfig{IsClient: true} },
}, "protocol", "settings")
@@ -338,6 +340,14 @@ func (c *OutboundDetourConfig) Build() (*core.OutboundHandlerConfig, error) {
return nil, err
}
if _, ok := ts.(*masque.ClientConfig); ok {
if ms := senderSettings.MultiplexSettings; ms != nil && ms.Enabled {
return nil, errors.New(`masque outbound does not support "mux"`)
}
} else if senderSettings.StreamSettings != nil && senderSettings.StreamSettings.ProtocolName == "masque" {
return nil, errors.New("the masque transport can only be used by the masque outbound")
}
if fc, ok := ts.(*freedom.Config); ok {
if senderSettings.StreamSettings != nil &&
senderSettings.StreamSettings.SocketSettings != nil &&
+3
View File
@@ -41,6 +41,7 @@ import (
_ "github.com/xtls/xray-core/proxy/freedom"
_ "github.com/xtls/xray-core/proxy/http"
_ "github.com/xtls/xray-core/proxy/loopback"
_ "github.com/xtls/xray-core/proxy/masque"
_ "github.com/xtls/xray-core/proxy/shadowsocks"
_ "github.com/xtls/xray-core/proxy/socks"
_ "github.com/xtls/xray-core/proxy/trojan"
@@ -54,12 +55,14 @@ import (
_ "github.com/xtls/xray-core/transport/internet/grpc"
_ "github.com/xtls/xray-core/transport/internet/httpupgrade"
_ "github.com/xtls/xray-core/transport/internet/kcp"
_ "github.com/xtls/xray-core/transport/internet/masque"
_ "github.com/xtls/xray-core/transport/internet/reality"
_ "github.com/xtls/xray-core/transport/internet/splithttp"
_ "github.com/xtls/xray-core/transport/internet/tcp"
_ "github.com/xtls/xray-core/transport/internet/tls"
_ "github.com/xtls/xray-core/transport/internet/udp"
_ "github.com/xtls/xray-core/transport/internet/websocket"
_ "github.com/xtls/xray-core/transport/internet/xdrive"
// Transport headers
_ "github.com/xtls/xray-core/transport/internet/headers/http"
+3 -3
View File
@@ -236,14 +236,14 @@ type UDPReader struct {
func (r *UDPReader) ReadFrom(p []byte) (n int, addr *net.Destination, err error) {
for {
var buf [hysteria.MaxDatagramFrameSize]byte
var packet [1500]byte
n, err := r.reader.Read(buf[:])
n, err := r.reader.Read(packet[:])
if err != nil {
return 0, nil, err
}
msg, err := ParseUDPMessage(buf[:n])
msg, err := ParseUDPMessage(packet[:n])
if err != nil {
continue
}
+328
View File
@@ -0,0 +1,328 @@
package masque
import (
"context"
go_errors "errors"
"io"
"net/netip"
"slices"
"sync"
"sync/atomic"
"time"
"golang.zx2c4.com/wireguard/tun"
"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/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/core"
"github.com/xtls/xray-core/features/policy"
"github.com/xtls/xray-core/proxy/wireguard"
"github.com/xtls/xray-core/transport"
"github.com/xtls/xray-core/transport/internet"
"github.com/xtls/xray-core/transport/internet/masque"
"github.com/xtls/xray-core/transport/internet/stat"
"github.com/xtls/xray-core/transport/internet/tls"
)
const (
establishTimeout = 10 * time.Second
retryInterval = time.Second
)
type Client struct {
server *protocol.ServerSpec
policyManager policy.Manager
remoteDNS []netip.Addr
ctx context.Context
cancel context.CancelFunc
tunnel atomic.Pointer[tunnel]
mu sync.Mutex
lastErr error
lastErrAt time.Time
}
func NewClient(ctx context.Context, config *ClientConfig) (*Client, error) {
v := core.MustFromContext(ctx)
p := v.GetFeature(policy.ManagerType()).(policy.Manager)
streamSettings := session.StreamSettingsFromContext(ctx).(*internet.MemoryStreamConfig)
if _, ok := streamSettings.ProtocolSettings.(*masque.Config); !ok {
return nil, errors.New("not masque transport")
}
if tls.ConfigFromStreamSettings(streamSettings) == nil {
return nil, errors.New(`MASQUE requires "security": "tls"`)
}
if config.Server == nil {
return nil, errors.New(`no target server found`)
}
server, err := protocol.NewServerSpecFromPB(config.Server)
if err != nil {
return nil, errors.New("failed to get server spec").Base(err)
}
dns := config.RemoteDns
if len(dns) == 0 {
dns = []string{"1.1.1.1", "1.0.0.1", "2606:4700:4700::1111", "2606:4700:4700::1001"}
}
remoteDNS := make([]netip.Addr, 0, len(dns))
for _, s := range dns {
addr, err := netip.ParseAddr(s)
if err != nil {
return nil, errors.New("invalid remote DNS server ", s).Base(err)
}
remoteDNS = append(remoteDNS, addr)
}
c := &Client{
server: server,
policyManager: p,
remoteDNS: remoteDNS,
}
c.ctx, c.cancel = context.WithCancel(context.Background())
return c, nil
}
func (c *Client) Process(ctx context.Context, link *transport.Link, dialer internet.Dialer) error {
outbounds := session.OutboundsFromContext(ctx)
ob := outbounds[len(outbounds)-1]
if !ob.Target.IsValid() {
return errors.New("target not specified")
}
ob.Name = "masque"
ob.CanSpliceCopy = 3
t, err := c.getTunnel(ctx, dialer)
if err != nil {
return errors.New("failed to establish CONNECT-IP tunnel").Base(err)
}
var newCtx context.Context
var newCancel context.CancelFunc
if session.TimeoutOnlyFromContext(ctx) {
newCtx, newCancel = context.WithCancel(context.Background())
}
sessionPolicy := c.policyManager.ForLevel(0)
ctx, cancel := context.WithCancel(ctx)
timer := signal.CancelAfterInactivity(ctx, func() {
cancel()
if newCancel != nil {
newCancel()
}
}, sessionPolicy.Timeouts.ConnectionIdle)
if newCtx != nil {
ctx = newCtx
}
var reader buf.Reader
var writer buf.Writer
switch ob.Target.Network {
case net.Network_TCP:
var conn net.Conn
var err error
if sessionPolicy.Timeouts.Handshake != 0 {
timeoutCtx, timeoutCancel := context.WithTimeout(ctx, sessionPolicy.Timeouts.Handshake)
conn, err = t.tnet.DialContext(timeoutCtx, "tcp", ob.Target.NetAddr())
timeoutCancel()
} else {
conn, err = t.tnet.Dial("tcp", ob.Target.NetAddr())
}
if err != nil {
return errors.New("failed to create TCP connection").Base(err)
}
defer conn.Close()
reader = buf.NewReader(conn)
writer = buf.NewWriter(conn)
case net.Network_UDP:
conn, err := t.tnet.Dial("udp", ob.Target.NetAddr())
if err != nil {
return errors.New("failed to create UDP connection").Base(err)
}
defer conn.Close()
uc := &wireguard.UDPConnClient{
PacketConn: conn.(*internet.PacketConnWrapper).PacketConn,
Dest: conn.RemoteAddr().(*net.UDPAddr),
}
reader = uc
writer = uc
default:
panic(ob.Target.Network)
}
requestFunc := func() error {
defer timer.SetTimeout(sessionPolicy.Timeouts.DownlinkOnly)
return buf.Copy(link.Reader, writer, buf.UpdateActivity(timer))
}
responseFunc := func() error {
defer timer.SetTimeout(sessionPolicy.Timeouts.UplinkOnly)
return buf.Copy(reader, link.Writer, buf.UpdateActivity(timer))
}
responseDonePost := task.OnSuccess(responseFunc, task.Close(link.Writer))
if err := task.Run(ctx, requestFunc, responseDonePost); err != nil {
common.Interrupt(link.Reader)
common.Interrupt(link.Writer)
return errors.New("connection ends").Base(err)
}
return nil
}
func (c *Client) getTunnel(ctx context.Context, dialer internet.Dialer) (*tunnel, error) {
c.mu.Lock()
defer c.mu.Unlock()
if c.ctx.Err() != nil {
return nil, errors.New("closed")
}
if t := c.tunnel.Load(); t != nil {
select {
case <-t.done:
default:
return t, nil
}
}
if err := ctx.Err(); err != nil {
return nil, err
}
if c.lastErr != nil && time.Since(c.lastErrAt) < retryInterval {
return nil, c.lastErr
}
t, err := c.establish(ctx, dialer)
if err != nil {
c.lastErr, c.lastErrAt = err, time.Now()
return nil, err
}
c.lastErr = nil
c.tunnel.Store(t)
if c.ctx.Err() != nil {
if c.tunnel.CompareAndSwap(t, nil) {
t.close()
}
return nil, errors.New("closed")
}
return t, nil
}
func (c *Client) establish(ctx context.Context, dialer internet.Dialer) (*tunnel, error) {
ctx, cancel := context.WithTimeout(context.WithoutCancel(ctx), establishTimeout)
defer cancel()
defer context.AfterFunc(c.ctx, cancel)()
conn, err := dialer.Dial(ctx, c.server.Destination)
if err != nil {
return nil, err
}
mconn, ok := stat.TryUnwrapStatsConn(conn).(*masque.Conn)
if !ok {
conn.Close()
return nil, errors.New("not a CONNECT-IP connection")
}
t, err := newTunnel(conn, mconn.LocalAddrs(), c.remoteDNS)
if err != nil {
conn.Close()
return nil, err
}
errors.LogInfo(ctx, "MASQUE: tunnel established from ", mconn.LocalAddrs())
return t, nil
}
func (c *Client) Close() error {
c.cancel()
if t := c.tunnel.Swap(nil); t != nil {
t.close()
}
return nil
}
type tunnel struct {
conn stat.Connection
dev tun.Device
tnet *wireguard.Net
done chan struct{}
closeOnce sync.Once
}
func newTunnel(conn stat.Connection, local []netip.Addr, remoteDNS []netip.Addr) (*tunnel, error) {
var dns []netip.Addr
for _, addr := range remoteDNS {
if slices.ContainsFunc(local, func(l netip.Addr) bool { return l.Is4() == addr.Is4() }) {
dns = append(dns, addr)
}
}
if len(dns) == 0 {
errors.LogWarning(context.Background(), "MASQUE: no remote DNS server is reachable from the assigned addresses ", local, ", domain names will fail to resolve")
dns = remoteDNS
}
dev, tnet, _, err := wireguard.CreateNetTUN(local, dns, masque.MinPacketSize, true)
if err != nil {
return nil, err
}
t := &tunnel{
conn: conn,
dev: dev,
tnet: tnet,
done: make(chan struct{}),
}
go t.readFromTunnel()
go t.writeToTunnel()
return t, nil
}
func (t *tunnel) readFromTunnel() {
defer t.close()
b := make([]byte, buf.Size)
for {
n, err := t.conn.Read(b)
if err != nil {
if go_errors.Is(err, io.ErrShortBuffer) {
continue
}
errors.LogInfoInner(context.Background(), err, "MASQUE: tunnel closed")
return
}
t.dev.Write([][]byte{b[:n]}, 0)
}
}
func (t *tunnel) writeToTunnel() {
bufs := [][]byte{make([]byte, masque.MinPacketSize)}
sizes := []int{0}
for {
if _, err := t.dev.Read(bufs, sizes, 0); err != nil {
return
}
if _, err := t.conn.Write(bufs[0][:sizes[0]]); err != nil {
var ptb *masque.PacketTooBigError
if go_errors.As(err, &ptb) {
go t.dev.Write([][]byte{ptb.ICMP}, 0)
}
}
}
}
func (t *tunnel) close() {
t.closeOnce.Do(func() {
close(t.done)
t.conn.Close()
t.dev.Close()
})
}
func init() {
common.Must(common.RegisterConfig((*ClientConfig)(nil), func(ctx context.Context, config interface{}) (interface{}, error) {
return NewClient(ctx, config.(*ClientConfig))
}))
}
+136
View File
@@ -0,0 +1,136 @@
// Code generated by protoc-gen-go. DO NOT EDIT.
// versions:
// protoc-gen-go v1.36.11
// protoc v6.33.5
// source: proxy/masque/config.proto
package masque
import (
protocol "github.com/xtls/xray-core/common/protocol"
protoreflect "google.golang.org/protobuf/reflect/protoreflect"
protoimpl "google.golang.org/protobuf/runtime/protoimpl"
reflect "reflect"
sync "sync"
unsafe "unsafe"
)
const (
// Verify that this generated code is sufficiently up-to-date.
_ = protoimpl.EnforceVersion(20 - protoimpl.MinVersion)
// Verify that runtime/protoimpl is sufficiently up-to-date.
_ = protoimpl.EnforceVersion(protoimpl.MaxVersion - 20)
)
type ClientConfig struct {
state protoimpl.MessageState `protogen:"open.v1"`
Server *protocol.ServerEndpoint `protobuf:"bytes,1,opt,name=server,proto3" json:"server,omitempty"`
RemoteDns []string `protobuf:"bytes,2,rep,name=remote_dns,json=remoteDns,proto3" json:"remote_dns,omitempty"`
unknownFields protoimpl.UnknownFields
sizeCache protoimpl.SizeCache
}
func (x *ClientConfig) Reset() {
*x = ClientConfig{}
mi := &file_proxy_masque_config_proto_msgTypes[0]
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
ms.StoreMessageInfo(mi)
}
func (x *ClientConfig) String() string {
return protoimpl.X.MessageStringOf(x)
}
func (*ClientConfig) ProtoMessage() {}
func (x *ClientConfig) ProtoReflect() protoreflect.Message {
mi := &file_proxy_masque_config_proto_msgTypes[0]
if x != nil {
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
if ms.LoadMessageInfo() == nil {
ms.StoreMessageInfo(mi)
}
return ms
}
return mi.MessageOf(x)
}
// Deprecated: Use ClientConfig.ProtoReflect.Descriptor instead.
func (*ClientConfig) Descriptor() ([]byte, []int) {
return file_proxy_masque_config_proto_rawDescGZIP(), []int{0}
}
func (x *ClientConfig) GetServer() *protocol.ServerEndpoint {
if x != nil {
return x.Server
}
return nil
}
func (x *ClientConfig) GetRemoteDns() []string {
if x != nil {
return x.RemoteDns
}
return nil
}
var File_proxy_masque_config_proto protoreflect.FileDescriptor
const file_proxy_masque_config_proto_rawDesc = "" +
"\n" +
"\x19proxy/masque/config.proto\x12\x11xray.proxy.masque\x1a!common/protocol/server_spec.proto\"k\n" +
"\fClientConfig\x12<\n" +
"\x06server\x18\x01 \x01(\v2$.xray.common.protocol.ServerEndpointR\x06server\x12\x1d\n" +
"\n" +
"remote_dns\x18\x02 \x03(\tR\tremoteDnsBU\n" +
"\x15com.xray.proxy.masqueP\x01Z&github.com/xtls/xray-core/proxy/masque\xaa\x02\x11Xray.Proxy.Masqueb\x06proto3"
var (
file_proxy_masque_config_proto_rawDescOnce sync.Once
file_proxy_masque_config_proto_rawDescData []byte
)
func file_proxy_masque_config_proto_rawDescGZIP() []byte {
file_proxy_masque_config_proto_rawDescOnce.Do(func() {
file_proxy_masque_config_proto_rawDescData = protoimpl.X.CompressGZIP(unsafe.Slice(unsafe.StringData(file_proxy_masque_config_proto_rawDesc), len(file_proxy_masque_config_proto_rawDesc)))
})
return file_proxy_masque_config_proto_rawDescData
}
var file_proxy_masque_config_proto_msgTypes = make([]protoimpl.MessageInfo, 1)
var file_proxy_masque_config_proto_goTypes = []any{
(*ClientConfig)(nil), // 0: xray.proxy.masque.ClientConfig
(*protocol.ServerEndpoint)(nil), // 1: xray.common.protocol.ServerEndpoint
}
var file_proxy_masque_config_proto_depIdxs = []int32{
1, // 0: xray.proxy.masque.ClientConfig.server:type_name -> xray.common.protocol.ServerEndpoint
1, // [1:1] is the sub-list for method output_type
1, // [1:1] is the sub-list for method input_type
1, // [1:1] is the sub-list for extension type_name
1, // [1:1] is the sub-list for extension extendee
0, // [0:1] is the sub-list for field type_name
}
func init() { file_proxy_masque_config_proto_init() }
func file_proxy_masque_config_proto_init() {
if File_proxy_masque_config_proto != nil {
return
}
type x struct{}
out := protoimpl.TypeBuilder{
File: protoimpl.DescBuilder{
GoPackagePath: reflect.TypeOf(x{}).PkgPath(),
RawDescriptor: unsafe.Slice(unsafe.StringData(file_proxy_masque_config_proto_rawDesc), len(file_proxy_masque_config_proto_rawDesc)),
NumEnums: 0,
NumMessages: 1,
NumExtensions: 0,
NumServices: 0,
},
GoTypes: file_proxy_masque_config_proto_goTypes,
DependencyIndexes: file_proxy_masque_config_proto_depIdxs,
MessageInfos: file_proxy_masque_config_proto_msgTypes,
}.Build()
File_proxy_masque_config_proto = out.File
file_proxy_masque_config_proto_goTypes = nil
file_proxy_masque_config_proto_depIdxs = nil
}
+14
View File
@@ -0,0 +1,14 @@
syntax = "proto3";
package xray.proxy.masque;
option csharp_namespace = "Xray.Proxy.Masque";
option go_package = "github.com/xtls/xray-core/proxy/masque";
option java_package = "com.xray.proxy.masque";
option java_multiple_files = true;
import "common/protocol/server_spec.proto";
message ClientConfig {
xray.common.protocol.ServerEndpoint server = 1;
repeated string remote_dns = 2;
}
+33 -4
View File
@@ -37,6 +37,25 @@ type Handler struct {
downlinkCounter stats.Counter
}
type tunUDPStatsWriter struct {
writer buf.Writer
counter stats.Counter
}
func (w *tunUDPStatsWriter) WriteMultiBuffer(mb buf.MultiBuffer) error {
for len(mb) > 0 {
remaining, packet := buf.SplitFirst(mb)
packetSize := packet.Len()
if err := w.writer.WriteMultiBuffer(buf.MultiBuffer{packet}); err != nil {
buf.ReleaseMulti(remaining)
return err
}
w.counter.Add(int64(packetSize))
mb = remaining
}
return nil
}
// ConnectionHandler interface with the only method that stack is going to push new connections to
type ConnectionHandler interface {
HandleConnection(conn net.Conn, destination net.Destination)
@@ -104,7 +123,7 @@ func (t *Handler) Start() error {
iface := updater.Get()
if iface == nil {
errors.LogInfo(context.Background(), "[tun] falied to set interface > iface == nil")
return nil
return errors.New("iface not found")
}
return c.Control(func(fd uintptr) {
addrPort, _ := netip.ParseAddrPort(address)
@@ -171,7 +190,8 @@ func (t *Handler) HandleConnection(conn net.Conn, destination net.Destination) {
return
}
source := net.DestinationFromAddr(remote)
if t.uplinkCounter != nil || t.downlinkCounter != nil {
isUDP := destination.Network == net.Network_UDP
if !isUDP && (t.uplinkCounter != nil || t.downlinkCounter != nil) {
conn = &stat.CounterConnection{
Connection: conn,
ReadCounter: t.uplinkCounter,
@@ -203,9 +223,18 @@ func (t *Handler) HandleConnection(conn net.Conn, destination net.Destination) {
})
errors.LogInfo(ctx, "processing from ", source, " to ", destination)
reader := &buf.TimeoutWrapperReader{Reader: buf.NewReader(conn)}
writer := buf.NewWriter(conn)
if isUDP {
reader.Counter = t.uplinkCounter
if t.downlinkCounter != nil {
writer = &tunUDPStatsWriter{writer: writer, counter: t.downlinkCounter}
}
}
link := &transport.Link{
Reader: &buf.TimeoutWrapperReader{Reader: buf.NewReader(conn)},
Writer: buf.NewWriter(conn),
Reader: reader,
Writer: writer,
}
if err := t.dispatcher.DispatchLink(ctx, destination, link); err != nil {
errors.LogError(ctx, errors.New("connection closed").Base(err))
+123 -138
View File
@@ -3,7 +3,6 @@ package wireguard
import (
"context"
"fmt"
gonet "net"
"net/netip"
"reflect"
"strings"
@@ -28,14 +27,10 @@ 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"
)
type entry struct {
got []net.IP
time time.Time
}
type Handler struct {
conf *DeviceConfig
policyManager policy.Manager
@@ -49,11 +44,6 @@ type Handler struct {
tnet *Net
dev *device.Device
mu sync.Mutex
// TODO: cache cleanup loop
local bool
cache map[string]entry
cacheMu sync.Mutex
}
func NewClient(ctx context.Context, conf *DeviceConfig) (*Handler, error) {
@@ -109,15 +99,10 @@ func NewClient(ctx context.Context, conf *DeviceConfig) (*Handler, error) {
return nil, err
}
local := false
dns := conf.DNS
if len(dns) == 0 {
dns = []string{"1.1.1.1", "1.0.0.1", "2606:4700:4700::1111", "2606:4700:4700::1001"}
}
if len(dns) == 1 && dns[0] == "local" {
local = true
dns = nil
}
dnses := make([]netip.Addr, 0, len(dns))
for _, dns := range dns {
dnses = append(dnses, netip.MustParseAddr(dns))
@@ -151,9 +136,6 @@ func NewClient(ctx context.Context, conf *DeviceConfig) (*Handler, error) {
tun: tun,
tnet: tnet,
local: local,
cache: make(map[string]entry),
}, nil
}
@@ -172,22 +154,6 @@ func (h *Handler) Process(ctx context.Context, link *transport.Link, dialer inte
return err
}
var addr netip.Addr
if ob.Target.Address.Family().IsDomain() {
ip, err := h.resolveRemote(ob.Target.Address.String())
if err != nil {
return errors.New("failed to resolve domain").Base(err)
}
addr, _ = netip.AddrFromSlice(ip)
} else {
addr, _ = netip.AddrFromSlice(ob.Target.Address.IP())
}
addrPort := netip.AddrPortFrom(addr, ob.Target.Port.Value())
if !addrPort.IsValid() {
return errors.New("invalid target ", ob.Target)
}
var newCtx context.Context
var newCancel context.CancelFunc
if session.TimeoutOnlyFromContext(ctx) {
@@ -216,10 +182,10 @@ func (h *Handler) Process(ctx context.Context, link *transport.Link, dialer inte
var err error
if sessionPolicy.Timeouts.Handshake != 0 {
timeoutCtx, timeoutCancel := context.WithTimeout(ctx, sessionPolicy.Timeouts.Handshake)
conn, err = h.tnet.DialContextTCPAddrPort(timeoutCtx, addrPort)
conn, err = h.tnet.DialContext(timeoutCtx, "tcp", ob.Target.NetAddr())
timeoutCancel()
} else {
conn, err = h.tnet.DialContextTCPAddrPort(ctx, addrPort)
conn, err = h.tnet.Dial("tcp", ob.Target.NetAddr())
}
if err != nil {
return errors.New("failed to create TCP connection").Base(err)
@@ -228,15 +194,14 @@ func (h *Handler) Process(ctx context.Context, link *transport.Link, dialer inte
reader = buf.NewReader(conn)
writer = buf.NewWriter(conn)
case net.Network_UDP:
conn, err := h.tnet.DialUDPAddrPort(netip.AddrPort{}, addrPort)
conn, err := h.tnet.Dial("udp", ob.Target.NetAddr())
if err != nil {
return errors.New("failed to create UDP connection").Base(err)
}
defer conn.Close()
c := &udpConnClient{
PacketConn: conn.(*internet.PacketConnWrapper).PacketConn,
resolveFunc: h.resolveRemote,
dest: gonet.UDPAddrFromAddrPort(addrPort),
c := &UDPConnClient{
PacketConn: conn.(*internet.PacketConnWrapper).PacketConn,
Dest: conn.RemoteAddr().(*net.UDPAddr),
}
reader = c
writer = c
@@ -293,26 +258,26 @@ func (h *Handler) init(ctx context.Context) error {
if err != nil {
return nil, err
}
conn, err := internet.DialSystem(ctx, dest, h.streamSettings.SocketSettings)
if err != nil {
return nil, err
}
var pktConn net.PacketConn
switch c := conn.(type) {
case *internet.PacketConnWrapper:
pktConn = c.PacketConn
case *cnc.Connection:
pktConn = &internet.FakePacketConn{Conn: c}
default:
panic(reflect.TypeOf(c))
}
if h.streamSettings.UdpmaskManager != nil {
newConn, err := h.streamSettings.UdpmaskManager.WrapPacketConnClient(pktConn)
if h.streamSettings.FinalMask != nil {
conn, err := h.streamSettings.FinalMask.DialUDP(ctx, dest)
if err != nil {
pktConn.Close()
return nil, errors.New("mask err").Base(err)
return nil, errors.New("failed to dial to dest").Base(err)
}
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 *internet.PacketConnWrapper:
pktConn = c.PacketConn
case *cnc.Connection:
pktConn = &internet.FakePacketConn{Conn: c}
default:
panic(reflect.TypeOf(c))
}
pktConn = newConn
}
if h.uplinkCounter != nil || h.downlinkCounter != nil {
pktConn = &PacketCounterConnection{
@@ -371,90 +336,54 @@ func (h *Handler) init(ctx context.Context) error {
}
func (h *Handler) resolveLocal(host string) (net.IP, error) {
return h.resolveDomain(host, h.conf.DomainStrategy, func(host string) ([]net.IP, uint32, error) {
return h.dns.LookupIP(host, dns.IPOption{IPv4Enable: true, IPv6Enable: true})
})
}
func (h *Handler) resolveRemote(host string) (net.IP, error) {
return h.resolveDomain(host, h.conf.DomainStrategy, func(host string) ([]net.IP, uint32, error) {
if h.local {
return h.dns.LookupIP(host, dns.IPOption{IPv4Enable: true, IPv6Enable: true})
}
return h.tnet.LookupHost(host)
})
}
func (h *Handler) resolveDomain(host string, strategy DeviceConfig_DomainStrategy, lookupIP func(host string) ([]net.IP, uint32, error)) (net.IP, error) {
if ip := net.ParseIP(host); ip != nil {
return ip, nil
}
h.cacheMu.Lock()
if entry, ok := h.cache[host]; ok {
if time.Now().Before(entry.time) {
h.cacheMu.Unlock()
return entry.got[dice.Roll(len(entry.got))], nil
}
delete(h.cache, host)
}
h.cacheMu.Unlock()
ips, ttl, err := lookupIP(host)
ips, _, err := h.dns.LookupIP(host, dns.IPOption{IPv4Enable: true, IPv6Enable: true})
if err != nil {
return nil, err
}
if len(ips) == 0 {
return nil, dns.ErrEmptyResponse
}
var got4, got6 []net.IP
for _, ip := range ips {
if ip.To4() != nil {
got4 = append(got4, ip)
} else {
got6 = append(got6, ip)
got := ips
if h.streamSettings.SocketSettings != nil {
var got4, got6 []net.IP
for _, ip := range ips {
if ip.To4() != nil {
got4 = append(got4, ip)
} else {
got6 = append(got6, ip)
}
}
}
var got []net.IP
switch strategy {
case DeviceConfig_FORCE_IP:
got = ips
return ips[dice.Roll(len(ips))], nil
case DeviceConfig_FORCE_IP4:
got = got4
case DeviceConfig_FORCE_IP6:
got = got6
case DeviceConfig_FORCE_IP46:
got = got4
if len(got) == 0 {
got = got6
}
case DeviceConfig_FORCE_IP64:
got = got6
if len(got) == 0 {
switch h.streamSettings.SocketSettings.DomainStrategy {
case internet.DomainStrategy_AS_IS, internet.DomainStrategy_USE_IP, internet.DomainStrategy_FORCE_IP:
got = ips
case internet.DomainStrategy_USE_IP4, internet.DomainStrategy_FORCE_IP4:
got = got4
case internet.DomainStrategy_USE_IP6, internet.DomainStrategy_FORCE_IP6:
got = got6
case internet.DomainStrategy_USE_IP46, internet.DomainStrategy_FORCE_IP46:
got = got4
if len(got) == 0 {
got = got6
}
case internet.DomainStrategy_USE_IP64, internet.DomainStrategy_FORCE_IP64:
got = got6
if len(got) == 0 {
got = got4
}
}
if len(got) == 0 {
return nil, dns.ErrEmptyResponse
}
default:
panic(strategy)
}
if len(got) == 0 {
return nil, dns.ErrEmptyResponse
}
entry := entry{
got: got,
time: time.Now().Add(time.Duration(ttl) * time.Second),
}
h.cacheMu.Lock()
h.cache[host] = entry
h.cacheMu.Unlock()
return got[dice.Roll(len(got))], nil
}
type udpConnClient struct {
type UDPConnClient struct {
net.PacketConn
resolveFunc func(host string) (net.IP, error)
dest *net.UDPAddr
Dest *net.UDPAddr
}
func (c *udpConnClient) ReadMultiBuffer() (buf.MultiBuffer, error) {
func (c *UDPConnClient) ReadMultiBuffer() (buf.MultiBuffer, error) {
b := buf.New()
b.Resize(0, buf.Size)
n, addr, err := c.PacketConn.ReadFrom(b.Bytes())
@@ -473,20 +402,13 @@ func (c *udpConnClient) ReadMultiBuffer() (buf.MultiBuffer, error) {
return buf.MultiBuffer{b}, nil
}
func (c *udpConnClient) WriteMultiBuffer(mb buf.MultiBuffer) error {
func (c *UDPConnClient) WriteMultiBuffer(mb buf.MultiBuffer) error {
for i, b := range mb {
dst := c.dest
dst := c.Dest
if b.UDP != nil {
if b.UDP.Address.Family().IsDomain() {
ip, err := c.resolveFunc(b.UDP.Address.String())
if err != nil {
errors.LogErrorInner(context.Background(), err, "drop packet to ", b.UDP, " with size ", len(b.Bytes()))
b.Release()
continue
}
dst = &net.UDPAddr{
IP: ip,
Port: int(b.UDP.Port),
if b.UDP.Port != net.Port(dst.Port) {
dst = &net.UDPAddr{IP: dst.IP, Port: int(b.UDP.Port)}
}
} else {
dst = b.UDP.RawNetAddr().(*net.UDPAddr)
@@ -523,3 +445,66 @@ func (c *PacketCounterConnection) WriteTo(p []byte, addr net.Addr) (n int, err e
}
return
}
type entry struct {
saddr []string
deadline time.Time
}
type cache struct {
running bool
m map[string]entry
mu sync.Mutex
}
func (c *cache) run() {
if c.running {
return
}
c.running = true
if c.m == nil {
c.m = make(map[string]entry)
}
go c.gc()
}
func (c *cache) gc() {
ticker := time.NewTicker(time.Minute)
defer ticker.Stop()
for now := range ticker.C {
c.mu.Lock()
for key, entry := range c.m {
if now.After(entry.deadline) {
delete(c.m, key)
}
}
if len(c.m) == 0 {
c.running = false
c.mu.Unlock()
return
}
c.mu.Unlock()
}
}
func (c *cache) LookupHost(host string) []string {
c.mu.Lock()
defer c.mu.Unlock()
c.run()
if entry, ok := c.m[host]; ok {
if time.Now().Before(entry.deadline) {
return entry.saddr
}
delete(c.m, host)
}
return nil
}
func (c *cache) Cache(host string, saddr []string, ttl uint32) {
c.mu.Lock()
defer c.mu.Unlock()
c.m[host] = entry{
saddr: saddr,
deadline: time.Now().Add(time.Second * time.Duration(ttl)),
}
}
+26 -102
View File
@@ -22,61 +22,6 @@ const (
_ = protoimpl.EnforceVersion(protoimpl.MaxVersion - 20)
)
type DeviceConfig_DomainStrategy int32
const (
DeviceConfig_FORCE_IP DeviceConfig_DomainStrategy = 0
DeviceConfig_FORCE_IP4 DeviceConfig_DomainStrategy = 1
DeviceConfig_FORCE_IP6 DeviceConfig_DomainStrategy = 2
DeviceConfig_FORCE_IP46 DeviceConfig_DomainStrategy = 3
DeviceConfig_FORCE_IP64 DeviceConfig_DomainStrategy = 4
)
// Enum value maps for DeviceConfig_DomainStrategy.
var (
DeviceConfig_DomainStrategy_name = map[int32]string{
0: "FORCE_IP",
1: "FORCE_IP4",
2: "FORCE_IP6",
3: "FORCE_IP46",
4: "FORCE_IP64",
}
DeviceConfig_DomainStrategy_value = map[string]int32{
"FORCE_IP": 0,
"FORCE_IP4": 1,
"FORCE_IP6": 2,
"FORCE_IP46": 3,
"FORCE_IP64": 4,
}
)
func (x DeviceConfig_DomainStrategy) Enum() *DeviceConfig_DomainStrategy {
p := new(DeviceConfig_DomainStrategy)
*p = x
return p
}
func (x DeviceConfig_DomainStrategy) String() string {
return protoimpl.X.EnumStringOf(x.Descriptor(), protoreflect.EnumNumber(x))
}
func (DeviceConfig_DomainStrategy) Descriptor() protoreflect.EnumDescriptor {
return file_proxy_wireguard_config_proto_enumTypes[0].Descriptor()
}
func (DeviceConfig_DomainStrategy) Type() protoreflect.EnumType {
return &file_proxy_wireguard_config_proto_enumTypes[0]
}
func (x DeviceConfig_DomainStrategy) Number() protoreflect.EnumNumber {
return protoreflect.EnumNumber(x)
}
// Deprecated: Use DeviceConfig_DomainStrategy.Descriptor instead.
func (DeviceConfig_DomainStrategy) EnumDescriptor() ([]byte, []int) {
return file_proxy_wireguard_config_proto_rawDescGZIP(), []int{1, 0}
}
type PeerConfig struct {
state protoimpl.MessageState `protogen:"open.v1"`
PublicKey string `protobuf:"bytes,1,opt,name=public_key,json=publicKey,proto3" json:"public_key,omitempty"`
@@ -154,19 +99,18 @@ func (x *PeerConfig) GetAllowedIps() []string {
}
type DeviceConfig struct {
state protoimpl.MessageState `protogen:"open.v1"`
SecretKey string `protobuf:"bytes,1,opt,name=secret_key,json=secretKey,proto3" json:"secret_key,omitempty"`
Endpoint []string `protobuf:"bytes,2,rep,name=endpoint,proto3" json:"endpoint,omitempty"`
Peers []*PeerConfig `protobuf:"bytes,3,rep,name=peers,proto3" json:"peers,omitempty"`
Users []*protocol.User `protobuf:"bytes,5,rep,name=users,proto3" json:"users,omitempty"`
Mtu int32 `protobuf:"varint,4,opt,name=mtu,proto3" json:"mtu,omitempty"`
Reserved []byte `protobuf:"bytes,6,opt,name=reserved,proto3" json:"reserved,omitempty"`
DomainStrategy DeviceConfig_DomainStrategy `protobuf:"varint,7,opt,name=domain_strategy,json=domainStrategy,proto3,enum=xray.proxy.wireguard.DeviceConfig_DomainStrategy" json:"domain_strategy,omitempty"`
IsClient bool `protobuf:"varint,8,opt,name=is_client,json=isClient,proto3" json:"is_client,omitempty"`
NoKernelTun bool `protobuf:"varint,9,opt,name=no_kernel_tun,json=noKernelTun,proto3" json:"no_kernel_tun,omitempty"`
DNS []string `protobuf:"bytes,10,rep,name=DNS,proto3" json:"DNS,omitempty"`
unknownFields protoimpl.UnknownFields
sizeCache protoimpl.SizeCache
state protoimpl.MessageState `protogen:"open.v1"`
SecretKey string `protobuf:"bytes,1,opt,name=secret_key,json=secretKey,proto3" json:"secret_key,omitempty"`
Endpoint []string `protobuf:"bytes,2,rep,name=endpoint,proto3" json:"endpoint,omitempty"`
Peers []*PeerConfig `protobuf:"bytes,3,rep,name=peers,proto3" json:"peers,omitempty"`
Users []*protocol.User `protobuf:"bytes,5,rep,name=users,proto3" json:"users,omitempty"`
Mtu int32 `protobuf:"varint,4,opt,name=mtu,proto3" json:"mtu,omitempty"`
Reserved []byte `protobuf:"bytes,6,opt,name=reserved,proto3" json:"reserved,omitempty"`
IsClient bool `protobuf:"varint,8,opt,name=is_client,json=isClient,proto3" json:"is_client,omitempty"`
NoKernelTun bool `protobuf:"varint,9,opt,name=no_kernel_tun,json=noKernelTun,proto3" json:"no_kernel_tun,omitempty"`
DNS []string `protobuf:"bytes,10,rep,name=DNS,proto3" json:"DNS,omitempty"`
unknownFields protoimpl.UnknownFields
sizeCache protoimpl.SizeCache
}
func (x *DeviceConfig) Reset() {
@@ -241,13 +185,6 @@ func (x *DeviceConfig) GetReserved() []byte {
return nil
}
func (x *DeviceConfig) GetDomainStrategy() DeviceConfig_DomainStrategy {
if x != nil {
return x.DomainStrategy
}
return DeviceConfig_FORCE_IP
}
func (x *DeviceConfig) GetIsClient() bool {
if x != nil {
return x.IsClient
@@ -283,7 +220,7 @@ const file_proxy_wireguard_config_proto_rawDesc = "" +
"\n" +
"keep_alive\x18\x04 \x01(\tR\tkeepAlive\x12\x1f\n" +
"\vallowed_ips\x18\x05 \x03(\tR\n" +
"allowedIps\"\xee\x03\n" +
"allowedIps\"\xb4\x02\n" +
"\fDeviceConfig\x12\x1d\n" +
"\n" +
"secret_key\x18\x01 \x01(\tR\tsecretKey\x12\x1a\n" +
@@ -291,20 +228,11 @@ const file_proxy_wireguard_config_proto_rawDesc = "" +
"\x05peers\x18\x03 \x03(\v2 .xray.proxy.wireguard.PeerConfigR\x05peers\x120\n" +
"\x05users\x18\x05 \x03(\v2\x1a.xray.common.protocol.UserR\x05users\x12\x10\n" +
"\x03mtu\x18\x04 \x01(\x05R\x03mtu\x12\x1a\n" +
"\breserved\x18\x06 \x01(\fR\breserved\x12Z\n" +
"\x0fdomain_strategy\x18\a \x01(\x0e21.xray.proxy.wireguard.DeviceConfig.DomainStrategyR\x0edomainStrategy\x12\x1b\n" +
"\breserved\x18\x06 \x01(\fR\breserved\x12\x1b\n" +
"\tis_client\x18\b \x01(\bR\bisClient\x12\"\n" +
"\rno_kernel_tun\x18\t \x01(\bR\vnoKernelTun\x12\x10\n" +
"\x03DNS\x18\n" +
" \x03(\tR\x03DNS\"\\\n" +
"\x0eDomainStrategy\x12\f\n" +
"\bFORCE_IP\x10\x00\x12\r\n" +
"\tFORCE_IP4\x10\x01\x12\r\n" +
"\tFORCE_IP6\x10\x02\x12\x0e\n" +
"\n" +
"FORCE_IP46\x10\x03\x12\x0e\n" +
"\n" +
"FORCE_IP64\x10\x04B^\n" +
" \x03(\tR\x03DNSB^\n" +
"\x18com.xray.proxy.wireguardP\x01Z)github.com/xtls/xray-core/proxy/wireguard\xaa\x02\x14Xray.Proxy.WireGuardb\x06proto3"
var (
@@ -319,23 +247,20 @@ func file_proxy_wireguard_config_proto_rawDescGZIP() []byte {
return file_proxy_wireguard_config_proto_rawDescData
}
var file_proxy_wireguard_config_proto_enumTypes = make([]protoimpl.EnumInfo, 1)
var file_proxy_wireguard_config_proto_msgTypes = make([]protoimpl.MessageInfo, 2)
var file_proxy_wireguard_config_proto_goTypes = []any{
(DeviceConfig_DomainStrategy)(0), // 0: xray.proxy.wireguard.DeviceConfig.DomainStrategy
(*PeerConfig)(nil), // 1: xray.proxy.wireguard.PeerConfig
(*DeviceConfig)(nil), // 2: xray.proxy.wireguard.DeviceConfig
(*protocol.User)(nil), // 3: xray.common.protocol.User
(*PeerConfig)(nil), // 0: xray.proxy.wireguard.PeerConfig
(*DeviceConfig)(nil), // 1: xray.proxy.wireguard.DeviceConfig
(*protocol.User)(nil), // 2: xray.common.protocol.User
}
var file_proxy_wireguard_config_proto_depIdxs = []int32{
1, // 0: xray.proxy.wireguard.DeviceConfig.peers:type_name -> xray.proxy.wireguard.PeerConfig
3, // 1: xray.proxy.wireguard.DeviceConfig.users:type_name -> xray.common.protocol.User
0, // 2: xray.proxy.wireguard.DeviceConfig.domain_strategy:type_name -> xray.proxy.wireguard.DeviceConfig.DomainStrategy
3, // [3:3] is the sub-list for method output_type
3, // [3:3] is the sub-list for method input_type
3, // [3:3] is the sub-list for extension type_name
3, // [3:3] is the sub-list for extension extendee
0, // [0:3] is the sub-list for field type_name
0, // 0: xray.proxy.wireguard.DeviceConfig.peers:type_name -> xray.proxy.wireguard.PeerConfig
2, // 1: xray.proxy.wireguard.DeviceConfig.users:type_name -> xray.common.protocol.User
2, // [2:2] is the sub-list for method output_type
2, // [2:2] is the sub-list for method input_type
2, // [2:2] is the sub-list for extension type_name
2, // [2:2] is the sub-list for extension extendee
0, // [0:2] is the sub-list for field type_name
}
func init() { file_proxy_wireguard_config_proto_init() }
@@ -348,14 +273,13 @@ func file_proxy_wireguard_config_proto_init() {
File: protoimpl.DescBuilder{
GoPackagePath: reflect.TypeOf(x{}).PkgPath(),
RawDescriptor: unsafe.Slice(unsafe.StringData(file_proxy_wireguard_config_proto_rawDesc), len(file_proxy_wireguard_config_proto_rawDesc)),
NumEnums: 1,
NumEnums: 0,
NumMessages: 2,
NumExtensions: 0,
NumServices: 0,
},
GoTypes: file_proxy_wireguard_config_proto_goTypes,
DependencyIndexes: file_proxy_wireguard_config_proto_depIdxs,
EnumInfos: file_proxy_wireguard_config_proto_enumTypes,
MessageInfos: file_proxy_wireguard_config_proto_msgTypes,
}.Build()
File_proxy_wireguard_config_proto = out.File
-8
View File
@@ -17,13 +17,6 @@ message PeerConfig {
}
message DeviceConfig {
enum DomainStrategy {
FORCE_IP = 0;
FORCE_IP4 = 1;
FORCE_IP6 = 2;
FORCE_IP46 = 3;
FORCE_IP64 = 4;
}
string secret_key = 1;
repeated string endpoint = 2;
repeated PeerConfig peers = 3;
@@ -31,7 +24,6 @@ message DeviceConfig {
int32 mtu = 4;
bytes reserved = 6;
DomainStrategy domain_strategy = 7;
bool is_client = 8;
bool no_kernel_tun = 9;
repeated string DNS = 10;
+159 -14
View File
@@ -15,6 +15,8 @@ import (
"net"
"net/netip"
"os"
"regexp"
"strconv"
"strings"
"syscall"
"time"
@@ -42,6 +44,7 @@ type netTun struct {
events chan tun.Event
notifyHandle *channel.NotificationHandle
incomingPacket chan *buffer.View
closed chan struct{}
mtu int
dnsServers []netip.Addr
hasV4, hasV6 bool
@@ -58,6 +61,7 @@ func CreateNetTUN(localAddresses, dnsServers []netip.Addr, mtu int, handleLocal
stack: stack.New(opts),
events: make(chan tun.Event, 10),
incomingPacket: make(chan *buffer.View),
closed: make(chan struct{}),
dnsServers: dnsServers,
mtu: mtu,
}
@@ -124,12 +128,15 @@ func (tun *netTun) Events() <-chan tun.Event {
}
func (tun *netTun) Read(buf [][]byte, sizes []int, offset int) (int, error) {
view, ok := <-tun.incomingPacket
if !ok {
var view *buffer.View
select {
case view = <-tun.incomingPacket:
case <-tun.closed:
return 0, os.ErrClosed
}
n, err := view.Read(buf[0][offset:])
view.Release()
if err != nil {
return 0, err
}
@@ -166,7 +173,11 @@ func (tun *netTun) WriteNotify() {
view := pkt.ToView()
pkt.DecRef()
tun.incomingPacket <- view
select {
case tun.incomingPacket <- view:
case <-tun.closed:
view.Release()
}
}
func (tun *netTun) Close() error {
@@ -179,8 +190,9 @@ func (tun *netTun) Close() error {
close(tun.events)
}
if tun.incomingPacket != nil {
close(tun.incomingPacket)
// we don't close incomingPacket, because WriteNotify may be mid-send on it (DNS lookup) and would panic.
if tun.closed != nil {
close(tun.closed)
}
return nil
@@ -219,6 +231,7 @@ type Net struct {
DialUDPAddrPort func(laddr, raddr netip.AddrPort) (net.Conn, error)
dnsServers []netip.Addr
hasV4, hasV6 bool
cache cache
}
func convertToFullAddr(endpoint netip.AddrPort) (tcpip.FullAddress, tcpip.NetworkProtocolNumber) {
@@ -246,9 +259,12 @@ var (
errServerTemporarilyMisbehaving = errors.New("server misbehaving")
errCanceled = errors.New("operation was canceled")
errTimeout = errors.New("i/o timeout")
errNumericPort = errors.New("port must be numeric")
errNoSuitableAddress = errors.New("no suitable address found")
errMissingAddress = errors.New("missing address")
)
func (net *Net) LookupHost(host string) (addrs []net.IP, ttl uint32, err error) {
func (net *Net) LookupHost(host string) (addrs []string, err error) {
return net.LookupContextHost(context.Background(), host)
}
@@ -567,9 +583,12 @@ func (tnet *Net) tryOneName(ctx context.Context, name string, qtype dnsmessage.T
return dnsmessage.Parser{}, "", lastErr
}
func (tnet *Net) LookupContextHost(ctx context.Context, host string) ([]net.IP, uint32, error) {
func (tnet *Net) LookupContextHost(ctx context.Context, host string) ([]string, error) {
if saddr := tnet.cache.LookupHost(host); saddr != nil {
return saddr, nil
}
if host == "" || (!tnet.hasV6 && !tnet.hasV4) {
return nil, 0, &net.DNSError{Err: errNoSuchHost.Error(), Name: host, IsNotFound: true}
return nil, &net.DNSError{Err: errNoSuchHost.Error(), Name: host, IsNotFound: true}
}
zlen := len(host)
if strings.IndexByte(host, ':') != -1 {
@@ -578,11 +597,11 @@ func (tnet *Net) LookupContextHost(ctx context.Context, host string) ([]net.IP,
}
}
if ip, err := netip.ParseAddr(host[:zlen]); err == nil {
return []net.IP{ip.AsSlice()}, 0, nil
return []string{ip.String()}, nil
}
if !isDomainName(host) {
return nil, 0, &net.DNSError{Err: errNoSuchHost.Error(), Name: host, IsNotFound: true}
return nil, &net.DNSError{Err: errNoSuchHost.Error(), Name: host, IsNotFound: true}
}
type result struct {
p dnsmessage.Parser
@@ -683,11 +702,137 @@ func (tnet *Net) LookupContextHost(ctx context.Context, host string) ([]net.IP,
}
if len(addrs) == 0 && lastErr != nil {
return nil, 0, lastErr
return nil, lastErr
}
ips := make([]net.IP, 0, len(addrs))
saddrs := make([]string, 0, len(addrs))
for _, ip := range addrs {
ips = append(ips, ip.AsSlice())
saddrs = append(saddrs, ip.String())
}
return ips, ttl, nil
tnet.cache.Cache(host, saddrs, ttl)
return saddrs, nil
}
func partialDeadline(now, deadline time.Time, addrsRemaining int) (time.Time, error) {
if deadline.IsZero() {
return deadline, nil
}
timeRemaining := deadline.Sub(now)
if timeRemaining <= 0 {
return time.Time{}, errTimeout
}
timeout := timeRemaining / time.Duration(addrsRemaining)
const saneMinimum = 2 * time.Second
if timeout < saneMinimum {
if timeRemaining < saneMinimum {
timeout = timeRemaining
} else {
timeout = saneMinimum
}
}
return now.Add(timeout), nil
}
var protoSplitter = regexp.MustCompile(`^(tcp|udp|ping)(4|6)?$`)
func (tnet *Net) DialContext(ctx context.Context, network, address string) (net.Conn, error) {
if ctx == nil {
panic("nil context")
}
var acceptV4, acceptV6 bool
matches := protoSplitter.FindStringSubmatch(network)
if matches == nil {
return nil, &net.OpError{Op: "dial", Err: net.UnknownNetworkError(network)}
} else if len(matches[2]) == 0 {
acceptV4 = true
acceptV6 = true
} else {
acceptV4 = matches[2][0] == '4'
acceptV6 = !acceptV4
}
var host string
var port int
if matches[1] == "ping" {
host = address
} else {
var sport string
var err error
host, sport, err = net.SplitHostPort(address)
if err != nil {
return nil, &net.OpError{Op: "dial", Err: err}
}
port, err = strconv.Atoi(sport)
if err != nil || port < 0 || port > 65535 {
return nil, &net.OpError{Op: "dial", Err: errNumericPort}
}
}
allAddr, err := tnet.LookupContextHost(ctx, host)
if err != nil {
return nil, &net.OpError{Op: "dial", Err: err}
}
var addrs []netip.AddrPort
for _, addr := range allAddr {
ip, err := netip.ParseAddr(addr)
if err == nil && ((ip.Is4() && acceptV4) || (ip.Is6() && acceptV6)) {
addrs = append(addrs, netip.AddrPortFrom(ip, uint16(port)))
}
}
if len(addrs) == 0 && len(allAddr) != 0 {
return nil, &net.OpError{Op: "dial", Err: errNoSuitableAddress}
}
var firstErr error
for i, addr := range addrs {
select {
case <-ctx.Done():
err := ctx.Err()
if err == context.Canceled {
err = errCanceled
} else if err == context.DeadlineExceeded {
err = errTimeout
}
return nil, &net.OpError{Op: "dial", Err: err}
default:
}
dialCtx := ctx
if deadline, hasDeadline := ctx.Deadline(); hasDeadline {
partialDeadline, err := partialDeadline(time.Now(), deadline, len(addrs)-i)
if err != nil {
if firstErr == nil {
firstErr = &net.OpError{Op: "dial", Err: err}
}
break
}
if partialDeadline.Before(deadline) {
var cancel context.CancelFunc
dialCtx, cancel = context.WithDeadline(ctx, partialDeadline)
defer cancel()
}
}
var c net.Conn
switch matches[1] {
case "tcp":
c, err = tnet.DialContextTCPAddrPort(dialCtx, addr)
case "udp":
c, err = tnet.DialUDPAddrPort(netip.AddrPort{}, addr)
case "ping":
err = errors.New("not support")
// c, err = tnet.DialPingAddr(netip.Addr{}, addr.Addr())
}
if err == nil {
return c, nil
}
if firstErr == nil {
firstErr = err
}
}
if firstErr == nil {
firstErr = &net.OpError{Op: "dial", Err: errMissingAddress}
}
return nil, firstErr
}
func (tnet *Net) Dial(network, address string) (net.Conn, error) {
return tnet.DialContext(context.Background(), network, address)
}
+7 -9
View File
@@ -258,18 +258,16 @@ func (s *Server) Start() error {
return errors.New("address is domain")
}
listenFunc := func() (net.PacketConn, error) {
pktConn, err := internet.ListenSystemPacket(context.Background(), &net.UDPAddr{IP: s.src.Address.IP(), Port: int(s.src.Port)}, s.streamSettings.SocketSettings)
var pktConn net.PacketConn
var err error
if s.streamSettings.FinalMask != nil {
pktConn, err = s.streamSettings.FinalMask.ListenPacket(context.Background(), &net.UDPAddr{IP: s.src.Address.IP(), Port: int(s.src.Port)})
} else {
pktConn, err = internet.ListenSystemPacket(context.Background(), &net.UDPAddr{IP: s.src.Address.IP(), Port: int(s.src.Port)}, s.streamSettings.SocketSettings)
}
if err != nil {
return nil, err
}
if s.streamSettings.UdpmaskManager != nil {
newConn, err := s.streamSettings.UdpmaskManager.WrapPacketConnServer(pktConn)
if err != nil {
pktConn.Close()
return nil, errors.New("mask err").Base(err)
}
pktConn = newConn
}
if s.uplinkCounter != nil || s.downlinkCounter != nil {
pktConn = &PacketCounterConnection{
PacketConn: pktConn,
+275
View File
@@ -0,0 +1,275 @@
package scenarios
import (
"context"
gotls "crypto/tls"
"crypto/x509"
go_errors "errors"
"io"
"net/http"
"net/netip"
"sync/atomic"
"testing"
"time"
"github.com/apernet/quic-go"
"github.com/apernet/quic-go/http3"
"golang.org/x/sync/errgroup"
"gvisor.dev/gvisor/pkg/tcpip"
"gvisor.dev/gvisor/pkg/tcpip/adapters/gonet"
"gvisor.dev/gvisor/pkg/tcpip/network/ipv4"
"gvisor.dev/gvisor/pkg/tcpip/network/ipv6"
"github.com/xtls/xray-core/app/log"
"github.com/xtls/xray-core/app/proxyman"
"github.com/xtls/xray-core/common"
clog "github.com/xtls/xray-core/common/log"
"github.com/xtls/xray-core/common/net"
"github.com/xtls/xray-core/common/protocol"
"github.com/xtls/xray-core/common/protocol/tls/cert"
"github.com/xtls/xray-core/common/serial"
core "github.com/xtls/xray-core/core"
"github.com/xtls/xray-core/proxy/dokodemo"
"github.com/xtls/xray-core/proxy/masque"
"github.com/xtls/xray-core/proxy/wireguard"
"github.com/xtls/xray-core/testing/servers/tcp"
"github.com/xtls/xray-core/testing/servers/udp"
"github.com/xtls/xray-core/transport/internet"
transmasque "github.com/xtls/xray-core/transport/internet/masque"
"github.com/xtls/xray-core/transport/internet/masque/connectip"
"github.com/xtls/xray-core/transport/internet/tls"
)
var (
masqueServerV4 = netip.MustParseAddr("10.13.0.1")
masqueServerV6 = netip.MustParseAddr("fd13::1")
masqueClientV4 = netip.MustParsePrefix("10.13.0.2/32")
masqueClientV6 = netip.MustParsePrefix("fd13::2/128")
)
const (
masqueEchoPort = 7
masqueAuthorization = "Basic dTpw"
)
func startMasqueServer(t *testing.T) (net.Port, [32]byte) {
dev, _, gstack, err := wireguard.CreateNetTUN([]netip.Addr{masqueServerV4, masqueServerV6}, nil, transmasque.MinPacketSize, false)
common.Must(err)
t.Cleanup(func() { dev.Close() })
for _, addr := range []netip.Addr{masqueServerV4, masqueServerV6} {
proto := ipv4.ProtocolNumber
if addr.Is6() {
proto = ipv6.ProtocolNumber
}
local := tcpip.FullAddress{NIC: 1, Addr: tcpip.AddrFromSlice(addr.AsSlice()), Port: masqueEchoPort}
l, err := gonet.ListenTCP(gstack, local, proto)
common.Must(err)
go func() {
for {
c, err := l.Accept()
if err != nil {
return
}
go func() {
defer c.Close()
b := make([]byte, 2048)
for {
n, err := c.Read(b)
if err != nil {
return
}
if _, err := c.Write(xor(b[:n])); err != nil {
return
}
}
}()
}
}()
u, err := gonet.DialUDP(gstack, &local, nil, proto)
common.Must(err)
go func() {
b := make([]byte, 2048)
for {
n, addr, err := u.ReadFrom(b)
if err != nil {
return
}
u.WriteTo(xor(b[:n]), addr)
}
}()
}
var current atomic.Pointer[connectip.Conn]
go func() {
bufs := [][]byte{make([]byte, transmasque.MinPacketSize)}
sizes := []int{0}
for {
if _, err := dev.Read(bufs, sizes, 0); err != nil {
return
}
if conn := current.Load(); conn != nil {
if icmp, _ := conn.WritePacket(bufs[0][:sizes[0]]); len(icmp) > 0 {
go dev.Write([][]byte{icmp}, 0)
}
}
}
}()
handler := func(w http.ResponseWriter, r *http.Request) {
if r.URL.Path != transmasque.DefaultPath {
w.WriteHeader(http.StatusNotFound)
return
}
if r.Header.Get("Authorization") != masqueAuthorization {
w.WriteHeader(http.StatusUnauthorized)
return
}
req, err := connectip.ParseProxyRequest(r)
if err != nil {
var perr *connectip.ProxyRequestParseError
if go_errors.As(err, &perr) {
w.WriteHeader(perr.HTTPStatus)
}
return
}
conn, err := (&connectip.Proxy{}).Proxy(w, req)
if err != nil {
return
}
defer conn.Close()
common.Must(conn.AssignAddresses([]netip.Prefix{masqueClientV4, masqueClientV6}))
common.Must(conn.AdvertiseRoute([]connectip.IPRoute{
{StartIP: netip.IPv4Unspecified(), EndIP: netip.AddrFrom4([4]byte{255, 255, 255, 255})},
{StartIP: netip.IPv6Unspecified(), EndIP: netip.AddrFrom16([16]byte{0: 0xff, 1: 0xff, 2: 0xff, 3: 0xff, 4: 0xff, 5: 0xff, 6: 0xff, 7: 0xff, 8: 0xff, 9: 0xff, 10: 0xff, 11: 0xff, 12: 0xff, 13: 0xff, 14: 0xff, 15: 0xff})},
}))
go func() {
for {
ar, err := conn.ReceiveAddressRequest(context.Background())
if err != nil {
return
}
assigned := make([]netip.Prefix, len(ar.Prefixes))
for i, p := range ar.Prefixes {
if p.Addr().Is4() {
assigned[i] = masqueClientV4
} else {
assigned[i] = masqueClientV6
}
}
ar.Respond(assigned, nil)
}
}()
current.Store(conn)
b := make([]byte, 2048)
for {
n, err := conn.ReadPacket(b)
if err != nil {
if go_errors.Is(err, io.ErrShortBuffer) {
continue
}
return
}
dev.Write([][]byte{b[:n]}, 0)
}
}
certificate, certHash := cert.MustGenerate(nil, cert.CommonName("localhost"))
key := common.Must2(x509.ParsePKCS8PrivateKey(certificate.PrivateKey))
tlsConfig := &gotls.Config{
Certificates: []gotls.Certificate{{Certificate: [][]byte{certificate.Certificate}, PrivateKey: key}},
NextProtos: []string{http3.NextProtoH3},
}
pktConn := common.Must2(net.ListenUDP("udp", &net.UDPAddr{IP: net.LocalHostIP.IP()}))
tr := &quic.Transport{Conn: pktConn}
ln := common.Must2(tr.ListenEarly(tlsConfig, &quic.Config{EnableDatagrams: true, InitialPacketSize: 1350}))
server := &http3.Server{Handler: http.HandlerFunc(handler), EnableDatagrams: true}
go server.ServeListener(ln)
t.Cleanup(func() {
server.Close()
ln.Close()
tr.Close()
pktConn.Close()
})
return net.Port(pktConn.LocalAddr().(*net.UDPAddr).Port), certHash
}
func TestMasque(t *testing.T) {
serverPort, certHash := startMasqueServer(t)
tcpPort := tcp.PickPort()
tcp6Port := tcp.PickPort()
udpPort := udp.PickPort()
dokodemoTo := func(port net.Port, addr netip.Addr, network net.Network) *core.InboundHandlerConfig {
return &core.InboundHandlerConfig{
ReceiverSettings: serial.ToTypedMessage(&proxyman.ReceiverConfig{
PortList: &net.PortList{Range: []*net.PortRange{net.SinglePortRange(port)}},
Listen: net.NewIPOrDomain(net.LocalHostIP),
}),
ProxySettings: serial.ToTypedMessage(&dokodemo.Config{
RewriteAddress: net.NewIPOrDomain(net.IPAddress(addr.AsSlice())),
RewritePort: masqueEchoPort,
AllowedNetworks: []net.Network{network},
}),
}
}
clientConfig := &core.Config{
App: []*serial.TypedMessage{
serial.ToTypedMessage(&log.Config{
ErrorLogLevel: clog.Severity_Debug,
ErrorLogType: log.LogType_Console,
}),
},
Inbound: []*core.InboundHandlerConfig{
dokodemoTo(tcpPort, masqueServerV4, net.Network_TCP),
dokodemoTo(tcp6Port, masqueServerV6, net.Network_TCP),
dokodemoTo(udpPort, masqueServerV4, net.Network_UDP),
},
Outbound: []*core.OutboundHandlerConfig{
{
ProxySettings: serial.ToTypedMessage(&masque.ClientConfig{
Server: &protocol.ServerEndpoint{
Address: net.NewIPOrDomain(net.LocalHostIP),
Port: uint32(serverPort),
},
}),
SenderSettings: serial.ToTypedMessage(&proxyman.SenderConfig{
StreamSettings: &internet.StreamConfig{
ProtocolName: "masque",
TransportSettings: []*internet.TransportConfig{
{
ProtocolName: "masque",
Settings: serial.ToTypedMessage(&transmasque.Config{
Path: transmasque.DefaultPath,
Headers: map[string]string{"Authorization": masqueAuthorization},
}),
},
},
SecurityType: serial.GetMessageType(&tls.Config{}),
SecuritySettings: []*serial.TypedMessage{
serial.ToTypedMessage(&tls.Config{
ServerName: "localhost",
PinnedPeerCertSha256: [][]byte{certHash[:]},
}),
},
},
}),
},
},
}
servers, err := InitializeServerConfigs(clientConfig)
common.Must(err)
defer CloseAllServers(servers)
var errg errgroup.Group
for range 3 {
errg.Go(testTCPConn(tcpPort, 1024*1024, time.Second*20))
}
errg.Go(testTCPConn(tcp6Port, 1024*1024, time.Second*20))
errg.Go(testUDPConn(udpPort, 1024, time.Second*5))
if err := errg.Wait(); err != nil {
t.Error(err)
}
}
+2
View File
@@ -65,6 +65,7 @@ func TestWireguard(t *testing.T) {
ProxySettings: serial.ToTypedMessage(&freedom.Config{
FinalRules: []*freedom.FinalRuleConfig{{Action: freedom.RuleAction_Allow}},
}),
SenderSettings: serial.ToTypedMessage(&proxyman.SenderConfig{}),
},
},
}
@@ -104,6 +105,7 @@ func TestWireguard(t *testing.T) {
AllowedIps: []string{"0.0.0.0/0", "::0/0"},
}},
}),
SenderSettings: serial.ToTypedMessage(&proxyman.SenderConfig{}),
},
},
}
+263 -120
View File
@@ -2,103 +2,291 @@ package finalmask
import (
"context"
"net"
"fmt"
"slices"
"github.com/xtls/xray-core/common/buf"
"github.com/xtls/xray-core/common/errors"
"github.com/xtls/xray-core/common/net"
)
type Udpmask interface {
WrapPacketConnClient(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error)
WrapPacketConnServer(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error)
type Dialer struct {
DialTCP func(net.Destination) (net.Conn, error)
DialUDP func(net.Destination) (net.Conn, error)
}
type UdpmaskManager struct {
udpmasks []Udpmask
type ListenConfig struct {
Listen func(net.Addr) (net.Listener, error)
ListenPacket func(net.Addr) (net.PacketConn, error)
}
func NewUdpmaskManager(udpmasks []Udpmask) *UdpmaskManager {
slices.Reverse(udpmasks)
return &UdpmaskManager{udpmasks: udpmasks}
type TCPMask interface {
WrapConnClient(net.Conn, *net.Destination, *Dialer) (net.Conn, error)
WrapConnServer(net.Conn) (net.Conn, error)
// Listen(net.Listener) (net.Listener, error)
}
func (m *UdpmaskManager) WrapPacketConnClient(raw net.PacketConn) (net.PacketConn, error) {
var sizes []int
var conns []net.PacketConn
for i, mask := range m.udpmasks {
if _, ok := mask.(headerConn); ok {
conn, err := mask.WrapPacketConnClient(nil, i, len(m.udpmasks)-1)
if err != nil {
return nil, err
}
sizes = append(sizes, conn.(headerSize).Size())
conns = append(conns, conn)
} else {
if len(conns) > 0 {
raw = &headerManagerConn{sizes: sizes, conns: conns, PacketConn: raw}
sizes = nil
conns = nil
}
var err error
raw, err = mask.WrapPacketConnClient(raw, i, len(m.udpmasks)-1)
if err != nil {
return nil, err
type UDPMask interface {
WrapPacketConnClient(net.PacketConn, *net.Destination, *Dialer) (net.PacketConn, error)
WrapPacketConnServer(net.PacketConn, net.Addr, *ListenConfig) (net.PacketConn, error)
}
type FinalMask struct {
tcpMasks []TCPMask
udpMasks []UDPMask
dialTCP func(context.Context, net.Destination) (net.Conn, error)
listen func(context.Context, net.Addr) (net.Listener, error)
dialUDP func(context.Context, net.Destination) (net.PacketConn, net.Addr, error)
listenPacket func(context.Context, net.Addr) (net.PacketConn, error)
}
func NewFinalMask(tcpMasks []TCPMask, udpMasks []UDPMask, dialTCP func(context.Context, net.Destination) (net.Conn, error), listen func(context.Context, net.Addr) (net.Listener, error), dialUDP func(context.Context, net.Destination) (net.PacketConn, net.Addr, error), listenPacket func(context.Context, net.Addr) (net.PacketConn, error)) *FinalMask {
slices.Reverse(tcpMasks)
slices.Reverse(udpMasks)
return &FinalMask{
tcpMasks: tcpMasks,
udpMasks: udpMasks,
dialTCP: dialTCP,
dialUDP: dialUDP,
listen: listen,
listenPacket: listenPacket,
}
}
func (fm *FinalMask) DialTCP(ctx context.Context, dest net.Destination) (net.Conn, error) {
if len(fm.tcpMasks) == 0 {
return fm.dialTCP(ctx, dest)
}
for i := range fm.tcpMasks {
if i > 0 {
if _, ok := fm.tcpMasks[i].(interface{ HandleDial() }); ok {
return nil, fmt.Errorf("incorrect index: %d %T", i, fm.tcpMasks[i])
}
}
}
if len(conns) > 0 {
raw = &headerManagerConn{sizes: sizes, conns: conns, PacketConn: raw}
sizes = nil
conns = nil
var conn net.Conn
var err error
if _, ok := fm.tcpMasks[0].(interface{ HandleDial() }); !ok {
conn, err = fm.dialTCP(ctx, dest)
if err != nil {
return nil, err
}
}
return raw, nil
dialer := &Dialer{
DialTCP: func(dest net.Destination) (net.Conn, error) {
return fm.dialTCP(ctx, dest)
},
DialUDP: func(dest net.Destination) (net.Conn, error) {
conn, addr, err := fm.dialUDP(ctx, dest)
if err != nil {
return nil, err
}
return &PacketConnWrapper{PacketConn: conn, udpAddr: addr}, err
},
}
for i := range fm.tcpMasks {
var newConn net.Conn
newConn, err = fm.tcpMasks[i].WrapConnClient(conn, &dest, dialer)
if err != nil {
_ = conn.Close()
return nil, err
}
conn = newConn
}
return conn, nil
}
func (m *UdpmaskManager) WrapPacketConnServer(raw net.PacketConn) (net.PacketConn, error) {
var sizes []int
var conns []net.PacketConn
for i, mask := range m.udpmasks {
if _, ok := mask.(headerConn); ok {
conn, err := mask.WrapPacketConnServer(nil, i, len(m.udpmasks)-1)
if err != nil {
return nil, err
func (fm *FinalMask) Listen(ctx context.Context, addr net.Addr) (net.Listener, error) {
if len(fm.tcpMasks) == 0 {
return fm.listen(ctx, addr)
}
off := 0
listener, err := fm.listen(ctx, addr)
if err != nil {
return nil, err
}
for i := range fm.tcpMasks {
if _, ok := fm.tcpMasks[i].(interface {
Listen(net.Listener) (net.Listener, error)
}); ok {
if i-off == 0 {
l, err := fm.tcpMasks[i].(interface {
Listen(net.Listener) (net.Listener, error)
}).Listen(listener)
if err != nil {
listener.Close()
return nil, err
}
listener = l
} else {
l, err := fm.tcpMasks[i].(interface {
Listen(net.Listener) (net.Listener, error)
}).Listen(&TCPListener{Listener: listener, tcpMasks: fm.tcpMasks[off:i]})
if err != nil {
listener.Close()
return nil, err
}
listener = l
}
sizes = append(sizes, conn.(headerSize).Size())
conns = append(conns, conn)
} else {
if len(conns) > 0 {
raw = &headerManagerConn{sizes: sizes, conns: conns, PacketConn: raw}
sizes = nil
conns = nil
}
var err error
raw, err = mask.WrapPacketConnServer(raw, i, len(m.udpmasks)-1)
if err != nil {
return nil, err
off = i + 1
}
}
if off < len(fm.tcpMasks) {
return &TCPListener{Listener: listener, tcpMasks: fm.tcpMasks[off:]}, nil
}
return listener, nil
}
func (fm *FinalMask) DialUDP(ctx context.Context, dest net.Destination) (net.Conn, error) {
if len(fm.udpMasks) == 0 {
conn, addr, err := fm.dialUDP(ctx, dest)
if err != nil {
return nil, err
}
return &PacketConnWrapper{PacketConn: conn, udpAddr: addr}, nil
}
for i := range fm.udpMasks {
if i > 0 {
if _, ok := fm.udpMasks[i].(interface{ HandleDial() }); ok {
return nil, fmt.Errorf("incorrect index: %d %T", i, fm.udpMasks[i])
}
}
}
var conn net.PacketConn
var addr net.Addr
var err error
if _, ok := fm.udpMasks[0].(interface{ HandleDial() }); !ok {
conn, addr, err = fm.dialUDP(ctx, dest)
if err != nil {
return nil, err
}
}
dialer := &Dialer{
DialTCP: func(dest net.Destination) (net.Conn, error) {
return fm.dialTCP(ctx, dest)
},
DialUDP: func(dest net.Destination) (net.Conn, error) {
conn, addr, err := fm.dialUDP(ctx, dest)
if err != nil {
return nil, err
}
return &PacketConnWrapper{PacketConn: conn, udpAddr: addr}, err
},
}
var sizes []int
var conns []net.PacketConn
for i := range fm.udpMasks {
var newConn net.PacketConn
if _, ok := fm.udpMasks[i].(interface{ HeaderConn() }); ok {
newConn, err = fm.udpMasks[i].WrapPacketConnClient(nil, nil, nil)
if err != nil {
_ = conn.Close()
return nil, err
}
sizes = append(sizes, newConn.(interface{ Size() int }).Size())
conns = append(conns, newConn)
} else {
if len(conns) > 0 {
conn = &headerManagerConn{PacketConn: conn, sizes: sizes, conns: conns}
sizes = nil
conns = nil
}
newConn, err = fm.udpMasks[i].WrapPacketConnClient(conn, &dest, dialer)
if err != nil {
_ = conn.Close()
return nil, err
}
conn = newConn
}
}
if len(conns) > 0 {
raw = &headerManagerConn{sizes: sizes, conns: conns, PacketConn: raw}
conn = &headerManagerConn{PacketConn: conn, sizes: sizes, conns: conns}
sizes = nil
conns = nil
}
return raw, nil
if addr == nil {
addr = &net.UDPAddr{IP: []byte{0, 0, 0, 0}}
}
return &PacketConnWrapper{PacketConn: conn, udpAddr: addr}, nil
}
func (fm *FinalMask) ListenPacket(ctx context.Context, addr net.Addr) (net.PacketConn, error) {
if len(fm.udpMasks) == 0 {
return fm.listenPacket(ctx, addr)
}
for i := range fm.udpMasks {
if i > 0 {
if _, ok := fm.udpMasks[i].(interface{ HandleListen() }); ok {
return nil, fmt.Errorf("incorrect index: %d %T", i, fm.udpMasks[i])
}
}
}
var conn net.PacketConn
var err error
if _, ok := fm.udpMasks[0].(interface{ HandleListen() }); !ok {
conn, err = fm.listenPacket(ctx, addr)
if err != nil {
return nil, err
}
}
lc := &ListenConfig{
Listen: func(addr net.Addr) (net.Listener, error) { return fm.listen(ctx, addr) },
ListenPacket: func(addr net.Addr) (net.PacketConn, error) { return fm.listenPacket(ctx, addr) },
}
var sizes []int
var conns []net.PacketConn
for i := range fm.udpMasks {
var newConn net.PacketConn
if _, ok := fm.udpMasks[i].(interface{ HeaderConn() }); ok {
newConn, err = fm.udpMasks[i].WrapPacketConnServer(nil, nil, nil)
if err != nil {
_ = conn.Close()
return nil, err
}
sizes = append(sizes, newConn.(interface{ Size() int }).Size())
conns = append(conns, newConn)
} else {
if len(conns) > 0 {
conn = &headerManagerConn{PacketConn: conn, sizes: sizes, conns: conns}
sizes = nil
conns = nil
}
newConn, err = fm.udpMasks[i].WrapPacketConnServer(conn, addr, lc)
if err != nil {
_ = conn.Close()
return nil, err
}
conn = newConn
}
}
if len(conns) > 0 {
conn = &headerManagerConn{PacketConn: conn, sizes: sizes, conns: conns}
sizes = nil
conns = nil
}
return conn, nil
}
const (
UDPSize = 4096
)
type headerConn interface {
HeaderConn()
type PacketConnWrapper struct {
net.PacketConn
udpAddr net.Addr
}
type headerSize interface {
Size() int
func (c *PacketConnWrapper) RemoteAddr() net.Addr {
return c.udpAddr
}
func (c *PacketConnWrapper) Read(b []byte) (n int, err error) {
n, _, err = c.PacketConn.ReadFrom(b)
return
}
func (c *PacketConnWrapper) Write(b []byte) (n int, err error) {
return c.PacketConn.WriteTo(b, c.udpAddr)
}
type headerManagerConn struct {
@@ -191,72 +379,27 @@ func (c *headerManagerConn) WriteTo(p []byte, addr net.Addr) (n int, err error)
return len(p), nil
}
type Tcpmask interface {
WrapConnClient(net.Conn) (net.Conn, error)
WrapConnServer(net.Conn) (net.Conn, error)
}
type TcpmaskManager struct {
tcpmasks []Tcpmask
}
func NewTcpmaskManager(tcpmasks []Tcpmask) *TcpmaskManager {
slices.Reverse(tcpmasks)
return &TcpmaskManager{tcpmasks: tcpmasks}
}
func (m *TcpmaskManager) WrapConnClient(raw net.Conn) (net.Conn, error) {
var err error
for _, mask := range m.tcpmasks {
raw, err = mask.WrapConnClient(raw)
if err != nil {
return nil, err
}
}
return raw, nil
}
func (m *TcpmaskManager) WrapConnServer(raw net.Conn) (net.Conn, error) {
var err error
for _, mask := range m.tcpmasks {
raw, err = mask.WrapConnServer(raw)
if err != nil {
return nil, err
}
}
return raw, nil
}
func (m *TcpmaskManager) WrapListener(l net.Listener) (net.Listener, error) {
return NewTcpListener(m, l)
}
type tcpListener struct {
m *TcpmaskManager
type TCPListener struct {
net.Listener
tcpMasks []TCPMask
}
func NewTcpListener(m *TcpmaskManager, l net.Listener) (net.Listener, error) {
return &tcpListener{
m: m,
Listener: l,
}, nil
}
func (l *tcpListener) Accept() (net.Conn, error) {
func (l *TCPListener) Accept() (net.Conn, error) {
conn, err := l.Listener.Accept()
if err != nil {
return conn, err
}
newConn, err := l.m.WrapConnServer(conn)
if err != nil {
errors.LogDebugInner(context.Background(), err, "mask err")
_ = conn.Close()
return nil, err
for i := range l.tcpMasks {
var newConn net.Conn
newConn, err = l.tcpMasks[i].WrapConnServer(conn)
if err != nil {
_ = conn.Close()
return nil, err
}
conn = newConn
}
return newConn, nil
return conn, nil
}
type TcpMaskConn interface {
@@ -1,11 +1,14 @@
package fragment
import "net"
import (
"github.com/xtls/xray-core/common/net"
"github.com/xtls/xray-core/transport/internet/finalmask"
)
func (c *Config) WrapConnClient(raw net.Conn) (net.Conn, error) {
return NewConnClient(c, raw, false)
func (c *Config) WrapConnClient(conn net.Conn, dest *net.Destination, dialer *finalmask.Dialer) (net.Conn, error) {
return NewConnClient(c, conn, false)
}
func (c *Config) WrapConnServer(raw net.Conn) (net.Conn, error) {
return NewConnServer(c, raw, true)
func (c *Config) WrapConnServer(conn net.Conn) (net.Conn, error) {
return NewConnServer(c, conn, true)
}
@@ -1,29 +1,30 @@
package custom
import (
"net"
"github.com/xtls/xray-core/common/net"
"github.com/xtls/xray-core/transport/internet/finalmask"
)
func (c *TCPConfig) WrapConnClient(raw net.Conn) (net.Conn, error) {
return NewConnClientTCP(c, raw)
func (c *TCPConfig) WrapConnClient(conn net.Conn, dest *net.Destination, dialer *finalmask.Dialer) (net.Conn, error) {
return NewConnClientTCP(c, conn)
}
func (c *TCPConfig) WrapConnServer(raw net.Conn) (net.Conn, error) {
return NewConnServerTCP(c, raw)
func (c *TCPConfig) WrapConnServer(conn net.Conn) (net.Conn, error) {
return NewConnServerTCP(c, conn)
}
func (c *UDPConfig) WrapPacketConnClient(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error) {
return NewConnClientUDP(c, raw)
func (c *UDPConfig) WrapPacketConnClient(conn net.PacketConn, dest *net.Destination, dialer *finalmask.Dialer) (net.PacketConn, error) {
return NewConnClientUDP(c, conn)
}
func (c *UDPConfig) WrapPacketConnServer(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error) {
return NewConnServerUDP(c, raw)
func (c *UDPConfig) WrapPacketConnServer(conn net.PacketConn, addr net.Addr, lc *finalmask.ListenConfig) (net.PacketConn, error) {
return NewConnServerUDP(c, conn)
}
func (c *UDPStandaloneConfig) WrapPacketConnClient(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error) {
return NewConnClientUDPStandalone(c, raw)
func (c *UDPStandaloneConfig) WrapPacketConnClient(conn net.PacketConn, dest *net.Destination, dialer *finalmask.Dialer) (net.PacketConn, error) {
return NewConnClientUDPStandalone(c, conn)
}
func (c *UDPStandaloneConfig) WrapPacketConnServer(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error) {
return NewConnServerUDPStandalone(c, raw)
func (c *UDPStandaloneConfig) WrapPacketConnServer(conn net.PacketConn, addr net.Addr, lc *finalmask.ListenConfig) (net.PacketConn, error) {
return NewConnServerUDPStandalone(c, conn)
}
@@ -9,8 +9,6 @@ import (
"strings"
"testing"
"time"
"github.com/xtls/xray-core/transport/internet/finalmask"
)
func TestMetadataEvaluatorRejectsUnknownName(t *testing.T) {
@@ -156,7 +154,7 @@ func TestMetadataUDPStandaloneWriteUsesRemotePort(t *testing.T) {
}
defer serverRaw.Close()
client, err := finalmask.NewUdpmaskManager([]finalmask.Udpmask{cfg}).WrapPacketConnClient(clientRaw)
client, err := cfg.WrapPacketConnClient(clientRaw, nil, nil)
if err != nil {
t.Fatal(err)
}
@@ -301,7 +299,7 @@ func TestMetadataTCPHandshakeUsesEndpointPorts(t *testing.T) {
}
defer serverRaw.Close()
client, err := clientCfg.WrapConnClient(clientRaw)
client, err := clientCfg.WrapConnClient(clientRaw, nil, nil)
if err != nil {
t.Fatal(err)
}
@@ -5,8 +5,6 @@ import (
"net"
"testing"
"time"
"github.com/xtls/xray-core/transport/internet/finalmask"
)
func mustSendRecvUDP(t *testing.T, from net.PacketConn, to net.PacketConn, msg []byte) {
@@ -48,7 +46,6 @@ func TestStateUDPResponseReusesPriorCapturedValues(t *testing.T) {
},
},
}
maskManager := finalmask.NewUdpmaskManager([]finalmask.Udpmask{cfg})
clientRaw, err := net.ListenPacket("udp", "127.0.0.1:0")
if err != nil {
@@ -62,11 +59,11 @@ func TestStateUDPResponseReusesPriorCapturedValues(t *testing.T) {
}
defer serverRaw.Close()
client, err := maskManager.WrapPacketConnClient(clientRaw)
client, err := cfg.WrapPacketConnClient(clientRaw, nil, nil)
if err != nil {
t.Fatal(err)
}
server, err := maskManager.WrapPacketConnServer(serverRaw)
server, err := cfg.WrapPacketConnServer(serverRaw, nil, nil)
if err != nil {
t.Fatal(err)
}
@@ -37,7 +37,7 @@ func TestDSLTCPHandshakeReusesCapturedValue(t *testing.T) {
defer clientRaw.Close()
defer serverRaw.Close()
client, err := cfg.WrapConnClient(clientRaw)
client, err := cfg.WrapConnClient(clientRaw, nil, nil)
if err != nil {
t.Fatal(err)
}
@@ -117,7 +117,7 @@ func TestDSLTCPClientRejectsMismatchedResponseSequence(t *testing.T) {
defer clientRaw.Close()
defer serverRaw.Close()
client, err := clientCfg.WrapConnClient(clientRaw)
client, err := clientCfg.WrapConnClient(clientRaw, nil, nil)
if err != nil {
t.Fatal(err)
}
@@ -1,15 +1,16 @@
package aes128gcm
import (
"net"
"github.com/xtls/xray-core/common/net"
"github.com/xtls/xray-core/transport/internet/finalmask"
)
func (c *Config) HeaderConn() {}
func (c *Config) WrapPacketConnClient(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error) {
return NewConnClient(c, raw)
func (c *Config) WrapPacketConnClient(conn net.PacketConn, dest *net.Destination, dialer *finalmask.Dialer) (net.PacketConn, error) {
return NewConnClient(c, conn)
}
func (c *Config) WrapPacketConnServer(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error) {
return NewConnServer(c, raw)
func (c *Config) WrapPacketConnServer(conn net.PacketConn, addr net.Addr, lc *finalmask.ListenConfig) (net.PacketConn, error) {
return NewConnServer(c, conn)
}
@@ -1,15 +1,16 @@
package header
import (
"net"
"github.com/xtls/xray-core/common/net"
"github.com/xtls/xray-core/transport/internet/finalmask"
)
func (c *Config) HeaderConn() {}
func (c *Config) WrapPacketConnClient(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error) {
return NewConnClient(c, raw)
func (c *Config) WrapPacketConnClient(conn net.PacketConn, dest *net.Destination, dialer *finalmask.Dialer) (net.PacketConn, error) {
return NewConnClient(c, conn)
}
func (c *Config) WrapPacketConnServer(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error) {
return NewConnServer(c, raw)
func (c *Config) WrapPacketConnServer(conn net.PacketConn, addr net.Addr, lc *finalmask.ListenConfig) (net.PacketConn, error) {
return NewConnServer(c, conn)
}
@@ -1,15 +1,16 @@
package original
import (
"net"
"github.com/xtls/xray-core/common/net"
"github.com/xtls/xray-core/transport/internet/finalmask"
)
func (c *Config) HeaderConn() {}
func (c *Config) WrapPacketConnClient(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error) {
return NewConnClient(c, raw)
func (c *Config) WrapPacketConnClient(conn net.PacketConn, dest *net.Destination, dialer *finalmask.Dialer) (net.PacketConn, error) {
return NewConnClient(c, conn)
}
func (c *Config) WrapPacketConnServer(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error) {
return NewConnServer(c, raw)
func (c *Config) WrapPacketConnServer(conn net.PacketConn, addr net.Addr, lc *finalmask.ListenConfig) (net.PacketConn, error) {
return NewConnServer(c, conn)
}
+8 -5
View File
@@ -1,11 +1,14 @@
package noise
import "net"
import (
"github.com/xtls/xray-core/common/net"
"github.com/xtls/xray-core/transport/internet/finalmask"
)
func (c *Config) WrapPacketConnClient(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error) {
return NewConnClient(c, raw)
func (c *Config) WrapPacketConnClient(conn net.PacketConn, dest *net.Destination, dialer *finalmask.Dialer) (net.PacketConn, error) {
return NewConnClient(c, conn)
}
func (c *Config) WrapPacketConnServer(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error) {
return NewConnServer(c, raw)
func (c *Config) WrapPacketConnServer(conn net.PacketConn, addr net.Addr, lc *finalmask.ListenConfig) (net.PacketConn, error) {
return NewConnServer(c, conn)
}
+6 -15
View File
@@ -1,23 +1,14 @@
package realm
import (
"net"
"github.com/xtls/xray-core/common/errors"
"github.com/xtls/xray-core/transport/internet"
"github.com/xtls/xray-core/common/net"
"github.com/xtls/xray-core/transport/internet/finalmask"
)
func (c *Config) WrapPacketConnClient(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error) {
_, ok1 := raw.(*internet.FakePacketConn)
if level != 0 || ok1 {
return nil, errors.New("realm requires being at the outermost level")
}
return NewConnClient(c, raw)
func (c *Config) WrapPacketConnClient(conn net.PacketConn, dest *net.Destination, dialer *finalmask.Dialer) (net.PacketConn, error) {
return NewConnClient(c, conn)
}
func (c *Config) WrapPacketConnServer(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error) {
if level != 0 {
return nil, errors.New("realm requires being at the outermost level")
}
return NewConnServer(c, raw)
func (c *Config) WrapPacketConnServer(conn net.PacketConn, addr net.Addr, lc *finalmask.ListenConfig) (net.PacketConn, error) {
return NewConnServer(c, conn)
}
@@ -1,23 +1,24 @@
package salamander
import (
"net"
"github.com/xtls/xray-core/common/net"
"github.com/xtls/xray-core/transport/internet/finalmask"
)
func (c *Config) HeaderConn() {}
func (c *Config) WrapPacketConnClient(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error) {
return NewSalamanderConnClient(c, raw)
func (c *Config) WrapPacketConnClient(conn net.PacketConn, dest *net.Destination, dialer *finalmask.Dialer) (net.PacketConn, error) {
return NewSalamanderConnClient(c, conn)
}
func (c *Config) WrapPacketConnServer(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error) {
return NewSalamanderConnServer(c, raw)
func (c *Config) WrapPacketConnServer(conn net.PacketConn, addr net.Addr, lc *finalmask.ListenConfig) (net.PacketConn, error) {
return NewSalamanderConnServer(c, conn)
}
func (c *GeckoConfig) WrapPacketConnClient(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error) {
return NewGeckoConnClient(c, raw)
func (c *GeckoConfig) WrapPacketConnClient(conn net.PacketConn, dest *net.Destination, dialer *finalmask.Dialer) (net.PacketConn, error) {
return NewGeckoConnClient(c, conn)
}
func (c *GeckoConfig) WrapPacketConnServer(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error) {
return NewGeckoConnServer(c, raw)
func (c *GeckoConfig) WrapPacketConnServer(conn net.PacketConn, addr net.Addr, lc *finalmask.ListenConfig) (net.PacketConn, error) {
return NewGeckoConnServer(c, conn)
}
+10 -17
View File
@@ -1,19 +1,18 @@
package sudoku
import (
"net"
"github.com/xtls/xray-core/common/errors"
"github.com/xtls/xray-core/common/net"
"github.com/xtls/xray-core/transport/internet/finalmask"
)
// Sudoku in finalmask mode is a pure appearance transform with no standalone handshake.
// TCP always keeps classic sudoku on uplink and uses packed downlink optimization on server writes.
func (c *Config) WrapConnClient(raw net.Conn) (net.Conn, error) {
return newPackedDirectionalConn(raw, c, true)
func (c *Config) WrapConnClient(conn net.Conn, dest *net.Destination, dialer *finalmask.Dialer) (net.Conn, error) {
return newPackedDirectionalConn(conn, c, true)
}
func (c *Config) WrapConnServer(raw net.Conn) (net.Conn, error) {
return newPackedDirectionalConn(raw, c, false)
func (c *Config) WrapConnServer(conn net.Conn) (net.Conn, error) {
return newPackedDirectionalConn(conn, c, false)
}
func newPackedDirectionalConn(raw net.Conn, config *Config, readPacked bool) (net.Conn, error) {
@@ -36,16 +35,10 @@ func newPackedDirectionalConn(raw net.Conn, config *Config, readPacked bool) (ne
return newWrappedConn(raw, reader, writer), nil
}
func (c *Config) WrapPacketConnClient(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error) {
if level != levelCount {
return nil, errors.New("sudoku udp mask must be the innermost mask in chain")
}
return NewUDPConn(raw, c)
func (c *Config) WrapPacketConnClient(conn net.PacketConn, dest *net.Destination, dialer *finalmask.Dialer) (net.PacketConn, error) {
return NewUDPConn(conn, c)
}
func (c *Config) WrapPacketConnServer(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error) {
if level != levelCount {
return nil, errors.New("sudoku udp mask must be the innermost mask in chain")
}
return NewUDPConn(raw, c)
func (c *Config) WrapPacketConnServer(conn net.PacketConn, addr net.Addr, lc *finalmask.ListenConfig) (net.PacketConn, error) {
return NewUDPConn(conn, c)
}
+55 -53
View File
@@ -2,12 +2,14 @@ package finalmask_test
import (
"bytes"
"context"
"io"
"net"
gonet "net"
"strings"
"testing"
"time"
"github.com/xtls/xray-core/common/net"
"github.com/xtls/xray-core/transport/internet/finalmask"
"github.com/xtls/xray-core/transport/internet/finalmask/header/custom"
)
@@ -20,11 +22,14 @@ func mustSendRecvTcp(
) {
t.Helper()
waitCh := make(chan error)
go func() {
_, err := from.Write(msg)
if err != nil {
t.Error(err)
t.Fatal(err)
}
close(waitCh)
}()
buf := make([]byte, 1024)
@@ -40,18 +45,23 @@ func mustSendRecvTcp(
if !bytes.Equal(buf[:n], msg) {
t.Fatalf("unexpected data %q", buf[:n])
}
<-waitCh
}
type layerMaskTcp struct {
name string
mask finalmask.Tcpmask
mask finalmask.TCPMask
}
type failingWrapMask struct{}
func (failingWrapMask) TCP() {}
func (f failingWrapMask) WrapConnClient(raw net.Conn) (net.Conn, error) { return raw, nil }
func (f failingWrapMask) WrapConnServer(raw net.Conn) (net.Conn, error) {
func (failingWrapMask) TCP() {}
func (f failingWrapMask) WrapConnClient(conn net.Conn, dest *net.Destination, dialer *finalmask.Dialer) (net.Conn, error) {
return conn, nil
}
func (f failingWrapMask) WrapConnServer(conn net.Conn) (net.Conn, error) {
return nil, io.ErrClosedPipe
}
@@ -92,32 +102,31 @@ func TestConnReadWrite(t *testing.T) {
t.Run(c.name, func(t *testing.T) {
mask := c.mask
maskManager := finalmask.NewTcpmaskManager([]finalmask.Tcpmask{mask})
dialTCP := func(ctx context.Context, dest net.Destination) (net.Conn, error) {
return net.Dial("tcp", dest.NetAddr())
}
listen := func(ctx context.Context, addr net.Addr) (net.Listener, error) {
return net.Listen("tcp", addr.String())
}
finalMask := finalmask.NewFinalMask([]finalmask.TCPMask{mask}, nil, dialTCP, listen, nil, nil)
ln, err := net.Listen("tcp", "127.0.0.1:0")
listener, err := finalMask.Listen(context.Background(), &net.TCPAddr{IP: net.LocalHostIP.IP()})
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { listener.Close() })
client, err := net.Dial("tcp", ln.Addr().String())
client, err := finalMask.DialTCP(context.Background(), net.TCPDestination(net.IPAddress(listener.Addr().(*net.TCPAddr).IP), net.Port(listener.Addr().(*net.TCPAddr).Port)))
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { client.Close() })
client, err = maskManager.WrapConnClient(client)
if err != nil {
t.Fatal(err)
}
server, err := ln.Accept()
if err != nil {
t.Fatal(err)
}
server, err = maskManager.WrapConnServer(server)
server, err := listener.Accept()
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { server.Close() })
_ = client.SetDeadline(time.Now().Add(time.Second))
_ = server.SetDeadline(time.Now().Add(time.Second))
@@ -150,34 +159,32 @@ func TestTCPcustomStaticHandshakeRoundTrip(t *testing.T) {
},
},
}
maskManager := finalmask.NewTcpmaskManager([]finalmask.Tcpmask{cfg})
ln, err := net.Listen("tcp", "127.0.0.1:0")
if err != nil {
t.Fatal(err)
dialTCP := func(ctx context.Context, dest net.Destination) (net.Conn, error) {
return net.Dial("tcp", dest.NetAddr())
}
defer ln.Close()
listen := func(ctx context.Context, addr net.Addr) (net.Listener, error) {
return net.Listen("tcp", addr.String())
}
finalMask := finalmask.NewFinalMask([]finalmask.TCPMask{cfg}, nil, dialTCP, listen, nil, nil)
clientRaw, err := net.Dial("tcp", ln.Addr().String())
listener, err := finalMask.Listen(context.Background(), &net.TCPAddr{IP: net.LocalHostIP.IP()})
if err != nil {
t.Fatal(err)
}
defer clientRaw.Close()
defer listener.Close()
serverRaw, err := ln.Accept()
client, err := finalMask.DialTCP(context.Background(), net.TCPDestination(net.IPAddress(listener.Addr().(*net.TCPAddr).IP), net.Port(listener.Addr().(*net.TCPAddr).Port)))
if err != nil {
t.Fatal(err)
}
defer serverRaw.Close()
defer client.Close()
client, err := maskManager.WrapConnClient(clientRaw)
if err != nil {
t.Fatal(err)
}
server, err := maskManager.WrapConnServer(serverRaw)
server, err := listener.Accept()
if err != nil {
t.Fatal(err)
}
defer server.Close()
_ = client.SetDeadline(time.Now().Add(time.Second))
_ = server.SetDeadline(time.Now().Add(time.Second))
@@ -220,11 +227,11 @@ func TestTCPcustomClientRejectsMismatchedServerSequence(t *testing.T) {
},
}
clientRaw, serverRaw := net.Pipe()
clientRaw, serverRaw := gonet.Pipe()
defer clientRaw.Close()
defer serverRaw.Close()
client, err := clientCfg.WrapConnClient(clientRaw)
client, err := clientCfg.WrapConnClient(clientRaw, nil, nil)
if err != nil {
t.Fatal(err)
}
@@ -257,42 +264,37 @@ func TestTCPcustomClientRejectsMismatchedServerSequence(t *testing.T) {
}
func TestTCPWrapListenerRejectsImmediateWrapErrors(t *testing.T) {
clientManager := finalmask.NewTcpmaskManager([]finalmask.Tcpmask{failingWrapMask{}})
serverManager := finalmask.NewTcpmaskManager([]finalmask.Tcpmask{failingWrapMask{}})
dialTCP := func(ctx context.Context, dest net.Destination) (net.Conn, error) {
return net.Dial("tcp", dest.NetAddr())
}
listen := func(ctx context.Context, addr net.Addr) (net.Listener, error) {
return net.Listen("tcp", addr.String())
}
finalMask := finalmask.NewFinalMask([]finalmask.TCPMask{failingWrapMask{}}, nil, dialTCP, listen, nil, nil)
rawLn, err := net.Listen("tcp", "127.0.0.1:0")
if err != nil {
t.Fatal(err)
}
defer rawLn.Close()
ln, err := serverManager.WrapListener(rawLn)
listener, err := finalMask.Listen(context.Background(), &net.TCPAddr{IP: net.LocalHostIP.IP()})
if err != nil {
t.Fatal(err)
}
defer listener.Close()
accepted := make(chan struct {
conn net.Conn
err error
}, 1)
go func() {
conn, err := ln.Accept()
conn, err := listener.Accept()
accepted <- struct {
conn net.Conn
err error
}{conn: conn, err: err}
}()
clientRaw, err := net.Dial("tcp", rawLn.Addr().String())
if err != nil {
t.Fatal(err)
}
defer clientRaw.Close()
client, err := clientManager.WrapConnClient(clientRaw)
client, err := finalMask.DialTCP(context.Background(), net.TCPDestination(net.IPAddress(listener.Addr().(*net.TCPAddr).IP), net.Port(listener.Addr().(*net.TCPAddr).Port)))
if err != nil {
t.Fatal(err)
}
defer client.Close()
_ = client.SetDeadline(time.Now().Add(time.Second))
+52 -59
View File
@@ -2,13 +2,15 @@ package finalmask_test
import (
"bytes"
"context"
"encoding/binary"
"io"
"net"
gonet "net"
"sync/atomic"
"testing"
"time"
"github.com/xtls/xray-core/common/net"
"github.com/xtls/xray-core/proxy"
"github.com/xtls/xray-core/transport/internet/finalmask"
"github.com/xtls/xray-core/transport/internet/finalmask/header/custom"
@@ -51,7 +53,7 @@ func mustSendRecv(
type layerMask struct {
name string
mask finalmask.Udpmask
mask finalmask.UDPMask
layers int
}
@@ -213,25 +215,23 @@ func newStandaloneStunLikeUDPServerConfig() *custom.UDPStandaloneConfig {
func newUDPClientServerPair(t *testing.T, cfg *custom.UDPStandaloneConfig) (net.PacketConn, net.PacketConn, net.PacketConn, net.PacketConn) {
t.Helper()
clientRaw, err := net.ListenPacket("udp", "127.0.0.1:0")
clientRaw, err := gonet.ListenPacket("udp", "127.0.0.1:0")
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { _ = clientRaw.Close() })
serverRaw, err := net.ListenPacket("udp", "127.0.0.1:0")
serverRaw, err := gonet.ListenPacket("udp", "127.0.0.1:0")
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { _ = serverRaw.Close() })
maskManager := finalmask.NewUdpmaskManager([]finalmask.Udpmask{cfg})
client, err := maskManager.WrapPacketConnClient(clientRaw)
client, err := cfg.WrapPacketConnClient(clientRaw, nil, nil)
if err != nil {
t.Fatal(err)
}
server, err := maskManager.WrapPacketConnServer(serverRaw)
server, err := cfg.WrapPacketConnServer(serverRaw, nil, nil)
if err != nil {
t.Fatal(err)
}
@@ -348,31 +348,39 @@ func TestPacketConnReadWrite(t *testing.T) {
if layers <= 0 {
layers = 1
}
masks := make([]finalmask.Udpmask, 0, layers)
masks := make([]finalmask.UDPMask, 0, layers)
for i := 0; i < layers; i++ {
masks = append(masks, mask)
}
maskManager := finalmask.NewUdpmaskManager(masks)
client, err := net.ListenPacket("udp", "127.0.0.1:0")
dialUDP := func(ctx context.Context, dest net.Destination) (net.PacketConn, net.Addr, error) {
udpAddr, err := net.ResolveUDPAddr("udp", dest.NetAddr())
if err != nil {
return nil, nil, err
}
conn, err := gonet.ListenPacket("udp", "127.0.0.1:0")
if err != nil {
return nil, nil, err
}
return conn, udpAddr, nil
}
listenPacket := func(ctx context.Context, addr net.Addr) (net.PacketConn, error) {
return gonet.ListenPacket(addr.Network(), addr.String())
}
finalMask := finalmask.NewFinalMask(nil, masks, nil, nil, dialUDP, listenPacket)
server, err := finalMask.ListenPacket(context.Background(), &net.UDPAddr{IP: net.LocalHostIP.IP()})
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { server.Close() })
client, err = maskManager.WrapPacketConnClient(client)
if err != nil {
t.Fatal(err)
}
server, err := net.ListenPacket("udp", "127.0.0.1:0")
if err != nil {
t.Fatal(err)
}
server, err = maskManager.WrapPacketConnServer(server)
clientConn, err := finalMask.DialUDP(context.Background(), net.UDPDestination(net.IPAddress(server.LocalAddr().(*net.UDPAddr).IP), net.Port(server.LocalAddr().(*net.UDPAddr).Port)))
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { clientConn.Close() })
client := clientConn.(*finalmask.PacketConnWrapper).PacketConn
_ = client.SetDeadline(time.Now().Add(time.Second))
_ = server.SetDeadline(time.Now().Add(time.Second))
@@ -397,21 +405,20 @@ func TestUDPcustomStaticHeaderWireShape(t *testing.T) {
{Rand: 1, RandMin: 0x30, RandMax: 0x40},
},
}
maskManager := finalmask.NewUdpmaskManager([]finalmask.Udpmask{cfg})
clientRaw, err := net.ListenPacket("udp", "127.0.0.1:0")
clientRaw, err := gonet.ListenPacket("udp", "127.0.0.1:0")
if err != nil {
t.Fatal(err)
}
defer clientRaw.Close()
serverRaw, err := net.ListenPacket("udp", "127.0.0.1:0")
serverRaw, err := gonet.ListenPacket("udp", "127.0.0.1:0")
if err != nil {
t.Fatal(err)
}
defer serverRaw.Close()
client, err := maskManager.WrapPacketConnClient(clientRaw)
client, err := cfg.WrapPacketConnClient(clientRaw, nil, nil)
if err != nil {
t.Fatal(err)
}
@@ -642,11 +649,11 @@ func TestSudokuBDD(t *testing.T) {
Ascii: "prefer_ascii",
}
clientRaw, serverRaw := net.Pipe()
clientRaw, serverRaw := gonet.Pipe()
defer clientRaw.Close()
defer serverRaw.Close()
clientConn, err := cfg.WrapConnClient(clientRaw)
clientConn, err := cfg.WrapConnClient(clientRaw, nil, nil)
if err != nil {
t.Fatal(err)
}
@@ -683,11 +690,11 @@ func TestSudokuBDD(t *testing.T) {
PaddingMax: 0,
}
clientRaw, serverRaw := net.Pipe()
clientRaw, serverRaw := gonet.Pipe()
defer clientRaw.Close()
defer serverRaw.Close()
clientConn, err := cfg.WrapConnClient(clientRaw)
clientConn, err := cfg.WrapConnClient(clientRaw, nil, nil)
if err != nil {
t.Fatal(err)
}
@@ -738,10 +745,10 @@ func TestSudokuBDD(t *testing.T) {
countWireBytes := func(wrapServer func(net.Conn, *sudoku.Config) (net.Conn, error), cfg *sudoku.Config) int64 {
t.Helper()
clientRaw, serverRaw := net.Pipe()
clientRaw, serverRaw := gonet.Pipe()
watchedServerRaw := &countingConn{Conn: serverRaw}
clientConn, err := cfg.WrapConnClient(clientRaw)
clientConn, err := cfg.WrapConnClient(clientRaw, nil, nil)
if err != nil {
t.Fatal(err)
}
@@ -793,11 +800,11 @@ func TestSudokuBDD(t *testing.T) {
CustomTables: []string{"xpxvvpvv", "vxpvxvvp"},
}
clientRaw, serverRaw := net.Pipe()
clientRaw, serverRaw := gonet.Pipe()
defer clientRaw.Close()
defer serverRaw.Close()
clientConn, err := cfg.WrapConnClient(clientRaw)
clientConn, err := cfg.WrapConnClient(clientRaw, nil, nil)
if err != nil {
t.Fatal(err)
}
@@ -835,11 +842,11 @@ func TestSudokuBDD(t *testing.T) {
PaddingMax: 0,
}
clientRaw, serverRaw := net.Pipe()
clientRaw, serverRaw := gonet.Pipe()
defer clientRaw.Close()
defer serverRaw.Close()
clientConn, err := cfg.WrapConnClient(clientRaw)
clientConn, err := cfg.WrapConnClient(clientRaw, nil, nil)
if err != nil {
t.Fatal(err)
}
@@ -868,19 +875,6 @@ func TestSudokuBDD(t *testing.T) {
}
})
t.Run("GivenSudokuUDPMask_WhenNotInnermost_ThenWrapFails", func(t *testing.T) {
cfg := &sudoku.Config{Password: "sudoku-udp"}
raw, err := net.ListenPacket("udp", "127.0.0.1:0")
if err != nil {
t.Fatal(err)
}
defer raw.Close()
if _, err := cfg.WrapPacketConnClient(raw, 0, 1); err == nil {
t.Fatal("expected innermost check failure")
}
})
t.Run("GivenSudokuMultiTableUDPMask_WhenClientSendsMultipleDatagrams_ThenPayloadMatches", func(t *testing.T) {
cfg := &sudoku.Config{
Password: "sudoku-udp-multi",
@@ -889,25 +883,24 @@ func TestSudokuBDD(t *testing.T) {
PaddingMin: 0,
PaddingMax: 0,
}
maskManager := finalmask.NewUdpmaskManager([]finalmask.Udpmask{cfg})
clientRaw, err := net.ListenPacket("udp", "127.0.0.1:0")
clientRaw, err := gonet.ListenPacket("udp", "127.0.0.1:0")
if err != nil {
t.Fatal(err)
}
defer clientRaw.Close()
serverRaw, err := net.ListenPacket("udp", "127.0.0.1:0")
serverRaw, err := gonet.ListenPacket("udp", "127.0.0.1:0")
if err != nil {
t.Fatal(err)
}
defer serverRaw.Close()
client, err := maskManager.WrapPacketConnClient(clientRaw)
client, err := cfg.WrapPacketConnClient(clientRaw, nil, nil)
if err != nil {
t.Fatal(err)
}
server, err := maskManager.WrapPacketConnServer(serverRaw)
server, err := cfg.WrapPacketConnServer(serverRaw, nil, nil)
if err != nil {
t.Fatal(err)
}
@@ -961,7 +954,7 @@ func TestSudokuBDD(t *testing.T) {
}
defer serverRaw.Close()
clientConn, err := cfg.WrapConnClient(clientRaw)
clientConn, err := cfg.WrapConnClient(clientRaw, nil, nil)
if err != nil {
t.Fatal(err)
}
@@ -1008,11 +1001,11 @@ func TestSudokuBDD(t *testing.T) {
Ascii: "prefer_entropy",
}
clientRaw, serverRaw := net.Pipe()
clientRaw, serverRaw := gonet.Pipe()
defer clientRaw.Close()
defer serverRaw.Close()
clientConn, err := cfg.WrapConnClient(clientRaw)
clientConn, err := cfg.WrapConnClient(clientRaw, nil, nil)
if err != nil {
t.Fatal(err)
}
@@ -1032,11 +1025,11 @@ func TestSudokuBDD(t *testing.T) {
Ascii: "prefer_entropy",
}
clientRaw, serverRaw := net.Pipe()
clientRaw, serverRaw := gonet.Pipe()
defer clientRaw.Close()
defer serverRaw.Close()
clientConn, err := cfg.WrapConnClient(clientRaw)
clientConn, err := cfg.WrapConnClient(clientRaw, nil, nil)
if err != nil {
t.Fatal(err)
}
+7 -10
View File
@@ -1,20 +1,17 @@
package udphop
import (
"net"
"github.com/xtls/xray-core/common/errors"
"github.com/xtls/xray-core/transport/internet"
"github.com/xtls/xray-core/common/net"
"github.com/xtls/xray-core/transport/internet/finalmask"
)
func (c *Config) WrapPacketConnClient(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error) {
_, ok1 := raw.(*internet.FakePacketConn)
if level != 0 || ok1 {
return nil, errors.New("udphop requires being at the outermost level")
}
return NewUDPHopConn(c, raw)
func (c *Config) HandleDial() {}
func (c *Config) WrapPacketConnClient(conn net.PacketConn, dest *net.Destination, dialer *finalmask.Dialer) (net.PacketConn, error) {
return NewUDPHopConn(c, dest, dialer)
}
func (c *Config) WrapPacketConnServer(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error) {
func (c *Config) WrapPacketConnServer(conn net.PacketConn, addr net.Addr, lc *finalmask.ListenConfig) (net.PacketConn, error) {
return nil, errors.New("udphop: client only")
}
@@ -7,7 +7,6 @@
package udphop
import (
internet "github.com/xtls/xray-core/transport/internet"
protoreflect "google.golang.org/protobuf/reflect/protoreflect"
protoimpl "google.golang.org/protobuf/runtime/protoimpl"
reflect "reflect"
@@ -24,14 +23,13 @@ const (
type Config struct {
state protoimpl.MessageState `protogen:"open.v1"`
Sockopt *internet.SocketConfig `protobuf:"bytes,1,opt,name=sockopt,proto3" json:"sockopt,omitempty"`
Local bool `protobuf:"varint,2,opt,name=local,proto3" json:"local,omitempty"`
Remote bool `protobuf:"varint,3,opt,name=remote,proto3" json:"remote,omitempty"`
RemoteOnce bool `protobuf:"varint,4,opt,name=remote_once,json=remoteOnce,proto3" json:"remote_once,omitempty"`
IntervalMin int64 `protobuf:"varint,5,opt,name=interval_min,json=intervalMin,proto3" json:"interval_min,omitempty"`
IntervalMax int64 `protobuf:"varint,6,opt,name=interval_max,json=intervalMax,proto3" json:"interval_max,omitempty"`
RemotePorts []uint32 `protobuf:"varint,7,rep,packed,name=remote_ports,json=remotePorts,proto3" json:"remote_ports,omitempty"`
RemoteIPs []string `protobuf:"bytes,8,rep,name=remoteIPs,proto3" json:"remoteIPs,omitempty"`
RemoteIPs []string `protobuf:"bytes,7,rep,name=remoteIPs,proto3" json:"remoteIPs,omitempty"`
RemotePorts []uint32 `protobuf:"varint,8,rep,packed,name=remote_ports,json=remotePorts,proto3" json:"remote_ports,omitempty"`
unknownFields protoimpl.UnknownFields
sizeCache protoimpl.SizeCache
}
@@ -66,13 +64,6 @@ func (*Config) Descriptor() ([]byte, []int) {
return file_transport_internet_finalmask_udphop_config_proto_rawDescGZIP(), []int{0}
}
func (x *Config) GetSockopt() *internet.SocketConfig {
if x != nil {
return x.Sockopt
}
return nil
}
func (x *Config) GetLocal() bool {
if x != nil {
return x.Local
@@ -108,16 +99,16 @@ func (x *Config) GetIntervalMax() int64 {
return 0
}
func (x *Config) GetRemotePorts() []uint32 {
func (x *Config) GetRemoteIPs() []string {
if x != nil {
return x.RemotePorts
return x.RemoteIPs
}
return nil
}
func (x *Config) GetRemoteIPs() []string {
func (x *Config) GetRemotePorts() []uint32 {
if x != nil {
return x.RemoteIPs
return x.RemotePorts
}
return nil
}
@@ -126,17 +117,16 @@ var File_transport_internet_finalmask_udphop_config_proto protoreflect.FileDescr
const file_transport_internet_finalmask_udphop_config_proto_rawDesc = "" +
"\n" +
"0transport/internet/finalmask/udphop/config.proto\x12(xray.transport.internet.finalmask.udphop\x1a\x1ftransport/internet/config.proto\"\x9f\x02\n" +
"\x06Config\x12?\n" +
"\asockopt\x18\x01 \x01(\v2%.xray.transport.internet.SocketConfigR\asockopt\x12\x14\n" +
"0transport/internet/finalmask/udphop/config.proto\x12(xray.transport.internet.finalmask.udphop\"\xe4\x01\n" +
"\x06Config\x12\x14\n" +
"\x05local\x18\x02 \x01(\bR\x05local\x12\x16\n" +
"\x06remote\x18\x03 \x01(\bR\x06remote\x12\x1f\n" +
"\vremote_once\x18\x04 \x01(\bR\n" +
"remoteOnce\x12!\n" +
"\finterval_min\x18\x05 \x01(\x03R\vintervalMin\x12!\n" +
"\finterval_max\x18\x06 \x01(\x03R\vintervalMax\x12!\n" +
"\fremote_ports\x18\a \x03(\rR\vremotePorts\x12\x1c\n" +
"\tremoteIPs\x18\b \x03(\tR\tremoteIPsB\x9a\x01\n" +
"\finterval_max\x18\x06 \x01(\x03R\vintervalMax\x12\x1c\n" +
"\tremoteIPs\x18\a \x03(\tR\tremoteIPs\x12!\n" +
"\fremote_ports\x18\b \x03(\rR\vremotePortsJ\x04\b\x01\x10\x02B\x9a\x01\n" +
",com.xray.transport.internet.finalmask.udphopP\x01Z=github.com/xtls/xray-core/transport/internet/finalmask/udphop\xaa\x02(Xray.Transport.Internet.Finalmask.Udphopb\x06proto3"
var (
@@ -153,16 +143,14 @@ func file_transport_internet_finalmask_udphop_config_proto_rawDescGZIP() []byte
var file_transport_internet_finalmask_udphop_config_proto_msgTypes = make([]protoimpl.MessageInfo, 1)
var file_transport_internet_finalmask_udphop_config_proto_goTypes = []any{
(*Config)(nil), // 0: xray.transport.internet.finalmask.udphop.Config
(*internet.SocketConfig)(nil), // 1: xray.transport.internet.SocketConfig
(*Config)(nil), // 0: xray.transport.internet.finalmask.udphop.Config
}
var file_transport_internet_finalmask_udphop_config_proto_depIdxs = []int32{
1, // 0: xray.transport.internet.finalmask.udphop.Config.sockopt:type_name -> xray.transport.internet.SocketConfig
1, // [1:1] is the sub-list for method output_type
1, // [1:1] is the sub-list for method input_type
1, // [1:1] is the sub-list for extension type_name
1, // [1:1] is the sub-list for extension extendee
0, // [0:1] is the sub-list for field type_name
0, // [0:0] is the sub-list for method output_type
0, // [0:0] is the sub-list for method input_type
0, // [0:0] is the sub-list for extension type_name
0, // [0:0] is the sub-list for extension extendee
0, // [0:0] is the sub-list for field type_name
}
func init() { file_transport_internet_finalmask_udphop_config_proto_init() }
@@ -6,16 +6,14 @@ option go_package = "github.com/xtls/xray-core/transport/internet/finalmask/udph
option java_package = "com.xray.transport.internet.finalmask.udphop";
option java_multiple_files = true;
import "transport/internet/config.proto";
message Config {
xray.transport.internet.SocketConfig sockopt = 1;
reserved 1;
bool local = 2;
bool remote = 3;
bool remote_once = 4;
int64 interval_min = 5;
int64 interval_max = 6;
repeated uint32 remote_ports = 7;
repeated string remoteIPs = 8;
repeated string remoteIPs = 7;
repeated uint32 remote_ports = 8;
}
+81 -97
View File
@@ -6,9 +6,7 @@ import (
goerrors "errors"
"io"
mrand "math/rand"
gonet "net"
"net/netip"
"reflect"
"sync"
"time"
@@ -16,8 +14,6 @@ import (
"github.com/xtls/xray-core/common/crypto"
"github.com/xtls/xray-core/common/errors"
"github.com/xtls/xray-core/common/net"
"github.com/xtls/xray-core/common/net/cnc"
"github.com/xtls/xray-core/transport/internet"
"github.com/xtls/xray-core/transport/internet/finalmask"
)
@@ -34,16 +30,14 @@ type packet struct {
}
type udpHopConn struct {
conn net.PacketConn
sockopt *internet.SocketConfig
local bool
remote bool
remoteOnce bool
dialer *finalmask.Dialer
local bool
remote bool
intervalMin int64
intervalMax int64
remotePorts []uint32
remoteIPs []netip.Prefix
remotePorts []uint32
deadline time.Time
readDeadline time.Time
@@ -55,10 +49,10 @@ type udpHopConn struct {
readCh chan packet
closeCh chan struct{}
wg sync.WaitGroup
mu sync.Mutex
mu sync.RWMutex
}
func NewUDPHopConn(c *Config, raw net.PacketConn) (net.PacketConn, error) {
func NewUDPHopConn(c *Config, dest *net.Destination, dialer *finalmask.Dialer) (net.PacketConn, error) {
if c.IntervalMin < 5 || c.IntervalMax < 5 {
return nil, errors.New("invalid interval")
}
@@ -66,22 +60,40 @@ func NewUDPHopConn(c *Config, raw net.PacketConn) (net.PacketConn, error) {
for _, ip := range c.RemoteIPs {
remoteIPs = append(remoteIPs, netip.MustParsePrefix(ip))
}
conn := &udpHopConn{
conn: raw,
sockopt: c.Sockopt,
local: c.Local,
remote: c.Remote,
remoteOnce: c.RemoteOnce,
remotePorts := c.RemotePorts
if c.Remote || c.RemoteOnce {
if len(remoteIPs) > 0 {
dest.Address = net.IPAddress(randPrefix(remoteIPs[mrand.Intn(len(remoteIPs))]))
}
if len(remotePorts) > 0 {
dest.Port = net.Port(remotePorts[mrand.Intn(len(remotePorts))])
}
}
conn, err := dialer.DialUDP(*dest)
if err != nil {
return nil, err
}
cur := conn.(*finalmask.PacketConnWrapper).PacketConn
addr := conn.RemoteAddr().(*net.UDPAddr)
client := &udpHopConn{
dialer: dialer,
local: c.Local,
remote: c.Remote,
intervalMin: c.IntervalMin,
intervalMax: c.IntervalMax,
remotePorts: c.RemotePorts,
remoteIPs: remoteIPs,
remotePorts: remotePorts,
cur: cur,
addr: addr,
readCh: make(chan packet),
closeCh: make(chan struct{}),
}
return conn, nil
go client.run()
client.wg.Add(1)
go client.recv(client.cur)
return client, nil
}
func (c *udpHopConn) closed() bool {
@@ -93,61 +105,67 @@ func (c *udpHopConn) closed() bool {
}
}
func (c *udpHopConn) hop(addr *net.UDPAddr) {
func (c *udpHopConn) run() {
ticker := time.NewTicker(time.Second * time.Duration(crypto.RandBetween(c.intervalMin, c.intervalMax+1)))
defer ticker.Stop()
for {
select {
case <-c.closeCh:
return
case <-ticker.C:
ticker.Reset(time.Second * time.Duration(crypto.RandBetween(c.intervalMin, c.intervalMax+1)))
c.hop()
}
}
}
func (c *udpHopConn) hop() {
c.mu.Lock()
defer c.mu.Unlock()
if c.closed() {
return
}
newAddr := &net.UDPAddr{IP: addr.IP, Port: addr.Port}
newConn := c.conn
if c.remote || c.remoteOnce && c.addr == nil {
if len(c.remotePorts) > 0 {
newAddr.Port = int(c.remotePorts[mrand.Intn(len(c.remotePorts))])
}
oldIP := c.addr.IP
oldPort := c.addr.Port
if c.remote {
if len(c.remoteIPs) > 0 {
newAddr.IP = randPrefix(c.remoteIPs[mrand.Intn(len(c.remoteIPs))])
c.addr.IP = randPrefix(c.remoteIPs[mrand.Intn(len(c.remoteIPs))])
}
if len(c.remotePorts) > 0 {
c.addr.Port = int(c.remotePorts[mrand.Intn(len(c.remotePorts))])
}
}
if c.local {
raw, err := internet.DialSystem(context.Background(), net.UDPDestination(net.IPAddress(newAddr.IP), net.Port(newAddr.Port)), c.sockopt)
conn, err := c.dialer.DialUDP(net.UDPDestination(net.IPAddress(c.addr.IP), net.Port(c.addr.Port)))
if err != nil {
c.addr.IP = oldIP
c.addr.Port = oldPort
errors.LogErrorInner(context.Background(), err, "hop err")
return
}
switch c := raw.(type) {
case *internet.PacketConnWrapper:
newConn = c.PacketConn
case *cnc.Connection:
newConn = &internet.FakePacketConn{Conn: c}
default:
panic(reflect.TypeOf(c))
}
newConn.SetDeadline(c.deadline)
newConn.SetReadDeadline(c.readDeadline)
newConn.SetWriteDeadline(c.writeDeadline)
conn.SetDeadline(c.deadline)
conn.SetReadDeadline(c.readDeadline)
conn.SetWriteDeadline(c.writeDeadline)
if c.pre != nil {
_ = c.pre.Close()
}
c.pre = c.cur
c.cur = conn.(*finalmask.PacketConnWrapper).PacketConn
c.wg.Add(1)
go c.recv(newConn)
go c.recv(c.cur)
}
c.addr = newAddr
c.cur = newConn
}
func (c *udpHopConn) recv(conn net.PacketConn) {
defer c.wg.Done()
for {
if c.closed() {
return
}
p := pool.Get().([]byte)
n, addr, err := conn.ReadFrom(p)
if err != nil {
pool.Put(p[:cap(p)])
if goerrors.Is(err, io.EOF) || goerrors.Is(err, io.ErrClosedPipe) || goerrors.Is(err, gonet.ErrClosed) {
break
if c.closed() {
return
}
var netErr net.Error
if goerrors.As(err, &netErr) && netErr.Timeout() {
@@ -156,9 +174,10 @@ func (c *udpHopConn) recv(conn net.PacketConn) {
case <-c.closeCh:
return
}
continue
}
errors.LogErrorInner(context.Background(), err, "recv err")
continue
return
}
select {
case c.readCh <- packet{p: p[:n], addr: addr}:
@@ -169,22 +188,6 @@ func (c *udpHopConn) recv(conn net.PacketConn) {
}
}
func (c *udpHopConn) hopLoop() {
ticker := time.NewTicker(time.Second * time.Duration(crypto.RandBetween(c.intervalMin, c.intervalMax+1)))
defer ticker.Stop()
for {
select {
case <-ticker.C:
ticker.Reset(time.Second * time.Duration(crypto.RandBetween(c.intervalMin, c.intervalMax+1)))
c.mu.Lock()
c.hop(c.addr)
c.mu.Unlock()
case <-c.closeCh:
return
}
}
}
func (c *udpHopConn) ReadFrom(p []byte) (n int, addr net.Addr, err error) {
packet, ok := <-c.readCh
if ok {
@@ -194,21 +197,12 @@ func (c *udpHopConn) ReadFrom(p []byte) (n int, addr net.Addr, err error) {
}
return n, packet.addr, packet.err
}
return 0, nil, io.EOF
return 0, nil, io.ErrClosedPipe
}
func (c *udpHopConn) WriteTo(p []byte, addr net.Addr) (n int, err error) {
c.mu.Lock()
defer c.mu.Unlock()
if c.cur == nil {
c.hop(addr.(*net.UDPAddr))
if c.cur == nil {
return 0, nil
}
go c.hopLoop()
}
c.mu.RLock()
defer c.mu.RUnlock()
_, err = c.cur.WriteTo(p, c.addr)
if err != nil {
errors.LogErrorInner(context.Background(), err, "send err")
@@ -227,15 +221,12 @@ func (c *udpHopConn) Close() error {
if c.pre != nil {
_ = c.pre.Close()
}
if c.cur != nil {
_ = c.cur.Close()
}
_ = c.conn.Close()
_ = c.cur.Close()
c.wg.Wait()
select {
case p := <-c.readCh:
if p.p != nil {
pool.Put(p.p[:cap(p.p)])
case packet := <-c.readCh:
if packet.p != nil {
pool.Put(packet.p[:cap(packet.p)])
}
default:
}
@@ -244,7 +235,9 @@ func (c *udpHopConn) Close() error {
}
func (c *udpHopConn) LocalAddr() net.Addr {
return c.conn.LocalAddr()
c.mu.RLock()
defer c.mu.RUnlock()
return c.cur.LocalAddr()
}
func (c *udpHopConn) SetDeadline(t time.Time) error {
@@ -254,10 +247,7 @@ func (c *udpHopConn) SetDeadline(t time.Time) error {
if c.pre != nil {
_ = c.pre.SetDeadline(t)
}
if c.cur != nil {
_ = c.cur.SetDeadline(t)
}
return nil
return c.cur.SetDeadline(t)
}
func (c *udpHopConn) SetReadDeadline(t time.Time) error {
@@ -267,10 +257,7 @@ func (c *udpHopConn) SetReadDeadline(t time.Time) error {
if c.pre != nil {
_ = c.pre.SetReadDeadline(t)
}
if c.cur != nil {
_ = c.cur.SetReadDeadline(t)
}
return nil
return c.cur.SetReadDeadline(t)
}
func (c *udpHopConn) SetWriteDeadline(t time.Time) error {
@@ -280,10 +267,7 @@ func (c *udpHopConn) SetWriteDeadline(t time.Time) error {
if c.pre != nil {
_ = c.pre.SetWriteDeadline(t)
}
if c.cur != nil {
_ = c.cur.SetWriteDeadline(t)
}
return nil
return c.cur.SetWriteDeadline(t)
}
func randPrefix(p netip.Prefix) []byte {
+6 -13
View File
@@ -1,21 +1,14 @@
package xdns
import (
"net"
"github.com/xtls/xray-core/common/net"
"github.com/xtls/xray-core/transport/internet/finalmask"
)
func (c *Config) WrapPacketConnClient(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error) {
// _, ok1 := raw.(*internet.FakePacketConn)
// _, ok2 := raw.(*udphop.UdpHopPacketConn)
// if level != 0 || ok1 || ok2 {
// return nil, errors.New("xdns requires being at the outermost level")
// }
return NewConnClient(c, raw)
func (c *Config) WrapPacketConnClient(conn net.PacketConn, dest *net.Destination, dialer *finalmask.Dialer) (net.PacketConn, error) {
return NewConnClient(c, conn)
}
func (c *Config) WrapPacketConnServer(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error) {
// if level != 0 {
// return nil, errors.New("xdns requires being at the outermost level")
// }
return NewConnServer(c, raw)
func (c *Config) WrapPacketConnServer(conn net.PacketConn, addr net.Addr, lc *finalmask.ListenConfig) (net.PacketConn, error) {
return NewConnServer(c, conn)
}
+28 -34
View File
@@ -8,8 +8,7 @@ import (
goerrors "errors"
"fmt"
"io"
mathrand "math/rand"
"net"
mrand "math/rand"
"net/netip"
"sync"
"time"
@@ -17,6 +16,7 @@ import (
"github.com/xtls/xray-core/common"
"github.com/xtls/xray-core/common/errors"
"github.com/xtls/xray-core/common/net"
"github.com/xtls/xray-core/transport/internet/finalmask"
"golang.org/x/net/icmp"
"golang.org/x/net/ipv4"
@@ -36,11 +36,11 @@ type packet struct {
}
type xicmpConnClient struct {
conn net.PacketConn
icmp4 *icmp.PacketConn
icmp6 *icmp.PacketConn
udp bool
ips []netip.Addr
ip net.IP
clientID [8]byte
id int
seq int
@@ -50,7 +50,7 @@ type xicmpConnClient struct {
mu sync.Mutex
}
func NewConnClient(c *Config, raw net.PacketConn) (net.PacketConn, error) {
func NewConnClient(c *Config, dest *net.Destination) (net.PacketConn, error) {
var icmp4, icmp6 *icmp.PacketConn
var err4, err6 error
if c.DGRAM {
@@ -69,17 +69,24 @@ func NewConnClient(c *Config, raw net.PacketConn) (net.PacketConn, error) {
ips = append(ips, netip.MustParseAddr(ip))
}
var ip net.IP
if len(ips) > 0 {
ip = ips[mrand.Intn(len(ips))].AsSlice()
} else {
ip = dest.Address.IP()
}
var clientID [8]byte
common.Must2(rand.Read(clientID[:]))
conn := &xicmpConnClient{
conn: raw,
icmp4: icmp4,
icmp6: icmp6,
udp: c.DGRAM,
ips: ips,
ip: ip,
clientID: clientID,
id: mathrand.Intn(65536),
id: mrand.Intn(65536),
seq: 1,
readCh: make(chan packet),
closeCh: make(chan struct{}),
@@ -92,10 +99,6 @@ func NewConnClient(c *Config, raw net.PacketConn) (net.PacketConn, error) {
return conn, nil
}
func (c *xicmpConnClient) ring(a, b uint16) uint16 {
return min(a-b, b-a)
}
func (c *xicmpConnClient) closed() bool {
select {
case <-c.closeCh:
@@ -110,12 +113,11 @@ func (c *xicmpConnClient) recv4() {
var b [finalmask.UDPSize]byte
for {
if c.closed() {
return
}
n, addr, err := c.icmp4.ReadFrom(b[:])
if err != nil {
if c.closed() {
return
}
var netErr net.Error
if goerrors.As(err, &netErr) && netErr.Timeout() {
select {
@@ -125,9 +127,10 @@ func (c *xicmpConnClient) recv4() {
case <-c.closeCh:
return
}
continue
}
errors.LogErrorInner(context.Background(), err, "recv4 err")
continue
errors.LogErrorInner(context.Background(), err, "recv err 4")
return
}
msg, err := icmp.ParseMessage(1, b[:n])
@@ -150,10 +153,6 @@ func (c *xicmpConnClient) recv4() {
continue
}
if c.ring(uint16(echo.Seq), uint16(c.seq)) > 1000 {
continue
}
if len(echo.Data) > 8 && bytes.Equal(echo.Data[:8], c.clientID[:]) {
continue
}
@@ -182,12 +181,11 @@ func (c *xicmpConnClient) recv6() {
var b [finalmask.UDPSize]byte
for {
if c.closed() {
return
}
n, addr, err := c.icmp6.ReadFrom(b[:])
if err != nil {
if c.closed() {
return
}
var netErr net.Error
if goerrors.As(err, &netErr) && netErr.Timeout() {
select {
@@ -197,9 +195,10 @@ func (c *xicmpConnClient) recv6() {
case <-c.closeCh:
return
}
continue
}
errors.LogErrorInner(context.Background(), err, "recv6 err")
continue
errors.LogErrorInner(context.Background(), err, "recv err 6")
return
}
msg, err := icmp.ParseMessage(58, b[:n])
@@ -222,10 +221,6 @@ func (c *xicmpConnClient) recv6() {
continue
}
if c.ring(uint16(echo.Seq), uint16(c.seq)) > 1000 {
continue
}
if len(echo.Data) > 8 && bytes.Equal(echo.Data[:8], c.clientID[:]) {
continue
}
@@ -273,9 +268,9 @@ func (c *xicmpConnClient) WriteTo(p []byte, addr net.Addr) (n int, err error) {
c.seq %= 65536
c.mu.Unlock()
ip := addr.(*net.UDPAddr).IP
ip := c.ip
if len(c.ips) > 0 {
ip = c.ips[mathrand.Intn(len(c.ips))].AsSlice()
ip = c.ips[mrand.Intn(len(c.ips))].AsSlice()
}
if c.udp {
@@ -314,7 +309,6 @@ func (c *xicmpConnClient) Close() error {
close(c.closeCh)
_ = c.icmp4.Close()
_ = c.icmp6.Close()
_ = c.conn.Close()
c.wg.Wait()
select {
case p := <-c.readCh:
@@ -328,7 +322,7 @@ func (c *xicmpConnClient) Close() error {
}
func (c *xicmpConnClient) LocalAddr() net.Addr {
return c.conn.LocalAddr()
return &net.UDPAddr{IP: []byte{0, 0, 0, 0}}
}
func (c *xicmpConnClient) SetDeadline(t time.Time) error {
+13 -13
View File
@@ -1,23 +1,23 @@
package xicmp
import (
"net"
"errors"
"github.com/xtls/xray-core/common/errors"
"github.com/xtls/xray-core/transport/internet"
"github.com/xtls/xray-core/common/net"
"github.com/xtls/xray-core/transport/internet/finalmask"
)
func (c *Config) WrapPacketConnClient(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error) {
_, ok1 := raw.(*internet.FakePacketConn)
if level != 0 || ok1 {
return nil, errors.New("xicmp requires being at the outermost level")
func (c *Config) HandleDial() {}
func (c *Config) HandleListen() {}
func (c *Config) WrapPacketConnClient(conn net.PacketConn, dest *net.Destination, dialer *finalmask.Dialer) (net.PacketConn, error) {
if dest.Address.Family().IsDomain() && len(c.IPs) == 0 {
return nil, errors.New("empty ip addresses")
}
return NewConnClient(c, raw)
return NewConnClient(c, dest)
}
func (c *Config) WrapPacketConnServer(raw net.PacketConn, level int, levelCount int) (net.PacketConn, error) {
if level != 0 {
return nil, errors.New("xicmp requires being at the outermost level")
}
return NewConnServer(c, raw)
func (c *Config) WrapPacketConnServer(conn net.PacketConn, addr net.Addr, lc *finalmask.ListenConfig) (net.PacketConn, error) {
return NewConnServer(c)
}
+14 -17
View File
@@ -37,7 +37,6 @@ type record struct {
}
type xicmpConnServer struct {
conn net.PacketConn
icmp4 *icmp.PacketConn
icmp6 *icmp.PacketConn
ips map[netip.Addr]struct{}
@@ -48,7 +47,7 @@ type xicmpConnServer struct {
mu sync.Mutex
}
func NewConnServer(c *Config, raw net.PacketConn) (net.PacketConn, error) {
func NewConnServer(c *Config) (net.PacketConn, error) {
icmp4, err := icmp.ListenPacket("ip4:icmp", "0.0.0.0")
if err != nil {
return nil, err
@@ -64,7 +63,6 @@ func NewConnServer(c *Config, raw net.PacketConn) (net.PacketConn, error) {
}
conn := &xicmpConnServer{
conn: raw,
icmp4: icmp4,
icmp6: icmp6,
ips: ips,
@@ -115,12 +113,11 @@ func (c *xicmpConnServer) recv4() {
var b [finalmask.UDPSize]byte
for {
if c.closed() {
return
}
n, addr, err := c.icmp4.ReadFrom(b[:])
if err != nil {
if c.closed() {
return
}
var netErr net.Error
if goerrors.As(err, &netErr) && netErr.Timeout() {
select {
@@ -130,9 +127,10 @@ func (c *xicmpConnServer) recv4() {
case <-c.closeCh:
return
}
continue
}
errors.LogErrorInner(context.Background(), err, "recv4 err")
continue
errors.LogErrorInner(context.Background(), err, "recv err 4")
return
}
msg, err := icmp.ParseMessage(1, b[:n])
@@ -195,12 +193,11 @@ func (c *xicmpConnServer) recv6() {
var b [finalmask.UDPSize]byte
for {
if c.closed() {
return
}
n, addr, err := c.icmp6.ReadFrom(b[:])
if err != nil {
if c.closed() {
return
}
var netErr net.Error
if goerrors.As(err, &netErr) && netErr.Timeout() {
select {
@@ -210,9 +207,10 @@ func (c *xicmpConnServer) recv6() {
case <-c.closeCh:
return
}
continue
}
errors.LogErrorInner(context.Background(), err, "recv6 err")
continue
errors.LogErrorInner(context.Background(), err, "recv err 6")
return
}
msg, err := icmp.ParseMessage(58, b[:n])
@@ -330,7 +328,6 @@ func (c *xicmpConnServer) Close() error {
close(c.closeCh)
_ = c.icmp4.Close()
_ = c.icmp6.Close()
_ = c.conn.Close()
c.wg.Wait()
select {
case p := <-c.readCh:
@@ -344,7 +341,7 @@ func (c *xicmpConnServer) Close() error {
}
func (c *xicmpConnServer) LocalAddr() net.Addr {
return c.conn.LocalAddr()
return &net.UDPAddr{IP: []byte{0, 0, 0, 0}}
}
func (c *xicmpConnServer) SetDeadline(t time.Time) error {
@@ -39,7 +39,6 @@ type record struct {
}
type xicmpConnServer struct {
conn net.PacketConn
icmp4 *icmp.PacketConn
icmp6 *icmp.PacketConn
ipv4PC *ipv4.PacketConn
@@ -52,7 +51,7 @@ type xicmpConnServer struct {
mu sync.Mutex
}
func NewConnServer(c *Config, raw net.PacketConn) (net.PacketConn, error) {
func NewConnServer(c *Config) (net.PacketConn, error) {
icmp4, err := icmp.ListenPacket("ip4:icmp", "0.0.0.0")
if err != nil {
return nil, err
@@ -68,7 +67,6 @@ func NewConnServer(c *Config, raw net.PacketConn) (net.PacketConn, error) {
}
conn := &xicmpConnServer{
conn: raw,
icmp4: icmp4,
icmp6: icmp6,
ipv4PC: icmp4.IPv4PacketConn(),
@@ -124,12 +122,11 @@ func (c *xicmpConnServer) recv4() {
var b [finalmask.UDPSize]byte
for {
if c.closed() {
return
}
n, cm, addr, err := c.ipv4PC.ReadFrom(b[:])
if err != nil {
if c.closed() {
return
}
var netErr net.Error
if goerrors.As(err, &netErr) && netErr.Timeout() {
select {
@@ -139,9 +136,10 @@ func (c *xicmpConnServer) recv4() {
case <-c.closeCh:
return
}
continue
}
errors.LogErrorInner(context.Background(), err, "recv4 err")
continue
errors.LogErrorInner(context.Background(), err, "recv err 4")
return
}
msg, err := icmp.ParseMessage(1, b[:n])
@@ -205,12 +203,11 @@ func (c *xicmpConnServer) recv6() {
var b [finalmask.UDPSize]byte
for {
if c.closed() {
return
}
n, cm, addr, err := c.ipv6PC.ReadFrom(b[:])
if err != nil {
if c.closed() {
return
}
var netErr net.Error
if goerrors.As(err, &netErr) && netErr.Timeout() {
select {
@@ -220,9 +217,10 @@ func (c *xicmpConnServer) recv6() {
case <-c.closeCh:
return
}
continue
}
errors.LogErrorInner(context.Background(), err, "recv6 err")
continue
errors.LogErrorInner(context.Background(), err, "recv err 6")
return
}
msg, err := icmp.ParseMessage(58, b[:n])
@@ -341,7 +339,6 @@ func (c *xicmpConnServer) Close() error {
close(c.closeCh)
_ = c.icmp4.Close()
_ = c.icmp6.Close()
_ = c.conn.Close()
c.wg.Wait()
select {
case p := <-c.readCh:
@@ -355,7 +352,7 @@ func (c *xicmpConnServer) Close() error {
}
func (c *xicmpConnServer) LocalAddr() net.Addr {
return c.conn.LocalAddr()
return &net.UDPAddr{IP: []byte{0, 0, 0, 0}}
}
func (c *xicmpConnServer) SetDeadline(t time.Time) error {
+4 -2
View File
@@ -2,10 +2,12 @@ package xmc
import (
"fmt"
"net"
"github.com/xtls/xray-core/common/net"
"github.com/xtls/xray-core/transport/internet/finalmask"
)
func (c *Config) WrapConnClient(conn net.Conn) (net.Conn, error) {
func (c *Config) WrapConnClient(conn net.Conn, dest *net.Destination, dialer *finalmask.Dialer) (net.Conn, error) {
profiles, err := profilesFromConfig(c.Profiles)
if err != nil {
return nil, fmt.Errorf("minecraft finalmask: %w", err)
+6 -11
View File
@@ -83,7 +83,6 @@ func getGrpcClient(ctx context.Context, dest net.Destination, streamSettings *in
}
tlsConfig := tls.ConfigFromStreamSettings(streamSettings)
realityConfig := reality.ConfigFromStreamSettings(streamSettings)
sockopt := streamSettings.SocketSettings
grpcSettings := streamSettings.ProtocolSettings.(*Config)
if client, found := globalDialerMap[dialerConf{dest, streamSettings}]; found && client.GetState() != connectivity.Shutdown {
@@ -124,17 +123,13 @@ func getGrpcClient(ctx context.Context, dest net.Destination, streamSettings *in
gctx = session.ContextWithOutbounds(gctx, session.OutboundsFromContext(ctx))
gctx = session.ContextWithTimeoutOnly(gctx, true)
c, err := internet.DialSystem(gctx, net.TCPDestination(address, port), sockopt)
var c net.Conn
if streamSettings.FinalMask != nil {
c, err = streamSettings.FinalMask.DialTCP(gctx, net.TCPDestination(address, port))
} else {
c, err = internet.DialSystem(ctx, dest, streamSettings.SocketSettings)
}
if err == nil {
if streamSettings.TcpmaskManager != nil {
newConn, err := streamSettings.TcpmaskManager.WrapConnClient(c)
if err != nil {
c.Close()
return nil, errors.New("mask err").Base(err)
}
c = newConn
}
if tlsConfig != nil {
config := tlsConfig.GetTLSConfig(tls.WithDestination(dest))
if fingerprint := tls.GetFingerprint(tlsConfig.Fingerprint); fingerprint != nil {
+11 -19
View File
@@ -104,28 +104,20 @@ func Listen(ctx context.Context, address net.Address, port net.Port, settings *i
go func() {
var streamListener net.Listener
var err error
var addr net.Addr
if port == net.Port(0) { // unix
streamListener, err = internet.ListenSystem(ctx, &net.UnixAddr{
Name: address.Domain(),
Net: "unix",
}, settings.SocketSettings)
if err != nil {
errors.LogErrorInner(ctx, err, "failed to listen on ", address)
return
}
addr = &net.UnixAddr{Name: address.Domain(), Net: "unix"}
} else { // tcp
streamListener, err = internet.ListenSystem(ctx, &net.TCPAddr{
IP: address.IP(),
Port: int(port),
}, settings.SocketSettings)
if err != nil {
errors.LogErrorInner(ctx, err, "failed to listen on ", address, ":", port)
return
}
addr = &net.TCPAddr{IP: address.IP(), Port: int(port)}
}
if settings.TcpmaskManager != nil {
streamListener, _ = settings.TcpmaskManager.WrapListener(streamListener)
if settings.FinalMask != nil {
streamListener, err = settings.FinalMask.Listen(ctx, addr)
} else {
streamListener, err = internet.ListenSystem(ctx, addr, settings.SocketSettings)
}
if err != nil {
errors.LogErrorInner(ctx, err, "failed to listen on ", address, ":", port)
return
}
errors.LogDebug(ctx, "gRPC listen for service name `"+grpcSettings.getServiceName()+"` tun `"+grpcSettings.getTunStreamName()+"` multi tun `"+grpcSettings.getTunMultiStreamName()+"`")
+7 -10
View File
@@ -46,21 +46,18 @@ func (c *ConnRF) Read(b []byte) (int, error) {
func dialhttpUpgrade(ctx context.Context, dest net.Destination, streamSettings *internet.MemoryStreamConfig) (net.Conn, error) {
transportConfiguration := streamSettings.ProtocolSettings.(*Config)
pconn, err := internet.DialSystem(ctx, dest, streamSettings.SocketSettings)
var pconn net.Conn
var err error
if streamSettings.FinalMask != nil {
pconn, err = streamSettings.FinalMask.DialTCP(ctx, dest)
} else {
pconn, err = internet.DialSystem(ctx, dest, streamSettings.SocketSettings)
}
if err != nil {
errors.LogErrorInner(ctx, err, "failed to dial to ", dest)
return nil, err
}
if streamSettings.TcpmaskManager != nil {
newConn, err := streamSettings.TcpmaskManager.WrapConnClient(pconn)
if err != nil {
pconn.Close()
return nil, errors.New("mask err").Base(err)
}
pconn = newConn
}
var conn net.Conn
var requestURL url.URL
tConfig := tls.ConfigFromStreamSettings(streamSettings)
+11 -19
View File
@@ -124,29 +124,21 @@ func ListenHTTPUpgrade(ctx context.Context, address net.Address, port net.Port,
}
var listener net.Listener
var err error
var addr net.Addr
if port == net.Port(0) { // unix
listener, err = internet.ListenSystem(ctx, &net.UnixAddr{
Name: address.Domain(),
Net: "unix",
}, streamSettings.SocketSettings)
if err != nil {
return nil, errors.New("failed to listen unix domain socket(for HttpUpgrade) on ", address).Base(err)
}
errors.LogInfo(ctx, "listening unix domain socket(for HttpUpgrade) on ", address)
addr = &net.UnixAddr{Name: address.Domain(), Net: "unix"}
} else { // tcp
listener, err = internet.ListenSystem(ctx, &net.TCPAddr{
IP: address.IP(),
Port: int(port),
}, streamSettings.SocketSettings)
if err != nil {
return nil, errors.New("failed to listen TCP(for HttpUpgrade) on ", address, ":", port).Base(err)
}
errors.LogInfo(ctx, "listening TCP(for HttpUpgrade) on ", address, ":", port)
addr = &net.TCPAddr{IP: address.IP(), Port: int(port)}
}
if streamSettings.TcpmaskManager != nil {
listener, _ = streamSettings.TcpmaskManager.WrapListener(listener)
if streamSettings.FinalMask != nil {
listener, err = streamSettings.FinalMask.Listen(ctx, addr)
} else {
listener, err = internet.ListenSystem(ctx, addr, streamSettings.SocketSettings)
}
if err != nil {
return nil, errors.New("failed to listen ", addr.Network(), "(for HttpUpgrade) on ", address, ":", port).Base(err)
}
errors.LogInfo(ctx, "listening ", addr.Network(), "(for HttpUpgrade) on ", address, ":", port)
if streamSettings.SocketSettings != nil && streamSettings.SocketSettings.AcceptProxyProtocol {
errors.LogWarning(ctx, "accepting PROXY protocol")
+35 -36
View File
@@ -2,7 +2,7 @@ package hysteria
import (
"context"
go_tls "crypto/tls"
gotls "crypto/tls"
"net/http"
"net/url"
"reflect"
@@ -28,12 +28,12 @@ import (
type client struct {
sync.Mutex
dest net.Destination
config *Config
tlsConfig *go_tls.Config
socketConfig *internet.SocketConfig
udpmaskManager *finalmask.UdpmaskManager
quicParams *internet.QuicParams
dest net.Destination
config *Config
tlsConfig *gotls.Config
socketConfig *internet.SocketConfig
finalMask *finalmask.FinalMask
quicParams *internet.QuicParams
conn *quic.Conn
tr *quic.Transport
@@ -113,30 +113,29 @@ func (c *client) dial(ctx context.Context) error {
// }
var pktConn net.PacketConn
var udpAddr *net.UDPAddr
raw, err := internet.DialSystem(ctx, c.dest, c.socketConfig)
if err != nil {
return errors.New("failed to dial to dest").Base(err)
}
switch c := raw.(type) {
case *internet.PacketConnWrapper:
pktConn = c.PacketConn
udpAddr = raw.RemoteAddr().(*net.UDPAddr)
case *cnc.Connection:
pktConn = &internet.FakePacketConn{Conn: c}
udpAddr = &net.UDPAddr{IP: c.RemoteAddr().(*net.TCPAddr).IP, Port: c.RemoteAddr().(*net.TCPAddr).Port}
default:
panic(reflect.TypeOf(c))
}
if c.udpmaskManager != nil {
newConn, err := c.udpmaskManager.WrapPacketConnClient(pktConn)
var udpAddr net.Addr
if c.finalMask != nil {
conn, err := c.finalMask.DialUDP(ctx, c.dest)
if err != nil {
pktConn.Close()
return errors.New("mask err").Base(err)
return errors.New("failed to dial to dest").Base(err)
}
pktConn = conn.(*finalmask.PacketConnWrapper).PacketConn
udpAddr = conn.RemoteAddr()
} else {
conn, err := internet.DialSystem(ctx, c.dest, c.socketConfig)
if err != nil {
return errors.New("failed to dial to dest").Base(err)
}
switch c := conn.(type) {
case *internet.PacketConnWrapper:
pktConn = c.PacketConn
udpAddr = c.RemoteAddr()
case *cnc.Connection:
pktConn = &internet.FakePacketConn{Conn: c}
udpAddr = &net.UDPAddr{IP: []byte{0, 0, 0, 0}}
default:
panic(reflect.TypeOf(c))
}
pktConn = newConn
}
tr := &quic.Transport{Conn: pktConn, DisableGSO: quicParams.DisableGSO}
@@ -150,7 +149,7 @@ func (c *client) dial(ctx context.Context) error {
rt := &http3.Transport{
TLSClientConfig: c.tlsConfig,
QUICConfig: quicConfig,
Dial: func(ctx context.Context, _ string, tlsCfg *go_tls.Config, cfg *quic.Config) (*quic.Conn, error) {
Dial: func(ctx context.Context, _ string, tlsCfg *gotls.Config, cfg *quic.Config) (*quic.Conn, error) {
qc, err := tr.DialEarly(ctx, udpAddr, tlsCfg, cfg)
if err != nil {
return nil, err
@@ -316,12 +315,12 @@ func Dial(ctx context.Context, dest net.Destination, streamSettings *internet.Me
c = manager.m[dialerConf{dest, streamSettings}]
if c == nil {
c = &client{
dest: dest,
config: streamSettings.ProtocolSettings.(*Config),
tlsConfig: tlsConfig.GetTLSConfig(tls.WithDestination(dest)),
socketConfig: streamSettings.SocketSettings,
udpmaskManager: streamSettings.UdpmaskManager,
quicParams: streamSettings.QuicParams,
dest: dest,
config: streamSettings.ProtocolSettings.(*Config),
tlsConfig: tlsConfig.GetTLSConfig(tls.WithDestination(dest)),
socketConfig: streamSettings.SocketSettings,
finalMask: streamSettings.FinalMask,
quicParams: streamSettings.QuicParams,
}
manager.m[dialerConf{dest, streamSettings}] = c
}
+7 -10
View File
@@ -316,20 +316,17 @@ func Listen(ctx context.Context, address net.Address, port net.Port, streamSetti
quicConfig.MaxIncomingStreams = 1024
}
pktConn, err := internet.ListenSystemPacket(context.Background(), &net.UDPAddr{IP: address.IP(), Port: int(port)}, streamSettings.SocketSettings)
var pktConn net.PacketConn
var err error
if streamSettings.FinalMask != nil {
pktConn, err = streamSettings.FinalMask.ListenPacket(context.Background(), &net.UDPAddr{IP: address.IP(), Port: int(port)})
} else {
pktConn, err = internet.ListenSystemPacket(context.Background(), &net.UDPAddr{IP: address.IP(), Port: int(port)}, streamSettings.SocketSettings)
}
if err != nil {
return nil, err
}
if streamSettings.UdpmaskManager != nil {
newConn, err := streamSettings.UdpmaskManager.WrapPacketConnServer(pktConn)
if err != nil {
pktConn.Close()
return nil, errors.New("mask err").Base(err)
}
pktConn = newConn
}
var k *quic.StatelessResetKey
if !quicParams.DisableStatelessReset {
k = &quic.StatelessResetKey{}
@@ -0,0 +1,165 @@
package hysteria
import (
"context"
"crypto/tls"
"crypto/x509"
"errors"
"net"
"runtime"
"testing"
"time"
"github.com/apernet/quic-go"
"github.com/xtls/xray-core/common"
"github.com/xtls/xray-core/common/protocol/tls/cert"
)
func TestDatagram(t *testing.T) {
run := func() (addr net.Addr, recv chan int64, cancel func()) {
cert, _ := cert.MustGenerate(nil)
Certificate := [][]byte{cert.Certificate}
PrivateKey := common.Must2(x509.ParsePKCS8PrivateKey(cert.PrivateKey))
tlsConf := &tls.Config{
Certificates: []tls.Certificate{
{
Certificate: Certificate,
PrivateKey: PrivateKey,
},
},
NextProtos: []string{"h3"},
}
quicConf := &quic.Config{
InitialStreamReceiveWindow: 8388608,
MaxStreamReceiveWindow: 8388608,
InitialConnectionReceiveWindow: 8388608 * 5 / 2,
MaxConnectionReceiveWindow: 8388608 * 5 / 2,
MaxIdleTimeout: 30 * time.Second,
MaxIncomingStreams: 1024,
DisablePathMTUDiscovery: runtime.GOOS != "linux" && runtime.GOOS != "windows" && runtime.GOOS != "darwin",
EnableDatagrams: true,
MaxDatagramFrameSize: MaxDatagramFrameSize,
AssumePeerMaxDatagramFrameSize: MaxDatagramFrameSize,
DisablePathManager: true,
}
pktConn := common.Must2(net.ListenPacket("udp", "127.0.0.1:0"))
tr := &quic.Transport{Conn: pktConn}
l := common.Must2(tr.Listen(tlsConf, quicConf))
recv = make(chan int64)
ctx, cancel := context.WithCancel(context.Background())
go func() {
defer pktConn.Close()
defer tr.Close()
defer l.Close()
defer close(recv)
var buf [1500]byte
for {
conn, err := l.Accept(ctx)
if err != nil {
if !errors.Is(err, context.Canceled) {
t.Error(err)
}
break
}
err = conn.SendDatagram(buf[:])
var qErr *quic.DatagramTooLargeError
if !errors.As(err, &qErr) {
t.Error(err)
}
recv <- qErr.MaxDatagramPayloadSize
defer conn.CloseWithError(0, "")
}
}()
return l.Addr(), recv, cancel
}
addr, recv, cancel := run()
t.Run("With ChromeParrot", func(t *testing.T) {
tlsConf := &tls.Config{
InsecureSkipVerify: true,
}
quicConf := &quic.Config{
InitialStreamReceiveWindow: 8388608,
MaxStreamReceiveWindow: 8388608,
InitialConnectionReceiveWindow: 8388608 * 5 / 2,
MaxConnectionReceiveWindow: 8388608 * 5 / 2,
MaxIdleTimeout: 30 * time.Second,
KeepAlivePeriod: 10 * time.Second,
DisablePathMTUDiscovery: runtime.GOOS != "linux" && runtime.GOOS != "windows" && runtime.GOOS != "darwin",
ChromeParrot: true,
EnableDatagrams: true,
MaxDatagramFrameSize: MaxDatagramFrameSize,
OmitMaxDatagramFrameSize: true,
DisablePathManager: true,
}
pktConn := common.Must2(net.ListenPacket("udp", "127.0.0.1:0"))
tr := &quic.Transport{Conn: pktConn, ConnectionIDGenerator: quic.ZeroLengthConnectionIDGenerator{}}
conn := common.Must2(tr.DialEarly(context.Background(), addr, tlsConf, quicConf))
defer pktConn.Close()
defer tr.Close()
defer conn.CloseWithError(0, "")
var buf [1500]byte
err := conn.SendDatagram(buf[:])
var qErr *quic.DatagramTooLargeError
if !errors.As(err, &qErr) || qErr.MaxDatagramPayloadSize != 1197 {
t.Error(err)
}
if server := <-recv; server != 1243 {
t.Error(server)
}
})
t.Run("Without ChromeParrot", func(t *testing.T) {
tlsConf := &tls.Config{
InsecureSkipVerify: true,
NextProtos: []string{"h3"},
}
quicConf := &quic.Config{
InitialStreamReceiveWindow: 8388608,
MaxStreamReceiveWindow: 8388608,
InitialConnectionReceiveWindow: 8388608 * 5 / 2,
MaxConnectionReceiveWindow: 8388608 * 5 / 2,
MaxIdleTimeout: 30 * time.Second,
KeepAlivePeriod: 10 * time.Second,
DisablePathMTUDiscovery: runtime.GOOS != "linux" && runtime.GOOS != "windows" && runtime.GOOS != "darwin",
ChromeParrot: false,
EnableDatagrams: true,
MaxDatagramFrameSize: MaxDatagramFrameSize,
OmitMaxDatagramFrameSize: true,
DisablePathManager: true,
}
pktConn := common.Must2(net.ListenPacket("udp", "127.0.0.1:0"))
tr := &quic.Transport{Conn: pktConn}
conn := common.Must2(tr.DialEarly(context.Background(), addr, tlsConf, quicConf))
defer pktConn.Close()
defer tr.Close()
defer conn.CloseWithError(0, "")
var buf [1500]byte
err := conn.SendDatagram(buf[:])
var qErr *quic.DatagramTooLargeError
if !errors.As(err, &qErr) || qErr.MaxDatagramPayloadSize != 1197 {
t.Error(err)
}
if server := <-recv; server != 1197 {
t.Error(server)
}
})
cancel()
}
+7 -28
View File
@@ -3,7 +3,6 @@ package kcp
import (
"context"
"io"
reflect "reflect"
"sync/atomic"
"github.com/xtls/xray-core/common"
@@ -11,7 +10,6 @@ import (
"github.com/xtls/xray-core/common/dice"
"github.com/xtls/xray-core/common/errors"
"github.com/xtls/xray-core/common/net"
"github.com/xtls/xray-core/common/net/cnc"
"github.com/xtls/xray-core/transport/internet"
"github.com/xtls/xray-core/transport/internet/stat"
"github.com/xtls/xray-core/transport/internet/tls"
@@ -51,36 +49,17 @@ func DialKCP(ctx context.Context, dest net.Destination, streamSettings *internet
dest.Network = net.Network_UDP
errors.LogInfo(ctx, "dialing mKCP to ", dest)
conn, err := internet.DialSystem(ctx, dest, streamSettings.SocketSettings)
var conn net.Conn
var err error
if streamSettings.FinalMask != nil {
conn, err = streamSettings.FinalMask.DialUDP(ctx, dest)
} else {
conn, err = internet.DialSystem(ctx, dest, streamSettings.SocketSettings)
}
if err != nil {
return nil, errors.New("failed to dial to dest: ", err).AtWarning().Base(err)
}
if streamSettings.UdpmaskManager != nil {
var pktConn net.PacketConn
var udpAddr *net.UDPAddr
switch c := conn.(type) {
case *internet.PacketConnWrapper:
pktConn = c.PacketConn
udpAddr = c.RemoteAddr().(*net.UDPAddr)
case *cnc.Connection:
pktConn = &internet.FakePacketConn{Conn: c}
udpAddr = &net.UDPAddr{IP: c.RemoteAddr().(*net.TCPAddr).IP, Port: c.RemoteAddr().(*net.TCPAddr).Port}
default:
panic(reflect.TypeOf(c))
}
newConn, err := streamSettings.UdpmaskManager.WrapPacketConnClient(pktConn)
if err != nil {
pktConn.Close()
return nil, errors.New("mask err").Base(err)
}
pktConn = newConn
conn = &internet.PacketConnWrapper{
PacketConn: pktConn,
Dest: udpAddr,
}
}
kcpSettings := streamSettings.ProtocolSettings.(*Config)
reader := &KCPPacketReader{}
+18
View File
@@ -0,0 +1,18 @@
package masque
import (
"github.com/xtls/xray-core/common"
"github.com/xtls/xray-core/transport/internet"
)
const protocolName = "masque"
const DefaultPath = "/.well-known/masque/ip/*/*/"
func init() {
common.Must(internet.RegisterProtocolConfigCreator(protocolName, func() interface{} {
return &Config{
Path: DefaultPath,
}
}))
}
+146
View File
@@ -0,0 +1,146 @@
// Code generated by protoc-gen-go. DO NOT EDIT.
// versions:
// protoc-gen-go v1.36.11
// protoc v6.33.5
// source: transport/internet/masque/config.proto
package masque
import (
protoreflect "google.golang.org/protobuf/reflect/protoreflect"
protoimpl "google.golang.org/protobuf/runtime/protoimpl"
reflect "reflect"
sync "sync"
unsafe "unsafe"
)
const (
// Verify that this generated code is sufficiently up-to-date.
_ = protoimpl.EnforceVersion(20 - protoimpl.MinVersion)
// Verify that runtime/protoimpl is sufficiently up-to-date.
_ = protoimpl.EnforceVersion(protoimpl.MaxVersion - 20)
)
type Config struct {
state protoimpl.MessageState `protogen:"open.v1"`
Host string `protobuf:"bytes,1,opt,name=host,proto3" json:"host,omitempty"`
Path string `protobuf:"bytes,2,opt,name=path,proto3" json:"path,omitempty"`
Headers map[string]string `protobuf:"bytes,3,rep,name=headers,proto3" json:"headers,omitempty" protobuf_key:"bytes,1,opt,name=key" protobuf_val:"bytes,2,opt,name=value"`
unknownFields protoimpl.UnknownFields
sizeCache protoimpl.SizeCache
}
func (x *Config) Reset() {
*x = Config{}
mi := &file_transport_internet_masque_config_proto_msgTypes[0]
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
ms.StoreMessageInfo(mi)
}
func (x *Config) String() string {
return protoimpl.X.MessageStringOf(x)
}
func (*Config) ProtoMessage() {}
func (x *Config) ProtoReflect() protoreflect.Message {
mi := &file_transport_internet_masque_config_proto_msgTypes[0]
if x != nil {
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
if ms.LoadMessageInfo() == nil {
ms.StoreMessageInfo(mi)
}
return ms
}
return mi.MessageOf(x)
}
// Deprecated: Use Config.ProtoReflect.Descriptor instead.
func (*Config) Descriptor() ([]byte, []int) {
return file_transport_internet_masque_config_proto_rawDescGZIP(), []int{0}
}
func (x *Config) GetHost() string {
if x != nil {
return x.Host
}
return ""
}
func (x *Config) GetPath() string {
if x != nil {
return x.Path
}
return ""
}
func (x *Config) GetHeaders() map[string]string {
if x != nil {
return x.Headers
}
return nil
}
var File_transport_internet_masque_config_proto protoreflect.FileDescriptor
const file_transport_internet_masque_config_proto_rawDesc = "" +
"\n" +
"&transport/internet/masque/config.proto\x12\x1exray.transport.internet.masque\"\xbb\x01\n" +
"\x06Config\x12\x12\n" +
"\x04host\x18\x01 \x01(\tR\x04host\x12\x12\n" +
"\x04path\x18\x02 \x01(\tR\x04path\x12M\n" +
"\aheaders\x18\x03 \x03(\v23.xray.transport.internet.masque.Config.HeadersEntryR\aheaders\x1a:\n" +
"\fHeadersEntry\x12\x10\n" +
"\x03key\x18\x01 \x01(\tR\x03key\x12\x14\n" +
"\x05value\x18\x02 \x01(\tR\x05value:\x028\x01B|\n" +
"\"com.xray.transport.internet.masqueP\x01Z3github.com/xtls/xray-core/transport/internet/masque\xaa\x02\x1eXray.Transport.Internet.Masqueb\x06proto3"
var (
file_transport_internet_masque_config_proto_rawDescOnce sync.Once
file_transport_internet_masque_config_proto_rawDescData []byte
)
func file_transport_internet_masque_config_proto_rawDescGZIP() []byte {
file_transport_internet_masque_config_proto_rawDescOnce.Do(func() {
file_transport_internet_masque_config_proto_rawDescData = protoimpl.X.CompressGZIP(unsafe.Slice(unsafe.StringData(file_transport_internet_masque_config_proto_rawDesc), len(file_transport_internet_masque_config_proto_rawDesc)))
})
return file_transport_internet_masque_config_proto_rawDescData
}
var file_transport_internet_masque_config_proto_msgTypes = make([]protoimpl.MessageInfo, 2)
var file_transport_internet_masque_config_proto_goTypes = []any{
(*Config)(nil), // 0: xray.transport.internet.masque.Config
nil, // 1: xray.transport.internet.masque.Config.HeadersEntry
}
var file_transport_internet_masque_config_proto_depIdxs = []int32{
1, // 0: xray.transport.internet.masque.Config.headers:type_name -> xray.transport.internet.masque.Config.HeadersEntry
1, // [1:1] is the sub-list for method output_type
1, // [1:1] is the sub-list for method input_type
1, // [1:1] is the sub-list for extension type_name
1, // [1:1] is the sub-list for extension extendee
0, // [0:1] is the sub-list for field type_name
}
func init() { file_transport_internet_masque_config_proto_init() }
func file_transport_internet_masque_config_proto_init() {
if File_transport_internet_masque_config_proto != nil {
return
}
type x struct{}
out := protoimpl.TypeBuilder{
File: protoimpl.DescBuilder{
GoPackagePath: reflect.TypeOf(x{}).PkgPath(),
RawDescriptor: unsafe.Slice(unsafe.StringData(file_transport_internet_masque_config_proto_rawDesc), len(file_transport_internet_masque_config_proto_rawDesc)),
NumEnums: 0,
NumMessages: 2,
NumExtensions: 0,
NumServices: 0,
},
GoTypes: file_transport_internet_masque_config_proto_goTypes,
DependencyIndexes: file_transport_internet_masque_config_proto_depIdxs,
MessageInfos: file_transport_internet_masque_config_proto_msgTypes,
}.Build()
File_transport_internet_masque_config_proto = out.File
file_transport_internet_masque_config_proto_goTypes = nil
file_transport_internet_masque_config_proto_depIdxs = nil
}
+13
View File
@@ -0,0 +1,13 @@
syntax = "proto3";
package xray.transport.internet.masque;
option csharp_namespace = "Xray.Transport.Internet.Masque";
option go_package = "github.com/xtls/xray-core/transport/internet/masque";
option java_package = "com.xray.transport.internet.masque";
option java_multiple_files = true;
message Config {
string host = 1;
string path = 2;
map<string, string> headers = 3;
}
+115
View File
@@ -0,0 +1,115 @@
package masque
import (
"context"
go_errors "errors"
"net/netip"
"slices"
"sync"
"time"
"github.com/apernet/quic-go"
"github.com/apernet/quic-go/http3"
"github.com/xtls/xray-core/common/errors"
"github.com/xtls/xray-core/common/net"
"github.com/xtls/xray-core/transport/internet/masque/connectip"
)
type PacketTooBigError struct {
ICMP []byte
}
func (e *PacketTooBigError) Error() string {
return "packet too big for the tunnel"
}
type Conn struct {
ipConn *connectip.Conn
quicConn *quic.Conn
local []netip.Addr
closeOnce sync.Once
}
func (c *Conn) LocalAddrs() []netip.Addr {
return c.local
}
func (c *Conn) Read(b []byte) (int, error) {
return c.ipConn.ReadPacket(b)
}
func (c *Conn) Write(b []byte) (int, error) {
icmp, err := c.ipConn.WritePacket(b)
if err != nil {
if go_errors.Is(err, connectip.ErrMTUTooSmall) {
errors.LogWarning(context.Background(), "MASQUE: closing the tunnel as it cannot carry ", MinPacketSize, "-byte packets")
} else {
errors.LogInfoInner(context.Background(), err, "MASQUE: closing the tunnel as sending failed")
}
c.Close()
return 0, err
}
if len(icmp) > 0 {
return 0, &PacketTooBigError{ICMP: icmp}
}
return len(b), nil
}
func (c *Conn) Close() error {
c.closeOnce.Do(func() {
c.ipConn.Close()
c.quicConn.CloseWithError(quic.ApplicationErrorCode(http3.ErrCodeNoError), "")
})
return nil
}
func (c *Conn) LocalAddr() net.Addr {
return c.quicConn.LocalAddr()
}
func (c *Conn) RemoteAddr() net.Addr {
return c.quicConn.RemoteAddr()
}
func (c *Conn) SetDeadline(time.Time) error {
return nil
}
func (c *Conn) SetReadDeadline(time.Time) error {
return nil
}
func (c *Conn) SetWriteDeadline(time.Time) error {
return nil
}
func (c *Conn) serveAddressAssignments() {
for {
assigned, err := c.ipConn.ReceiveAddressAssignment(context.Background())
if err != nil {
return
}
for _, addr := range c.local {
if !slices.ContainsFunc(assigned, func(a connectip.AssignedAddress) bool { return !a.Rejected() && a.IPPrefix.Contains(addr) }) {
errors.LogInfo(context.Background(), "MASQUE: closing the tunnel as the proxy withdrew ", addr)
c.Close()
return
}
}
if len(localAddrs(assigned)) > len(c.local) {
errors.LogInfo(context.Background(), "MASQUE: the proxy assigned another IP family, which is used once the tunnel is set up again")
}
}
}
func (c *Conn) serveAddressRequests() {
for {
req, err := c.ipConn.ReceiveAddressRequest(context.Background())
if err != nil {
return
}
if err := req.Respond(make([]netip.Prefix, len(req.Prefixes)), nil); err != nil {
return
}
}
}
@@ -0,0 +1,7 @@
Copyright 2024 Marten Seemann
Permission is hereby granted, free of charge, to any person obtaining a copy of this software and associated documentation files (the "Software"), to deal in the Software without restriction, including without limitation the rights to use, copy, modify, merge, publish, distribute, sublicense, and/or sell copies of the Software, and to permit persons to whom the Software is furnished to do so, subject to the following conditions:
The above copyright notice and this permission notice shall be included in all copies or substantial portions of the Software.
THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE SOFTWARE.
@@ -0,0 +1,86 @@
/* SPDX-License-Identifier: MIT
*
* Copyright 2024 Marten Seemann
* Adapted from github.com/quic-go/connect-ip-go (commit a0c35fa).
*/
package connectip
import (
"errors"
"fmt"
"net/netip"
"slices"
"sync/atomic"
)
var (
rejectedIPv4Prefix = netip.PrefixFrom(netip.IPv4Unspecified(), 32)
rejectedIPv6Prefix = netip.PrefixFrom(netip.IPv6Unspecified(), 128)
)
type AddressRequestID uint64
type AddressRequest struct {
Prefixes []netip.Prefix
conn *Conn
requested *addressRequestCapsule
responded *atomic.Bool
}
func newAddressRequest(conn *Conn, requested *addressRequestCapsule) *AddressRequest {
return &AddressRequest{
Prefixes: slices.Clone(requested.Prefixes),
conn: conn,
requested: requested,
responded: &atomic.Bool{},
}
}
func (r *AddressRequest) Respond(assignments, additional []netip.Prefix) error {
if r.conn == nil {
return errors.New("connect-ip: invalid address request")
}
if len(assignments) != len(r.requested.RequestIDs) {
return fmt.Errorf(
"connect-ip: expected %d address assignments, got %d",
len(r.requested.RequestIDs),
len(assignments),
)
}
capsule := &addressAssignCapsule{
AssignedAddresses: make([]AssignedAddress, 0, len(assignments)+len(additional)),
}
var zeroPrefix netip.Prefix
for i, p := range assignments {
if p == zeroPrefix {
if r.requested.Prefixes[i].Addr().Is4() {
p = rejectedIPv4Prefix
} else {
p = rejectedIPv6Prefix
}
} else if !p.IsValid() || p != p.Masked() {
return fmt.Errorf("connect-ip: invalid assigned prefix %d: %s", i, p)
}
capsule.AssignedAddresses = append(
capsule.AssignedAddresses,
AssignedAddress{RequestID: r.requested.RequestIDs[i], IPPrefix: p},
)
}
for i, p := range additional {
if !p.IsValid() || p != p.Masked() {
return fmt.Errorf("connect-ip: invalid additional prefix %d: %s", i, p)
}
capsule.AssignedAddresses = append(capsule.AssignedAddresses, AssignedAddress{IPPrefix: p})
}
if !r.responded.CompareAndSwap(false, true) {
return errors.New("connect-ip: address request already answered")
}
restrictPeer := slices.ContainsFunc(capsule.AssignedAddresses, func(a AssignedAddress) bool { return !a.Rejected() })
if err := r.conn.sendAddressAssignment(capsule, restrictPeer); err != nil {
r.responded.Store(false)
return err
}
return nil
}
@@ -0,0 +1,85 @@
/* SPDX-License-Identifier: MIT
*
* Copyright 2024 Marten Seemann
* Adapted from github.com/quic-go/connect-ip-go (commit a0c35fa).
*/
package connectip
import (
"context"
"net/netip"
"testing"
"time"
"github.com/stretchr/testify/require"
)
func TestAddressRequests(t *testing.T) {
client, server := setupConns(t)
ctx, cancel := context.WithTimeout(context.Background(), time.Second)
defer cancel()
prefixes := []netip.Prefix{
netip.MustParsePrefix("0.0.0.0/32"),
netip.MustParsePrefix("0.0.0.0/32"),
netip.MustParsePrefix("::/64"),
}
ids, err := client.RequestAddresses(prefixes)
require.NoError(t, err)
require.Equal(t, []AddressRequestID{1, 2, 3}, ids)
req, err := server.ReceiveAddressRequest(ctx)
require.NoError(t, err)
require.Equal(t, prefixes, req.Prefixes)
assignments := []netip.Prefix{netip.MustParsePrefix("192.0.2.1/32"), {}, {}}
additional := []netip.Prefix{netip.MustParsePrefix("2001:db8::/64")}
require.NoError(t, req.Respond(assignments, additional))
received, err := client.ReceiveAddressAssignment(ctx)
require.NoError(t, err)
require.Len(t, received, 4)
require.Equal(t, AssignedAddress{RequestID: ids[0], IPPrefix: assignments[0]}, received[0])
require.Equal(t, ids[1], received[1].RequestID)
require.True(t, received[1].Rejected())
require.Equal(t, ids[2], received[2].RequestID)
require.True(t, received[2].Rejected())
require.Equal(t, AssignedAddress{IPPrefix: additional[0]}, received[3])
ids, err = client.RequestAddresses(prefixes[:1])
require.NoError(t, err)
require.Equal(t, []AddressRequestID{4}, ids)
}
func TestAddressRequestValidation(t *testing.T) {
conn := newProxiedConn(&mockStream{})
defer conn.Close()
for _, prefixes := range [][]netip.Prefix{
nil,
{{}},
{netip.MustParsePrefix("192.0.2.1/24")},
{netip.MustParsePrefix("2001:db8::1/64")},
} {
ids, err := conn.RequestAddresses(prefixes)
require.Error(t, err)
require.Nil(t, ids)
}
}
func TestAddressResponseValidation(t *testing.T) {
conn := newProxiedConn(&mockStream{})
defer conn.Close()
prefixes := []netip.Prefix{netip.MustParsePrefix("192.0.2.1/32")}
req := newAddressRequest(conn, &addressRequestCapsule{RequestIDs: []AddressRequestID{1}, Prefixes: prefixes})
require.ErrorContains(t, (&AddressRequest{}).Respond(nil, nil), "invalid address request")
require.ErrorContains(t, req.Respond(nil, nil), "expected 1 address assignments")
require.ErrorContains(t, req.Respond(prefixes, []netip.Prefix{{}}), "invalid additional prefix")
require.ErrorContains(t,
req.Respond([]netip.Prefix{netip.MustParsePrefix("192.0.2.1/24")}, nil),
"invalid assigned prefix",
)
copied := *req
require.NoError(t, req.Respond(prefixes, nil))
require.ErrorContains(t, copied.Respond(prefixes, nil), "already answered")
}
@@ -0,0 +1,308 @@
/* SPDX-License-Identifier: MIT
*
* Copyright 2024 Marten Seemann
* Adapted from github.com/quic-go/connect-ip-go (commit a0c35fa).
*/
package connectip
import (
"cmp"
"encoding/binary"
"errors"
"fmt"
"io"
"net/netip"
"github.com/apernet/quic-go/http3"
"github.com/apernet/quic-go/quicvarint"
)
const (
capsuleTypeDatagram http3.CapsuleType = 0
capsuleTypeAddressAssign http3.CapsuleType = 1
capsuleTypeAddressRequest http3.CapsuleType = 2
capsuleTypeRouteAdvertisement http3.CapsuleType = 3
)
const (
maxAddressesPerCapsule = 8192
maxRoutesPerCapsule = 8192
)
type addressAssignCapsule struct {
AssignedAddresses []AssignedAddress
}
type AssignedAddress struct {
RequestID AddressRequestID
IPPrefix netip.Prefix
}
func (a AssignedAddress) Rejected() bool {
return a.IPPrefix == rejectedIPv4Prefix || a.IPPrefix == rejectedIPv6Prefix
}
func (a AssignedAddress) len() int {
return quicvarint.Len(uint64(a.RequestID)) + 1 + a.IPPrefix.Addr().BitLen()/8 + 1
}
type addressRequestCapsule struct {
RequestIDs []AddressRequestID
Prefixes []netip.Prefix
}
func parseAddressAssignCapsule(r http3.CapsuleReader) (*addressAssignCapsule, error) {
var assignedAddresses []AssignedAddress
for r.Remaining() > 0 {
if len(assignedAddresses) >= maxAddressesPerCapsule {
return nil, fmt.Errorf("%w: ADDRESS_ASSIGN capsule contains too many addresses (maximum %d)", errCapsuleLimit, maxAddressesPerCapsule)
}
requestID, prefix, err := parseAddress(r)
if err != nil {
return nil, err
}
assignedAddresses = append(assignedAddresses, AssignedAddress{RequestID: AddressRequestID(requestID), IPPrefix: prefix})
}
return &addressAssignCapsule{AssignedAddresses: assignedAddresses}, nil
}
func (c *addressAssignCapsule) append(b []byte) []byte {
totalLen := 0
for _, addr := range c.AssignedAddresses {
totalLen += addr.len()
}
b = quicvarint.Append(b, uint64(capsuleTypeAddressAssign))
b = quicvarint.Append(b, uint64(totalLen))
for _, addr := range c.AssignedAddresses {
b = quicvarint.Append(b, uint64(addr.RequestID))
if addr.IPPrefix.Addr().Is4() {
b = append(b, 4)
} else {
b = append(b, 6)
}
b = append(b, addr.IPPrefix.Addr().AsSlice()...)
b = append(b, byte(addr.IPPrefix.Bits()))
}
return b
}
func parseAddressRequestCapsule(r http3.CapsuleReader) (*addressRequestCapsule, error) {
if r.Remaining() == 0 {
return nil, errors.New("ADDRESS_REQUEST capsule contains no addresses")
}
capsule := &addressRequestCapsule{}
for r.Remaining() > 0 {
if len(capsule.Prefixes) >= maxAddressesPerCapsule {
return nil, fmt.Errorf("%w: ADDRESS_REQUEST capsule contains too many addresses (maximum %d)", errCapsuleLimit, maxAddressesPerCapsule)
}
requestID, prefix, err := parseAddress(r)
if err != nil {
return nil, err
}
if requestID == 0 {
return nil, errors.New("ADDRESS_REQUEST capsule contains a zero request ID")
}
capsule.RequestIDs = append(capsule.RequestIDs, AddressRequestID(requestID))
capsule.Prefixes = append(capsule.Prefixes, prefix)
}
return capsule, nil
}
func (c *addressRequestCapsule) append(b []byte) []byte {
var totalLen int
for i, p := range c.Prefixes {
totalLen += quicvarint.Len(uint64(c.RequestIDs[i])) + 1 + p.Addr().BitLen()/8 + 1
}
b = quicvarint.Append(b, uint64(capsuleTypeAddressRequest))
b = quicvarint.Append(b, uint64(totalLen))
for i, p := range c.Prefixes {
b = quicvarint.Append(b, uint64(c.RequestIDs[i]))
if p.Addr().Is4() {
b = append(b, 4)
} else {
b = append(b, 6)
}
b = append(b, p.Addr().AsSlice()...)
b = append(b, byte(p.Bits()))
}
return b
}
func parseAddress(r io.Reader) (requestID uint64, prefix netip.Prefix, _ error) {
vr := quicvarint.NewReader(r)
requestID, err := quicvarint.Read(vr)
if err != nil {
return 0, netip.Prefix{}, err
}
ipVersion, err := vr.ReadByte()
if err != nil {
return 0, netip.Prefix{}, err
}
var ip netip.Addr
switch ipVersion {
case 4:
var ipv4 [4]byte
if _, err := io.ReadFull(r, ipv4[:]); err != nil {
return 0, netip.Prefix{}, err
}
ip = netip.AddrFrom4(ipv4)
case 6:
var ipv6 [16]byte
if _, err := io.ReadFull(r, ipv6[:]); err != nil {
return 0, netip.Prefix{}, err
}
ip = netip.AddrFrom16(ipv6)
default:
return 0, netip.Prefix{}, fmt.Errorf("invalid IP version: %d", ipVersion)
}
prefixLen, err := vr.ReadByte()
if err != nil {
return 0, netip.Prefix{}, err
}
if int(prefixLen) > ip.BitLen() {
return 0, netip.Prefix{}, fmt.Errorf("prefix length %d exceeds IP address length (%d)", prefixLen, ip.BitLen())
}
prefix = netip.PrefixFrom(ip, int(prefixLen))
if prefix != prefix.Masked() {
return 0, netip.Prefix{}, errors.New("lower bits not covered by prefix length are not all zero")
}
return requestID, prefix, nil
}
type routeAdvertisementCapsule struct {
IPAddressRanges []IPRoute
}
type IPRoute struct {
StartIP netip.Addr
EndIP netip.Addr
IPProtocol uint8
}
func (r IPRoute) len() int { return 1 + r.StartIP.BitLen()/8 + r.EndIP.BitLen()/8 + 1 }
func (r IPRoute) Prefixes() []netip.Prefix { return rangeToPrefixes(r.StartIP, r.EndIP) }
func parseRouteAdvertisementCapsule(r http3.CapsuleReader) (*routeAdvertisementCapsule, error) {
var ranges []IPRoute
for r.Remaining() > 0 {
if len(ranges) >= maxRoutesPerCapsule {
return nil, fmt.Errorf("%w: ROUTE_ADVERTISEMENT capsule contains too many routes (maximum %d)", errCapsuleLimit, maxRoutesPerCapsule)
}
ipRange, err := parseIPAddressRange(r)
if err != nil {
return nil, err
}
if len(ranges) > 0 {
if err := checkRouteOrder(ranges[len(ranges)-1], ipRange); err != nil {
return nil, err
}
}
ranges = append(ranges, ipRange)
}
return &routeAdvertisementCapsule{IPAddressRanges: ranges}, nil
}
func (r IPRoute) validate() error {
if !r.StartIP.IsValid() || !r.EndIP.IsValid() || r.StartIP.Zone() != "" || r.EndIP.Zone() != "" {
return fmt.Errorf("invalid IP address range %s-%s", r.StartIP, r.EndIP)
}
if r.StartIP.Is4() != r.EndIP.Is4() {
return fmt.Errorf("IP address range %s-%s mixes IP versions", r.StartIP, r.EndIP)
}
if r.StartIP.Compare(r.EndIP) > 0 {
return fmt.Errorf("start IP %s is greater than end IP %s", r.StartIP, r.EndIP)
}
return nil
}
func checkRouteOrder(a, b IPRoute) error {
switch cmp.Or(
cmp.Compare(a.StartIP.BitLen(), b.StartIP.BitLen()),
cmp.Compare(a.IPProtocol, b.IPProtocol),
) {
case 1:
return fmt.Errorf("routes are not ordered by IP version and IP protocol: %s-%s (protocol %d) precedes %s-%s (protocol %d)",
a.StartIP, a.EndIP, a.IPProtocol, b.StartIP, b.EndIP, b.IPProtocol)
case 0:
if a.EndIP.Compare(b.StartIP) >= 0 {
return fmt.Errorf("IP address ranges %s-%s and %s-%s (protocol %d) overlap or are not in ascending order",
a.StartIP, a.EndIP, b.StartIP, b.EndIP, b.IPProtocol)
}
}
return nil
}
func (c *routeAdvertisementCapsule) append(b []byte) []byte {
var totalLen int
for _, ipRange := range c.IPAddressRanges {
totalLen += ipRange.len()
}
b = quicvarint.Append(b, uint64(capsuleTypeRouteAdvertisement))
b = quicvarint.Append(b, uint64(totalLen))
for _, ipRange := range c.IPAddressRanges {
if ipRange.StartIP.Is4() {
b = append(b, 4)
} else {
b = append(b, 6)
}
b = append(b, ipRange.StartIP.AsSlice()...)
b = append(b, ipRange.EndIP.AsSlice()...)
b = append(b, ipRange.IPProtocol)
}
return b
}
func parseIPAddressRange(r io.Reader) (IPRoute, error) {
var ipVersion uint8
if err := binary.Read(r, binary.LittleEndian, &ipVersion); err != nil {
return IPRoute{}, err
}
var startIP, endIP netip.Addr
switch ipVersion {
case 4:
var start, end [4]byte
if _, err := io.ReadFull(r, start[:]); err != nil {
return IPRoute{}, err
}
if _, err := io.ReadFull(r, end[:]); err != nil {
return IPRoute{}, err
}
startIP = netip.AddrFrom4(start)
endIP = netip.AddrFrom4(end)
case 6:
var start, end [16]byte
if _, err := io.ReadFull(r, start[:]); err != nil {
return IPRoute{}, err
}
if _, err := io.ReadFull(r, end[:]); err != nil {
return IPRoute{}, err
}
startIP = netip.AddrFrom16(start)
endIP = netip.AddrFrom16(end)
default:
return IPRoute{}, fmt.Errorf("invalid IP version: %d", ipVersion)
}
if startIP.Compare(endIP) > 0 {
return IPRoute{}, errors.New("start IP is greater than end IP")
}
var ipProtocol uint8
if err := binary.Read(r, binary.LittleEndian, &ipProtocol); err != nil {
return IPRoute{}, err
}
return IPRoute{
StartIP: startIP,
EndIP: endIP,
IPProtocol: ipProtocol,
}, nil
}
@@ -0,0 +1,449 @@
/* SPDX-License-Identifier: MIT
*
* Copyright 2024 Marten Seemann
* Adapted from github.com/quic-go/connect-ip-go (commit a0c35fa).
*/
package connectip
import (
"bytes"
"context"
"io"
"net"
"net/netip"
"testing"
"time"
"github.com/apernet/quic-go/http3"
"github.com/apernet/quic-go/quicvarint"
"github.com/stretchr/testify/require"
)
func newCapsuleReader(t *testing.T, typ http3.CapsuleType, payload []byte) http3.CapsuleReader {
t.Helper()
data := quicvarint.Append(nil, uint64(typ))
data = quicvarint.Append(data, uint64(len(payload)))
data = append(data, payload...)
parsedType, cr, err := http3.NewCapsuleParser(bytes.NewReader(data)).Next()
require.NoError(t, err)
require.Equal(t, typ, parsedType)
return cr
}
func testIncompleteCapsule(t *testing.T, data []byte, parse func(http3.CapsuleReader) error) {
t.Helper()
r := bytes.NewReader(data)
_, cr, err := http3.NewCapsuleParser(r).Next()
require.NoError(t, err)
require.NoError(t, parse(cr))
require.Zero(t, r.Len())
for i := range data {
_, cr, err := http3.NewCapsuleParser(bytes.NewReader(data[:i])).Next()
if err != nil {
if i == 0 {
require.ErrorIs(t, err, io.EOF)
} else {
require.ErrorIs(t, err, io.ErrUnexpectedEOF)
}
continue
}
require.ErrorIs(t, parse(cr), io.ErrUnexpectedEOF)
}
}
func testCapsuleEntryLimit[T any](t *testing.T, typ http3.CapsuleType, limit int, entry func(i int) []byte, parse func(http3.CapsuleReader) (*T, error)) {
t.Helper()
var payload []byte
for i := range limit {
payload = append(payload, entry(i)...)
}
r := newCapsuleReader(t, typ, payload)
_, err := parse(r)
require.NoError(t, err)
require.Zero(t, r.Remaining())
data := quicvarint.Append(nil, uint64(typ))
data = quicvarint.Append(data, uint64(len(payload)+1))
_, r, err = http3.NewCapsuleParser(bytes.NewReader(append(data, payload...))).Next()
require.NoError(t, err)
_, err = parse(r)
require.ErrorContains(t, err, "too many")
require.Equal(t, int64(1), r.Remaining())
}
func TestParseAddressAssignCapsule(t *testing.T) {
addr1 := quicvarint.Append(nil, 1337)
addr1 = append(addr1, 4)
addr1 = append(addr1, netip.AddrFrom4([4]byte{1, 2, 3, 0}).AsSlice()...)
addr1 = append(addr1, 24)
addr2 := quicvarint.Append(nil, 1338)
addr2 = append(addr2, 6)
addr2 = append(addr2, netip.MustParseAddr("2001:db8::1").AsSlice()...)
addr2 = append(addr2, 128)
data := quicvarint.Append(nil, uint64(capsuleTypeAddressAssign))
data = quicvarint.Append(data, uint64(len(addr1)+len(addr2)))
data = append(data, addr1...)
data = append(data, addr2...)
r := bytes.NewReader(data)
typ, cr, err := http3.NewCapsuleParser(r).Next()
require.NoError(t, err)
require.Equal(t, capsuleTypeAddressAssign, typ)
capsule, err := parseAddressAssignCapsule(cr)
require.NoError(t, err)
require.Equal(t,
[]AssignedAddress{
{RequestID: 1337, IPPrefix: netip.MustParsePrefix("1.2.3.0/24")},
{RequestID: 1338, IPPrefix: netip.MustParsePrefix("2001:db8::1/128")},
},
capsule.AssignedAddresses,
)
require.Zero(t, r.Len())
}
func TestParseAddressAssignCapsuleLimit(t *testing.T) {
entry := []byte{1, 4, 192, 0, 2, 1, 32}
testCapsuleEntryLimit(t, capsuleTypeAddressAssign, maxAddressesPerCapsule, func(int) []byte { return entry }, parseAddressAssignCapsule)
}
func TestAssignedAddressRejected(t *testing.T) {
for _, prefix := range []string{"0.0.0.0/32", "::/128"} {
require.True(t, (AssignedAddress{RequestID: 1, IPPrefix: netip.MustParsePrefix(prefix)}).Rejected())
}
for _, prefix := range []string{"0.0.0.0/0", "0.0.0.0/31", "::/0", "::/127", "192.0.2.1/32", "2001:db8::1/128"} {
require.False(t, (AssignedAddress{RequestID: 1, IPPrefix: netip.MustParsePrefix(prefix)}).Rejected())
}
require.False(t, (AssignedAddress{}).Rejected())
}
func TestWriteAddressAssignCapsule(t *testing.T) {
c := &addressAssignCapsule{
AssignedAddresses: []AssignedAddress{
{RequestID: 1337, IPPrefix: netip.MustParsePrefix("1.2.3.0/24")},
{RequestID: 1338, IPPrefix: netip.MustParsePrefix("2001:db8::1/128")},
},
}
data := c.append(nil)
r := bytes.NewReader(data)
typ, cr, err := http3.NewCapsuleParser(r).Next()
require.NoError(t, err)
require.Equal(t, capsuleTypeAddressAssign, typ)
parsed, err := parseAddressAssignCapsule(cr)
require.NoError(t, err)
require.Equal(t, c, parsed)
require.Zero(t, r.Len())
}
func TestParseAddressAssignCapsuleInvalid(t *testing.T) {
testParseAddressCapsuleInvalid(t, capsuleTypeAddressAssign, func(r http3.CapsuleReader) error {
_, err := parseAddressAssignCapsule(r)
return err
})
}
func testParseAddressCapsuleInvalid(t *testing.T, typ http3.CapsuleType, f func(r http3.CapsuleReader) error) {
t.Run("invalid IP version", func(t *testing.T) {
addr1 := quicvarint.Append(nil, 1337)
addr1 = append(addr1, 5)
addr1 = append(addr1, netip.AddrFrom4([4]byte{1, 2, 3, 4}).AsSlice()...)
addr1 = append(addr1, 32)
require.ErrorContains(t, f(newCapsuleReader(t, typ, addr1)), "invalid IP version: 5")
})
t.Run("invalid prefix length", func(t *testing.T) {
addr1 := quicvarint.Append(nil, 1337)
addr1 = append(addr1, 4)
addr1 = append(addr1, netip.AddrFrom4([4]byte{1, 2, 3, 4}).AsSlice()...)
addr1 = append(addr1, 33)
require.ErrorContains(t, f(newCapsuleReader(t, typ, addr1)), "prefix length 33 exceeds IP address length (32)")
})
t.Run("lower bits not covered by prefix length are not all zero", func(t *testing.T) {
addr1 := quicvarint.Append(nil, 1337)
addr1 = append(addr1, 4)
addr1 = append(addr1, netip.AddrFrom4([4]byte{1, 2, 3, 4}).AsSlice()...)
addr1 = append(addr1, 28)
require.ErrorContains(t, f(newCapsuleReader(t, typ, addr1)), "lower bits not covered by prefix length are not all zero")
})
t.Run("incomplete capsule", func(t *testing.T) {
var data []byte
switch typ {
case capsuleTypeAddressAssign:
data = (&addressAssignCapsule{
AssignedAddresses: []AssignedAddress{
{RequestID: 1337, IPPrefix: netip.MustParsePrefix("1.2.3.4/32")},
{RequestID: 1338, IPPrefix: netip.MustParsePrefix("2001:db8::1/128")},
},
}).append(nil)
case capsuleTypeAddressRequest:
data = (&addressRequestCapsule{
RequestIDs: []AddressRequestID{1337, 1338},
Prefixes: []netip.Prefix{netip.MustParsePrefix("1.2.3.4/32"), netip.MustParsePrefix("2001:db8::1/128")},
}).append(nil)
default:
t.Fatalf("unexpected capsule type: %d", typ)
}
testIncompleteCapsule(t, data, f)
})
}
func TestParseAddressRequestCapsule(t *testing.T) {
addr1 := quicvarint.Append(nil, 1337)
addr1 = append(addr1, 4)
addr1 = append(addr1, netip.AddrFrom4([4]byte{1, 2, 3, 0}).AsSlice()...)
addr1 = append(addr1, 24)
addr2 := quicvarint.Append(nil, 1338)
addr2 = append(addr2, 6)
addr2 = append(addr2, netip.MustParseAddr("2001:db8::1").AsSlice()...)
addr2 = append(addr2, 128)
data := quicvarint.Append(nil, uint64(capsuleTypeAddressRequest))
data = quicvarint.Append(data, uint64(len(addr1)+len(addr2)))
data = append(data, addr1...)
data = append(data, addr2...)
r := bytes.NewReader(data)
typ, cr, err := http3.NewCapsuleParser(r).Next()
require.NoError(t, err)
require.Equal(t, capsuleTypeAddressRequest, typ)
capsule, err := parseAddressRequestCapsule(cr)
require.NoError(t, err)
require.Equal(t, []AddressRequestID{1337, 1338}, capsule.RequestIDs)
require.Equal(t, []netip.Prefix{netip.MustParsePrefix("1.2.3.0/24"), netip.MustParsePrefix("2001:db8::1/128")}, capsule.Prefixes)
require.Zero(t, r.Len())
}
func TestParseAddressRequestCapsuleLimit(t *testing.T) {
entry := []byte{1, 4, 192, 0, 2, 1, 32}
testCapsuleEntryLimit(t, capsuleTypeAddressRequest, maxAddressesPerCapsule, func(int) []byte { return entry }, parseAddressRequestCapsule)
}
func TestWriteAddressRequestCapsule(t *testing.T) {
c := &addressRequestCapsule{
RequestIDs: []AddressRequestID{1337, 1338},
Prefixes: []netip.Prefix{netip.MustParsePrefix("1.2.3.0/24"), netip.MustParsePrefix("2001:db8::1/128")},
}
data := c.append(nil)
r := bytes.NewReader(data)
typ, cr, err := http3.NewCapsuleParser(r).Next()
require.NoError(t, err)
require.Equal(t, capsuleTypeAddressRequest, typ)
parsed, err := parseAddressRequestCapsule(cr)
require.NoError(t, err)
require.Equal(t, c, parsed)
require.Zero(t, r.Len())
}
func TestParseAddressRequestCapsuleInvalid(t *testing.T) {
t.Run("empty", func(t *testing.T) {
_, err := parseAddressRequestCapsule(newCapsuleReader(t, capsuleTypeAddressRequest, nil))
require.ErrorContains(t, err, "contains no addresses")
})
t.Run("zero request ID", func(t *testing.T) {
_, err := parseAddressRequestCapsule(newCapsuleReader(t, capsuleTypeAddressRequest, []byte{0, 4, 192, 0, 2, 1, 32}))
require.ErrorContains(t, err, "zero request ID")
})
testParseAddressCapsuleInvalid(t, capsuleTypeAddressRequest, func(r http3.CapsuleReader) error {
_, err := parseAddressRequestCapsule(r)
return err
})
}
func TestParseRouteAdvertisementCapsule(t *testing.T) {
iprange1 := []byte{4}
iprange1 = append(iprange1, netip.AddrFrom4([4]byte{1, 1, 1, 1}).AsSlice()...)
iprange1 = append(iprange1, netip.AddrFrom4([4]byte{1, 2, 3, 4}).AsSlice()...)
iprange1 = append(iprange1, 13)
iprange2 := []byte{6}
iprange2 = append(iprange2, netip.MustParseAddr("2001:db8::1").AsSlice()...)
iprange2 = append(iprange2, netip.MustParseAddr("2001:db8::100").AsSlice()...)
iprange2 = append(iprange2, 37)
data := quicvarint.Append(nil, uint64(capsuleTypeRouteAdvertisement))
data = quicvarint.Append(data, uint64(len(iprange1)+len(iprange2)))
data = append(data, iprange1...)
data = append(data, iprange2...)
r := bytes.NewReader(data)
typ, cr, err := http3.NewCapsuleParser(r).Next()
require.NoError(t, err)
require.Equal(t, capsuleTypeRouteAdvertisement, typ)
capsule, err := parseRouteAdvertisementCapsule(cr)
require.NoError(t, err)
require.Equal(t,
[]IPRoute{
{StartIP: netip.MustParseAddr("1.1.1.1"), EndIP: netip.MustParseAddr("1.2.3.4"), IPProtocol: 13},
{StartIP: netip.MustParseAddr("2001:db8::1"), EndIP: netip.MustParseAddr("2001:db8::100"), IPProtocol: 37},
},
capsule.IPAddressRanges,
)
require.Equal(t,
rangeToPrefixes(netip.MustParseAddr("1.1.1.1"), netip.MustParseAddr("1.2.3.4")),
capsule.IPAddressRanges[0].Prefixes(),
)
require.Equal(t,
rangeToPrefixes(netip.MustParseAddr("2001:db8::1"), netip.MustParseAddr("2001:db8::100")),
capsule.IPAddressRanges[1].Prefixes(),
)
require.Zero(t, r.Len())
}
func TestParseRouteAdvertisementCapsuleLimit(t *testing.T) {
entry := func(i int) []byte { return []byte{4, 10, 0, byte(i >> 8), byte(i), 10, 0, byte(i >> 8), byte(i), 0} }
testCapsuleEntryLimit(t, capsuleTypeRouteAdvertisement, maxRoutesPerCapsule, entry, parseRouteAdvertisementCapsule)
}
func TestWriteRouteAdvertisementCapsule(t *testing.T) {
c := &routeAdvertisementCapsule{
IPAddressRanges: []IPRoute{
{StartIP: netip.MustParseAddr("1.1.1.1"), EndIP: netip.MustParseAddr("1.2.3.4"), IPProtocol: 13},
{StartIP: netip.MustParseAddr("2001:db8::1"), EndIP: netip.MustParseAddr("2001:db8::100"), IPProtocol: 37},
},
}
data := c.append(nil)
r := bytes.NewReader(data)
typ, cr, err := http3.NewCapsuleParser(r).Next()
require.NoError(t, err)
require.Equal(t, capsuleTypeRouteAdvertisement, typ)
parsed, err := parseRouteAdvertisementCapsule(cr)
require.NoError(t, err)
require.Equal(t, c, parsed)
require.Zero(t, r.Len())
}
func TestParseRouteAdvertisementCapsuleInvalid(t *testing.T) {
t.Run("invalid IP version", func(t *testing.T) {
iprange1 := []byte{5}
iprange1 = append(iprange1, netip.AddrFrom4([4]byte{1, 1, 1, 1}).AsSlice()...)
iprange1 = append(iprange1, netip.AddrFrom4([4]byte{1, 1, 1, 2}).AsSlice()...)
iprange1 = append(iprange1, 13)
_, err := parseRouteAdvertisementCapsule(newCapsuleReader(t, capsuleTypeRouteAdvertisement, iprange1))
require.ErrorContains(t, err, "invalid IP version: 5")
})
t.Run("start IP is greater than end IP", func(t *testing.T) {
iprange1 := []byte{4}
iprange1 = append(iprange1, netip.AddrFrom4([4]byte{1, 2, 3, 4}).AsSlice()...)
iprange1 = append(iprange1, netip.AddrFrom4([4]byte{1, 1, 1, 1}).AsSlice()...)
iprange1 = append(iprange1, 13)
_, err := parseRouteAdvertisementCapsule(newCapsuleReader(t, capsuleTypeRouteAdvertisement, iprange1))
require.ErrorContains(t, err, "start IP is greater than end IP")
})
t.Run("incomplete capsule", func(t *testing.T) {
data := (&routeAdvertisementCapsule{
IPAddressRanges: []IPRoute{
{StartIP: netip.MustParseAddr("1.1.1.1"), EndIP: netip.MustParseAddr("2.2.2.2"), IPProtocol: 13},
{StartIP: netip.MustParseAddr("2001:db8::1"), EndIP: netip.MustParseAddr("2001:db8::100"), IPProtocol: 37},
},
}).append(nil)
testIncompleteCapsule(t, data, func(r http3.CapsuleReader) error {
_, err := parseRouteAdvertisementCapsule(r)
return err
})
})
}
var (
route4a = IPRoute{StartIP: netip.MustParseAddr("10.0.0.0"), EndIP: netip.MustParseAddr("10.0.0.9")}
route4b = IPRoute{StartIP: netip.MustParseAddr("10.0.0.10"), EndIP: netip.MustParseAddr("10.0.0.20")}
route4ab = IPRoute{StartIP: netip.MustParseAddr("10.0.0.9"), EndIP: netip.MustParseAddr("10.0.0.20")}
route6 = IPRoute{StartIP: netip.MustParseAddr("2001:db8::"), EndIP: netip.MustParseAddr("2001:db8::ffff")}
)
func withProtocol(r IPRoute, proto uint8) IPRoute {
r.IPProtocol = proto
return r
}
var routeOrderTests = []struct {
name string
routes []IPRoute
err string
}{
{name: "empty"},
{name: "adjacent ranges", routes: []IPRoute{route4a, route4b}},
{name: "IPv4 before IPv6 with a lower IP protocol", routes: []IPRoute{withProtocol(route4a, 17), route6}},
{name: "same range for different IP protocols", routes: []IPRoute{withProtocol(route4a, 6), withProtocol(route4a, 17)}},
{name: "IP protocol order before address order", routes: []IPRoute{withProtocol(route4b, 6), withProtocol(route4a, 17)}},
{name: "IPv6 before IPv4", routes: []IPRoute{route6, route4a}, err: "not ordered by IP version and IP protocol"},
{name: "descending IP protocols", routes: []IPRoute{withProtocol(route4a, 17), withProtocol(route4b, 6)}, err: "not ordered by IP version and IP protocol"},
{name: "descending ranges", routes: []IPRoute{route4b, route4a}, err: "overlap or are not in ascending order"},
{name: "overlapping ranges", routes: []IPRoute{route4a, route4ab}, err: "overlap or are not in ascending order"},
{name: "duplicate range", routes: []IPRoute{route6, route6}, err: "overlap or are not in ascending order"},
}
func TestParseRouteAdvertisementCapsuleOrder(t *testing.T) {
for _, tc := range routeOrderTests {
t.Run(tc.name, func(t *testing.T) {
data := (&routeAdvertisementCapsule{IPAddressRanges: tc.routes}).append(nil)
_, cr, err := http3.NewCapsuleParser(bytes.NewReader(data)).Next()
require.NoError(t, err)
capsule, err := parseRouteAdvertisementCapsule(cr)
if tc.err != "" {
require.ErrorContains(t, err, tc.err)
return
}
require.NoError(t, err)
require.Equal(t, tc.routes, capsule.IPAddressRanges)
})
}
}
func TestAdvertiseRouteValidation(t *testing.T) {
tests := []struct {
name string
routes []IPRoute
err string
}{
{name: "invalid start IP", routes: []IPRoute{{EndIP: route4a.EndIP}}, err: "invalid IP address range"},
{name: "invalid end IP", routes: []IPRoute{{StartIP: route4a.StartIP}}, err: "invalid IP address range"},
{
name: "IPv6 zone",
routes: []IPRoute{{StartIP: netip.MustParseAddr("fe80::1%eth0"), EndIP: netip.MustParseAddr("fe80::2%eth0")}},
err: "invalid IP address range",
},
{name: "mixed IP versions", routes: []IPRoute{{StartIP: route4a.StartIP, EndIP: route6.EndIP}}, err: "mixes IP versions"},
{
name: "IPv4 and IPv4-mapped IPv6",
routes: []IPRoute{{StartIP: netip.MustParseAddr("10.0.0.1"), EndIP: netip.MustParseAddr("::ffff:10.0.0.2")}},
err: "mixes IP versions",
},
{name: "start after end", routes: []IPRoute{route4a, {StartIP: route4b.EndIP, EndIP: route4b.StartIP}}, err: "invalid route 1: start IP 10.0.0.20 is greater than end IP 10.0.0.10"},
}
tests = append(tests, routeOrderTests...)
for _, tc := range tests {
t.Run(tc.name, func(t *testing.T) {
conn := newProxiedConn(&mockStream{})
t.Cleanup(func() { conn.Close() })
err := conn.AdvertiseRoute(tc.routes)
if tc.err != "" {
require.ErrorContains(t, err, tc.err)
conn.mu.Lock()
defer conn.mu.Unlock()
require.Empty(t, conn.queuedWrites)
require.Nil(t, conn.localRoutes)
return
}
require.NoError(t, err)
})
}
}
func TestReceiveMisorderedRouteAdvertisement(t *testing.T) {
toRead := make(chan []byte, 1)
conn := newProxiedConn(&mockStream{toRead: toRead})
t.Cleanup(func() { conn.Close() })
toRead <- (&routeAdvertisementCapsule{IPAddressRanges: []IPRoute{route6, route4a}}).append(nil)
ctx, cancel := context.WithTimeout(t.Context(), time.Second)
defer cancel()
_, err := conn.Routes(ctx)
require.ErrorIs(t, err, net.ErrClosed)
}
@@ -0,0 +1,95 @@
/* SPDX-License-Identifier: MIT
*
* Copyright 2024 Marten Seemann
* Adapted from github.com/quic-go/connect-ip-go (commit a0c35fa).
*/
package connectip
import (
"crypto/rand"
"crypto/rsa"
"crypto/tls"
"crypto/x509"
"crypto/x509/pkix"
"log"
"math/big"
"time"
"github.com/apernet/quic-go/http3"
)
var (
tlsConf *tls.Config
certPool *x509.CertPool
)
func generateCA() (*x509.Certificate, *rsa.PrivateKey, error) {
certTempl := &x509.Certificate{
SerialNumber: big.NewInt(2019),
Subject: pkix.Name{},
NotBefore: time.Now(),
NotAfter: time.Now().Add(24 * time.Hour),
IsCA: true,
ExtKeyUsage: []x509.ExtKeyUsage{x509.ExtKeyUsageClientAuth, x509.ExtKeyUsageServerAuth},
KeyUsage: x509.KeyUsageDigitalSignature | x509.KeyUsageCertSign,
BasicConstraintsValid: true,
}
caPrivateKey, err := rsa.GenerateKey(rand.Reader, 2048)
if err != nil {
return nil, nil, err
}
caBytes, err := x509.CreateCertificate(rand.Reader, certTempl, certTempl, &caPrivateKey.PublicKey, caPrivateKey)
if err != nil {
return nil, nil, err
}
ca, err := x509.ParseCertificate(caBytes)
if err != nil {
return nil, nil, err
}
return ca, caPrivateKey, nil
}
func generateLeafCert(ca *x509.Certificate, caPrivateKey *rsa.PrivateKey) (*x509.Certificate, *rsa.PrivateKey, error) {
certTempl := &x509.Certificate{
SerialNumber: big.NewInt(1),
DNSNames: []string{"localhost", "127.0.0.1"},
NotBefore: time.Now(),
NotAfter: time.Now().Add(24 * time.Hour),
ExtKeyUsage: []x509.ExtKeyUsage{x509.ExtKeyUsageClientAuth, x509.ExtKeyUsageServerAuth},
KeyUsage: x509.KeyUsageDigitalSignature,
}
privKey, err := rsa.GenerateKey(rand.Reader, 2048)
if err != nil {
return nil, nil, err
}
certBytes, err := x509.CreateCertificate(rand.Reader, certTempl, ca, &privKey.PublicKey, caPrivateKey)
if err != nil {
return nil, nil, err
}
cert, err := x509.ParseCertificate(certBytes)
if err != nil {
return nil, nil, err
}
return cert, privKey, nil
}
func init() {
ca, caPrivateKey, err := generateCA()
if err != nil {
log.Fatal("failed to generate CA certificate:", err)
}
leafCert, leafPrivateKey, err := generateLeafCert(ca, caPrivateKey)
if err != nil {
log.Fatal("failed to generate leaf certificate:", err)
}
certPool = x509.NewCertPool()
certPool.AddCert(ca)
tlsConf = &tls.Config{
Certificates: []tls.Certificate{{
Certificate: [][]byte{leafCert.Raw},
PrivateKey: leafPrivateKey,
}},
NextProtos: []string{http3.NextProtoH3},
}
}
@@ -0,0 +1,23 @@
/* SPDX-License-Identifier: MIT
*
* Copyright 2024 Marten Seemann
* Adapted from github.com/quic-go/connect-ip-go (commit a0c35fa).
*/
package connectip
import "encoding/binary"
func calculateIPv4Checksum(header []byte) uint16 {
var sum uint32
for i := 0; i < len(header); i += 2 {
if i == 10 {
continue
}
sum += uint32(binary.BigEndian.Uint16(header[i : i+2]))
}
for (sum >> 16) > 0 {
sum = (sum & 0xffff) + (sum >> 16)
}
return ^uint16(sum)
}
@@ -0,0 +1,27 @@
/* SPDX-License-Identifier: MIT
*
* Copyright 2024 Marten Seemann
* Adapted from github.com/quic-go/connect-ip-go (commit a0c35fa).
*/
package connectip
import (
"testing"
"github.com/stretchr/testify/require"
)
func TestIPv4ChecksumTestVector(t *testing.T) {
data := []byte{0x45, 0x00, 0x00, 0x73, 0x00, 0x00, 0x40, 0x00, 0x40, 0x11, 0xb8, 0x61, 0xc0, 0xa8, 0x00, 0x01, 0xc0, 0xa8, 0x00, 0xc7}
checksum := calculateIPv4Checksum(data)
require.Equal(t, uint16(0xb861), checksum)
}
func TestIPv4ChecksumWithOptions(t *testing.T) {
data := []byte{0x46, 0x00, 0x00, 0x77, 0x00, 0x00, 0x40, 0x00, 0x40, 0x11, 0x00, 0x00, 0xc0, 0xa8, 0x00, 0x01, 0xc0, 0xa8, 0x00, 0xc7, 0x94, 0x04, 0x00, 0x00}
checksum := calculateIPv4Checksum(data)
data[10], data[11] = byte(checksum>>8), byte(checksum)
require.True(t, ipv4ChecksumValid(data))
require.NotEqual(t, checksum, calculateIPv4Checksum(data[:20]), "the options must be covered")
}

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