mirror of
https://github.com/XTLS/Xray-core.git
synced 2026-10-04 14:56:49 +00:00
Compare commits
25
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
7ab6930f27 | ||
|
|
73fb3e8f4a | ||
|
|
745526f14c | ||
|
|
399563b6d9 | ||
|
|
e38794ed88 | ||
|
|
2610e57ecf | ||
|
|
5afe260f10 | ||
|
|
5d1d8200d9 | ||
|
|
2440f53cdd | ||
|
|
1c52c65872 | ||
|
|
5724db08f4 | ||
|
|
3d3306503d | ||
|
|
5e1bb92b98 | ||
|
|
459301d42e | ||
|
|
3982028a9c | ||
|
|
70b8e9a61d | ||
|
|
219f758060 | ||
|
|
9628003594 | ||
|
|
72d9ab50b9 | ||
|
|
235843c5d2 | ||
|
|
60e2a0c502 | ||
|
|
7d3e44fee2 | ||
|
|
9927942aaa | ||
|
|
a308ded2e6 | ||
|
|
7741e9e77e |
@@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
+23
-4
@@ -93,6 +93,7 @@ type NameServer struct {
|
||||
UnexpectedIp []*geodata.IPRule `protobuf:"bytes,13,rep,name=unexpected_ip,json=unexpectedIp,proto3" json:"unexpected_ip,omitempty"`
|
||||
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,6 +266,8 @@ type Config struct {
|
||||
DisableFallback bool `protobuf:"varint,10,opt,name=disableFallback,proto3" json:"disableFallback,omitempty"`
|
||||
DisableFallbackIfMatch bool `protobuf:"varint,11,opt,name=disableFallbackIfMatch,proto3" json:"disableFallbackIfMatch,omitempty"`
|
||||
EnableParallelQuery bool `protobuf:"varint,14,opt,name=enableParallelQuery,proto3" json:"enableParallelQuery,omitempty"`
|
||||
// Absolute path to the Lua DNS query script.
|
||||
Script string `protobuf:"bytes,15,opt,name=script,proto3" json:"script,omitempty"`
|
||||
unknownFields protoimpl.UnknownFields
|
||||
sizeCache protoimpl.SizeCache
|
||||
}
|
||||
@@ -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" +
|
||||
|
||||
@@ -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;
|
||||
}
|
||||
|
||||
@@ -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 {
|
||||
|
||||
+165
@@ -0,0 +1,165 @@
|
||||
package dns
|
||||
|
||||
import (
|
||||
"context"
|
||||
"strings"
|
||||
|
||||
"github.com/xtls/xray-core/common/errors"
|
||||
xlua "github.com/xtls/xray-core/common/lua"
|
||||
"github.com/xtls/xray-core/common/net"
|
||||
featureDNS "github.com/xtls/xray-core/features/dns"
|
||||
"github.com/xtls/xray-core/features/dns/localdns"
|
||||
lua "github.com/yuin/gopher-lua"
|
||||
)
|
||||
|
||||
// luaDNSServer adapts configured and local DNS to the same Lua API.
|
||||
type luaDNSServer struct {
|
||||
id string
|
||||
name string
|
||||
query func(context.Context, string, featureDNS.IPOption) ([]net.IP, uint32, error)
|
||||
}
|
||||
|
||||
// RegisterLua makes xray.dns available to scripts backed by client.
|
||||
func RegisterLua(L *lua.LState, client featureDNS.Client) {
|
||||
var servers []luaDNSServer
|
||||
switch client := client.(type) {
|
||||
case *DNS:
|
||||
servers = luaServers(client)
|
||||
case *localdns.Client:
|
||||
servers = []luaDNSServer{{
|
||||
id: "localhost",
|
||||
name: "localhost",
|
||||
query: func(_ context.Context, domain string, option featureDNS.IPOption) ([]net.IP, uint32, error) {
|
||||
return client.LookupIP(domain, option)
|
||||
},
|
||||
}}
|
||||
}
|
||||
registerLua(L, servers, client)
|
||||
}
|
||||
|
||||
// registerLua makes xray.dns available to DNS scripts.
|
||||
func (s *DNS) registerLua(L *lua.LState) {
|
||||
registerLua(L, luaServers(s), nil)
|
||||
}
|
||||
|
||||
func luaServers(s *DNS) []luaDNSServer {
|
||||
servers := make([]luaDNSServer, len(s.clients))
|
||||
for i, client := range s.clients {
|
||||
servers[i] = luaDNSServer{id: client.id, name: client.Name(), query: client.QueryIP}
|
||||
}
|
||||
return servers
|
||||
}
|
||||
|
||||
func registerLua(L *lua.LState, servers []luaDNSServer, client featureDNS.Client) {
|
||||
L.PreloadModule("xray.dns", func(L *lua.LState) int {
|
||||
serverList := L.NewTable()
|
||||
for i, client := range servers {
|
||||
server := L.NewTable()
|
||||
|
||||
server.RawSetString("ID", lua.LString(client.id))
|
||||
|
||||
server.RawSetString("Query", L.NewFunction(func(L *lua.LState) int {
|
||||
domain, ok := L.Get(2).(lua.LString)
|
||||
if !ok {
|
||||
L.RaiseError("server:Query requires a domain")
|
||||
return 0
|
||||
}
|
||||
option := featureDNS.IPOption{
|
||||
IPv4Enable: L.CheckBool(3),
|
||||
IPv6Enable: L.CheckBool(4),
|
||||
FakeEnable: L.CheckBool(5),
|
||||
}
|
||||
ctx := L.Context()
|
||||
if ctx == nil {
|
||||
L.RaiseError("server:Query requires an active DNS query")
|
||||
return 0
|
||||
}
|
||||
var ips []net.IP
|
||||
var ttl uint32
|
||||
var err error
|
||||
if !option.FakeEnable && strings.EqualFold(client.name, "FakeDNS") {
|
||||
err = featureDNS.ErrEmptyResponse
|
||||
} else {
|
||||
ips, ttl, err = client.query(ctx, string(domain), option)
|
||||
}
|
||||
xlua.PushUserData(L, ips)
|
||||
xlua.PushNumber(L, ttl)
|
||||
xlua.PushError(L, err)
|
||||
return 3
|
||||
}))
|
||||
serverList.RawSetInt(i+1, server)
|
||||
}
|
||||
|
||||
module := L.NewTable()
|
||||
if servers != nil {
|
||||
module.RawSetString("Servers", serverList)
|
||||
}
|
||||
if client != nil {
|
||||
module.RawSetString("Query", newLuaClientQuery(L, client))
|
||||
}
|
||||
L.Push(module)
|
||||
return 1
|
||||
})
|
||||
}
|
||||
|
||||
func newLuaClientQuery(L *lua.LState, client featureDNS.Client) *lua.LFunction {
|
||||
return L.NewFunction(func(L *lua.LState) int {
|
||||
domain, ok := L.Get(1).(lua.LString)
|
||||
if !ok {
|
||||
L.RaiseError("dns.Query requires a domain")
|
||||
return 0
|
||||
}
|
||||
option := featureDNS.IPOption{
|
||||
IPv4Enable: L.CheckBool(2),
|
||||
IPv6Enable: L.CheckBool(3),
|
||||
FakeEnable: L.CheckBool(4),
|
||||
}
|
||||
if L.Context() == nil {
|
||||
L.RaiseError("dns.Query requires an active DNS query")
|
||||
return 0
|
||||
}
|
||||
ips, ttl, err := client.LookupIP(string(domain), option)
|
||||
xlua.PushUserData(L, ips)
|
||||
xlua.PushNumber(L, ttl)
|
||||
xlua.PushError(L, err)
|
||||
return 3
|
||||
})
|
||||
}
|
||||
|
||||
// callLuaHook invokes HandleDNSQuery in the supplied state.
|
||||
// Returned slices and IP bytes may share storage with DNS caches or matcher inputs.
|
||||
func (s *DNS) callLuaHook(L *lua.LState, domain string, option featureDNS.IPOption) ([]net.IP, uint32, error) {
|
||||
top := L.GetTop()
|
||||
defer L.SetTop(top)
|
||||
fn := L.GetGlobal("HandleDNSQuery")
|
||||
if fn.Type() != lua.LTFunction {
|
||||
return nil, 0, errors.New("DNS script must define HandleDNSQuery(...)")
|
||||
}
|
||||
if err := L.CallByParam(lua.P{Fn: fn, NRet: 3, Protect: true},
|
||||
lua.LString(strings.ToLower(domain)), lua.LBool(option.IPv4Enable),
|
||||
lua.LBool(option.IPv6Enable), lua.LBool(option.FakeEnable)); err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
return readLuaDNSResult(L.Get(-3), L.Get(-2), L.Get(-1))
|
||||
}
|
||||
|
||||
func readLuaDNSResult(addresses, ttlValue, errorValue lua.LValue) ([]net.IP, uint32, error) {
|
||||
if err := xlua.ReadError(errorValue, "DNS script error must be an error or string"); err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
ttl, err := xlua.ReadUint32(ttlValue, "DNS script returned invalid TTL")
|
||||
if err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
if addresses == lua.LNil {
|
||||
return nil, 0, featureDNS.ErrEmptyResponse
|
||||
}
|
||||
ips, err := xlua.ReadUserData[[]net.IP](addresses, "DNS script IPs must be native IP slice userdata")
|
||||
if err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
if len(ips) == 0 {
|
||||
return nil, 0, featureDNS.ErrEmptyResponse
|
||||
}
|
||||
return ips, ttl, nil
|
||||
}
|
||||
@@ -0,0 +1,294 @@
|
||||
package dns
|
||||
|
||||
import (
|
||||
"context"
|
||||
go_errors "errors"
|
||||
"math"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/xtls/xray-core/common/geodata"
|
||||
"github.com/xtls/xray-core/common/net"
|
||||
featureDNS "github.com/xtls/xray-core/features/dns"
|
||||
"github.com/xtls/xray-core/features/dns/localdns"
|
||||
lua "github.com/yuin/gopher-lua"
|
||||
)
|
||||
|
||||
func TestReadLuaDNSResult(t *testing.T) {
|
||||
L := lua.NewState()
|
||||
defer L.Close()
|
||||
want := []net.IP{net.ParseIP("8.8.8.8"), {127, 0, 0, 1}, net.ParseIP("::1")}
|
||||
addresses := L.NewUserData()
|
||||
addresses.Value = want
|
||||
ips, ttl, err := readLuaDNSResult(addresses, lua.LNumber(45), lua.LNil)
|
||||
if err != nil || ttl != 45 || len(ips) != len(want) {
|
||||
t.Fatalf("readLuaDNSResult() = %v, TTL %d, %v", ips, ttl, err)
|
||||
}
|
||||
for i := range want {
|
||||
if !ips[i].Equal(want[i]) {
|
||||
t.Fatalf("IP %d = %v, want %v", i, ips[i], want[i])
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestReadLuaDNSResultValidation(t *testing.T) {
|
||||
L := lua.NewState()
|
||||
defer L.Close()
|
||||
|
||||
for _, tc := range []struct {
|
||||
name string
|
||||
change func(*[3]lua.LValue)
|
||||
want string
|
||||
}{
|
||||
{"fractional TTL", func(v *[3]lua.LValue) { v[1] = lua.LNumber(1.5) }, "invalid TTL"},
|
||||
{"oversized TTL", func(v *[3]lua.LValue) { v[1] = lua.LNumber(4294967296) }, "invalid TTL"},
|
||||
{"negative TTL", func(v *[3]lua.LValue) { v[1] = lua.LNumber(-1) }, "invalid TTL"},
|
||||
{"NaN TTL", func(v *[3]lua.LValue) { v[1] = lua.LNumber(math.NaN()) }, "invalid TTL"},
|
||||
{"missing TTL", func(v *[3]lua.LValue) { v[1] = lua.LNil }, "invalid TTL"},
|
||||
{"string IPs", func(v *[3]lua.LValue) { v[0] = lua.LString("127.0.0.1") }, "native IP slice"},
|
||||
{"wrong userdata", func(v *[3]lua.LValue) { v[0].(*lua.LUserData).Value = net.ParseIP("127.0.0.1") }, "native IP slice"},
|
||||
{"script error", func(v *[3]lua.LValue) { v[2] = lua.LString("blocked by script") }, "blocked by script"},
|
||||
{"invalid error", func(v *[3]lua.LValue) { v[2] = lua.LTrue }, "error or string"},
|
||||
} {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
addresses := L.NewUserData()
|
||||
addresses.Value = []net.IP{net.ParseIP("127.0.0.1")}
|
||||
values := [3]lua.LValue{addresses, lua.LNumber(60), lua.LNil}
|
||||
tc.change(&values)
|
||||
_, _, err := readLuaDNSResult(values[0], values[1], values[2])
|
||||
if err == nil || !strings.Contains(err.Error(), tc.want) {
|
||||
t.Fatalf("readLuaDNSResult error = %v, want %q", err, tc.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
addresses := L.NewUserData()
|
||||
addresses.Value = []net.IP(nil)
|
||||
for _, empty := range []lua.LValue{addresses, lua.LNil} {
|
||||
if _, _, err := readLuaDNSResult(empty, lua.LNumber(0), lua.LNil); !go_errors.Is(err, featureDNS.ErrEmptyResponse) {
|
||||
t.Fatalf("empty result error = %v, want ErrEmptyResponse", err)
|
||||
}
|
||||
}
|
||||
wantErr := go_errors.New("upstream failed")
|
||||
errorValue := L.NewUserData()
|
||||
errorValue.Value = wantErr
|
||||
if _, _, err := readLuaDNSResult(lua.LNil, lua.LNil, errorValue); err != wantErr {
|
||||
t.Fatalf("upstream error = %v, want original error %v", err, wantErr)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCallLuaHookCancellation(t *testing.T) {
|
||||
L := lua.NewState()
|
||||
defer L.Close()
|
||||
if err := L.DoString(`function HandleDNSQuery(domain, ipv4, ipv6, fake) while true do end end`); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
cancel()
|
||||
L.SetContext(ctx)
|
||||
_, _, err := (&DNS{}).callLuaHook(L, "example.com", featureDNS.IPOption{IPv4Enable: true})
|
||||
if err == nil {
|
||||
t.Fatal("CallLuaHook did not stop after context cancellation")
|
||||
}
|
||||
if L.Context() != ctx {
|
||||
t.Fatal("CallLuaHook changed the Lua state's context")
|
||||
}
|
||||
}
|
||||
|
||||
func TestCallLuaHookNormalizesDomain(t *testing.T) {
|
||||
L := lua.NewState()
|
||||
defer L.Close()
|
||||
addresses := L.NewUserData()
|
||||
addresses.Value = []net.IP{net.ParseIP("127.0.0.1")}
|
||||
L.SetGlobal("ips", addresses)
|
||||
if err := L.DoString(`
|
||||
function HandleDNSQuery(domain, ipv4, ipv6, fake)
|
||||
assert(domain == "example.com")
|
||||
assert(ipv4 and not ipv6 and not fake)
|
||||
return ips, 60, nil
|
||||
end
|
||||
`); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
s := &DNS{}
|
||||
if _, _, err := s.callLuaHook(L, "ExAmPlE.CoM", featureDNS.IPOption{IPv4Enable: true}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCallLuaHookRestoresStack(t *testing.T) {
|
||||
for _, tc := range []struct {
|
||||
name string
|
||||
body string
|
||||
wantErr bool
|
||||
}{
|
||||
{"success", `return ips, 60`, false},
|
||||
{"error", `error("failed")`, true},
|
||||
} {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
L := lua.NewState()
|
||||
defer L.Close()
|
||||
addresses := L.NewUserData()
|
||||
addresses.Value = []net.IP{net.ParseIP("127.0.0.1")}
|
||||
L.SetGlobal("ips", addresses)
|
||||
if err := L.DoString("function HandleDNSQuery() " + tc.body + " end"); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
L.Push(lua.LTrue)
|
||||
_, _, err := (&DNS{}).callLuaHook(L, "example.com", featureDNS.IPOption{IPv4Enable: true})
|
||||
if (err != nil) != tc.wantErr {
|
||||
t.Fatalf("hook error = %v, want error %t", err, tc.wantErr)
|
||||
}
|
||||
if L.GetTop() != 1 || L.Get(1) != lua.LTrue {
|
||||
t.Fatal("hook did not restore the stack")
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestLuaDNSServerQuery(t *testing.T) {
|
||||
L := lua.NewState()
|
||||
defer L.Close()
|
||||
geodata.RegisterLua(L)
|
||||
option := featureDNS.IPOption{IPv4Enable: true}
|
||||
ips := []net.IP{net.ParseIP("127.0.0.1"), net.ParseIP("8.8.8.8")}
|
||||
server := &DNS{clients: []*Client{{server: &benchmarkLuaNameServer{ips: ips}, ipOption: &option, timeoutMs: time.Second}}}
|
||||
server.registerLua(L)
|
||||
if err := L.DoString(`
|
||||
local server = require("xray.dns").Servers[1]
|
||||
local matcher = require("xray.geodata").BuildIPMatcher("127.0.0.0/8")
|
||||
function HandleDNSQuery(domain, ipv4, ipv6, fake)
|
||||
local ips, ttl, err = server:Query(domain, ipv4, ipv6, fake)
|
||||
assert(type(ips) == "userdata" and not err)
|
||||
assert(matcher:AnyMatch(ips))
|
||||
local matched = matcher:FilterIPs(ips)
|
||||
return matched, ttl, err
|
||||
end
|
||||
`); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
L.SetContext(context.Background())
|
||||
got, ttl, err := server.callLuaHook(L, "example.com", option)
|
||||
if err != nil || ttl != 60 || len(got) != 1 || !got[0].Equal(ips[0]) {
|
||||
t.Fatalf("server query = %v, TTL %d, %v", got, ttl, err)
|
||||
}
|
||||
}
|
||||
|
||||
type luaDNSClient struct {
|
||||
featureDNS.Client
|
||||
lookup func(string, featureDNS.IPOption) ([]net.IP, uint32, error)
|
||||
}
|
||||
|
||||
func (c *luaDNSClient) LookupIP(domain string, option featureDNS.IPOption) ([]net.IP, uint32, error) {
|
||||
return c.lookup(domain, option)
|
||||
}
|
||||
|
||||
func TestLuaDNSClientQuery(t *testing.T) {
|
||||
L := lua.NewState()
|
||||
defer L.Close()
|
||||
L.SetContext(context.Background())
|
||||
geodata.RegisterLua(L)
|
||||
want := []net.IP{{127, 0, 0, 1}}
|
||||
client := &luaDNSClient{lookup: func(domain string, option featureDNS.IPOption) ([]net.IP, uint32, error) {
|
||||
if domain != "MiXeD.Example." || !option.IPv4Enable || option.IPv6Enable || !option.FakeEnable {
|
||||
t.Fatalf("dns.Query arguments = %q, %+v", domain, option)
|
||||
}
|
||||
return want, 42, nil
|
||||
}}
|
||||
RegisterLua(L, client)
|
||||
if err := L.DoString(`
|
||||
local dns = require("xray.dns")
|
||||
local matcher = require("xray.geodata").BuildIPMatcher("127.0.0.1")
|
||||
assert(dns.Servers == nil)
|
||||
ips, ttl, err = dns.Query("MiXeD.Example.", true, false, true)
|
||||
assert(not err and ttl == 42 and matcher:AnyMatch(ips))
|
||||
`); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
got := L.GetGlobal("ips").(*lua.LUserData).Value.([]net.IP)
|
||||
if &got[0] != &want[0] {
|
||||
t.Fatal("dns.Query copied the IP slice")
|
||||
}
|
||||
}
|
||||
|
||||
func TestLuaDNSLocalClient(t *testing.T) {
|
||||
L := lua.NewState()
|
||||
defer L.Close()
|
||||
L.SetContext(context.Background())
|
||||
RegisterLua(L, localdns.New())
|
||||
if err := L.DoString(`
|
||||
local dns = require("xray.dns")
|
||||
assert(dns.Servers[1].ID == "localhost")
|
||||
serverIPs, _, serverErr = dns.Servers[1]:Query("127.0.0.1", true, false, false)
|
||||
clientIPs, _, clientErr = dns.Query("127.0.0.1", true, false, false)
|
||||
assert(not serverErr and not clientErr)
|
||||
`); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
for _, name := range []string{"serverIPs", "clientIPs"} {
|
||||
ips := L.GetGlobal(name).(*lua.LUserData).Value.([]net.IP)
|
||||
if len(ips) != 1 || !ips[0].Equal(net.ParseIP("127.0.0.1")) {
|
||||
t.Fatalf("%s = %v", name, ips)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
type benchmarkLuaNameServer struct {
|
||||
ips []net.IP
|
||||
}
|
||||
|
||||
func (*benchmarkLuaNameServer) Name() string { return "benchmark" }
|
||||
func (*benchmarkLuaNameServer) IsDisableCache() bool { return true }
|
||||
func (s *benchmarkLuaNameServer) QueryIP(context.Context, string, featureDNS.IPOption) ([]net.IP, uint32, error) {
|
||||
return s.ips, 60, nil
|
||||
}
|
||||
|
||||
// BenchmarkLuaDNSHookCall isolates a preloaded Lua hook and its server:Query bridge.
|
||||
// The direct case measures the same DNS client without Lua.
|
||||
func BenchmarkLuaDNSHookCall(b *testing.B) {
|
||||
option := featureDNS.IPOption{IPv4Enable: true}
|
||||
ip := net.ParseIP("127.0.0.1")
|
||||
upstream := &benchmarkLuaNameServer{ips: []net.IP{ip}}
|
||||
client := &Client{server: upstream, ipOption: &option, timeoutMs: time.Second}
|
||||
server := &DNS{clients: []*Client{client}}
|
||||
L := lua.NewState()
|
||||
defer L.Close()
|
||||
server.registerLua(L)
|
||||
if err := L.DoString(`
|
||||
local server = require("xray.dns").Servers[1]
|
||||
function HandleDNSQuery(domain, ipv4, ipv6, fake)
|
||||
return server:Query(domain, ipv4, ipv6, fake)
|
||||
end
|
||||
`); err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
|
||||
ctx := context.Background()
|
||||
L.SetContext(ctx)
|
||||
for _, bench := range []struct {
|
||||
name string
|
||||
query func() ([]net.IP, uint32, error)
|
||||
}{
|
||||
{"direct", func() ([]net.IP, uint32, error) { return client.QueryIP(ctx, "example.com", option) }},
|
||||
{"lua_hook", func() ([]net.IP, uint32, error) { return server.callLuaHook(L, "example.com", option) }},
|
||||
} {
|
||||
b.Run(bench.name, func(b *testing.B) {
|
||||
b.ReportAllocs()
|
||||
b.ResetTimer()
|
||||
var ips []net.IP
|
||||
var ttl uint32
|
||||
var err error
|
||||
for i := 0; i < b.N; i++ {
|
||||
ips, ttl, err = bench.query()
|
||||
if err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
}
|
||||
b.StopTimer()
|
||||
if ttl != 60 || len(ips) != 1 || !ips[0].Equal(ip) {
|
||||
b.Fatalf("query() = %v, TTL %d; want %v, TTL 60", ips, ttl, ip)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -29,6 +29,7 @@ type Server interface {
|
||||
|
||||
// Client is the interface for DNS client.
|
||||
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)
|
||||
|
||||
@@ -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}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,59 @@
|
||||
package dns
|
||||
|
||||
import (
|
||||
"time"
|
||||
|
||||
"github.com/xtls/xray-core/common/errors"
|
||||
"github.com/xtls/xray-core/common/geodata"
|
||||
"github.com/xtls/xray-core/common/log"
|
||||
xlua "github.com/xtls/xray-core/common/lua"
|
||||
"github.com/xtls/xray-core/common/net"
|
||||
"github.com/xtls/xray-core/features/dns"
|
||||
lua "github.com/yuin/gopher-lua"
|
||||
)
|
||||
|
||||
const scriptExecutionTimeout = 6 * time.Second
|
||||
|
||||
type scriptEngine struct {
|
||||
dns *DNS
|
||||
pool *xlua.Pool
|
||||
}
|
||||
|
||||
func newScriptEngine(path string, server *DNS) (*scriptEngine, error) {
|
||||
program, err := xlua.CompileFile(path)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
e := &scriptEngine{dns: server}
|
||||
e.pool, err = xlua.NewPool(server.ctx, scriptExecutionTimeout, program.NewStateFactory(
|
||||
scriptExecutionTimeout*20,
|
||||
func(L *lua.LState) {
|
||||
geodata.RegisterLua(L)
|
||||
log.RegisterLua(L)
|
||||
server.registerLua(L)
|
||||
},
|
||||
func(L *lua.LState) error {
|
||||
if L.GetGlobal("HandleDNSQuery").Type() != lua.LTFunction {
|
||||
return errors.New("DNS script must define HandleDNSQuery(...)")
|
||||
}
|
||||
return nil
|
||||
}))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
errors.LogInfo(server.ctx, "DNS script initialized from ", path)
|
||||
return e, nil
|
||||
}
|
||||
|
||||
func (e *scriptEngine) close() {
|
||||
e.pool.Close()
|
||||
}
|
||||
|
||||
func (e *scriptEngine) query(domain string, option dns.IPOption) (ips []net.IP, ttl uint32, err error) {
|
||||
err = e.pool.WithState(nil, 0, func(L *lua.LState) error {
|
||||
var hookErr error
|
||||
ips, ttl, hookErr = e.dns.callLuaHook(L, domain, option)
|
||||
return hookErr
|
||||
})
|
||||
return
|
||||
}
|
||||
@@ -0,0 +1,197 @@
|
||||
package dns
|
||||
|
||||
import (
|
||||
"context"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/xtls/xray-core/common/net"
|
||||
featureDNS "github.com/xtls/xray-core/features/dns"
|
||||
)
|
||||
|
||||
type geoIPScriptNameServer struct {
|
||||
name string
|
||||
answers map[string]net.IP
|
||||
ttl uint32
|
||||
calls int
|
||||
}
|
||||
|
||||
func (s *geoIPScriptNameServer) Name() string { return s.name }
|
||||
func (s *geoIPScriptNameServer) IsDisableCache() bool { return true }
|
||||
|
||||
func (s *geoIPScriptNameServer) QueryIP(ctx context.Context, domain string, _ featureDNS.IPOption) ([]net.IP, uint32, error) {
|
||||
if err := ctx.Err(); err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
s.calls++
|
||||
ip, ok := s.answers[domain]
|
||||
if !ok {
|
||||
return nil, 0, featureDNS.ErrEmptyResponse
|
||||
}
|
||||
return []net.IP{ip}, s.ttl, nil
|
||||
}
|
||||
|
||||
func TestDNSScriptGeoIPFallback(t *testing.T) {
|
||||
t.Setenv("xray.location.asset", filepath.Join("..", "..", "resources"))
|
||||
script := `
|
||||
local servers = require("xray.dns").Servers
|
||||
local us_ips = require("xray.geodata").BuildIPMatcher("geoip:us")
|
||||
|
||||
local by_id = {}
|
||||
for _, server in ipairs(servers) do
|
||||
by_id[server.ID] = server
|
||||
end
|
||||
assert(by_id.primary and by_id.fallback, "primary and fallback DNS servers are required")
|
||||
|
||||
function HandleDNSQuery(domain, ipv4, ipv6, fake)
|
||||
local ips, ttl, err = by_id.primary:Query(domain, ipv4, ipv6, fake)
|
||||
if not err and us_ips:AnyMatch(ips) then
|
||||
return ips, ttl, nil
|
||||
end
|
||||
return by_id.fallback:Query(domain, ipv4, ipv6, fake)
|
||||
end
|
||||
`
|
||||
scriptPath := filepath.Join(t.TempDir(), "geoip_fallback.lua")
|
||||
if err := os.WriteFile(scriptPath, []byte(script), 0o600); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
primary := &geoIPScriptNameServer{
|
||||
name: "primary",
|
||||
answers: map[string]net.IP{
|
||||
"us.example": net.ParseIP("2001:4860:4860::8888"),
|
||||
"other.example": net.ParseIP("127.0.0.1"),
|
||||
},
|
||||
ttl: 30,
|
||||
}
|
||||
fallback := &geoIPScriptNameServer{
|
||||
name: "fallback",
|
||||
answers: map[string]net.IP{"other.example": net.ParseIP("9.9.9.9")},
|
||||
ttl: 60,
|
||||
}
|
||||
option := featureDNS.IPOption{IPv4Enable: true, IPv6Enable: true}
|
||||
hosts, err := NewStaticHosts(nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
server := &DNS{
|
||||
ctx: context.Background(),
|
||||
hosts: hosts,
|
||||
ipOption: &option,
|
||||
scriptPath: scriptPath,
|
||||
clients: []*Client{
|
||||
{id: "primary", server: primary, ipOption: &option, timeoutMs: 2 * time.Second},
|
||||
{id: "fallback", server: fallback, ipOption: &option, timeoutMs: 2 * time.Second},
|
||||
},
|
||||
}
|
||||
if err := server.Start(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer server.Close()
|
||||
|
||||
for _, tc := range []struct {
|
||||
domain string
|
||||
ip net.IP
|
||||
ttl uint32
|
||||
}{
|
||||
{"Us.Example.", net.ParseIP("2001:4860:4860::8888"), 30},
|
||||
{"other.example", net.ParseIP("9.9.9.9"), 60},
|
||||
} {
|
||||
ips, ttl, err := server.LookupIP(tc.domain, option)
|
||||
if err != nil {
|
||||
t.Fatalf("LookupIP(%q): %v", tc.domain, err)
|
||||
}
|
||||
if ttl != tc.ttl || len(ips) != 1 || !ips[0].Equal(tc.ip) {
|
||||
t.Fatalf("LookupIP(%q) = %v, TTL %d; want %v, TTL %d", tc.domain, ips, ttl, tc.ip, tc.ttl)
|
||||
}
|
||||
}
|
||||
if primary.calls != 2 || fallback.calls != 1 {
|
||||
t.Fatalf("upstream calls: primary %d, fallback %d; want 2 and 1", primary.calls, fallback.calls)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDNSScriptRejectsInvalidStartup(t *testing.T) {
|
||||
for _, tc := range []struct {
|
||||
name string
|
||||
script string
|
||||
}{
|
||||
{"syntax", "function HandleDNSQuery("},
|
||||
{"missing hook", "value = 1"},
|
||||
{"top-level error", `error("setup failed")`},
|
||||
} {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
path := filepath.Join(t.TempDir(), "script.lua")
|
||||
if err := os.WriteFile(path, []byte(tc.script), 0o600); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
server := &DNS{ctx: context.Background(), scriptPath: path}
|
||||
if err := server.Start(); err == nil {
|
||||
t.Fatal("Start accepted an invalid DNS script")
|
||||
}
|
||||
if server.script != nil {
|
||||
t.Fatal("Start retained a script engine after failure")
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestDNSScriptHookErrorAndFakeDNSOption(t *testing.T) {
|
||||
path := filepath.Join(t.TempDir(), "script.lua")
|
||||
script := `
|
||||
local server = require("xray.dns").Servers[1]
|
||||
local log = require("xray.log")
|
||||
log.Info("DNS script loaded")
|
||||
function HandleDNSQuery(domain, ipv4, ipv6, fake)
|
||||
log.Debug("DNS query: ", domain)
|
||||
if domain == "bad.example" then error("script failure") end
|
||||
local ips, ttl, err = server:Query(domain, ipv4, ipv6, fake)
|
||||
if err then log.Error("DNS failed: ", err) end
|
||||
return ips, ttl, err
|
||||
end
|
||||
`
|
||||
if err := os.WriteFile(path, []byte(script), 0o600); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
option := featureDNS.IPOption{IPv4Enable: true}
|
||||
hosts, err := NewStaticHosts(nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
upstream := &geoIPScriptNameServer{
|
||||
name: "FakeDNS",
|
||||
answers: map[string]net.IP{"good.example": net.ParseIP("198.18.0.1")},
|
||||
ttl: 30,
|
||||
}
|
||||
server := &DNS{
|
||||
ctx: context.Background(),
|
||||
hosts: hosts,
|
||||
ipOption: &option,
|
||||
scriptPath: path,
|
||||
clients: []*Client{{id: "fake", server: upstream, ipOption: &option, timeoutMs: time.Second}},
|
||||
}
|
||||
if err := server.Start(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer server.Close()
|
||||
|
||||
if _, _, err := server.LookupIP("bad.example", option); err == nil || !strings.Contains(err.Error(), "script failure") {
|
||||
t.Fatalf("hook failure = %v, want script failure", err)
|
||||
}
|
||||
if _, _, err := server.LookupIP("good.example", option); err != featureDNS.ErrEmptyResponse {
|
||||
t.Fatalf("FakeDNS without FakeEnable = %v, want ErrEmptyResponse", err)
|
||||
}
|
||||
if upstream.calls != 0 {
|
||||
t.Fatalf("FakeDNS was queried without FakeEnable: %d calls", upstream.calls)
|
||||
}
|
||||
withFake := featureDNS.IPOption{IPv4Enable: true, FakeEnable: true}
|
||||
ips, ttl, err := server.LookupIP("good.example", withFake)
|
||||
if err != nil || ttl != 30 || len(ips) != 1 || !ips[0].Equal(net.ParseIP("198.18.0.1")) {
|
||||
t.Fatalf("FakeDNS with FakeEnable = %v, TTL %d, %v", ips, ttl, err)
|
||||
}
|
||||
if upstream.calls != 1 {
|
||||
t.Fatalf("FakeDNS query count = %d, want 1", upstream.calls)
|
||||
}
|
||||
}
|
||||
+12
-2
@@ -587,6 +587,8 @@ type Config struct {
|
||||
DomainStrategy Config_DomainStrategy `protobuf:"varint,1,opt,name=domain_strategy,json=domainStrategy,proto3,enum=xray.app.router.Config_DomainStrategy" json:"domain_strategy,omitempty"`
|
||||
Rule []*RoutingRule `protobuf:"bytes,2,rep,name=rule,proto3" json:"rule,omitempty"`
|
||||
BalancingRule []*BalancingRule `protobuf:"bytes,3,rep,name=balancing_rule,json=balancingRule,proto3" json:"balancing_rule,omitempty"`
|
||||
// Absolute path to the Lua routing script.
|
||||
Script string `protobuf:"bytes,4,opt,name=script,proto3" json:"script,omitempty"`
|
||||
unknownFields protoimpl.UnknownFields
|
||||
sizeCache protoimpl.SizeCache
|
||||
}
|
||||
@@ -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" +
|
||||
|
||||
@@ -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;
|
||||
}
|
||||
|
||||
@@ -0,0 +1,167 @@
|
||||
package router
|
||||
|
||||
import (
|
||||
"runtime"
|
||||
"strings"
|
||||
|
||||
"github.com/xtls/xray-core/common/errors"
|
||||
xlua "github.com/xtls/xray-core/common/lua"
|
||||
"github.com/xtls/xray-core/common/net"
|
||||
"github.com/xtls/xray-core/features/routing"
|
||||
lua "github.com/yuin/gopher-lua"
|
||||
)
|
||||
|
||||
const (
|
||||
luaContextType = "xray.router.Context"
|
||||
luaAttributesType = "xray.router.Attributes"
|
||||
)
|
||||
|
||||
// RegisterLua makes xray.router available to routing scripts.
|
||||
func (r *Router) RegisterLua(L *lua.LState) {
|
||||
registerLuaContext(L)
|
||||
|
||||
L.PreloadModule("xray.router", func(L *lua.LState) int {
|
||||
module := L.NewTable()
|
||||
|
||||
module.RawSetString("NetworkUnknown", lua.LNumber(net.Network_Unknown))
|
||||
module.RawSetString("NetworkTCP", lua.LNumber(net.Network_TCP))
|
||||
module.RawSetString("NetworkUDP", lua.LNumber(net.Network_UDP))
|
||||
module.RawSetString("NetworkUNIX", lua.LNumber(net.Network_UNIX))
|
||||
module.RawSetString("LocalOS", lua.LString(runtime.GOOS))
|
||||
|
||||
module.RawSetString("PickOutbound", L.NewFunction(func(L *lua.LState) int {
|
||||
tag, ok := L.Get(2).(lua.LString)
|
||||
if !ok {
|
||||
L.ArgError(2, "balancer tag must be a string")
|
||||
return 0
|
||||
}
|
||||
balancer, found := (*r.balancers.Load())[string(tag)]
|
||||
if !found {
|
||||
xlua.PushNil(L)
|
||||
xlua.PushError(L, errors.New("balancer ", tag, " not found"))
|
||||
return 2
|
||||
}
|
||||
outboundTag, err := balancer.PickOutbound()
|
||||
xlua.PushString(L, outboundTag)
|
||||
xlua.PushError(L, err)
|
||||
return 2
|
||||
}))
|
||||
|
||||
module.RawSetString("FindProcess", L.NewFunction(func(L *lua.LState) int {
|
||||
pid, name, path, err := findProcess(checkLuaContext(L), net.FindProcess)
|
||||
xlua.PushNumber(L, pid)
|
||||
xlua.PushString(L, name)
|
||||
xlua.PushString(L, path)
|
||||
xlua.PushError(L, err)
|
||||
return 4
|
||||
}))
|
||||
|
||||
L.Push(module)
|
||||
return 1
|
||||
})
|
||||
}
|
||||
|
||||
func registerLuaContext(L *lua.LState) {
|
||||
attributes := L.NewTypeMetatable(luaAttributesType)
|
||||
L.SetField(attributes, "__index", L.NewFunction(func(L *lua.LState) int {
|
||||
values := L.CheckUserData(1).Value.(map[string]string)
|
||||
key := L.CheckString(2)
|
||||
if value, found := values[key]; found {
|
||||
xlua.PushString(L, value)
|
||||
} else {
|
||||
xlua.PushNil(L)
|
||||
}
|
||||
return 1
|
||||
}))
|
||||
methods := L.NewTable()
|
||||
L.SetFuncs(methods, map[string]lua.LGFunction{
|
||||
"GetSourceIPs": func(L *lua.LState) int {
|
||||
xlua.PushUserData(L, checkLuaContext(L).GetSourceIPs())
|
||||
return 1
|
||||
},
|
||||
"GetTargetIPs": func(L *lua.LState) int {
|
||||
xlua.PushUserData(L, checkLuaContext(L).GetTargetIPs())
|
||||
return 1
|
||||
},
|
||||
"GetLocalIPs": func(L *lua.LState) int {
|
||||
xlua.PushUserData(L, checkLuaContext(L).GetLocalIPs())
|
||||
return 1
|
||||
},
|
||||
"GetAttributes": func(L *lua.LState) int {
|
||||
values := L.NewUserData()
|
||||
values.Value = checkLuaContext(L).GetAttributes()
|
||||
L.SetMetatable(values, attributes)
|
||||
L.Push(values)
|
||||
return 1
|
||||
},
|
||||
})
|
||||
L.SetField(L.NewTypeMetatable(luaContextType), "__index", methods)
|
||||
}
|
||||
|
||||
func checkLuaContext(L *lua.LState) routing.Context {
|
||||
ctx, ok := L.CheckUserData(1).Value.(routing.Context)
|
||||
if !ok {
|
||||
L.ArgError(1, "routing context expected")
|
||||
}
|
||||
return ctx
|
||||
}
|
||||
|
||||
// callLuaHook invokes HandleRoute in the supplied state.
|
||||
func (r *Router) callLuaHook(L *lua.LState, routeCtx routing.Context) (string, string, error) {
|
||||
top := L.GetTop()
|
||||
defer L.SetTop(top)
|
||||
fn := L.GetGlobal("HandleRoute")
|
||||
if fn.Type() != lua.LTFunction {
|
||||
return "", "", errors.New("routing script must define HandleRoute(...)")
|
||||
}
|
||||
value := L.NewUserData()
|
||||
value.Value = routeCtx
|
||||
L.SetMetatable(value, L.GetTypeMetatable(luaContextType))
|
||||
if err := L.CallByParam(lua.P{Fn: fn, NRet: 3, Protect: true},
|
||||
value, lua.LString(routeCtx.GetInboundTag()), lua.LNumber(routeCtx.GetSourcePort()),
|
||||
lua.LNumber(routeCtx.GetTargetPort()), lua.LNumber(routeCtx.GetLocalPort()),
|
||||
lua.LString(strings.ToLower(routeCtx.GetTargetDomain())), lua.LNumber(routeCtx.GetNetwork()),
|
||||
lua.LString(routeCtx.GetProtocol()), lua.LString(routeCtx.GetUser()),
|
||||
lua.LNumber(routeCtx.GetVlessRoute()), lua.LBool(routeCtx.GetSkipDNSResolve())); err != nil {
|
||||
return "", "", err
|
||||
}
|
||||
return readLuaRouteResult(L.Get(-3), L.Get(-2), L.Get(-1))
|
||||
}
|
||||
|
||||
func readLuaRouteResult(tagValue, ruleValue, errorValue lua.LValue) (string, string, error) {
|
||||
if err := xlua.ReadError(errorValue, "routing script error must be an error or string"); err != nil {
|
||||
return "", "", err
|
||||
}
|
||||
tag, err := xlua.ReadOptionalString(tagValue, "routing script outboundTag must be a string or nil")
|
||||
if err != nil || tag == "" {
|
||||
return "", "", err
|
||||
}
|
||||
ruleTag, err := xlua.ReadOptionalString(ruleValue, "routing script ruleTag must be a string")
|
||||
if err != nil {
|
||||
return "", "", err
|
||||
}
|
||||
return tag, ruleTag, nil
|
||||
}
|
||||
|
||||
type processFinder func(string, string, uint16, string, uint16) (int, string, string, error)
|
||||
|
||||
func findProcess(ctx routing.Context, finder processFinder) (int, string, string, error) {
|
||||
sources := ctx.GetSourceIPs()
|
||||
if len(sources) == 0 {
|
||||
return 0, "", "", errors.New("process lookup requires a source IP")
|
||||
}
|
||||
var network string
|
||||
switch ctx.GetNetwork() {
|
||||
case net.Network_TCP:
|
||||
network = "tcp"
|
||||
case net.Network_UDP:
|
||||
network = "udp"
|
||||
default:
|
||||
return 0, "", "", errors.New("process lookup requires TCP or UDP")
|
||||
}
|
||||
targetIP, targetPort := "", uint16(0)
|
||||
if targets := ctx.GetTargetIPs(); len(targets) > 0 {
|
||||
targetIP, targetPort = targets[0].String(), uint16(ctx.GetTargetPort())
|
||||
}
|
||||
return finder(network, sources[0].String(), uint16(ctx.GetSourcePort()), targetIP, targetPort)
|
||||
}
|
||||
@@ -0,0 +1,300 @@
|
||||
package router
|
||||
|
||||
import (
|
||||
"context"
|
||||
go_errors "errors"
|
||||
"runtime"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/xtls/xray-core/common/geodata"
|
||||
"github.com/xtls/xray-core/common/net"
|
||||
"github.com/xtls/xray-core/common/protocol"
|
||||
"github.com/xtls/xray-core/common/session"
|
||||
"github.com/xtls/xray-core/features/routing"
|
||||
routing_session "github.com/xtls/xray-core/features/routing/session"
|
||||
lua "github.com/yuin/gopher-lua"
|
||||
)
|
||||
|
||||
type luaRouteTestContext struct {
|
||||
*routing_session.Context
|
||||
sourceIPs, targetIPs, localIPs []net.IP
|
||||
}
|
||||
|
||||
func (c *luaRouteTestContext) GetSourceIPs() []net.IP { return c.sourceIPs }
|
||||
func (c *luaRouteTestContext) GetTargetIPs() []net.IP { return c.targetIPs }
|
||||
func (c *luaRouteTestContext) GetLocalIPs() []net.IP { return c.localIPs }
|
||||
|
||||
func newLuaRouteTestContext() *luaRouteTestContext {
|
||||
return &luaRouteTestContext{
|
||||
Context: &routing_session.Context{
|
||||
Inbound: &session.Inbound{
|
||||
Tag: "in", VlessRoute: 4321,
|
||||
Source: net.TCPDestination(net.LocalHostIP, 1234),
|
||||
Local: net.TCPDestination(net.LocalHostIP, 5678),
|
||||
User: &protocol.MemoryUser{Email: "user@example.com"},
|
||||
},
|
||||
Outbound: &session.Outbound{
|
||||
Target: net.TCPDestination(net.LocalHostIP, 443),
|
||||
RouteTarget: net.TCPDestination(net.DomainAddress("MiXeD.Example."), 443),
|
||||
},
|
||||
Content: &session.Content{
|
||||
Protocol: "tls", Attributes: map[string]string{"key": "value"}, SkipDNSResolve: true,
|
||||
},
|
||||
},
|
||||
sourceIPs: []net.IP{{127, 0, 0, 2}},
|
||||
targetIPs: []net.IP{{127, 0, 0, 3}},
|
||||
localIPs: []net.IP{{127, 0, 0, 1}},
|
||||
}
|
||||
}
|
||||
|
||||
func newLuaRouterState(t *testing.T, script string) (*Router, *lua.LState) {
|
||||
t.Helper()
|
||||
r := new(Router)
|
||||
if err := r.Init(context.Background(), &Config{}, nil, nil, nil); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
L := lua.NewState()
|
||||
t.Cleanup(L.Close)
|
||||
r.RegisterLua(L)
|
||||
geodata.RegisterLua(L)
|
||||
if err := L.DoString(script); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
return r, L
|
||||
}
|
||||
|
||||
func TestLuaRouteBinding(t *testing.T) {
|
||||
r, L := newLuaRouterState(t, `
|
||||
local router = require("xray.router")
|
||||
local matcher = require("xray.geodata").BuildIPMatcher("127.0.0.0/8")
|
||||
assert(router.NetworkUnknown == 0 and router.NetworkTCP == 2)
|
||||
assert(router.NetworkUDP == 3 and router.NetworkUNIX == 4)
|
||||
assert(router.BuildIPMatcher == nil and router.BuildDomainMatcher == nil)
|
||||
function HandleRoute(ctx, inboundTag, sourcePort, targetPort, localPort,
|
||||
targetDomain, network, protocol, user, vlessRoute, skipDNSResolve, ...)
|
||||
assert(select("#", ...) == 0)
|
||||
assert(inboundTag == "in" and sourcePort == 1234 and targetPort == 443 and localPort == 5678)
|
||||
assert(targetDomain == "mixed.example." and network == router.NetworkTCP)
|
||||
assert(protocol == "tls" and user == "user@example.com" and vlessRoute == 4321 and skipDNSResolve)
|
||||
assert(ctx.GetNetwork == nil and ctx.Context == nil)
|
||||
savedContext = ctx
|
||||
sourceIPs, targetIPs, localIPs = ctx:GetSourceIPs(), ctx:GetTargetIPs(), ctx:GetLocalIPs()
|
||||
attributes = ctx:GetAttributes()
|
||||
assert(matcher:AnyMatch(sourceIPs) and matcher:AnyMatch(targetIPs) and matcher:AnyMatch(localIPs))
|
||||
assert(attributes.key == "value" and attributes.missing == nil)
|
||||
assert(not pcall(function() attributes.key = "changed" end))
|
||||
return "out", "rule"
|
||||
end`)
|
||||
|
||||
ctx := newLuaRouteTestContext()
|
||||
tag, rule, err := r.callLuaHook(L, ctx)
|
||||
if err != nil || tag != "out" || rule != "rule" {
|
||||
t.Fatalf("hook = %q, %q, %v", tag, rule, err)
|
||||
}
|
||||
if L.GetGlobal("savedContext").(*lua.LUserData).Value != ctx {
|
||||
t.Fatal("routing context was copied")
|
||||
}
|
||||
for _, tc := range []struct {
|
||||
name string
|
||||
want []net.IP
|
||||
}{
|
||||
{"sourceIPs", ctx.sourceIPs},
|
||||
{"targetIPs", ctx.targetIPs},
|
||||
{"localIPs", ctx.localIPs},
|
||||
} {
|
||||
got := L.GetGlobal(tc.name).(*lua.LUserData).Value.([]net.IP)
|
||||
if &got[0] != &tc.want[0] {
|
||||
t.Fatalf("%s storage was copied", tc.name)
|
||||
}
|
||||
}
|
||||
ctx.Content.Attributes["key"] = "updated"
|
||||
L.SetGlobal("expectedOS", lua.LString(runtime.GOOS))
|
||||
if err := L.DoString(`
|
||||
assert(attributes.key == "updated")
|
||||
assert(require("xray.router").LocalOS == expectedOS)`); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestLuaRouteResult(t *testing.T) {
|
||||
nativeErr := go_errors.New("native failure")
|
||||
for _, tc := range []struct {
|
||||
name, body, tag, rule, wantErr string
|
||||
native bool
|
||||
}{
|
||||
{name: "route", body: `return "out", "rule"`, tag: "out", rule: "rule"},
|
||||
{name: "no match", body: `return nil`},
|
||||
{name: "empty tag", body: `return ""`},
|
||||
{name: "no match ignores rule", body: `return nil, false`},
|
||||
{name: "empty tag ignores rule", body: `return "", false`},
|
||||
{name: "missing rule", body: `return "out"`, tag: "out"},
|
||||
{name: "invalid tag", body: `return 1`, wantErr: "outboundTag"},
|
||||
{name: "invalid rule", body: `return "out", false`, wantErr: "ruleTag"},
|
||||
{name: "string error", body: `return nil, nil, "script failure"`, wantErr: "script failure"},
|
||||
{name: "native error", body: `return nil, nil, nativeError`, native: true},
|
||||
{name: "error overrides invalid tags", body: `return false, false, nativeError`, native: true},
|
||||
{name: "invalid error", body: `return "out", "rule", false`, wantErr: "error or string"},
|
||||
{name: "wrong error userdata", body: `return "out", "rule", wrongError`, wantErr: "error or string"},
|
||||
{name: "runtime error", body: `error("runtime failure")`, wantErr: "runtime failure"},
|
||||
} {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
r, L := newLuaRouterState(t, "function HandleRoute() "+tc.body+" end")
|
||||
value := L.NewUserData()
|
||||
value.Value = nativeErr
|
||||
L.SetGlobal("nativeError", value)
|
||||
wrong := L.NewUserData()
|
||||
wrong.Value = "not a native error"
|
||||
L.SetGlobal("wrongError", wrong)
|
||||
L.Push(lua.LTrue)
|
||||
|
||||
tag, rule, err := r.callLuaHook(L, &routing_session.Context{})
|
||||
if tag != tc.tag || rule != tc.rule {
|
||||
t.Fatalf("result = %q, %q, %v", tag, rule, err)
|
||||
}
|
||||
switch {
|
||||
case tc.native:
|
||||
if err != nativeErr {
|
||||
t.Fatalf("error = %v, want original error", err)
|
||||
}
|
||||
case tc.wantErr != "":
|
||||
if err == nil || !strings.Contains(err.Error(), tc.wantErr) {
|
||||
t.Fatalf("error = %v, want %q", err, tc.wantErr)
|
||||
}
|
||||
case err != nil:
|
||||
t.Fatal(err)
|
||||
}
|
||||
if L.GetTop() != 1 || L.Get(1) != lua.LTrue {
|
||||
t.Fatal("hook did not restore the stack")
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestLuaRouteCancellation(t *testing.T) {
|
||||
r, L := newLuaRouterState(t, `function HandleRoute() while true do end end`)
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
cancel()
|
||||
L.SetContext(ctx)
|
||||
if _, _, err := r.callLuaHook(L, &routing_session.Context{}); err == nil {
|
||||
t.Fatal("CallLuaHook did not stop after context cancellation")
|
||||
}
|
||||
if L.Context() != ctx || L.GetTop() != 0 {
|
||||
t.Fatal("CallLuaHook did not restore the Lua state")
|
||||
}
|
||||
}
|
||||
|
||||
func TestFindProcess(t *testing.T) {
|
||||
for _, tc := range []struct {
|
||||
name, network, target string
|
||||
targetPort uint16
|
||||
modify func(*luaRouteTestContext)
|
||||
wantErr bool
|
||||
}{
|
||||
{name: "TCP", network: "tcp", target: "127.0.0.3", targetPort: 443},
|
||||
{name: "UDP", network: "udp", target: "127.0.0.3", targetPort: 443, modify: func(c *luaRouteTestContext) {
|
||||
c.Outbound.Target.Network = net.Network_UDP
|
||||
}},
|
||||
{name: "domain target", network: "tcp", modify: func(c *luaRouteTestContext) { c.targetIPs = nil }},
|
||||
{name: "missing source", modify: func(c *luaRouteTestContext) { c.sourceIPs = nil }, wantErr: true},
|
||||
{name: "unsupported network", modify: func(c *luaRouteTestContext) {
|
||||
c.Outbound.Target.Network = net.Network_UNIX
|
||||
}, wantErr: true},
|
||||
} {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
ctx := newLuaRouteTestContext()
|
||||
if tc.modify != nil {
|
||||
tc.modify(ctx)
|
||||
}
|
||||
called := false
|
||||
pid, name, path, err := findProcess(ctx, func(network, source string, sourcePort uint16, target string, targetPort uint16) (int, string, string, error) {
|
||||
called = true
|
||||
if network != tc.network || source != "127.0.0.2" || sourcePort != 1234 || target != tc.target || targetPort != tc.targetPort {
|
||||
t.Fatalf("endpoints = %s %s:%d -> %s:%d", network, source, sourcePort, target, targetPort)
|
||||
}
|
||||
return 42, "process", "/path/process", nil
|
||||
})
|
||||
if tc.wantErr {
|
||||
if err == nil || called {
|
||||
t.Fatalf("findProcess = %d, %q, %q, %v", pid, name, path, err)
|
||||
}
|
||||
return
|
||||
}
|
||||
if err != nil || !called || pid != 42 || name != "process" || path != "/path/process" {
|
||||
t.Fatalf("findProcess = %d, %q, %q, %v", pid, name, path, err)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// BenchmarkLuaRouteHookCall isolates a preloaded Lua hook and its routing context bridge.
|
||||
// The direct case runs an equivalent native routing rule.
|
||||
func BenchmarkLuaRouteHookCall(b *testing.B) {
|
||||
r := new(Router)
|
||||
if err := r.Init(context.Background(), &Config{Rule: []*RoutingRule{{
|
||||
TargetTag: &RoutingRule_Tag{Tag: "out"},
|
||||
RuleTag: "rule",
|
||||
InboundTag: []string{"in"},
|
||||
Networks: []net.Network{net.Network_TCP},
|
||||
Ip: []*geodata.IPRule{{
|
||||
Value: &geodata.IPRule_Custom{Custom: &geodata.CIDRRule{
|
||||
Cidr: &geodata.CIDR{Ip: []byte{127, 0, 0, 0}, Prefix: 8},
|
||||
}},
|
||||
}},
|
||||
}}}, nil, nil, nil); err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
L := lua.NewState()
|
||||
defer L.Close()
|
||||
r.RegisterLua(L)
|
||||
geodata.RegisterLua(L)
|
||||
if err := L.DoString(`
|
||||
local router = require("xray.router")
|
||||
local matcher = require("xray.geodata").BuildIPMatcher("127.0.0.0/8")
|
||||
function HandleRoute(ctx, inboundTag, sourcePort, targetPort, localPort,
|
||||
targetDomain, network, protocol, user, vlessRoute, skipDNSResolve)
|
||||
if inboundTag == "in" and network == router.NetworkTCP and matcher:AnyMatch(ctx:GetTargetIPs()) then
|
||||
return "out", "rule"
|
||||
end
|
||||
end
|
||||
`); err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
|
||||
L.SetContext(context.Background())
|
||||
routeCtx := newLuaRouteTestContext()
|
||||
for _, benchmark := range []struct {
|
||||
name string
|
||||
route func() (string, string, error)
|
||||
}{
|
||||
{"direct", func() (string, string, error) {
|
||||
route, err := r.PickRoute(routeCtx)
|
||||
if err != nil {
|
||||
return "", "", err
|
||||
}
|
||||
return route.GetOutboundTag(), route.GetRuleTag(), nil
|
||||
}},
|
||||
{"lua_hook", func() (string, string, error) {
|
||||
return r.callLuaHook(L, routeCtx)
|
||||
}},
|
||||
} {
|
||||
b.Run(benchmark.name, func(b *testing.B) {
|
||||
b.ReportAllocs()
|
||||
b.ResetTimer()
|
||||
var tag, rule string
|
||||
var err error
|
||||
for i := 0; i < b.N; i++ {
|
||||
tag, rule, err = benchmark.route()
|
||||
if err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
}
|
||||
b.StopTimer()
|
||||
if tag != "out" || rule != "rule" {
|
||||
b.Fatalf("route() = %q, %q; want out, rule", tag, rule)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
var _ routing.Context = (*luaRouteTestContext)(nil)
|
||||
@@ -20,6 +20,8 @@ import (
|
||||
type Router struct {
|
||||
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())
|
||||
|
||||
@@ -0,0 +1,68 @@
|
||||
package router
|
||||
|
||||
import (
|
||||
"time"
|
||||
|
||||
"github.com/xtls/xray-core/app/dns"
|
||||
"github.com/xtls/xray-core/common"
|
||||
"github.com/xtls/xray-core/common/errors"
|
||||
"github.com/xtls/xray-core/common/geodata"
|
||||
"github.com/xtls/xray-core/common/log"
|
||||
xlua "github.com/xtls/xray-core/common/lua"
|
||||
"github.com/xtls/xray-core/features/routing"
|
||||
lua "github.com/yuin/gopher-lua"
|
||||
)
|
||||
|
||||
const scriptExecutionTimeout = 6 * time.Second
|
||||
|
||||
type scriptEngine struct {
|
||||
router *Router
|
||||
pool *xlua.Pool
|
||||
}
|
||||
|
||||
func newScriptEngine(path string, router *Router) (*scriptEngine, error) {
|
||||
program, err := xlua.CompileFile(path)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
e := &scriptEngine{router: router}
|
||||
e.pool, err = xlua.NewPool(router.ctx, scriptExecutionTimeout, program.NewStateFactory(
|
||||
scriptExecutionTimeout*20,
|
||||
func(L *lua.LState) {
|
||||
geodata.RegisterLua(L)
|
||||
log.RegisterLua(L)
|
||||
router.RegisterLua(L)
|
||||
dns.RegisterLua(L, router.dns)
|
||||
},
|
||||
func(L *lua.LState) error {
|
||||
if L.GetGlobal("HandleRoute").Type() != lua.LTFunction {
|
||||
return errors.New("routing script must define HandleRoute(...)")
|
||||
}
|
||||
return nil
|
||||
}))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
errors.LogInfo(router.ctx, "routing script initialized from ", path)
|
||||
return e, nil
|
||||
}
|
||||
|
||||
func (e *scriptEngine) close() {
|
||||
e.pool.Close()
|
||||
}
|
||||
|
||||
func (e *scriptEngine) pickRoute(ctx routing.Context) (routing.Route, error) {
|
||||
var tag, ruleTag string
|
||||
err := e.pool.WithState(nil, 0, func(L *lua.LState) error {
|
||||
var hookErr error
|
||||
tag, ruleTag, hookErr = e.router.callLuaHook(L, ctx)
|
||||
return hookErr
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if tag == "" {
|
||||
return nil, common.ErrNoClue
|
||||
}
|
||||
return &Route{Context: ctx, outboundTag: tag, ruleTag: ruleTag}, nil
|
||||
}
|
||||
@@ -0,0 +1,372 @@
|
||||
package router
|
||||
|
||||
import (
|
||||
"context"
|
||||
stdnet "net"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
wireDNS "github.com/miekg/dns"
|
||||
"github.com/xtls/xray-core/app/dispatcher"
|
||||
appdns "github.com/xtls/xray-core/app/dns"
|
||||
"github.com/xtls/xray-core/app/proxyman"
|
||||
_ "github.com/xtls/xray-core/app/proxyman/outbound"
|
||||
"github.com/xtls/xray-core/common"
|
||||
"github.com/xtls/xray-core/common/net"
|
||||
"github.com/xtls/xray-core/common/serial"
|
||||
"github.com/xtls/xray-core/core"
|
||||
featureDNS "github.com/xtls/xray-core/features/dns"
|
||||
"github.com/xtls/xray-core/features/outbound"
|
||||
"github.com/xtls/xray-core/features/routing"
|
||||
routing_session "github.com/xtls/xray-core/features/routing/session"
|
||||
"github.com/xtls/xray-core/proxy/blackhole"
|
||||
"github.com/xtls/xray-core/proxy/freedom"
|
||||
)
|
||||
|
||||
type luaRouteDNSClient struct {
|
||||
featureDNS.Client
|
||||
lookup func(string, featureDNS.IPOption) ([]net.IP, uint32, error)
|
||||
}
|
||||
|
||||
func (d *luaRouteDNSClient) LookupIP(domain string, option featureDNS.IPOption) ([]net.IP, uint32, error) {
|
||||
return d.lookup(domain, option)
|
||||
}
|
||||
|
||||
type luaRouteOutboundManager struct{ outbound.Manager }
|
||||
|
||||
func (*luaRouteOutboundManager) Select(selectors []string) []string { return selectors }
|
||||
|
||||
func writeRouteScript(t *testing.T, script string) string {
|
||||
t.Helper()
|
||||
path := filepath.Join(t.TempDir(), "route.lua")
|
||||
if err := os.WriteFile(path, []byte(script), 0o600); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
return path
|
||||
}
|
||||
|
||||
func startLuaRouter(t *testing.T, script string, d featureDNS.Client, config *Config) *Router {
|
||||
t.Helper()
|
||||
if config == nil {
|
||||
config = &Config{}
|
||||
}
|
||||
config.Script = writeRouteScript(t, script)
|
||||
r := new(Router)
|
||||
if err := r.Init(context.Background(), config, d, &luaRouteOutboundManager{}, nil); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := r.Start(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
t.Cleanup(func() {
|
||||
if err := r.Close(); err != nil {
|
||||
t.Error(err)
|
||||
}
|
||||
})
|
||||
return r
|
||||
}
|
||||
|
||||
func TestRouterScriptStartup(t *testing.T) {
|
||||
for _, tc := range []struct{ name, script string }{
|
||||
{"syntax error", "function HandleRoute("},
|
||||
{"missing hook", "value = 1"},
|
||||
{"initialization error", `error("setup failed")`},
|
||||
} {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
r := new(Router)
|
||||
if err := r.Init(context.Background(), &Config{Script: writeRouteScript(t, tc.script)}, nil, nil, nil); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer r.Close()
|
||||
if err := r.Start(); err == nil {
|
||||
t.Fatal("Start accepted an invalid routing script")
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestRouterScriptRouting(t *testing.T) {
|
||||
var dnsCalls atomic.Int32
|
||||
d := &luaRouteDNSClient{lookup: func(string, featureDNS.IPOption) ([]net.IP, uint32, error) {
|
||||
dnsCalls.Add(1)
|
||||
return []net.IP{{1, 2, 3, 4}}, 60, nil
|
||||
}}
|
||||
r := startLuaRouter(t, `
|
||||
function HandleRoute(ctx, inbound)
|
||||
if inbound == "miss" then return nil end
|
||||
return "lua-out", "lua-rule"
|
||||
end`, d, &Config{
|
||||
DomainStrategy: Config_IpOnDemand,
|
||||
Rule: []*RoutingRule{{
|
||||
TargetTag: &RoutingRule_Tag{Tag: "json-out"},
|
||||
Networks: []net.Network{net.Network_TCP},
|
||||
}},
|
||||
})
|
||||
ctx := newLuaRouteTestContext()
|
||||
ctx.Content.SkipDNSResolve = false
|
||||
route, err := r.PickRoute(ctx)
|
||||
if err != nil || route.GetOutboundTag() != "lua-out" || route.GetRuleTag() != "lua-rule" || route.(*Route).Context != ctx {
|
||||
t.Fatalf("route = %v, %v", route, err)
|
||||
}
|
||||
ctx.Inbound.Tag = "miss"
|
||||
if route, err := r.PickRoute(ctx); route != nil || err != common.ErrNoClue {
|
||||
t.Fatalf("miss = %v, %v", route, err)
|
||||
}
|
||||
if dnsCalls.Load() != 0 {
|
||||
t.Fatal("script routing implicitly resolved DNS")
|
||||
}
|
||||
}
|
||||
|
||||
func TestRouterScriptModules(t *testing.T) {
|
||||
ips := []net.IP{{127, 0, 0, 7}}
|
||||
calls := 0
|
||||
d := &luaRouteDNSClient{lookup: func(domain string, option featureDNS.IPOption) ([]net.IP, uint32, error) {
|
||||
calls++
|
||||
if domain != "mixed.example." || !option.IPv4Enable || option.IPv6Enable || !option.FakeEnable {
|
||||
t.Fatalf("dns.Query arguments = %q, %+v", domain, option)
|
||||
}
|
||||
return ips, 17, nil
|
||||
}}
|
||||
r := startLuaRouter(t, `
|
||||
local dns = require("xray.dns")
|
||||
local matcher = require("xray.geodata").BuildIPMatcher("127.0.0.0/8")
|
||||
assert(dns.Servers == nil and type(dns.Query) == "function")
|
||||
assert(type(require("xray.log").Info) == "function")
|
||||
function HandleRoute(ctx, inbound, sourcePort, targetPort, localPort, domain)
|
||||
local ips, ttl, err = dns.Query(domain, true, false, true)
|
||||
assert(not err and ttl == 17)
|
||||
assert(matcher:AnyMatch(ips) and matcher:AnyMatch(ctx:GetTargetIPs()))
|
||||
return "out"
|
||||
end`, d, nil)
|
||||
if _, err := r.PickRoute(newLuaRouteTestContext()); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if calls != 1 {
|
||||
t.Fatalf("DNS calls = %d, want 1", calls)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRouterScriptBalancerReload(t *testing.T) {
|
||||
config := func(tag string) *Config {
|
||||
return &Config{BalancingRule: []*BalancingRule{{
|
||||
Tag: "balance", Strategy: "roundrobin", OutboundSelector: []string{tag},
|
||||
}}}
|
||||
}
|
||||
r := startLuaRouter(t, `
|
||||
local router = require("xray.router")
|
||||
function HandleRoute()
|
||||
local tag, err = router:PickOutbound("balance")
|
||||
return tag, "balanced", err
|
||||
end`, nil, config("old"))
|
||||
pick := func(want string) {
|
||||
t.Helper()
|
||||
route, err := r.PickRoute(&routing_session.Context{})
|
||||
if err != nil || route.GetOutboundTag() != want || route.GetRuleTag() != "balanced" {
|
||||
t.Fatalf("route = %v, %v, want %q", route, err, want)
|
||||
}
|
||||
}
|
||||
|
||||
pick("old")
|
||||
if err := r.SetOverrideTarget("balance", "override"); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
pick("override")
|
||||
if err := r.SetOverrideTarget("balance", ""); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := r.ReloadRules(config("new"), false); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
pick("new")
|
||||
}
|
||||
|
||||
func TestRouterScriptConcurrentBalancerReload(t *testing.T) {
|
||||
config := func(tag string) *Config {
|
||||
return &Config{BalancingRule: []*BalancingRule{{
|
||||
Tag: "balance", Strategy: "roundrobin", OutboundSelector: []string{tag},
|
||||
}}}
|
||||
}
|
||||
r := startLuaRouter(t, `
|
||||
local router = require("xray.router")
|
||||
function HandleRoute()
|
||||
local tag, err = router:PickOutbound("balance")
|
||||
return tag, nil, err
|
||||
end`, nil, config("a"))
|
||||
|
||||
var wg sync.WaitGroup
|
||||
for range 4 {
|
||||
wg.Go(func() {
|
||||
for range 20 {
|
||||
route, err := r.PickRoute(&routing_session.Context{})
|
||||
if err != nil {
|
||||
t.Errorf("PickRoute: %v", err)
|
||||
return
|
||||
}
|
||||
if tag := route.GetOutboundTag(); tag != "a" && tag != "b" {
|
||||
t.Errorf("unexpected tag %q", tag)
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
wg.Go(func() {
|
||||
for range 20 {
|
||||
for _, tag := range []string{"a", "b"} {
|
||||
if err := r.ReloadRules(config(tag), false); err != nil {
|
||||
t.Error(err)
|
||||
return
|
||||
}
|
||||
}
|
||||
}
|
||||
})
|
||||
wg.Wait()
|
||||
}
|
||||
|
||||
func TestRouterScriptStateReuse(t *testing.T) {
|
||||
r := startLuaRouter(t, `
|
||||
local calls = 0
|
||||
function HandleRoute(ctx, inbound)
|
||||
calls = calls + 1
|
||||
if inbound == "miss" then return nil end
|
||||
if inbound == "fail" then error("failed") end
|
||||
return tostring(calls)
|
||||
end`, nil, nil)
|
||||
ctx := newLuaRouteTestContext()
|
||||
pick := func(want string) {
|
||||
t.Helper()
|
||||
route, err := r.PickRoute(ctx)
|
||||
if err != nil || route.GetOutboundTag() != want {
|
||||
t.Fatalf("route = %v, %v, want %q", route, err, want)
|
||||
}
|
||||
}
|
||||
|
||||
pick("1")
|
||||
ctx.Inbound.Tag = "miss"
|
||||
if _, err := r.PickRoute(ctx); err != common.ErrNoClue {
|
||||
t.Fatalf("miss = %v", err)
|
||||
}
|
||||
ctx.Inbound.Tag = "in"
|
||||
pick("3")
|
||||
ctx.Inbound.Tag = "fail"
|
||||
if _, err := r.PickRoute(ctx); err == nil {
|
||||
t.Fatal("script error was ignored")
|
||||
}
|
||||
ctx.Inbound.Tag = "in"
|
||||
pick("1")
|
||||
}
|
||||
|
||||
func TestRouterScriptDNSDispatcherReentry(t *testing.T) {
|
||||
conn, err := stdnet.ListenPacket("udp4", "127.0.0.1:0")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
port := conn.LocalAddr().(*stdnet.UDPAddr).Port
|
||||
ready, stopped := make(chan struct{}), make(chan error, 1)
|
||||
var queries atomic.Int32
|
||||
server := &wireDNS.Server{
|
||||
PacketConn: conn,
|
||||
NotifyStartedFunc: func() {
|
||||
close(ready)
|
||||
},
|
||||
Handler: wireDNS.HandlerFunc(func(w wireDNS.ResponseWriter, query *wireDNS.Msg) {
|
||||
queries.Add(1)
|
||||
response := new(wireDNS.Msg).SetReply(query)
|
||||
for _, question := range query.Question {
|
||||
if question.Name == "nested.example." && question.Qtype == wireDNS.TypeA {
|
||||
response.Answer = append(response.Answer, &wireDNS.A{
|
||||
Hdr: wireDNS.RR_Header{Name: question.Name, Rrtype: wireDNS.TypeA, Class: wireDNS.ClassINET, Ttl: 60},
|
||||
A: stdnet.IP{127, 0, 0, 7},
|
||||
})
|
||||
}
|
||||
}
|
||||
if err := w.WriteMsg(response); err != nil {
|
||||
t.Error(err)
|
||||
}
|
||||
}),
|
||||
}
|
||||
go func() { stopped <- server.ActivateAndServe() }()
|
||||
defer func() {
|
||||
server.Shutdown()
|
||||
select {
|
||||
case err := <-stopped:
|
||||
if err != nil {
|
||||
t.Error(err)
|
||||
}
|
||||
case <-time.After(3 * time.Second):
|
||||
t.Error("DNS server did not stop")
|
||||
}
|
||||
}()
|
||||
select {
|
||||
case <-ready:
|
||||
case err := <-stopped:
|
||||
t.Fatalf("DNS server startup: %v", err)
|
||||
case <-time.After(3 * time.Second):
|
||||
t.Fatal("DNS server did not start")
|
||||
}
|
||||
|
||||
dnsScript := writeRouteScript(t, `
|
||||
local server = require("xray.dns").Servers[1]
|
||||
function HandleDNSQuery(domain, ipv4, ipv6, fake)
|
||||
return server:Query(domain, ipv4, ipv6, fake)
|
||||
end`)
|
||||
routerScript := writeRouteScript(t, `
|
||||
local router = require("xray.router")
|
||||
local dns = require("xray.dns")
|
||||
local matcher = require("xray.geodata").BuildIPMatcher("127.0.0.7")
|
||||
local active = false
|
||||
function HandleRoute(ctx, inbound, sourcePort, targetPort, localPort, domain, network,
|
||||
protocol, user, vlessRoute, skipDNSResolve)
|
||||
assert(not active, "borrowed Router VM reentered")
|
||||
if inbound == "dns" then
|
||||
assert(network == router.NetworkUDP and skipDNSResolve == false)
|
||||
return "direct", "dns-route"
|
||||
end
|
||||
active = true
|
||||
local ips, ttl, err = dns.Query("nested.example", true, false, false)
|
||||
assert(not err and matcher:AnyMatch(ips) and active)
|
||||
active = false
|
||||
return "direct", "outer-route"
|
||||
end`)
|
||||
instance, err := core.New(&core.Config{
|
||||
App: []*serial.TypedMessage{
|
||||
serial.ToTypedMessage(&appdns.Config{
|
||||
Tag: "dns", Script: dnsScript, DisableCache: true,
|
||||
NameServer: []*appdns.NameServer{{
|
||||
Id: "upstream", TimeoutMs: 1000,
|
||||
Address: &net.Endpoint{
|
||||
Network: net.Network_UDP,
|
||||
Address: &net.IPOrDomain{Address: &net.IPOrDomain_Ip{Ip: []byte{127, 0, 0, 1}}},
|
||||
Port: uint32(port),
|
||||
},
|
||||
}},
|
||||
}),
|
||||
serial.ToTypedMessage(&Config{Script: routerScript}),
|
||||
serial.ToTypedMessage(&dispatcher.Config{}),
|
||||
serial.ToTypedMessage(&proxyman.OutboundConfig{}),
|
||||
},
|
||||
Outbound: []*core.OutboundHandlerConfig{
|
||||
{Tag: "default", ProxySettings: serial.ToTypedMessage(&blackhole.Config{})},
|
||||
{Tag: "direct", ProxySettings: serial.ToTypedMessage(&freedom.Config{
|
||||
FinalRules: []*freedom.FinalRuleConfig{{Action: freedom.RuleAction_Allow}},
|
||||
})},
|
||||
},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer instance.Close()
|
||||
if err := instance.Start(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
r := instance.GetFeature(routing.RouterType()).(*Router)
|
||||
route, err := r.PickRoute(newLuaRouteTestContext())
|
||||
if err != nil || route.GetOutboundTag() != "direct" || route.GetRuleTag() != "outer-route" {
|
||||
t.Fatalf("nested DNS routing = %v, %v", route, err)
|
||||
}
|
||||
if queries.Load() == 0 {
|
||||
t.Fatal("DNS query did not pass through the dispatcher")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,58 @@
|
||||
package geodata
|
||||
|
||||
import (
|
||||
lua "github.com/yuin/gopher-lua"
|
||||
luar "layeh.com/gopher-luar"
|
||||
)
|
||||
|
||||
// RegisterLua makes xray.geodata available to require in an LState.
|
||||
func RegisterLua(L *lua.LState) {
|
||||
L.PreloadModule("xray.geodata", func(L *lua.LState) int {
|
||||
module := L.NewTable()
|
||||
|
||||
module.RawSetString("BuildDomainMatcher", L.NewFunction(func(L *lua.LState) int {
|
||||
parsed, err := ParseDomainRules(luaRules(L), Domain_Domain)
|
||||
if err != nil {
|
||||
L.RaiseError("%v", err)
|
||||
return 0
|
||||
}
|
||||
matcher, err := DomainReg.BuildDomainMatcher(parsed)
|
||||
if err != nil {
|
||||
L.RaiseError("%v", err)
|
||||
return 0
|
||||
}
|
||||
L.Push(luar.New(L, matcher))
|
||||
return 1
|
||||
}))
|
||||
|
||||
module.RawSetString("BuildIPMatcher", L.NewFunction(func(L *lua.LState) int {
|
||||
parsed, err := ParseIPRules(luaRules(L))
|
||||
if err != nil {
|
||||
L.RaiseError("%v", err)
|
||||
return 0
|
||||
}
|
||||
matcher, err := IPReg.BuildIPMatcher(parsed)
|
||||
if err != nil {
|
||||
L.RaiseError("%v", err)
|
||||
return 0
|
||||
}
|
||||
L.Push(luar.New(L, matcher))
|
||||
return 1
|
||||
}))
|
||||
L.Push(module)
|
||||
return 1
|
||||
})
|
||||
}
|
||||
|
||||
func luaRules(L *lua.LState) []string {
|
||||
rules := make([]string, L.GetTop())
|
||||
for i := range rules {
|
||||
value, ok := L.Get(i + 1).(lua.LString)
|
||||
if !ok {
|
||||
L.RaiseError("geodata rules must be strings")
|
||||
return nil
|
||||
}
|
||||
rules[i] = string(value)
|
||||
}
|
||||
return rules
|
||||
}
|
||||
@@ -0,0 +1,66 @@
|
||||
package geodata
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/xtls/xray-core/common/net"
|
||||
lua "github.com/yuin/gopher-lua"
|
||||
)
|
||||
|
||||
func TestLuaIPMatcher(t *testing.T) {
|
||||
L := lua.NewState()
|
||||
defer L.Close()
|
||||
RegisterLua(L)
|
||||
ip := L.NewUserData()
|
||||
ip.Value = net.ParseIP("127.0.0.1")
|
||||
L.SetGlobal("ip", ip)
|
||||
ips := L.NewUserData()
|
||||
ips.Value = []net.IP{ip.Value.(net.IP), net.ParseIP("8.8.8.8")}
|
||||
L.SetGlobal("ips", ips)
|
||||
if err := L.DoString(`
|
||||
local matcher = require("xray.geodata").BuildIPMatcher("127.0.0.0/8", "::1")
|
||||
assert(matcher:Match(ip))
|
||||
assert(matcher:AnyMatch(ips))
|
||||
assert(not matcher:Matches(ips))
|
||||
local matched, unmatched = matcher:FilterIPs(ips)
|
||||
assert(type(matched) == "userdata" and type(unmatched) == "userdata")
|
||||
assert(#matched == 1 and #unmatched == 1)
|
||||
`); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestLuaDomainMatcher(t *testing.T) {
|
||||
L := lua.NewState()
|
||||
defer L.Close()
|
||||
RegisterLua(L)
|
||||
if err := L.DoString(`
|
||||
local matcher = require("xray.geodata").BuildDomainMatcher("example.com", "full:other.com")
|
||||
assert(matcher:MatchAny("example.com"))
|
||||
assert(matcher:MatchAny("www.example.com"))
|
||||
assert(matcher:MatchAny("other.com"))
|
||||
assert(not matcher:MatchAny("www.other.com"))
|
||||
assert(#(matcher:Match("www.example.com")) == 1)
|
||||
`); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestLuaMatchersRejectInvalidRules(t *testing.T) {
|
||||
for _, tc := range []struct {
|
||||
name string
|
||||
script string
|
||||
}{
|
||||
{"IP rule", `require("xray.geodata").BuildIPMatcher("not-an-ip")`},
|
||||
{"non-string domain rule", `require("xray.geodata").BuildDomainMatcher("example.com", true)`},
|
||||
} {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
L := lua.NewState()
|
||||
defer L.Close()
|
||||
RegisterLua(L)
|
||||
if err := L.DoString(tc.script); err == nil {
|
||||
t.Fatal("invalid geodata rule was accepted")
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
})
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,3 @@
|
||||
// Package lua provides shared GopherLua programs, state management, and value
|
||||
// conversion and validation helpers for Xray scripts.
|
||||
package lua
|
||||
@@ -0,0 +1,150 @@
|
||||
package lua
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
glua "github.com/yuin/gopher-lua"
|
||||
)
|
||||
|
||||
const maxIdleStates = 16
|
||||
|
||||
// Pool lends each state to one caller at a time. It grows on contention and
|
||||
// keeps up to maxIdleStates idle states until Close. Acquire/Release callers
|
||||
// decide reusability; WithState uses its callback's error.
|
||||
type Pool struct {
|
||||
ctx context.Context
|
||||
cancel context.CancelFunc
|
||||
timeout time.Duration
|
||||
|
||||
factory LStateFactory
|
||||
idle []*glua.LState
|
||||
top int
|
||||
|
||||
mu sync.Mutex
|
||||
active sync.WaitGroup
|
||||
closed bool
|
||||
}
|
||||
|
||||
// NewPool tests the factory by creating one state during initialization.
|
||||
func NewPool(ctx context.Context, timeout time.Duration, factory LStateFactory) (*Pool, error) {
|
||||
if timeout <= 0 {
|
||||
return nil, errors.New("Lua pool timeout must be positive")
|
||||
}
|
||||
|
||||
poolCtx, cancel := context.WithCancel(ctx)
|
||||
|
||||
state, err := factory(poolCtx)
|
||||
if err != nil {
|
||||
cancel()
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return &Pool{ctx: poolCtx, cancel: cancel, timeout: timeout, factory: factory, idle: []*glua.LState{state}, top: state.GetTop()}, nil
|
||||
}
|
||||
|
||||
// Acquire returns an initialized exclusive state, growing the pool if necessary.
|
||||
// ctx is passed to the factory for state creation; nil uses the pool context.
|
||||
func (p *Pool) Acquire(ctx context.Context) (*glua.LState, error) {
|
||||
p.mu.Lock()
|
||||
if p.closed {
|
||||
p.mu.Unlock()
|
||||
return nil, errors.New("Lua pool is closed")
|
||||
}
|
||||
if err := p.ctx.Err(); err != nil {
|
||||
p.mu.Unlock()
|
||||
return nil, err
|
||||
}
|
||||
if ctx == nil {
|
||||
ctx = p.ctx
|
||||
} else if err := ctx.Err(); err != nil {
|
||||
p.mu.Unlock()
|
||||
return nil, err
|
||||
}
|
||||
|
||||
p.active.Add(1)
|
||||
|
||||
n := len(p.idle)
|
||||
if n != 0 {
|
||||
state := p.idle[n-1]
|
||||
p.idle = p.idle[:n-1]
|
||||
p.mu.Unlock()
|
||||
return state, nil
|
||||
}
|
||||
p.mu.Unlock()
|
||||
|
||||
// TODO: Limit the total number of states. When the limit is reached, wait
|
||||
// for a Release instead of creating another state; allow the wait to be
|
||||
// cancelled by the caller or by Close.
|
||||
state, err := p.factory(ctx)
|
||||
if err != nil {
|
||||
p.active.Done()
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return state, nil
|
||||
}
|
||||
|
||||
// WithState runs work on an exclusive state and releases it afterward.
|
||||
// Nil ctx and zero timeout use pool defaults. The timeout starts after acquisition.
|
||||
func (p *Pool) WithState(ctx context.Context, timeout time.Duration, work func(*glua.LState) error) error {
|
||||
state, err := p.Acquire(ctx)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if ctx == nil {
|
||||
ctx = p.ctx
|
||||
}
|
||||
if timeout == 0 {
|
||||
timeout = p.timeout
|
||||
}
|
||||
ctx, cancel := context.WithTimeout(ctx, timeout)
|
||||
state.SetContext(ctx)
|
||||
reusable := false
|
||||
defer func() {
|
||||
cancel()
|
||||
p.Release(state, reusable)
|
||||
}()
|
||||
err = work(state)
|
||||
reusable = err == nil
|
||||
return err
|
||||
}
|
||||
|
||||
// Release resets a state for reuse or closes it.
|
||||
func (p *Pool) Release(state *glua.LState, reusable bool) {
|
||||
if reusable {
|
||||
state.RemoveContext()
|
||||
state.SetTop(p.top)
|
||||
p.mu.Lock()
|
||||
if !p.closed && p.ctx.Err() == nil && len(p.idle) < maxIdleStates {
|
||||
p.idle = append(p.idle, state)
|
||||
} else {
|
||||
reusable = false
|
||||
}
|
||||
p.mu.Unlock()
|
||||
}
|
||||
|
||||
if !reusable {
|
||||
state.Close()
|
||||
}
|
||||
|
||||
p.active.Done()
|
||||
}
|
||||
|
||||
// Close cancels the pool context, closes idle states, and waits for borrowed states.
|
||||
func (p *Pool) Close() {
|
||||
p.mu.Lock()
|
||||
if !p.closed {
|
||||
p.closed = true
|
||||
p.cancel()
|
||||
for _, state := range p.idle {
|
||||
state.Close()
|
||||
}
|
||||
p.idle = nil
|
||||
}
|
||||
p.mu.Unlock()
|
||||
|
||||
p.active.Wait()
|
||||
}
|
||||
@@ -0,0 +1,466 @@
|
||||
package lua
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
glua "github.com/yuin/gopher-lua"
|
||||
)
|
||||
|
||||
func newTestPool(t testing.TB, ctx context.Context, timeout time.Duration, factory LStateFactory) *Pool {
|
||||
t.Helper()
|
||||
pool, err := NewPool(ctx, timeout, factory)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
t.Cleanup(pool.Close)
|
||||
return pool
|
||||
}
|
||||
|
||||
func assertPoolCloseBlocked(t *testing.T, done <-chan struct{}) {
|
||||
t.Helper()
|
||||
select {
|
||||
case <-done:
|
||||
t.Fatal("Close returned while work was still active")
|
||||
case <-time.After(20 * time.Millisecond):
|
||||
}
|
||||
}
|
||||
|
||||
func TestPoolTimeoutValidation(t *testing.T) {
|
||||
for _, tc := range []struct {
|
||||
name string
|
||||
timeout time.Duration
|
||||
wantErr bool
|
||||
}{
|
||||
{"zero", 0, true},
|
||||
{"negative", -time.Nanosecond, true},
|
||||
{"positive", time.Nanosecond, false},
|
||||
} {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
called := false
|
||||
pool, err := NewPool(context.Background(), tc.timeout, func(context.Context) (*glua.LState, error) {
|
||||
called = true
|
||||
return glua.NewState(), nil
|
||||
})
|
||||
if pool != nil {
|
||||
t.Cleanup(pool.Close)
|
||||
}
|
||||
if (err != nil) != tc.wantErr {
|
||||
t.Fatalf("NewPool error = %v, want error %t", err, tc.wantErr)
|
||||
}
|
||||
if tc.wantErr && (pool != nil || called) {
|
||||
t.Fatal("invalid timeout created a pool or called the factory")
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestPoolFactoryFailure(t *testing.T) {
|
||||
failure := errors.New("factory failed")
|
||||
_, err := NewPool(context.Background(), time.Second, func(context.Context) (*glua.LState, error) {
|
||||
return nil, failure
|
||||
})
|
||||
if !errors.Is(err, failure) {
|
||||
t.Fatalf("NewPool error = %v, want original factory error", err)
|
||||
}
|
||||
|
||||
calls := 0
|
||||
pool := newTestPool(t, context.Background(), time.Second, func(context.Context) (*glua.LState, error) {
|
||||
calls++
|
||||
if calls == 1 {
|
||||
return glua.NewState(), nil
|
||||
}
|
||||
return nil, failure
|
||||
})
|
||||
state, err := pool.Acquire(nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer pool.Release(state, true)
|
||||
err = pool.WithState(nil, 0, func(*glua.LState) error {
|
||||
t.Error("work ran after factory failure")
|
||||
return nil
|
||||
})
|
||||
if !errors.Is(err, failure) {
|
||||
t.Fatalf("WithState error = %v, want original factory error", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPoolReusesStatesAndLimitsIdle(t *testing.T) {
|
||||
created := 0
|
||||
pool := newTestPool(t, context.Background(), time.Second, func(context.Context) (*glua.LState, error) {
|
||||
created++
|
||||
return glua.NewState(), nil
|
||||
})
|
||||
var borrowed []*glua.LState
|
||||
defer func() {
|
||||
for _, state := range borrowed {
|
||||
pool.Release(state, false)
|
||||
}
|
||||
}()
|
||||
for range maxIdleStates + 3 {
|
||||
state, err := pool.Acquire(nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
borrowed = append(borrowed, state)
|
||||
state.SetContext(context.Background())
|
||||
}
|
||||
states := borrowed
|
||||
for _, state := range states {
|
||||
pool.Release(state, true)
|
||||
}
|
||||
borrowed = nil
|
||||
open := 0
|
||||
for _, state := range states {
|
||||
if !state.IsClosed() {
|
||||
if state.Context() != nil {
|
||||
t.Fatal("Release left a context on a reusable state")
|
||||
}
|
||||
open++
|
||||
}
|
||||
}
|
||||
if open != maxIdleStates {
|
||||
t.Fatalf("retained %d states, want %d", open, maxIdleStates)
|
||||
}
|
||||
if err := pool.WithState(nil, 0, func(*glua.LState) error { return nil }); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if created != len(states) {
|
||||
t.Fatalf("created %d states, want %d", created, len(states))
|
||||
}
|
||||
pool.Close()
|
||||
for _, state := range states {
|
||||
if !state.IsClosed() {
|
||||
t.Fatal("Close left an idle state open")
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestPoolWithStateOptions(t *testing.T) {
|
||||
key := struct{}{}
|
||||
parent := context.WithValue(context.Background(), key, "pool")
|
||||
caller := context.WithValue(context.Background(), key, "caller")
|
||||
pool := newTestPool(t, parent, time.Second, func(context.Context) (*glua.LState, error) {
|
||||
return glua.NewState(), nil
|
||||
})
|
||||
for _, tc := range []struct {
|
||||
name string
|
||||
ctx context.Context
|
||||
timeout time.Duration
|
||||
wantValue string
|
||||
wantTimeout time.Duration
|
||||
}{
|
||||
{"defaults", nil, 0, "pool", time.Second},
|
||||
{"context", caller, 0, "caller", time.Second},
|
||||
{"timeout", nil, 2 * time.Second, "pool", 2 * time.Second},
|
||||
{"both", caller, 2 * time.Second, "caller", 2 * time.Second},
|
||||
} {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
started := time.Now()
|
||||
err := pool.WithState(tc.ctx, tc.timeout, func(L *glua.LState) error {
|
||||
ctx := L.Context()
|
||||
if ctx.Value(key) != tc.wantValue {
|
||||
t.Errorf("context value = %v, want %q", ctx.Value(key), tc.wantValue)
|
||||
}
|
||||
deadline, ok := ctx.Deadline()
|
||||
if !ok || deadline.Before(started.Add(tc.wantTimeout)) || deadline.After(time.Now().Add(tc.wantTimeout)) {
|
||||
t.Errorf("deadline = %v, want timeout %v", deadline, tc.wantTimeout)
|
||||
}
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestPoolFactoryContext(t *testing.T) {
|
||||
caller, cancel := context.WithTimeout(context.Background(), time.Minute)
|
||||
defer cancel()
|
||||
for _, tc := range []struct {
|
||||
name string
|
||||
ctx context.Context
|
||||
}{
|
||||
{"default", nil},
|
||||
{"caller", caller},
|
||||
} {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
var contexts []context.Context
|
||||
pool := newTestPool(t, context.Background(), time.Second, func(ctx context.Context) (*glua.LState, error) {
|
||||
contexts = append(contexts, ctx)
|
||||
return glua.NewState(), nil
|
||||
})
|
||||
state, err := pool.Acquire(nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer pool.Release(state, true)
|
||||
if err := pool.WithState(tc.ctx, 2*time.Second, func(*glua.LState) error { return nil }); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
want := tc.ctx
|
||||
if want == nil {
|
||||
want = pool.ctx
|
||||
}
|
||||
if len(contexts) != 2 || contexts[0] != pool.ctx || contexts[1] != want {
|
||||
t.Fatal("factory did not receive the initialization and acquisition contexts unchanged")
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestPoolWithStateLifecycle(t *testing.T) {
|
||||
failure := errors.New("work failed")
|
||||
for _, tc := range []struct {
|
||||
name string
|
||||
work func(*glua.LState, context.CancelFunc) error
|
||||
reusable bool
|
||||
wantPanic bool
|
||||
wantErr error
|
||||
}{
|
||||
{"success", func(*glua.LState, context.CancelFunc) error { return nil }, true, false, nil},
|
||||
{"canceled success", func(_ *glua.LState, cancel context.CancelFunc) error {
|
||||
cancel()
|
||||
return nil
|
||||
}, true, false, nil},
|
||||
{"error", func(*glua.LState, context.CancelFunc) error { return failure }, false, false, failure},
|
||||
{"timeout", func(L *glua.LState, _ context.CancelFunc) error { return L.DoString("while true do end") }, false, false, nil},
|
||||
{"panic", func(*glua.LState, context.CancelFunc) error { panic(failure) }, false, true, nil},
|
||||
} {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
pool := newTestPool(t, context.Background(), 10*time.Millisecond, func(context.Context) (*glua.LState, error) {
|
||||
state := glua.NewState()
|
||||
state.Push(glua.LTrue)
|
||||
return state, nil
|
||||
})
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
defer cancel()
|
||||
var state *glua.LState
|
||||
var workCtx context.Context
|
||||
var recovered any
|
||||
err := func() (err error) {
|
||||
defer func() { recovered = recover() }()
|
||||
return pool.WithState(ctx, 0, func(L *glua.LState) error {
|
||||
state, workCtx = L, L.Context()
|
||||
L.Push(glua.LFalse)
|
||||
return tc.work(L, cancel)
|
||||
})
|
||||
}()
|
||||
if tc.wantPanic {
|
||||
if recovered != failure {
|
||||
t.Fatalf("panic = %v, want original panic", recovered)
|
||||
}
|
||||
} else {
|
||||
if recovered != nil || (err == nil) != tc.reusable {
|
||||
t.Fatalf("WithState error = %v, panic = %v", err, recovered)
|
||||
}
|
||||
if tc.wantErr != nil && !errors.Is(err, tc.wantErr) {
|
||||
t.Fatalf("WithState error = %v, want %v", err, tc.wantErr)
|
||||
}
|
||||
}
|
||||
if workCtx.Err() == nil {
|
||||
t.Fatal("WithState did not cancel the execution context")
|
||||
}
|
||||
if closed := state.IsClosed(); closed == tc.reusable {
|
||||
t.Fatalf("state closed = %t, want %t", closed, !tc.reusable)
|
||||
}
|
||||
if tc.reusable && (state.Context() != nil || state.GetTop() != 1 || state.Get(1) != glua.LTrue) {
|
||||
t.Fatal("WithState did not reset the state for reuse")
|
||||
}
|
||||
if err := pool.WithState(nil, 0, func(L *glua.LState) error {
|
||||
if (L == state) != tc.reusable {
|
||||
t.Error("unexpected state reuse")
|
||||
}
|
||||
return nil
|
||||
}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestPoolClose(t *testing.T) {
|
||||
pool := newTestPool(t, context.Background(), time.Second, func(context.Context) (*glua.LState, error) {
|
||||
return glua.NewState(), nil
|
||||
})
|
||||
finishCtx, finish := context.WithCancel(context.Background())
|
||||
t.Cleanup(finish)
|
||||
started, done := make(chan *glua.LState, 1), make(chan error, 1)
|
||||
var workCtx context.Context
|
||||
go func() {
|
||||
done <- pool.WithState(nil, 0, func(L *glua.LState) error {
|
||||
workCtx = L.Context()
|
||||
started <- L
|
||||
<-finishCtx.Done()
|
||||
return nil
|
||||
})
|
||||
}()
|
||||
var state *glua.LState
|
||||
select {
|
||||
case state = <-started:
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("WithState did not start")
|
||||
}
|
||||
closed := make(chan struct{})
|
||||
go func() {
|
||||
pool.Close()
|
||||
close(closed)
|
||||
}()
|
||||
select {
|
||||
case <-workCtx.Done():
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("Close did not cancel work using the pool context")
|
||||
}
|
||||
if !errors.Is(workCtx.Err(), context.Canceled) {
|
||||
t.Fatalf("work context error = %v, want context.Canceled", workCtx.Err())
|
||||
}
|
||||
assertPoolCloseBlocked(t, closed)
|
||||
finish()
|
||||
select {
|
||||
case err := <-done:
|
||||
if err != nil {
|
||||
t.Fatalf("successful work returned an error: %v", err)
|
||||
}
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("WithState did not finish")
|
||||
}
|
||||
select {
|
||||
case <-closed:
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("Close did not finish after WithState")
|
||||
}
|
||||
if !state.IsClosed() {
|
||||
t.Fatal("Release returned a state to a closed pool")
|
||||
}
|
||||
if state, err := pool.Acquire(nil); state != nil || err == nil || errors.Is(err, context.Canceled) {
|
||||
t.Fatalf("Acquire after Close = %v, %v; want closed pool error", state, err)
|
||||
}
|
||||
pool.Close()
|
||||
}
|
||||
|
||||
func TestPoolCloseWaitsForFactory(t *testing.T) {
|
||||
finishCtx, finish := context.WithCancel(context.Background())
|
||||
started, canceled := make(chan struct{}), make(chan struct{})
|
||||
first := true
|
||||
pool := newTestPool(t, context.Background(), time.Second, func(ctx context.Context) (*glua.LState, error) {
|
||||
if first {
|
||||
first = false
|
||||
return glua.NewState(), nil
|
||||
}
|
||||
close(started)
|
||||
<-ctx.Done()
|
||||
close(canceled)
|
||||
<-finishCtx.Done()
|
||||
return nil, ctx.Err()
|
||||
})
|
||||
t.Cleanup(finish)
|
||||
state, err := pool.Acquire(nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
pool.Release(state, false)
|
||||
acquireDone := make(chan error, 1)
|
||||
go func() {
|
||||
_, err := pool.Acquire(nil)
|
||||
acquireDone <- err
|
||||
}()
|
||||
select {
|
||||
case <-started:
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("state creation did not start")
|
||||
}
|
||||
closed := make(chan struct{})
|
||||
go func() {
|
||||
pool.Close()
|
||||
close(closed)
|
||||
}()
|
||||
select {
|
||||
case <-canceled:
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("Close did not cancel state creation")
|
||||
}
|
||||
assertPoolCloseBlocked(t, closed)
|
||||
finish()
|
||||
select {
|
||||
case err := <-acquireDone:
|
||||
if !errors.Is(err, context.Canceled) {
|
||||
t.Fatalf("Acquire error = %v, want context.Canceled", err)
|
||||
}
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("state creation did not finish")
|
||||
}
|
||||
select {
|
||||
case <-closed:
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("Close did not finish after state creation")
|
||||
}
|
||||
}
|
||||
|
||||
func TestPoolCloseWaitsForCallerContext(t *testing.T) {
|
||||
pool := newTestPool(t, context.Background(), time.Minute, func(context.Context) (*glua.LState, error) {
|
||||
return glua.NewState(), nil
|
||||
})
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
t.Cleanup(cancel)
|
||||
started, done := make(chan context.Context, 1), make(chan error, 1)
|
||||
go func() {
|
||||
done <- pool.WithState(ctx, 0, func(L *glua.LState) error {
|
||||
started <- L.Context()
|
||||
<-L.Context().Done()
|
||||
return L.Context().Err()
|
||||
})
|
||||
}()
|
||||
var workCtx context.Context
|
||||
select {
|
||||
case workCtx = <-started:
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("WithState did not start")
|
||||
}
|
||||
closed := make(chan struct{})
|
||||
go func() {
|
||||
pool.Close()
|
||||
close(closed)
|
||||
}()
|
||||
select {
|
||||
case <-pool.ctx.Done():
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("Close did not cancel the pool context")
|
||||
}
|
||||
assertPoolCloseBlocked(t, closed)
|
||||
if workCtx.Err() != nil || ctx.Err() != nil {
|
||||
t.Fatal("Close canceled the caller's execution context")
|
||||
}
|
||||
cancel()
|
||||
select {
|
||||
case err := <-done:
|
||||
if !errors.Is(err, context.Canceled) {
|
||||
t.Fatalf("WithState error = %v, want context.Canceled", err)
|
||||
}
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("WithState did not stop after caller cancellation")
|
||||
}
|
||||
select {
|
||||
case <-closed:
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("Close did not finish after WithState")
|
||||
}
|
||||
}
|
||||
|
||||
func BenchmarkPoolAcquireRelease(b *testing.B) {
|
||||
pool := newTestPool(b, context.Background(), time.Second, func(context.Context) (*glua.LState, error) {
|
||||
return glua.NewState(), nil
|
||||
})
|
||||
b.ReportAllocs()
|
||||
b.ResetTimer()
|
||||
for i := 0; i < b.N; i++ {
|
||||
state, err := pool.Acquire(nil)
|
||||
if err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
pool.Release(state, true)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,77 @@
|
||||
package lua
|
||||
|
||||
import (
|
||||
"bufio"
|
||||
"context"
|
||||
"os"
|
||||
"time"
|
||||
|
||||
glua "github.com/yuin/gopher-lua"
|
||||
"github.com/yuin/gopher-lua/parse"
|
||||
)
|
||||
|
||||
// Program holds immutable bytecode that can be run by independent LStates.
|
||||
type Program struct {
|
||||
proto *glua.FunctionProto
|
||||
}
|
||||
|
||||
// LStateFactory returns a fully initialized state or nil and an error.
|
||||
// Implementations must close partial states on failure; callers own successful states.
|
||||
type LStateFactory func(context.Context) (*glua.LState, error)
|
||||
|
||||
// CompileFile reads and compiles a Lua file once.
|
||||
func CompileFile(path string) (*Program, error) {
|
||||
f, err := os.Open(path)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer f.Close()
|
||||
chunk, err := parse.Parse(bufio.NewReader(f), path)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
proto, err := glua.Compile(chunk, path)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &Program{proto: proto}, nil
|
||||
}
|
||||
|
||||
// NewState creates a state, runs register, executes the program under ctx, and
|
||||
// runs validate. It removes the initialization context before returning a state
|
||||
// owned by the caller.
|
||||
func (p *Program) NewState(ctx context.Context, register func(*glua.LState), validate func(*glua.LState) error) (*glua.LState, error) {
|
||||
L := glua.NewState()
|
||||
valid := false
|
||||
defer func() {
|
||||
if !valid {
|
||||
L.Close()
|
||||
}
|
||||
}()
|
||||
L.SetContext(ctx)
|
||||
defer L.RemoveContext()
|
||||
if register != nil {
|
||||
register(L)
|
||||
}
|
||||
L.Push(L.NewFunctionFromProto(p.proto))
|
||||
// Execute the Lua script's top level.
|
||||
if err := L.PCall(0, 0, nil); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if validate != nil {
|
||||
if err := validate(L); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
valid = true
|
||||
return L, nil
|
||||
}
|
||||
|
||||
// NewStateFactory returns a factory that gives each state an initialization timeout.
|
||||
func (p *Program) NewStateFactory(initTimeout time.Duration, register func(*glua.LState), validate func(*glua.LState) error) LStateFactory {
|
||||
return func(ctx context.Context) (*glua.LState, error) {
|
||||
initCtx, cancel := context.WithTimeout(ctx, initTimeout)
|
||||
defer cancel()
|
||||
return p.NewState(initCtx, register, validate)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,76 @@
|
||||
package lua
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
|
||||
glua "github.com/yuin/gopher-lua"
|
||||
)
|
||||
|
||||
func TestProgramStatesAreIndependent(t *testing.T) {
|
||||
path := filepath.Join(t.TempDir(), "state.lua")
|
||||
if err := os.WriteFile(path, []byte("value = (value or 0) + 1"), 0o600); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
program, err := CompileFile(path)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
first, err := program.NewState(context.Background(), nil, nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer first.Close()
|
||||
first.SetGlobal("value", glua.LNumber(42))
|
||||
second, err := program.NewState(context.Background(), nil, nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer second.Close()
|
||||
if got := second.GetGlobal("value"); got != glua.LNumber(1) {
|
||||
t.Fatalf("second state value = %v, want 1", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestProgramInitializationObservesCancellation(t *testing.T) {
|
||||
path := filepath.Join(t.TempDir(), "loop.lua")
|
||||
if err := os.WriteFile(path, []byte("while true do end"), 0o600); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
program, err := CompileFile(path)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
cancel()
|
||||
state, err := program.NewState(ctx, nil, nil)
|
||||
if err == nil || state != nil {
|
||||
if state != nil {
|
||||
state.Close()
|
||||
}
|
||||
t.Fatalf("NewState with canceled context = %v, %v; want nil state and error", state, err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestNewStateClosesFailedValidation(t *testing.T) {
|
||||
path := filepath.Join(t.TempDir(), "state.lua")
|
||||
if err := os.WriteFile(path, []byte("value = 1"), 0o600); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
program, err := CompileFile(path)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
wantErr := errors.New("invalid script")
|
||||
var checked *glua.LState
|
||||
L, err := program.NewState(context.Background(), nil, func(L *glua.LState) error {
|
||||
checked = L
|
||||
return wantErr
|
||||
})
|
||||
if L != nil || !errors.Is(err, wantErr) || checked == nil || !checked.IsClosed() {
|
||||
t.Fatalf("state = %v, error = %v, checked state closed = %t", L, err, checked != nil && checked.IsClosed())
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,95 @@
|
||||
package lua
|
||||
|
||||
import (
|
||||
"math"
|
||||
|
||||
"github.com/xtls/xray-core/common/errors"
|
||||
glua "github.com/yuin/gopher-lua"
|
||||
)
|
||||
|
||||
type number interface {
|
||||
~int | ~int8 | ~int16 | ~int32 | ~int64 |
|
||||
~uint | ~uint8 | ~uint16 | ~uint32 | ~uint64 | ~uintptr |
|
||||
~float32 | ~float64
|
||||
}
|
||||
|
||||
// PushNumber converts a Go number to a Lua number and pushes it.
|
||||
func PushNumber[T number](L *glua.LState, value T) {
|
||||
L.Push(glua.LNumber(value))
|
||||
}
|
||||
|
||||
// PushString converts a Go string to a Lua string and pushes it.
|
||||
func PushString(L *glua.LState, value string) {
|
||||
L.Push(glua.LString(value))
|
||||
}
|
||||
|
||||
// PushNil pushes Lua nil.
|
||||
func PushNil(L *glua.LState) {
|
||||
L.Push(glua.LNil)
|
||||
}
|
||||
|
||||
// PushUserData pushes a native Go value without copying it.
|
||||
func PushUserData(L *glua.LState, value any) {
|
||||
ud := L.NewUserData()
|
||||
ud.Value = value
|
||||
L.Push(ud)
|
||||
}
|
||||
|
||||
// PushError pushes nil or the original Go error as userdata.
|
||||
func PushError(L *glua.LState, err error) {
|
||||
if err == nil {
|
||||
L.Push(glua.LNil)
|
||||
return
|
||||
}
|
||||
PushUserData(L, err)
|
||||
}
|
||||
|
||||
// ReadUserData reads a native Go value of type T without copying it.
|
||||
// Other Lua values or userdata containing a different type return invalidMessage.
|
||||
func ReadUserData[T any](value glua.LValue, invalidMessage string) (T, error) {
|
||||
if ud, ok := value.(*glua.LUserData); ok {
|
||||
if result, ok := ud.Value.(T); ok {
|
||||
return result, nil
|
||||
}
|
||||
}
|
||||
var zero T
|
||||
return zero, errors.New(invalidMessage)
|
||||
}
|
||||
|
||||
// ReadError accepts nil, a native Go error, or a Lua string.
|
||||
// Native errors retain their identity; other values return invalidMessage.
|
||||
func ReadError(value glua.LValue, invalidMessage string) error {
|
||||
if value == glua.LNil {
|
||||
return nil
|
||||
}
|
||||
if ud, ok := value.(*glua.LUserData); ok {
|
||||
if err, ok := ud.Value.(error); ok {
|
||||
return err
|
||||
}
|
||||
}
|
||||
if message, ok := value.(glua.LString); ok {
|
||||
return errors.New(string(message))
|
||||
}
|
||||
return errors.New(invalidMessage)
|
||||
}
|
||||
|
||||
// ReadUint32 accepts only integral Lua numbers in the uint32 range.
|
||||
func ReadUint32(value glua.LValue, invalidMessage string) (uint32, error) {
|
||||
number, ok := value.(glua.LNumber)
|
||||
if !ok || number < 0 || number > math.MaxUint32 || math.Trunc(float64(number)) != float64(number) {
|
||||
return 0, errors.New(invalidMessage)
|
||||
}
|
||||
return uint32(number), nil
|
||||
}
|
||||
|
||||
// ReadOptionalString accepts a Lua string or nil, which becomes an empty string.
|
||||
// It does not coerce other values to strings.
|
||||
func ReadOptionalString(value glua.LValue, invalidMessage string) (string, error) {
|
||||
if value == glua.LNil {
|
||||
return "", nil
|
||||
}
|
||||
if result, ok := value.(glua.LString); ok {
|
||||
return string(result), nil
|
||||
}
|
||||
return "", errors.New(invalidMessage)
|
||||
}
|
||||
@@ -0,0 +1,121 @@
|
||||
package lua
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"math"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
glua "github.com/yuin/gopher-lua"
|
||||
)
|
||||
|
||||
func TestReadUint32(t *testing.T) {
|
||||
for _, tc := range []struct {
|
||||
name string
|
||||
value glua.LValue
|
||||
want uint32
|
||||
wantErr bool
|
||||
}{
|
||||
{name: "zero", value: glua.LNumber(0)},
|
||||
{name: "integer", value: glua.LNumber(45), want: 45},
|
||||
{name: "maximum", value: glua.LNumber(math.MaxUint32), want: math.MaxUint32},
|
||||
{name: "fraction", value: glua.LNumber(1.5), wantErr: true},
|
||||
{name: "negative", value: glua.LNumber(-1), wantErr: true},
|
||||
{name: "overflow", value: glua.LNumber(math.MaxUint32 + 1), wantErr: true},
|
||||
{name: "NaN", value: glua.LNumber(math.NaN()), wantErr: true},
|
||||
{name: "positive infinity", value: glua.LNumber(math.Inf(1)), wantErr: true},
|
||||
{name: "negative infinity", value: glua.LNumber(math.Inf(-1)), wantErr: true},
|
||||
{name: "nil", value: glua.LNil, wantErr: true},
|
||||
{name: "numeric string", value: glua.LString("45"), wantErr: true},
|
||||
{name: "boolean", value: glua.LTrue, wantErr: true},
|
||||
} {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
got, err := ReadUint32(tc.value, "invalid number")
|
||||
if got != tc.want || (err != nil) != tc.wantErr {
|
||||
t.Fatalf("ReadUint32() = %d, %v; want %d, error %t", got, err, tc.want, tc.wantErr)
|
||||
}
|
||||
if err != nil && !strings.Contains(err.Error(), "invalid number") {
|
||||
t.Fatalf("error = %v, want invalid number", err)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestReadOptionalString(t *testing.T) {
|
||||
for _, tc := range []struct {
|
||||
name string
|
||||
value glua.LValue
|
||||
want string
|
||||
wantErr bool
|
||||
}{
|
||||
{name: "nil", value: glua.LNil},
|
||||
{name: "empty", value: glua.LString("")},
|
||||
{name: "string", value: glua.LString("out"), want: "out"},
|
||||
{name: "number", value: glua.LNumber(1), wantErr: true},
|
||||
{name: "boolean", value: glua.LFalse, wantErr: true},
|
||||
} {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
got, err := ReadOptionalString(tc.value, "invalid string")
|
||||
if got != tc.want || (err != nil) != tc.wantErr {
|
||||
t.Fatalf("ReadOptionalString() = %q, %v; want %q, error %t", got, err, tc.want, tc.wantErr)
|
||||
}
|
||||
if err != nil && !strings.Contains(err.Error(), "invalid string") {
|
||||
t.Fatalf("error = %v, want invalid string", err)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestUserDataRoundTrip(t *testing.T) {
|
||||
L := glua.NewState()
|
||||
defer L.Close()
|
||||
want := []int{1, 2}
|
||||
PushUserData(L, want)
|
||||
if L.GetTop() != 1 {
|
||||
t.Fatalf("stack top = %d, want 1", L.GetTop())
|
||||
}
|
||||
got, err := ReadUserData[[]int](L.Get(-1), "invalid userdata")
|
||||
if err != nil || len(got) != len(want) || &got[0] != &want[0] {
|
||||
t.Fatalf("userdata = %v, %v; want original slice", got, err)
|
||||
}
|
||||
PushUserData(L, []int(nil))
|
||||
if got, err := ReadUserData[[]int](L.Get(-1), "invalid userdata"); err != nil || got != nil {
|
||||
t.Fatalf("nil slice userdata = %v, %v", got, err)
|
||||
}
|
||||
for _, value := range []glua.LValue{glua.LNil, glua.LString("1"), L.NewTable(), L.Get(1)} {
|
||||
if got, err := ReadUserData[int](value, "invalid userdata"); got != 0 || err == nil || !strings.Contains(err.Error(), "invalid userdata") {
|
||||
t.Fatalf("ReadUserData(%v) = %d, %v; want invalid userdata", value, got, err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestErrorRoundTrip(t *testing.T) {
|
||||
L := glua.NewState()
|
||||
defer L.Close()
|
||||
want := errors.New("upstream failed")
|
||||
for _, err := range []error{nil, want} {
|
||||
PushError(L, err)
|
||||
if L.GetTop() != 1 {
|
||||
t.Fatalf("stack top = %d, want 1", L.GetTop())
|
||||
}
|
||||
if err == nil && L.Get(-1) != glua.LNil {
|
||||
t.Fatalf("nil error pushed as %v", L.Get(-1))
|
||||
}
|
||||
if got := ReadError(L.Get(-1), "invalid error"); got != err {
|
||||
t.Fatalf("ReadError() = %v, want original error %v", got, err)
|
||||
}
|
||||
L.Pop(1)
|
||||
}
|
||||
for _, message := range []string{"script failed", ""} {
|
||||
if err := ReadError(glua.LString(message), "invalid error"); err == nil || !strings.Contains(err.Error(), message) {
|
||||
t.Fatalf("string error = %v, want %q", err, message)
|
||||
}
|
||||
}
|
||||
wrong := L.NewUserData()
|
||||
wrong.Value = "not a native error"
|
||||
for _, value := range []glua.LValue{glua.LTrue, glua.LNumber(1), L.NewTable(), wrong, L.NewUserData()} {
|
||||
if err := ReadError(value, "invalid error"); err == nil || !strings.Contains(err.Error(), "invalid error") {
|
||||
t.Fatalf("ReadError(%v) = %v, want invalid error", value, err)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -23,6 +23,7 @@ require (
|
||||
github.com/stretchr/testify v1.12.1
|
||||
github.com/vishvananda/netlink v1.3.1
|
||||
github.com/xtls/reality v0.0.0-20260908062103-8cdf7bf9c7f0
|
||||
github.com/yuin/gopher-lua v1.1.2
|
||||
go4.org/netipx v0.0.0-20231129151722-fdeea329fbba
|
||||
golang.org/x/crypto v0.57.0
|
||||
golang.org/x/exp v0.0.0-20240506185415-9bf2ced13842
|
||||
@@ -31,11 +32,12 @@ require (
|
||||
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
|
||||
)
|
||||
@@ -60,6 +62,6 @@ require (
|
||||
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
|
||||
)
|
||||
|
||||
@@ -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=
|
||||
@@ -127,6 +115,7 @@ golang.org/x/sync v0.0.0-20190423024810-112230192c58/go.mod h1:RxMgew5VJxzue5/jJ
|
||||
golang.org/x/sync v0.0.0-20210220032951-036812b2e83c/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
||||
golang.org/x/sync v0.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=
|
||||
@@ -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,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())
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -36,6 +36,8 @@ 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:
|
||||
@@ -61,6 +63,7 @@ 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"`
|
||||
}
|
||||
@@ -195,6 +198,16 @@ 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 {
|
||||
|
||||
@@ -20,10 +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"
|
||||
)
|
||||
|
||||
@@ -786,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)
|
||||
|
||||
@@ -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 &&
|
||||
|
||||
@@ -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,6 +55,7 @@ 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"
|
||||
|
||||
@@ -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))
|
||||
}))
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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;
|
||||
}
|
||||
@@ -199,9 +199,9 @@ func (h *Handler) Process(ctx context.Context, link *transport.Link, dialer inte
|
||||
return errors.New("failed to create UDP connection").Base(err)
|
||||
}
|
||||
defer conn.Close()
|
||||
c := &udpConnClient{
|
||||
c := &UDPConnClient{
|
||||
PacketConn: conn.(*internet.PacketConnWrapper).PacketConn,
|
||||
dest: conn.RemoteAddr().(*net.UDPAddr),
|
||||
Dest: conn.RemoteAddr().(*net.UDPAddr),
|
||||
}
|
||||
reader = c
|
||||
writer = c
|
||||
@@ -336,6 +336,9 @@ func (h *Handler) init(ctx context.Context) error {
|
||||
}
|
||||
|
||||
func (h *Handler) resolveLocal(host string) (net.IP, error) {
|
||||
if ip := net.ParseIP(host); ip != nil {
|
||||
return ip, nil
|
||||
}
|
||||
ips, _, err := h.dns.LookupIP(host, dns.IPOption{IPv4Enable: true, IPv6Enable: true})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
@@ -375,12 +378,12 @@ func (h *Handler) resolveLocal(host string) (net.IP, error) {
|
||||
return got[dice.Roll(len(got))], nil
|
||||
}
|
||||
|
||||
type udpConnClient struct {
|
||||
type UDPConnClient struct {
|
||||
net.PacketConn
|
||||
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())
|
||||
@@ -399,9 +402,9 @@ 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() {
|
||||
if b.UDP.Port != net.Port(dst.Port) {
|
||||
@@ -459,20 +462,27 @@ func (c *cache) run() {
|
||||
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)
|
||||
for {
|
||||
now := <-ticker.C
|
||||
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()
|
||||
}
|
||||
}
|
||||
|
||||
@@ -136,6 +136,7 @@ func (tun *netTun) Read(buf [][]byte, sizes []int, offset int) (int, error) {
|
||||
}
|
||||
|
||||
n, err := view.Read(buf[0][offset:])
|
||||
view.Release()
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
@@ -175,6 +176,7 @@ func (tun *netTun) WriteNotify() {
|
||||
select {
|
||||
case tun.incomingPacket <- view:
|
||||
case <-tun.closed:
|
||||
view.Release()
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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,
|
||||
}
|
||||
}))
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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;
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
@@ -0,0 +1,75 @@
|
||||
/* SPDX-License-Identifier: MIT
|
||||
*
|
||||
* Copyright 2024 Marten Seemann
|
||||
* Adapted from github.com/quic-go/connect-ip-go (commit a0c35fa).
|
||||
*/
|
||||
|
||||
package connectip
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"net/http"
|
||||
|
||||
"github.com/apernet/quic-go"
|
||||
"github.com/apernet/quic-go/http3"
|
||||
)
|
||||
|
||||
type ClientConn struct {
|
||||
clientConn *http3.ClientConn
|
||||
}
|
||||
|
||||
func NewClientConn(conn *http3.ClientConn) *ClientConn {
|
||||
return &ClientConn{clientConn: conn}
|
||||
}
|
||||
|
||||
func (c *ClientConn) Dial(req *Request) (*Conn, *http.Response, error) {
|
||||
httpReq := req.httpRequest()
|
||||
if httpReq.URL == nil {
|
||||
return nil, nil, errors.New("connect-ip: request URL is nil")
|
||||
}
|
||||
if httpReq.Host == "" && httpReq.URL.Host == "" {
|
||||
return nil, nil, errors.New("connect-ip: request needs a host")
|
||||
}
|
||||
|
||||
select {
|
||||
case <-httpReq.Context().Done():
|
||||
return nil, nil, context.Cause(httpReq.Context())
|
||||
case <-c.clientConn.Context().Done():
|
||||
return nil, nil, context.Cause(c.clientConn.Context())
|
||||
case <-c.clientConn.ReceivedSettings():
|
||||
}
|
||||
|
||||
settings := c.clientConn.Settings()
|
||||
if !settings.EnableExtendedConnect {
|
||||
return nil, nil, errors.New("connect-ip: server didn't enable Extended CONNECT")
|
||||
}
|
||||
if !settings.EnableDatagrams {
|
||||
return nil, nil, errors.New("connect-ip: server didn't enable datagrams")
|
||||
}
|
||||
|
||||
rstr, err := c.clientConn.OpenRequestStream(httpReq.Context())
|
||||
if err != nil {
|
||||
return nil, nil, fmt.Errorf("connect-ip: failed to open request stream: %w", err)
|
||||
}
|
||||
var keepStream bool
|
||||
defer func() {
|
||||
if !keepStream {
|
||||
rstr.CancelRead(quic.StreamErrorCode(http3.ErrCodeNoError))
|
||||
rstr.CancelWrite(quic.StreamErrorCode(http3.ErrCodeNoError))
|
||||
}
|
||||
}()
|
||||
if err := rstr.SendRequestHeader(httpReq); err != nil {
|
||||
return nil, nil, fmt.Errorf("connect-ip: failed to send request: %w", err)
|
||||
}
|
||||
rsp, err := rstr.ReadResponse()
|
||||
if err != nil {
|
||||
return nil, nil, fmt.Errorf("connect-ip: failed to read response: %w", err)
|
||||
}
|
||||
if rsp.StatusCode < 200 || rsp.StatusCode > 299 {
|
||||
return nil, rsp, fmt.Errorf("connect-ip: server responded with %d", rsp.StatusCode)
|
||||
}
|
||||
keepStream = true
|
||||
return newProxiedConn(rstr), rsp, nil
|
||||
}
|
||||
@@ -0,0 +1,101 @@
|
||||
/* SPDX-License-Identifier: MIT
|
||||
*
|
||||
* Copyright 2024 Marten Seemann
|
||||
* Adapted from github.com/quic-go/connect-ip-go (commit a0c35fa).
|
||||
*/
|
||||
|
||||
package connectip
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net"
|
||||
"net/http"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/apernet/quic-go"
|
||||
"github.com/apernet/quic-go/http3"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestClientWaitForSettings(t *testing.T) {
|
||||
conn, err := net.ListenUDP("udp", &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1), Port: 0})
|
||||
require.NoError(t, err)
|
||||
ln, err := quic.Listen(conn, tlsConf, &quic.Config{EnableDatagrams: true})
|
||||
require.NoError(t, err)
|
||||
defer ln.Close()
|
||||
|
||||
h3conn := dialHTTP3(t, conn.LocalAddr().String())
|
||||
ctx, cancel := context.WithTimeout(t.Context(), 100*time.Millisecond)
|
||||
defer cancel()
|
||||
req, err := NewRequest(ctx, "https://example.org/.well-known/masque/ip/")
|
||||
require.NoError(t, err)
|
||||
_, _, err = NewClientConn(h3conn).Dial(req)
|
||||
require.ErrorIs(t, err, context.DeadlineExceeded)
|
||||
}
|
||||
|
||||
func TestClientDatagramCheck(t *testing.T) {
|
||||
s := http3.Server{
|
||||
TLSConfig: tlsConf,
|
||||
QUICConfig: &quic.Config{EnableDatagrams: true},
|
||||
EnableDatagrams: false,
|
||||
}
|
||||
ln, err := net.ListenUDP("udp", &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1), Port: 0})
|
||||
require.NoError(t, err)
|
||||
go func() { s.Serve(ln) }()
|
||||
defer s.Close()
|
||||
|
||||
h3conn := dialHTTP3(t, ln.LocalAddr().String())
|
||||
ctx, cancel := context.WithTimeout(t.Context(), 5*time.Second)
|
||||
defer cancel()
|
||||
req, err := NewRequest(ctx, "https://example.org/.well-known/masque/ip/")
|
||||
require.NoError(t, err)
|
||||
_, _, err = NewClientConn(h3conn).Dial(req)
|
||||
require.ErrorContains(t, err, "connect-ip: server didn't enable datagrams")
|
||||
}
|
||||
|
||||
func TestNewClientConnSharesHTTP3Connection(t *testing.T) {
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
|
||||
defer cancel()
|
||||
|
||||
ln, err := net.ListenUDP("udp4", &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1)})
|
||||
require.NoError(t, err)
|
||||
defer ln.Close()
|
||||
url := "https://" + ln.LocalAddr().String()
|
||||
|
||||
mux := http.NewServeMux()
|
||||
mux.HandleFunc("/connect-ip", func(w http.ResponseWriter, r *http.Request) {
|
||||
req, err := ParseProxyRequest(r)
|
||||
if !assert.NoError(t, err) {
|
||||
w.WriteHeader(http.StatusBadRequest)
|
||||
return
|
||||
}
|
||||
_, err = (&Proxy{}).Proxy(w, req)
|
||||
assert.NoError(t, err)
|
||||
})
|
||||
mux.HandleFunc("GET /hello", func(http.ResponseWriter, *http.Request) {})
|
||||
s := http3.Server{Handler: mux, TLSConfig: tlsConf, EnableDatagrams: true}
|
||||
go func() { s.Serve(ln) }()
|
||||
defer s.Close()
|
||||
|
||||
h3conn := dialHTTP3(t, ln.LocalAddr().String())
|
||||
httpClient := &http.Client{Transport: h3conn, Timeout: time.Second}
|
||||
checkHTTP := func() {
|
||||
t.Helper()
|
||||
rsp, err := httpClient.Get(url + "/hello")
|
||||
require.NoError(t, err)
|
||||
rsp.Body.Close()
|
||||
require.Equal(t, http.StatusOK, rsp.StatusCode)
|
||||
}
|
||||
|
||||
checkHTTP()
|
||||
req, err := NewRequest(ctx, url+"/connect-ip")
|
||||
require.NoError(t, err)
|
||||
tunnel, rsp, err := NewClientConn(h3conn).Dial(req)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, http.StatusOK, rsp.StatusCode)
|
||||
checkHTTP()
|
||||
require.NoError(t, tunnel.Close())
|
||||
checkHTTP()
|
||||
}
|
||||
@@ -0,0 +1,620 @@
|
||||
/* SPDX-License-Identifier: MIT
|
||||
*
|
||||
* Copyright 2024 Marten Seemann
|
||||
* Adapted from github.com/quic-go/connect-ip-go (commit a0c35fa).
|
||||
*/
|
||||
|
||||
package connectip
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/binary"
|
||||
goerrors "errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"net"
|
||||
"net/netip"
|
||||
"slices"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/apernet/quic-go"
|
||||
"github.com/apernet/quic-go/http3"
|
||||
"github.com/apernet/quic-go/quicvarint"
|
||||
"github.com/xtls/xray-core/common/errors"
|
||||
"golang.org/x/net/ipv4"
|
||||
"golang.org/x/net/ipv6"
|
||||
)
|
||||
|
||||
type CloseError struct {
|
||||
Remote bool
|
||||
}
|
||||
|
||||
func (e *CloseError) Error() string { return net.ErrClosed.Error() }
|
||||
func (e *CloseError) Is(target error) bool { return target == net.ErrClosed }
|
||||
|
||||
const (
|
||||
ipProtoICMP = 1
|
||||
ipProtoICMPv6 = 58
|
||||
)
|
||||
|
||||
type http3Stream interface {
|
||||
io.ReadWriteCloser
|
||||
StreamID() quic.StreamID
|
||||
ReceiveDatagram(context.Context) ([]byte, error)
|
||||
SendDatagram([]byte) error
|
||||
CancelRead(quic.StreamErrorCode)
|
||||
CancelWrite(quic.StreamErrorCode)
|
||||
SetWriteDeadline(time.Time) error
|
||||
}
|
||||
|
||||
var (
|
||||
_ http3Stream = &http3.Stream{}
|
||||
_ http3Stream = &http3.RequestStream{}
|
||||
)
|
||||
|
||||
const maxQueuedCapsules = 128
|
||||
|
||||
var errCapsuleLimit = goerrors.New("connect-ip: capsule limit exceeded")
|
||||
|
||||
type streamWrite struct {
|
||||
Data []byte
|
||||
Fin bool
|
||||
}
|
||||
|
||||
type Conn struct {
|
||||
str http3Stream
|
||||
writeNotify chan struct{}
|
||||
writeDone chan error
|
||||
|
||||
assignedAddressUpdates chan []AssignedAddress
|
||||
addressRequests chan *addressRequestCapsule
|
||||
availableRouteUpdates chan []IPRoute
|
||||
|
||||
mu sync.Mutex
|
||||
queuedWrites []streamWrite
|
||||
peerAddresses []netip.Prefix
|
||||
localRoutes []IPRoute
|
||||
assignedAddresses []netip.Prefix
|
||||
lastAddressRequestID AddressRequestID
|
||||
|
||||
closeChan chan struct{}
|
||||
closeErr error
|
||||
|
||||
closeOnce sync.Once
|
||||
closeResult error
|
||||
|
||||
datagramCapsuleOnce sync.Once
|
||||
}
|
||||
|
||||
func newProxiedConn(str http3Stream) *Conn {
|
||||
c := &Conn{
|
||||
str: str,
|
||||
writeNotify: make(chan struct{}, 1),
|
||||
writeDone: make(chan error, 1),
|
||||
assignedAddressUpdates: make(chan []AssignedAddress, maxQueuedCapsules),
|
||||
addressRequests: make(chan *addressRequestCapsule, maxQueuedCapsules),
|
||||
availableRouteUpdates: make(chan []IPRoute, 1),
|
||||
closeChan: make(chan struct{}),
|
||||
}
|
||||
go func() {
|
||||
err := c.readFromStream()
|
||||
c.mu.Lock()
|
||||
closing := c.closeErr != nil
|
||||
if !closing {
|
||||
c.closeErr = &CloseError{Remote: true}
|
||||
close(c.closeChan)
|
||||
if err != nil {
|
||||
code := http3.ErrCodeMessageError
|
||||
var streamErr *quic.StreamError
|
||||
var h3Err *http3.Error
|
||||
switch {
|
||||
case goerrors.Is(err, errCapsuleLimit):
|
||||
code = http3.ErrCodeExcessiveLoad
|
||||
case goerrors.As(err, &streamErr) && streamErr.Remote, goerrors.As(err, &h3Err) && h3Err.Remote:
|
||||
code = http3.ErrCodeRequestCanceled
|
||||
}
|
||||
c.str.CancelRead(quic.StreamErrorCode(code))
|
||||
c.str.CancelWrite(quic.StreamErrorCode(code))
|
||||
close(c.writeNotify)
|
||||
} else {
|
||||
c.queueFin()
|
||||
}
|
||||
}
|
||||
c.mu.Unlock()
|
||||
if err != nil && !closing {
|
||||
errors.LogInfoInner(context.Background(), err, "reading capsules failed")
|
||||
}
|
||||
}()
|
||||
go func() {
|
||||
err := c.writeToStream()
|
||||
if err != nil {
|
||||
c.mu.Lock()
|
||||
closing := c.closeErr != nil
|
||||
if !closing {
|
||||
c.closeErr = &CloseError{Remote: true}
|
||||
close(c.closeChan)
|
||||
c.str.CancelRead(quic.StreamErrorCode(http3.ErrCodeExcessiveLoad))
|
||||
c.str.CancelWrite(quic.StreamErrorCode(http3.ErrCodeExcessiveLoad))
|
||||
} else {
|
||||
c.str.CancelWrite(quic.StreamErrorCode(http3.ErrCodeNoError))
|
||||
}
|
||||
c.mu.Unlock()
|
||||
if !closing {
|
||||
errors.LogInfoInner(context.Background(), err, "writing capsules failed")
|
||||
}
|
||||
}
|
||||
c.writeDone <- err
|
||||
close(c.writeDone)
|
||||
}()
|
||||
return c
|
||||
}
|
||||
|
||||
func (c *Conn) AdvertiseRoute(routes []IPRoute) error {
|
||||
for i, route := range routes {
|
||||
err := route.validate()
|
||||
if err == nil && i > 0 {
|
||||
err = checkRouteOrder(routes[i-1], route)
|
||||
}
|
||||
if err != nil {
|
||||
return fmt.Errorf("connect-ip: invalid route %d: %w", i, err)
|
||||
}
|
||||
}
|
||||
|
||||
c.mu.Lock()
|
||||
if c.closeErr != nil {
|
||||
err := c.closeErr
|
||||
c.mu.Unlock()
|
||||
return err
|
||||
}
|
||||
routes = slices.Clone(routes)
|
||||
err := c.queueWrite(streamWrite{Data: (&routeAdvertisementCapsule{IPAddressRanges: routes}).append(nil)})
|
||||
if err == nil {
|
||||
c.localRoutes = routes
|
||||
}
|
||||
c.mu.Unlock()
|
||||
if err != nil {
|
||||
c.Close()
|
||||
return err
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (c *Conn) RequestAddresses(prefixes []netip.Prefix) ([]AddressRequestID, error) {
|
||||
if len(prefixes) == 0 {
|
||||
return nil, goerrors.New("connect-ip: address request must contain at least one prefix")
|
||||
}
|
||||
for i, p := range prefixes {
|
||||
if !p.IsValid() || p != p.Masked() {
|
||||
return nil, fmt.Errorf("connect-ip: invalid requested prefix %d: %s", i, p)
|
||||
}
|
||||
}
|
||||
|
||||
c.mu.Lock()
|
||||
if c.closeErr != nil {
|
||||
err := c.closeErr
|
||||
c.mu.Unlock()
|
||||
return nil, err
|
||||
}
|
||||
ids := make([]AddressRequestID, len(prefixes))
|
||||
for i := range ids {
|
||||
ids[i] = c.lastAddressRequestID + AddressRequestID(i) + 1
|
||||
}
|
||||
capsule := &addressRequestCapsule{RequestIDs: ids, Prefixes: prefixes}
|
||||
err := c.queueWrite(streamWrite{Data: capsule.append(nil)})
|
||||
if err == nil {
|
||||
c.lastAddressRequestID = ids[len(ids)-1]
|
||||
}
|
||||
c.mu.Unlock()
|
||||
if err != nil {
|
||||
c.Close()
|
||||
return nil, err
|
||||
}
|
||||
return ids, nil
|
||||
}
|
||||
|
||||
func (c *Conn) ReceiveAddressAssignment(ctx context.Context) ([]AssignedAddress, error) {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return nil, ctx.Err()
|
||||
case assignment := <-c.assignedAddressUpdates:
|
||||
return assignment, nil
|
||||
case <-c.closeChan:
|
||||
select {
|
||||
case assignment := <-c.assignedAddressUpdates:
|
||||
return assignment, nil
|
||||
default:
|
||||
return nil, c.closeErr
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (c *Conn) ReceiveAddressRequest(ctx context.Context) (*AddressRequest, error) {
|
||||
var requested *addressRequestCapsule
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return nil, ctx.Err()
|
||||
case requested = <-c.addressRequests:
|
||||
case <-c.closeChan:
|
||||
select {
|
||||
case requested = <-c.addressRequests:
|
||||
default:
|
||||
return nil, c.closeErr
|
||||
}
|
||||
}
|
||||
return newAddressRequest(c, requested), nil
|
||||
}
|
||||
|
||||
func (c *Conn) AssignAddresses(prefixes []netip.Prefix) error {
|
||||
capsule := &addressAssignCapsule{}
|
||||
if prefixes != nil {
|
||||
capsule.AssignedAddresses = make([]AssignedAddress, len(prefixes))
|
||||
for i, p := range prefixes {
|
||||
capsule.AssignedAddresses[i] = AssignedAddress{IPPrefix: p}
|
||||
}
|
||||
}
|
||||
return c.sendAddressAssignment(capsule, true)
|
||||
}
|
||||
|
||||
func (c *Conn) sendAddressAssignment(capsule *addressAssignCapsule, restrictPeer bool) error {
|
||||
c.mu.Lock()
|
||||
if c.closeErr != nil {
|
||||
err := c.closeErr
|
||||
c.mu.Unlock()
|
||||
return err
|
||||
}
|
||||
if err := c.queueWrite(streamWrite{Data: capsule.append(nil)}); err != nil {
|
||||
c.mu.Unlock()
|
||||
c.Close()
|
||||
return err
|
||||
}
|
||||
|
||||
if !restrictPeer && c.peerAddresses == nil {
|
||||
c.mu.Unlock()
|
||||
return nil
|
||||
}
|
||||
var prefixes []netip.Prefix
|
||||
if capsule.AssignedAddresses != nil {
|
||||
prefixes = make([]netip.Prefix, 0, len(capsule.AssignedAddresses))
|
||||
}
|
||||
for _, assigned := range capsule.AssignedAddresses {
|
||||
if !assigned.Rejected() {
|
||||
prefixes = append(prefixes, assigned.IPPrefix)
|
||||
}
|
||||
}
|
||||
c.peerAddresses = prefixes
|
||||
c.mu.Unlock()
|
||||
return nil
|
||||
}
|
||||
|
||||
func (c *Conn) queueWrite(w streamWrite) error {
|
||||
if len(c.queuedWrites) >= maxQueuedCapsules {
|
||||
c.closeErr = &CloseError{Remote: false}
|
||||
close(c.closeChan)
|
||||
c.str.CancelRead(quic.StreamErrorCode(http3.ErrCodeExcessiveLoad))
|
||||
c.str.CancelWrite(quic.StreamErrorCode(http3.ErrCodeExcessiveLoad))
|
||||
close(c.writeNotify)
|
||||
return goerrors.New("connect-ip: capsule queue full")
|
||||
}
|
||||
c.queuedWrites = append(c.queuedWrites, w)
|
||||
c.notifyWriter()
|
||||
return nil
|
||||
}
|
||||
|
||||
func (c *Conn) queueFin() {
|
||||
c.str.SetWriteDeadline(time.Now())
|
||||
c.queuedWrites = append(c.queuedWrites, streamWrite{Fin: true})
|
||||
c.notifyWriter()
|
||||
}
|
||||
|
||||
func (c *Conn) notifyWriter() {
|
||||
select {
|
||||
case c.writeNotify <- struct{}{}:
|
||||
default:
|
||||
}
|
||||
}
|
||||
|
||||
func queueLatest[T any](ch chan T, value T) {
|
||||
for {
|
||||
select {
|
||||
case ch <- value:
|
||||
return
|
||||
case <-ch:
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (c *Conn) Routes(ctx context.Context) ([]IPRoute, error) {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return nil, ctx.Err()
|
||||
case <-c.closeChan:
|
||||
return nil, c.closeErr
|
||||
case routes := <-c.availableRouteUpdates:
|
||||
return routes, nil
|
||||
}
|
||||
}
|
||||
|
||||
func (c *Conn) readFromStream() error {
|
||||
p := http3.NewCapsuleParser(c.str)
|
||||
for {
|
||||
t, cr, err := p.Next()
|
||||
if goerrors.Is(err, io.EOF) {
|
||||
return nil
|
||||
}
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
switch t {
|
||||
case capsuleTypeAddressAssign:
|
||||
capsule, err := parseAddressAssignCapsule(cr)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
prefixes := make([]netip.Prefix, 0, len(capsule.AssignedAddresses))
|
||||
for _, assigned := range capsule.AssignedAddresses {
|
||||
if !assigned.Rejected() {
|
||||
prefixes = append(prefixes, assigned.IPPrefix)
|
||||
}
|
||||
}
|
||||
c.mu.Lock()
|
||||
c.assignedAddresses = prefixes
|
||||
c.mu.Unlock()
|
||||
select {
|
||||
case c.assignedAddressUpdates <- capsule.AssignedAddresses:
|
||||
default:
|
||||
return fmt.Errorf("%w: address assignment queue full", errCapsuleLimit)
|
||||
}
|
||||
case capsuleTypeAddressRequest:
|
||||
capsule, err := parseAddressRequestCapsule(cr)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
select {
|
||||
case c.addressRequests <- capsule:
|
||||
default:
|
||||
return fmt.Errorf("%w: address request queue full", errCapsuleLimit)
|
||||
}
|
||||
case capsuleTypeRouteAdvertisement:
|
||||
capsule, err := parseRouteAdvertisementCapsule(cr)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
queueLatest(c.availableRouteUpdates, capsule.IPAddressRanges)
|
||||
case capsuleTypeDatagram:
|
||||
c.datagramCapsuleOnce.Do(func() {
|
||||
errors.LogWarning(context.Background(), "connect-ip: dropping IP packets sent in DATAGRAM capsules, only QUIC DATAGRAM frames are supported")
|
||||
})
|
||||
if err := cr.Discard(); err != nil {
|
||||
return err
|
||||
}
|
||||
default:
|
||||
if err := cr.Discard(); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (c *Conn) writeToStream() error {
|
||||
for range c.writeNotify {
|
||||
for {
|
||||
c.mu.Lock()
|
||||
if len(c.queuedWrites) == 0 {
|
||||
c.mu.Unlock()
|
||||
break
|
||||
}
|
||||
w := c.queuedWrites[0]
|
||||
c.queuedWrites[0] = streamWrite{}
|
||||
c.queuedWrites = c.queuedWrites[1:]
|
||||
c.mu.Unlock()
|
||||
|
||||
if w.Fin {
|
||||
return c.str.Close()
|
||||
}
|
||||
if _, err := c.str.Write(w.Data); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
}
|
||||
return c.closeErr
|
||||
}
|
||||
|
||||
func (c *Conn) ReadPacket(b []byte) (int, error) {
|
||||
for {
|
||||
select {
|
||||
case <-c.closeChan:
|
||||
return 0, c.closeErr
|
||||
default:
|
||||
}
|
||||
data, err := c.str.ReceiveDatagram(context.Background())
|
||||
if err != nil {
|
||||
select {
|
||||
case <-c.closeChan:
|
||||
return 0, c.closeErr
|
||||
default:
|
||||
return 0, err
|
||||
}
|
||||
}
|
||||
contextID, n, err := quicvarint.Parse(data)
|
||||
if err != nil {
|
||||
errors.LogDebugInner(context.Background(), err, "dropping malformed datagram")
|
||||
continue
|
||||
}
|
||||
if contextID != 0 {
|
||||
continue
|
||||
}
|
||||
packet := data[n:]
|
||||
if err := c.handleIncomingProxiedPacket(packet); err != nil {
|
||||
errors.LogDebugInner(context.Background(), err, "dropping proxied packet")
|
||||
continue
|
||||
}
|
||||
if len(packet) > len(b) {
|
||||
return 0, io.ErrShortBuffer
|
||||
}
|
||||
return copy(b, packet), nil
|
||||
}
|
||||
}
|
||||
|
||||
func (c *Conn) handleIncomingProxiedPacket(data []byte) error {
|
||||
if len(data) == 0 {
|
||||
return goerrors.New("connect-ip: empty packet")
|
||||
}
|
||||
var src, dst netip.Addr
|
||||
var ipProto uint8
|
||||
switch v := ipVersion(data); v {
|
||||
default:
|
||||
return fmt.Errorf("connect-ip: unknown IP versions: %d", v)
|
||||
case 4:
|
||||
if len(data) < ipv4.HeaderLen {
|
||||
return fmt.Errorf("connect-ip: malformed datagram: too short")
|
||||
}
|
||||
src = netip.AddrFrom4([4]byte(data[12:16]))
|
||||
dst = netip.AddrFrom4([4]byte(data[16:20]))
|
||||
ipProto = data[9]
|
||||
case 6:
|
||||
if len(data) < ipv6.HeaderLen {
|
||||
return fmt.Errorf("connect-ip: malformed datagram: too short")
|
||||
}
|
||||
src = netip.AddrFrom16([16]byte(data[8:24]))
|
||||
dst = netip.AddrFrom16([16]byte(data[24:40]))
|
||||
ipProto = data[6]
|
||||
}
|
||||
|
||||
c.mu.Lock()
|
||||
assignedAddresses := c.assignedAddresses
|
||||
localRoutes := c.localRoutes
|
||||
peerAddresses := c.peerAddresses
|
||||
c.mu.Unlock()
|
||||
|
||||
if peerAddresses != nil {
|
||||
if !slices.ContainsFunc(peerAddresses, func(p netip.Prefix) bool { return p.Contains(src) }) {
|
||||
return fmt.Errorf("connect-ip: datagram source address not allowed: %s", src)
|
||||
}
|
||||
}
|
||||
|
||||
var isAllowedDst bool
|
||||
if len(assignedAddresses) > 0 {
|
||||
isAllowedDst = slices.ContainsFunc(assignedAddresses, func(p netip.Prefix) bool { return p.Contains(dst) })
|
||||
}
|
||||
if !isAllowedDst {
|
||||
isAllowedDst = slices.ContainsFunc(localRoutes, func(r IPRoute) bool {
|
||||
if r.StartIP.Compare(dst) > 0 || dst.Compare(r.EndIP) > 0 {
|
||||
return false
|
||||
}
|
||||
if (ipVersion(data) == 4 && ipProto == ipProtoICMP) || (ipVersion(data) == 6 && ipProto == ipProtoICMPv6) {
|
||||
return true
|
||||
}
|
||||
return r.IPProtocol == 0 || r.IPProtocol == ipProto
|
||||
})
|
||||
}
|
||||
if !isAllowedDst {
|
||||
return fmt.Errorf("connect-ip: datagram destination address / protocol not allowed: %s (protocol: %d)", dst, ipProto)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (c *Conn) WritePacket(b []byte) (icmp []byte, err error) {
|
||||
select {
|
||||
case <-c.closeChan:
|
||||
return nil, c.closeErr
|
||||
default:
|
||||
}
|
||||
data, err := c.composeDatagram(b)
|
||||
if err != nil {
|
||||
errors.LogDebugInner(context.Background(), err, "dropping proxied packet (", len(b), " bytes) that can't be proxied")
|
||||
return nil, nil
|
||||
}
|
||||
if err := c.str.SendDatagram(data); err != nil {
|
||||
if tooLarge, ok := goerrors.AsType[*quic.DatagramTooLargeError](err); ok {
|
||||
icmpPacket, err := composeICMPTooLargePacket(b, int(tooLarge.MaxDatagramPayloadSize)-c.datagramOverhead())
|
||||
if err != nil {
|
||||
if goerrors.Is(err, ErrMTUTooSmall) {
|
||||
return nil, err
|
||||
}
|
||||
errors.LogDebugInner(context.Background(), err, "failed to compose ICMP Packet Too Big")
|
||||
}
|
||||
return icmpPacket, nil
|
||||
}
|
||||
select {
|
||||
case <-c.closeChan:
|
||||
return nil, c.closeErr
|
||||
default:
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
func (c *Conn) composeDatagram(b []byte) ([]byte, error) {
|
||||
if len(b) == 0 {
|
||||
return nil, goerrors.New("connect-ip: empty packet")
|
||||
}
|
||||
switch v := ipVersion(b); v {
|
||||
default:
|
||||
return nil, fmt.Errorf("connect-ip: unknown IP versions: %d", v)
|
||||
case 4:
|
||||
if len(b) < ipv4.HeaderLen {
|
||||
return nil, fmt.Errorf("connect-ip: IPv4 packet too short")
|
||||
}
|
||||
hdrLen := int(b[0]&0x0f) << 2
|
||||
totalLen := int(binary.BigEndian.Uint16(b[2:4]))
|
||||
if hdrLen < ipv4.HeaderLen || hdrLen > totalLen || totalLen > len(b) {
|
||||
return nil, fmt.Errorf("connect-ip: malformed IPv4 header: header length %d, total length %d, packet length %d", hdrLen, totalLen, len(b))
|
||||
}
|
||||
ttl := b[8]
|
||||
if ttl <= 1 {
|
||||
return nil, fmt.Errorf("connect-ip: datagram TTL too small: %d", ttl)
|
||||
}
|
||||
b[8]--
|
||||
binary.BigEndian.PutUint16(b[10:12], calculateIPv4Checksum(b[:hdrLen]))
|
||||
case 6:
|
||||
if len(b) < ipv6.HeaderLen {
|
||||
return nil, fmt.Errorf("connect-ip: IPv6 packet too short")
|
||||
}
|
||||
hopLimit := b[7]
|
||||
if hopLimit <= 1 {
|
||||
return nil, fmt.Errorf("connect-ip: datagram Hop Limit too small: %d", hopLimit)
|
||||
}
|
||||
b[7]--
|
||||
}
|
||||
data := make([]byte, 0, len(contextIDZero)+len(b))
|
||||
data = append(data, contextIDZero...)
|
||||
data = append(data, b...)
|
||||
return data, nil
|
||||
}
|
||||
|
||||
func (c *Conn) datagramOverhead() int {
|
||||
return quicvarint.Len(uint64(c.str.StreamID()/4)) + len(contextIDZero)
|
||||
}
|
||||
|
||||
func (c *Conn) MaxPacketSize() int {
|
||||
select {
|
||||
case <-c.closeChan:
|
||||
return 0
|
||||
default:
|
||||
}
|
||||
err := c.str.SendDatagram(make([]byte, 1<<16))
|
||||
tooLarge, ok := goerrors.AsType[*quic.DatagramTooLargeError](err)
|
||||
if !ok {
|
||||
return 0
|
||||
}
|
||||
return max(0, int(tooLarge.MaxDatagramPayloadSize)-c.datagramOverhead())
|
||||
}
|
||||
|
||||
func (c *Conn) Close() error {
|
||||
c.closeOnce.Do(func() {
|
||||
c.mu.Lock()
|
||||
if c.closeErr == nil {
|
||||
c.closeErr = &CloseError{Remote: false}
|
||||
close(c.closeChan)
|
||||
c.queueFin()
|
||||
}
|
||||
c.mu.Unlock()
|
||||
c.closeResult = <-c.writeDone
|
||||
c.str.CancelRead(quic.StreamErrorCode(http3.ErrCodeNoError))
|
||||
})
|
||||
return c.closeResult
|
||||
}
|
||||
|
||||
func ipVersion(b []byte) uint8 { return b[0] >> 4 }
|
||||
@@ -0,0 +1,706 @@
|
||||
/* SPDX-License-Identifier: MIT
|
||||
*
|
||||
* Copyright 2024 Marten Seemann
|
||||
* Adapted from github.com/quic-go/connect-ip-go (commit a0c35fa).
|
||||
*/
|
||||
|
||||
package connectip
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/binary"
|
||||
"errors"
|
||||
"io"
|
||||
"net"
|
||||
"net/netip"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/apernet/quic-go"
|
||||
"github.com/apernet/quic-go/http3"
|
||||
"github.com/apernet/quic-go/quicvarint"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"golang.org/x/net/icmp"
|
||||
"golang.org/x/net/ipv4"
|
||||
"golang.org/x/net/ipv6"
|
||||
)
|
||||
|
||||
var ipv6Header = []byte{
|
||||
0x60, 0x00, 0x00, 0x00,
|
||||
0x00, 0x20, 59, 64,
|
||||
0x20, 0x01, 0x0d, 0xb8, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x01,
|
||||
0x20, 0x01, 0x0d, 0xb8, 0x85, 0xa3, 0x08, 0xd3, 0x13, 0x19, 0x8a, 0x2e, 0x03, 0x70, 0x73, 0x48,
|
||||
}
|
||||
|
||||
var (
|
||||
testSrc4 = netip.MustParseAddr("192.0.2.1")
|
||||
testDst4 = netip.MustParseAddr("198.51.100.1")
|
||||
testSrc6 = netip.MustParseAddr("2001:db8::1")
|
||||
testDst6 = netip.MustParseAddr("2001:db8:1::1")
|
||||
)
|
||||
|
||||
func ipv4Packet(ttl, proto uint8, src, dst netip.Addr, options, payload []byte) []byte {
|
||||
hdrLen := ipv4.HeaderLen + len(options)
|
||||
b := make([]byte, hdrLen, hdrLen+len(payload))
|
||||
b[0] = 4<<4 | byte(hdrLen>>2)
|
||||
binary.BigEndian.PutUint16(b[2:4], uint16(hdrLen+len(payload)))
|
||||
b[8] = ttl
|
||||
b[9] = proto
|
||||
copy(b[12:16], src.AsSlice())
|
||||
copy(b[16:20], dst.AsSlice())
|
||||
copy(b[ipv4.HeaderLen:], options)
|
||||
return append(b, payload...)
|
||||
}
|
||||
|
||||
func ipv6Packet(hopLimit, nextHeader uint8, src, dst netip.Addr, payload []byte) []byte {
|
||||
b := make([]byte, ipv6.HeaderLen, ipv6.HeaderLen+len(payload))
|
||||
b[0] = 6 << 4
|
||||
binary.BigEndian.PutUint16(b[4:6], uint16(len(payload)))
|
||||
b[6] = nextHeader
|
||||
b[7] = hopLimit
|
||||
copy(b[8:24], src.AsSlice())
|
||||
copy(b[24:40], dst.AsSlice())
|
||||
return append(b, payload...)
|
||||
}
|
||||
|
||||
func ipv4ChecksumValid(header []byte) bool {
|
||||
var sum uint32
|
||||
for i := 0; i+1 < len(header); i += 2 {
|
||||
sum += uint32(binary.BigEndian.Uint16(header[i:]))
|
||||
}
|
||||
for sum > 0xffff {
|
||||
sum = sum&0xffff + sum>>16
|
||||
}
|
||||
return sum == 0xffff
|
||||
}
|
||||
|
||||
type mockStream struct {
|
||||
streamID quic.StreamID
|
||||
reading []byte
|
||||
toRead <-chan []byte
|
||||
datagrams <-chan []byte
|
||||
maxDatagramPayloadSize int
|
||||
sendDatagramErr error
|
||||
sent [][]byte
|
||||
writeStarted chan struct{}
|
||||
written chan<- []byte
|
||||
readErr error
|
||||
|
||||
mu sync.Mutex
|
||||
cancelWriteCodes []quic.StreamErrorCode
|
||||
}
|
||||
|
||||
func (m *mockStream) cancelWriteCode() (quic.StreamErrorCode, bool) {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
if len(m.cancelWriteCodes) == 0 {
|
||||
return 0, false
|
||||
}
|
||||
return m.cancelWriteCodes[0], true
|
||||
}
|
||||
|
||||
var _ http3Stream = &mockStream{}
|
||||
|
||||
func (m *mockStream) StreamID() quic.StreamID { return m.streamID }
|
||||
func (m *mockStream) Read(p []byte) (int, error) {
|
||||
if len(m.reading) == 0 && m.readErr != nil {
|
||||
return 0, m.readErr
|
||||
}
|
||||
if len(m.reading) == 0 {
|
||||
m.reading = <-m.toRead
|
||||
}
|
||||
n := copy(p, m.reading)
|
||||
m.reading = m.reading[n:]
|
||||
return n, nil
|
||||
}
|
||||
func (m *mockStream) CancelRead(quic.StreamErrorCode) {}
|
||||
func (m *mockStream) Write(p []byte) (int, error) {
|
||||
if m.writeStarted != nil {
|
||||
close(m.writeStarted)
|
||||
m.writeStarted = nil
|
||||
}
|
||||
if m.written != nil {
|
||||
m.written <- bytes.Clone(p)
|
||||
}
|
||||
return len(p), nil
|
||||
}
|
||||
func (m *mockStream) Close() error { return nil }
|
||||
func (m *mockStream) CancelWrite(code quic.StreamErrorCode) {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
m.cancelWriteCodes = append(m.cancelWriteCodes, code)
|
||||
}
|
||||
func (m *mockStream) SetWriteDeadline(time.Time) error { return nil }
|
||||
func (m *mockStream) SendDatagram(data []byte) error {
|
||||
if m.sendDatagramErr != nil {
|
||||
return m.sendDatagramErr
|
||||
}
|
||||
if size := quicvarint.Len(uint64(m.streamID/4)) + len(data); m.maxDatagramPayloadSize > 0 && size > m.maxDatagramPayloadSize {
|
||||
return &quic.DatagramTooLargeError{MaxDatagramPayloadSize: int64(m.maxDatagramPayloadSize)}
|
||||
}
|
||||
m.sent = append(m.sent, bytes.Clone(data))
|
||||
return nil
|
||||
}
|
||||
|
||||
func (m *mockStream) ReceiveDatagram(ctx context.Context) ([]byte, error) {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return nil, ctx.Err()
|
||||
case data, ok := <-m.datagrams:
|
||||
if !ok {
|
||||
return nil, io.EOF
|
||||
}
|
||||
return data, nil
|
||||
}
|
||||
}
|
||||
|
||||
func TestCapsuleWriteQueueLimit(t *testing.T) {
|
||||
writes := make(chan []byte)
|
||||
writeStarted := make(chan struct{})
|
||||
conn := newProxiedConn(&mockStream{
|
||||
writeStarted: writeStarted,
|
||||
written: writes,
|
||||
})
|
||||
t.Cleanup(func() { conn.Close() })
|
||||
|
||||
require.NoError(t, conn.AssignAddresses(nil))
|
||||
select {
|
||||
case <-writeStarted:
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("capsule write did not start")
|
||||
}
|
||||
|
||||
for range maxQueuedCapsules {
|
||||
require.NoError(t, conn.AssignAddresses(nil))
|
||||
}
|
||||
go func() {
|
||||
conn.Routes(context.Background())
|
||||
for range maxQueuedCapsules + 1 {
|
||||
<-writes
|
||||
}
|
||||
}()
|
||||
require.ErrorContains(t, conn.AssignAddresses(nil), "capsule queue full")
|
||||
require.ErrorIs(t, conn.AssignAddresses(nil), net.ErrClosed)
|
||||
}
|
||||
|
||||
func TestCapsuleReceiveQueueLimit(t *testing.T) {
|
||||
for _, name := range []string{"assignments", "requests"} {
|
||||
t.Run(name, func(t *testing.T) {
|
||||
var data []byte
|
||||
for i := range maxQueuedCapsules + 1 {
|
||||
if name == "assignments" {
|
||||
data = (&addressAssignCapsule{}).append(data)
|
||||
} else {
|
||||
data = (&addressRequestCapsule{
|
||||
RequestIDs: []AddressRequestID{AddressRequestID(i + 1)},
|
||||
Prefixes: []netip.Prefix{netip.MustParsePrefix("192.0.2.1/32")},
|
||||
}).append(data)
|
||||
}
|
||||
}
|
||||
conn := newProxiedConn(&mockStream{reading: data})
|
||||
t.Cleanup(func() { conn.Close() })
|
||||
ctx, cancel := context.WithTimeout(context.Background(), time.Second)
|
||||
defer cancel()
|
||||
_, err := conn.Routes(ctx)
|
||||
require.ErrorIs(t, err, net.ErrClosed)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestAbortErrorCode(t *testing.T) {
|
||||
var overflow []byte
|
||||
for range maxQueuedCapsules + 1 {
|
||||
overflow = (&addressAssignCapsule{}).append(overflow)
|
||||
}
|
||||
misordered := (&routeAdvertisementCapsule{IPAddressRanges: []IPRoute{
|
||||
{StartIP: netip.MustParseAddr("192.0.2.2"), EndIP: netip.MustParseAddr("192.0.2.1")},
|
||||
}}).append(nil)
|
||||
for _, c := range []struct {
|
||||
name string
|
||||
str *mockStream
|
||||
code http3.ErrCode
|
||||
}{
|
||||
{"malformed capsule", &mockStream{reading: misordered}, http3.ErrCodeMessageError},
|
||||
{"queue limit", &mockStream{reading: overflow}, http3.ErrCodeExcessiveLoad},
|
||||
{"reset by peer", &mockStream{readErr: &quic.StreamError{ErrorCode: quic.StreamErrorCode(http3.ErrCodeNoError), Remote: true}}, http3.ErrCodeRequestCanceled},
|
||||
{"reset by peer on a request stream", &mockStream{readErr: &http3.Error{ErrorCode: http3.ErrCodeNoError, Remote: true}}, http3.ErrCodeRequestCanceled},
|
||||
} {
|
||||
t.Run(c.name, func(t *testing.T) {
|
||||
conn := newProxiedConn(c.str)
|
||||
t.Cleanup(func() { conn.Close() })
|
||||
require.Eventually(t, func() bool {
|
||||
_, ok := c.str.cancelWriteCode()
|
||||
return ok
|
||||
}, time.Second, time.Millisecond)
|
||||
code, _ := c.str.cancelWriteCode()
|
||||
require.Equal(t, quic.StreamErrorCode(c.code), code)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestIncomingDatagrams(t *testing.T) {
|
||||
t.Run("empty packets", func(t *testing.T) {
|
||||
conn := newProxiedConn(&mockStream{})
|
||||
require.ErrorContains(t,
|
||||
conn.handleIncomingProxiedPacket([]byte{}),
|
||||
"connect-ip: empty packet",
|
||||
)
|
||||
})
|
||||
|
||||
t.Run("invalid IP version", func(t *testing.T) {
|
||||
conn := newProxiedConn(&mockStream{})
|
||||
data := make([]byte, 20)
|
||||
data[0] = 5 << 4
|
||||
require.ErrorContains(t,
|
||||
conn.handleIncomingProxiedPacket(data),
|
||||
"connect-ip: unknown IP versions: 5",
|
||||
)
|
||||
})
|
||||
|
||||
t.Run("IPv4 packet too short", func(t *testing.T) {
|
||||
conn := newProxiedConn(&mockStream{})
|
||||
data, err := (&ipv4.Header{
|
||||
Src: net.IPv4(1, 2, 3, 4),
|
||||
Dst: net.IPv4(159, 70, 42, 98),
|
||||
Len: 20,
|
||||
Checksum: 89,
|
||||
}).Marshal()
|
||||
require.NoError(t, err)
|
||||
require.ErrorContains(t,
|
||||
conn.handleIncomingProxiedPacket(data[:ipv4.HeaderLen-1]),
|
||||
"connect-ip: malformed datagram: too short",
|
||||
)
|
||||
})
|
||||
|
||||
t.Run("IPv6 packet too short", func(t *testing.T) {
|
||||
conn := newProxiedConn(&mockStream{})
|
||||
require.ErrorContains(t,
|
||||
conn.handleIncomingProxiedPacket(ipv6Header[:ipv6.HeaderLen-1]),
|
||||
"connect-ip: malformed datagram: too short",
|
||||
)
|
||||
})
|
||||
|
||||
t.Run("invalid source address", func(t *testing.T) {
|
||||
conn := newProxiedConn(&mockStream{})
|
||||
require.NoError(t, conn.AssignAddresses([]netip.Prefix{netip.MustParsePrefix("192.168.0.10/32")}))
|
||||
hdr := &ipv4.Header{
|
||||
Src: net.IPv4(192, 168, 0, 11),
|
||||
Dst: net.IPv4(159, 70, 42, 98),
|
||||
Len: 20,
|
||||
Checksum: 89,
|
||||
}
|
||||
data, err := hdr.Marshal()
|
||||
require.NoError(t, err)
|
||||
require.ErrorContains(t,
|
||||
conn.handleIncomingProxiedPacket(data),
|
||||
"connect-ip: datagram source address not allowed: 192.168.0.11",
|
||||
)
|
||||
})
|
||||
|
||||
t.Run("invalid destination address", func(t *testing.T) {
|
||||
conn := newProxiedConn(&mockStream{})
|
||||
require.NoError(t, conn.AssignAddresses([]netip.Prefix{netip.MustParsePrefix("192.168.0.10/32")}))
|
||||
require.NoError(t, conn.AdvertiseRoute([]IPRoute{
|
||||
{StartIP: netip.MustParseAddr("10.0.0.0"), EndIP: netip.MustParseAddr("10.1.2.3")},
|
||||
}))
|
||||
hdr := &ipv4.Header{
|
||||
Src: net.IPv4(192, 168, 0, 10),
|
||||
Dst: net.IPv4(10, 1, 2, 3),
|
||||
Len: 20,
|
||||
Checksum: 89,
|
||||
}
|
||||
data, err := hdr.Marshal()
|
||||
require.NoError(t, err)
|
||||
require.NoError(t, conn.handleIncomingProxiedPacket(data))
|
||||
|
||||
hdr.Dst = net.IPv4(10, 1, 2, 4)
|
||||
data, err = hdr.Marshal()
|
||||
require.NoError(t, err)
|
||||
require.ErrorContains(t,
|
||||
conn.handleIncomingProxiedPacket(data),
|
||||
"connect-ip: datagram destination address / protocol not allowed: 10.1.2.4 (protocol: 0)",
|
||||
)
|
||||
})
|
||||
|
||||
t.Run("invalid IP protocol", func(t *testing.T) {
|
||||
conn := newProxiedConn(&mockStream{})
|
||||
require.NoError(t, conn.AssignAddresses([]netip.Prefix{netip.MustParsePrefix("192.168.0.10/32")}))
|
||||
require.NoError(t, conn.AdvertiseRoute([]IPRoute{
|
||||
{StartIP: netip.MustParseAddr("10.0.0.0"), EndIP: netip.MustParseAddr("10.1.2.3"), IPProtocol: 42},
|
||||
}))
|
||||
hdr := &ipv4.Header{
|
||||
Src: net.IPv4(192, 168, 0, 10),
|
||||
Dst: net.IPv4(10, 1, 2, 3),
|
||||
Len: 20,
|
||||
Checksum: 89,
|
||||
Protocol: 42,
|
||||
}
|
||||
data, err := hdr.Marshal()
|
||||
require.NoError(t, err)
|
||||
require.NoError(t, conn.handleIncomingProxiedPacket(data))
|
||||
|
||||
hdr.Protocol = 41
|
||||
data, err = hdr.Marshal()
|
||||
require.NoError(t, err)
|
||||
require.ErrorContains(t,
|
||||
conn.handleIncomingProxiedPacket(data),
|
||||
"connect-ip: datagram destination address / protocol not allowed: 10.1.2.3 (protocol: 41)",
|
||||
)
|
||||
|
||||
hdr.Protocol = ipProtoICMP
|
||||
data, err = hdr.Marshal()
|
||||
require.NoError(t, err)
|
||||
require.NoError(t, conn.handleIncomingProxiedPacket(data))
|
||||
})
|
||||
|
||||
t.Run("packet from assigned address", func(t *testing.T) {
|
||||
readChan := make(chan []byte, 1)
|
||||
conn := newProxiedConn(&mockStream{toRead: readChan})
|
||||
|
||||
hdr := &ipv4.Header{
|
||||
Src: net.IPv4(159, 70, 42, 98),
|
||||
Dst: net.IPv4(192, 168, 0, 10),
|
||||
Len: 20,
|
||||
Checksum: 89,
|
||||
}
|
||||
data, err := hdr.Marshal()
|
||||
require.NoError(t, err)
|
||||
require.Error(t, conn.handleIncomingProxiedPacket(data), "connect-ip: datagram destination address")
|
||||
|
||||
readChan <- (&addressAssignCapsule{
|
||||
AssignedAddresses: []AssignedAddress{{IPPrefix: netip.MustParsePrefix("192.168.0.10/32")}},
|
||||
}).append(nil)
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), time.Second)
|
||||
defer cancel()
|
||||
_, err = conn.ReceiveAddressAssignment(ctx)
|
||||
require.NoError(t, err)
|
||||
require.NoError(t, conn.handleIncomingProxiedPacket(data))
|
||||
})
|
||||
}
|
||||
|
||||
func TestSkipUnknownCapsule(t *testing.T) {
|
||||
for _, typ := range []http3.CapsuleType{42, capsuleTypeDatagram} {
|
||||
readChan := make(chan []byte, 1)
|
||||
conn := newProxiedConn(&mockStream{toRead: readChan})
|
||||
|
||||
data := quicvarint.Append(nil, uint64(typ))
|
||||
data = quicvarint.Append(data, 3)
|
||||
data = append(data, "foo"...)
|
||||
data = (&addressAssignCapsule{
|
||||
AssignedAddresses: []AssignedAddress{{IPPrefix: netip.MustParsePrefix("192.168.0.10/32")}},
|
||||
}).append(data)
|
||||
readChan <- data
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), time.Second)
|
||||
assigned, err := conn.ReceiveAddressAssignment(ctx)
|
||||
cancel()
|
||||
require.NoError(t, err, "capsule type %d", typ)
|
||||
require.Equal(t, []AssignedAddress{{IPPrefix: netip.MustParsePrefix("192.168.0.10/32")}}, assigned)
|
||||
conn.Close()
|
||||
}
|
||||
}
|
||||
|
||||
func FuzzIncomingDatagram(f *testing.F) {
|
||||
conn := newProxiedConn(&mockStream{})
|
||||
require.NoError(f, conn.AssignAddresses([]netip.Prefix{
|
||||
netip.MustParsePrefix("192.168.0.0/16"),
|
||||
netip.MustParsePrefix("2001:db8::0/64"),
|
||||
}))
|
||||
require.NoError(f, conn.AdvertiseRoute([]IPRoute{
|
||||
{StartIP: netip.MustParseAddr("10.0.0.0"), EndIP: netip.MustParseAddr("10.1.2.3"), IPProtocol: 42},
|
||||
{StartIP: netip.MustParseAddr("2001:db8:1::"), EndIP: netip.MustParseAddr("2001:db8:1::ffff"), IPProtocol: 42},
|
||||
}))
|
||||
|
||||
ipv4Header, err := (&ipv4.Header{
|
||||
Src: net.IPv4(1, 2, 3, 4),
|
||||
Dst: net.IPv4(159, 70, 42, 98),
|
||||
Len: 20,
|
||||
Checksum: 89,
|
||||
}).Marshal()
|
||||
require.NoError(f, err)
|
||||
|
||||
f.Add(ipv4Header)
|
||||
f.Add(ipv6Header)
|
||||
|
||||
f.Fuzz(func(t *testing.T, data []byte) {
|
||||
conn.handleIncomingProxiedPacket(data)
|
||||
})
|
||||
}
|
||||
|
||||
func TestSendingDatagrams(t *testing.T) {
|
||||
t.Run("invalid IP version", func(t *testing.T) {
|
||||
conn := newProxiedConn(&mockStream{})
|
||||
data := make([]byte, 20)
|
||||
data[0] = 5 << 4
|
||||
_, err := conn.composeDatagram(data)
|
||||
require.ErrorContains(t, err, "connect-ip: unknown IP versions: 5")
|
||||
})
|
||||
|
||||
t.Run("IPv4 packet too short", func(t *testing.T) {
|
||||
conn := newProxiedConn(&mockStream{})
|
||||
data, err := (&ipv4.Header{
|
||||
Src: net.IPv4(1, 2, 3, 4),
|
||||
Dst: net.IPv4(159, 70, 42, 98),
|
||||
Len: 20,
|
||||
Checksum: 89,
|
||||
}).Marshal()
|
||||
require.NoError(t, err)
|
||||
_, err = conn.composeDatagram(data[:ipv4.HeaderLen-1])
|
||||
require.ErrorContains(t, err, "connect-ip: IPv4 packet too short")
|
||||
})
|
||||
|
||||
t.Run("IPv6 packet too short", func(t *testing.T) {
|
||||
conn := newProxiedConn(&mockStream{})
|
||||
_, err := conn.composeDatagram(ipv6Header[:ipv6.HeaderLen-1])
|
||||
require.ErrorContains(t, err, "connect-ip: IPv6 packet too short")
|
||||
})
|
||||
}
|
||||
|
||||
func TestWritePacketDropsWithoutSending(t *testing.T) {
|
||||
setIHL := func(b []byte, ihl byte) []byte { b[0] = 4<<4 | ihl; return b }
|
||||
setTotalLen := func(b []byte, l uint16) []byte { binary.BigEndian.PutUint16(b[2:4], l); return b }
|
||||
payload := make([]byte, 20)
|
||||
|
||||
for _, tc := range []struct {
|
||||
name string
|
||||
packet []byte
|
||||
}{
|
||||
{"nil", nil},
|
||||
{"empty", []byte{}},
|
||||
{"IPv4 TTL 1", ipv4Packet(1, 17, testSrc4, testDst4, nil, payload)},
|
||||
{"IPv4 TTL 0", ipv4Packet(0, 17, testSrc4, testDst4, nil, payload)},
|
||||
{"IPv6 Hop Limit 1", ipv6Packet(1, 17, testSrc6, testDst6, payload)},
|
||||
{"IPv6 Hop Limit 0", ipv6Packet(0, 17, testSrc6, testDst6, payload)},
|
||||
{"IPv4 IHL below 5", setIHL(ipv4Packet(64, 17, testSrc4, testDst4, nil, payload), 4)},
|
||||
{"IPv4 IHL beyond total length", setIHL(ipv4Packet(64, 17, testSrc4, testDst4, nil, payload[:8]), 8)},
|
||||
{"IPv4 IHL beyond packet", setTotalLen(setIHL(ipv4Packet(64, 17, testSrc4, testDst4, nil, payload), 15), 60)},
|
||||
{"IPv4 total length beyond packet", setTotalLen(ipv4Packet(64, 17, testSrc4, testDst4, nil, payload), 41)},
|
||||
{"IPv4 total length below header", setTotalLen(ipv4Packet(64, 17, testSrc4, testDst4, nil, payload), 19)},
|
||||
} {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
str := &mockStream{}
|
||||
conn := newProxiedConn(str)
|
||||
t.Cleanup(func() { conn.Close() })
|
||||
|
||||
orig := bytes.Clone(tc.packet)
|
||||
icmpPacket, err := conn.WritePacket(tc.packet)
|
||||
require.NoError(t, err)
|
||||
require.Nil(t, icmpPacket)
|
||||
require.Empty(t, str.sent)
|
||||
require.Equal(t, orig, tc.packet, "dropped packets must not be modified")
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestWritePacketIPv4Checksum(t *testing.T) {
|
||||
for _, tc := range []struct {
|
||||
name string
|
||||
options []byte
|
||||
}{
|
||||
{"no options", nil},
|
||||
{"Router Alert option", []byte{0x94, 0x04, 0x00, 0x00}},
|
||||
{"maximum header length", bytes.Repeat([]byte{0x01}, 40)},
|
||||
} {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
str := &mockStream{}
|
||||
conn := newProxiedConn(str)
|
||||
t.Cleanup(func() { conn.Close() })
|
||||
|
||||
packet := ipv4Packet(64, 17, testSrc4, testDst4, tc.options, []byte("foobar"))
|
||||
icmpPacket, err := conn.WritePacket(packet)
|
||||
require.NoError(t, err)
|
||||
require.Nil(t, icmpPacket)
|
||||
require.Len(t, str.sent, 1)
|
||||
require.Equal(t, contextIDZero, str.sent[0][:len(contextIDZero)])
|
||||
sent := str.sent[0][len(contextIDZero):]
|
||||
require.Len(t, sent, len(packet))
|
||||
require.Equal(t, uint8(63), sent[8])
|
||||
require.True(t, ipv4ChecksumValid(sent[:ipv4.HeaderLen+len(tc.options)]))
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestWritePacketTooLarge(t *testing.T) {
|
||||
for _, tc := range []struct {
|
||||
name string
|
||||
streamID quic.StreamID
|
||||
maxPayloadSize int
|
||||
ipv6 bool
|
||||
wantMTU int
|
||||
wantMTUTooSmall bool
|
||||
}{
|
||||
{name: "IPv4", maxPayloadSize: 1200, wantMTU: 1198},
|
||||
{name: "IPv4, 2-byte Quarter Stream ID", streamID: 4 * 64, maxPayloadSize: 1200, wantMTU: 1197},
|
||||
{name: "IPv4 minimum MTU", maxPayloadSize: 70, wantMTU: 68},
|
||||
{name: "IPv4 below minimum MTU", maxPayloadSize: 69, wantMTU: 67, wantMTUTooSmall: true},
|
||||
{name: "IPv6", maxPayloadSize: 1400, ipv6: true, wantMTU: 1398},
|
||||
{name: "IPv6, 4-byte Quarter Stream ID", streamID: 4 * 20000, maxPayloadSize: 1400, ipv6: true, wantMTU: 1395},
|
||||
{name: "IPv6 minimum MTU", maxPayloadSize: 1282, ipv6: true, wantMTU: 1280},
|
||||
{name: "IPv6 below minimum MTU", maxPayloadSize: 1281, ipv6: true, wantMTU: 1279, wantMTUTooSmall: true},
|
||||
} {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
str := &mockStream{streamID: tc.streamID, maxDatagramPayloadSize: tc.maxPayloadSize}
|
||||
conn := newProxiedConn(str)
|
||||
t.Cleanup(func() { conn.Close() })
|
||||
|
||||
packetOfSize := func(size int) []byte {
|
||||
if tc.ipv6 {
|
||||
return ipv6Packet(64, 17, testSrc6, testDst6, make([]byte, size-ipv6.HeaderLen))
|
||||
}
|
||||
return ipv4Packet(64, 17, testSrc4, testDst4, nil, make([]byte, size-ipv4.HeaderLen))
|
||||
}
|
||||
require.Equal(t, tc.wantMTU, conn.MaxPacketSize())
|
||||
|
||||
icmpPacket, err := conn.WritePacket(packetOfSize(tc.wantMTU))
|
||||
require.NoError(t, err)
|
||||
require.Nil(t, icmpPacket)
|
||||
require.Len(t, str.sent, 1)
|
||||
|
||||
icmpPacket, err = conn.WritePacket(packetOfSize(tc.wantMTU + 1))
|
||||
if tc.wantMTUTooSmall {
|
||||
require.ErrorIs(t, err, ErrMTUTooSmall)
|
||||
require.Nil(t, icmpPacket)
|
||||
return
|
||||
}
|
||||
require.NoError(t, err)
|
||||
if tc.ipv6 {
|
||||
msg, err := icmp.ParseMessage(ipProtoICMPv6, icmpPacket[ipv6.HeaderLen:])
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, ipv6.ICMPTypePacketTooBig, msg.Type)
|
||||
require.Equal(t, tc.wantMTU, msg.Body.(*icmp.PacketTooBig).MTU)
|
||||
} else {
|
||||
msg := icmpPacket[ipv4.HeaderLen:]
|
||||
require.Equal(t, []byte{3, 4}, msg[:2])
|
||||
require.Equal(t, uint16(tc.wantMTU), binary.BigEndian.Uint16(msg[6:8]))
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestMaxPacketSize(t *testing.T) {
|
||||
t.Run("sends nothing", func(t *testing.T) {
|
||||
str := &mockStream{streamID: 8, maxDatagramPayloadSize: 1350}
|
||||
conn := newProxiedConn(str)
|
||||
t.Cleanup(func() { conn.Close() })
|
||||
require.Equal(t, 1348, conn.MaxPacketSize())
|
||||
require.Empty(t, str.sent)
|
||||
})
|
||||
|
||||
t.Run("datagrams unsupported", func(t *testing.T) {
|
||||
conn := newProxiedConn(&mockStream{sendDatagramErr: errors.New("datagram support disabled")})
|
||||
t.Cleanup(func() { conn.Close() })
|
||||
require.Zero(t, conn.MaxPacketSize())
|
||||
})
|
||||
|
||||
t.Run("closed", func(t *testing.T) {
|
||||
conn := newProxiedConn(&mockStream{maxDatagramPayloadSize: 1350})
|
||||
require.NoError(t, conn.Close())
|
||||
require.Zero(t, conn.MaxPacketSize())
|
||||
})
|
||||
}
|
||||
|
||||
func TestReadPacketDropsMalformedDatagrams(t *testing.T) {
|
||||
packet := ipv4Packet(64, 17, testSrc4, testDst4, nil, []byte("foobar"))
|
||||
for _, tc := range []struct {
|
||||
name string
|
||||
datagram []byte
|
||||
}{
|
||||
{"empty", []byte{}},
|
||||
{"truncated Context ID", []byte{0x40}},
|
||||
{"unknown Context ID", append([]byte{0x02}, packet...)},
|
||||
{"empty IP packet", []byte{0x00}},
|
||||
{"invalid IP packet", []byte{0x00, 0x50}},
|
||||
} {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
datagrams := make(chan []byte, 2)
|
||||
conn := newProxiedConn(&mockStream{datagrams: datagrams})
|
||||
t.Cleanup(func() { conn.Close() })
|
||||
require.NoError(t, conn.AdvertiseRoute([]IPRoute{
|
||||
{StartIP: netip.IPv4Unspecified(), EndIP: netip.MustParseAddr("255.255.255.255")},
|
||||
}))
|
||||
|
||||
datagrams <- tc.datagram
|
||||
datagrams <- append(bytes.Clone(contextIDZero), packet...)
|
||||
b := make([]byte, 1500)
|
||||
n, err := conn.ReadPacket(b)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, packet, b[:n])
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestReadPacketShortBuffer(t *testing.T) {
|
||||
datagrams := make(chan []byte, 3)
|
||||
conn := newProxiedConn(&mockStream{datagrams: datagrams})
|
||||
t.Cleanup(func() { conn.Close() })
|
||||
require.NoError(t, conn.AdvertiseRoute([]IPRoute{
|
||||
{StartIP: netip.IPv4Unspecified(), EndIP: netip.MustParseAddr("255.255.255.255")},
|
||||
}))
|
||||
|
||||
packet := ipv4Packet(64, 17, testSrc4, testDst4, nil, []byte("foobar"))
|
||||
for range 3 {
|
||||
datagrams <- append(bytes.Clone(contextIDZero), packet...)
|
||||
}
|
||||
n, err := conn.ReadPacket(make([]byte, len(packet)-1))
|
||||
require.ErrorIs(t, err, io.ErrShortBuffer)
|
||||
require.Zero(t, n)
|
||||
|
||||
b := make([]byte, len(packet))
|
||||
n, err = conn.ReadPacket(b)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, packet, b[:n])
|
||||
n, err = conn.ReadPacket(make([]byte, 1500))
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, len(packet), n)
|
||||
}
|
||||
|
||||
func TestCloseConcurrently(t *testing.T) {
|
||||
for _, side := range []string{"client", "proxy"} {
|
||||
t.Run(side, func(t *testing.T) {
|
||||
client, server := setupConns(t)
|
||||
conn := client
|
||||
if side == "proxy" {
|
||||
conn = server
|
||||
}
|
||||
|
||||
readErr := make(chan error, 1)
|
||||
go func() {
|
||||
b := make([]byte, 1500)
|
||||
for {
|
||||
if _, err := conn.ReadPacket(b); err != nil {
|
||||
readErr <- err
|
||||
return
|
||||
}
|
||||
}
|
||||
}()
|
||||
writeErr := make(chan error, 1)
|
||||
go func() {
|
||||
for {
|
||||
if _, err := conn.WritePacket(ipv4Packet(64, 17, testSrc4, testDst4, nil, nil)); err != nil {
|
||||
writeErr <- err
|
||||
return
|
||||
}
|
||||
}
|
||||
}()
|
||||
|
||||
var wg sync.WaitGroup
|
||||
for range 4 {
|
||||
wg.Go(func() { assert.NoError(t, conn.Close()) })
|
||||
}
|
||||
wg.Wait()
|
||||
for _, errChan := range []chan error{readErr, writeErr} {
|
||||
select {
|
||||
case err := <-errChan:
|
||||
require.ErrorIs(t, err, net.ErrClosed)
|
||||
case <-time.After(5 * time.Second):
|
||||
t.Fatal("timeout")
|
||||
}
|
||||
}
|
||||
require.NoError(t, conn.Close())
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,94 @@
|
||||
/* SPDX-License-Identifier: MIT
|
||||
*
|
||||
* Copyright 2024 Marten Seemann
|
||||
* Adapted from github.com/quic-go/connect-ip-go (commit a0c35fa).
|
||||
*/
|
||||
|
||||
package connectip
|
||||
|
||||
import (
|
||||
"encoding/binary"
|
||||
"errors"
|
||||
"fmt"
|
||||
|
||||
"golang.org/x/net/icmp"
|
||||
"golang.org/x/net/ipv4"
|
||||
"golang.org/x/net/ipv6"
|
||||
)
|
||||
|
||||
const (
|
||||
ipv4MinMTU = 68
|
||||
ipv6MinMTU = 1280
|
||||
)
|
||||
|
||||
var ErrMTUTooSmall = errors.New("connect-ip: tunnel MTU below the minimum link MTU")
|
||||
|
||||
func composeICMPTooLargePacket(b []byte, mtu int) ([]byte, error) {
|
||||
if len(b) == 0 {
|
||||
return nil, errors.New("connect-ip: empty packet")
|
||||
}
|
||||
|
||||
var icmpMessage *icmp.Message
|
||||
var psh []byte
|
||||
switch v := ipVersion(b); v {
|
||||
case 4:
|
||||
if len(b) < ipv4.HeaderLen {
|
||||
return nil, errors.New("connect-ip: IPv4 packet too short")
|
||||
}
|
||||
if mtu < ipv4MinMTU {
|
||||
return nil, fmt.Errorf("%w: %d bytes", ErrMTUTooSmall, mtu)
|
||||
}
|
||||
icmpMessage = &icmp.Message{
|
||||
Type: ipv4.ICMPTypeDestinationUnreachable,
|
||||
Code: 4,
|
||||
Body: &icmp.PacketTooBig{
|
||||
MTU: mtu,
|
||||
Data: b[:min(len(b), max(ipv4.HeaderLen, int(b[0]&0x0f)<<2)+8)],
|
||||
},
|
||||
}
|
||||
case 6:
|
||||
if len(b) < ipv6.HeaderLen {
|
||||
return nil, errors.New("connect-ip: IPv6 packet too short")
|
||||
}
|
||||
if mtu < ipv6MinMTU {
|
||||
return nil, fmt.Errorf("%w: %d bytes", ErrMTUTooSmall, mtu)
|
||||
}
|
||||
icmpMessage = &icmp.Message{
|
||||
Type: ipv6.ICMPTypePacketTooBig,
|
||||
Body: &icmp.PacketTooBig{
|
||||
MTU: mtu,
|
||||
Data: b[:min(len(b), 1232)],
|
||||
},
|
||||
}
|
||||
psh = icmp.IPv6PseudoHeader(b[24:40], b[8:24])
|
||||
default:
|
||||
return nil, fmt.Errorf("connect-ip: unknown IP version: %d", v)
|
||||
}
|
||||
|
||||
icmp, err := icmpMessage.Marshal(psh)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("connect-ip: failed to marshal ICMP message: %w", err)
|
||||
}
|
||||
|
||||
if ipVersion(b) == 4 {
|
||||
var header [ipv4.HeaderLen]byte
|
||||
header[0] = 4<<4 | ipv4.HeaderLen>>2
|
||||
ipLen := ipv4.HeaderLen + len(icmp)
|
||||
binary.BigEndian.PutUint16(header[2:4], uint16(ipLen))
|
||||
header[8] = 64
|
||||
header[9] = 1
|
||||
copy(header[12:16], b[16:20])
|
||||
copy(header[16:20], b[12:16])
|
||||
binary.BigEndian.PutUint16(header[10:12], calculateIPv4Checksum(header[:]))
|
||||
return append(header[:], icmp...), nil
|
||||
}
|
||||
|
||||
var header [ipv6.HeaderLen]byte
|
||||
header[0] = 6 << 4
|
||||
binary.BigEndian.PutUint16(header[4:6], uint16(len(icmp)))
|
||||
header[6] = 58
|
||||
header[7] = 64
|
||||
copy(header[8:24], b[24:40])
|
||||
copy(header[24:40], b[8:24])
|
||||
return append(header[:], icmp...), nil
|
||||
}
|
||||
@@ -0,0 +1,158 @@
|
||||
/* SPDX-License-Identifier: MIT
|
||||
*
|
||||
* Copyright 2024 Marten Seemann
|
||||
* Adapted from github.com/quic-go/connect-ip-go (commit a0c35fa).
|
||||
*/
|
||||
|
||||
package connectip
|
||||
|
||||
import (
|
||||
"encoding/binary"
|
||||
"net"
|
||||
"net/netip"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
"golang.org/x/net/icmp"
|
||||
"golang.org/x/net/ipv4"
|
||||
"golang.org/x/net/ipv6"
|
||||
)
|
||||
|
||||
func TestICMPTooLargeIPv4(t *testing.T) {
|
||||
src := netip.MustParseAddr("192.168.1.1")
|
||||
dst := netip.MustParseAddr("8.8.8.8")
|
||||
origHdr := &ipv4.Header{
|
||||
Version: 4,
|
||||
Len: ipv4.HeaderLen,
|
||||
TotalLen: 60,
|
||||
TTL: 64,
|
||||
Protocol: 6,
|
||||
Src: src.AsSlice(),
|
||||
Dst: dst.AsSlice(),
|
||||
}
|
||||
origBytes, err := origHdr.Marshal()
|
||||
require.NoError(t, err)
|
||||
data, err := composeICMPTooLargePacket(origBytes, 1200)
|
||||
require.NoError(t, err)
|
||||
|
||||
hdr, err := ipv4.ParseHeader(data)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, 4, hdr.Version)
|
||||
require.Equal(t, ipProtoICMP, hdr.Protocol)
|
||||
require.Equal(t, dst.String(), hdr.Src.String())
|
||||
require.Equal(t, src.String(), hdr.Dst.String())
|
||||
require.Equal(t, uint16(hdr.Checksum), calculateIPv4Checksum(data[:ipv4.HeaderLen]))
|
||||
icmpMsg, err := icmp.ParseMessage(ipProtoICMP, data[ipv4.HeaderLen:])
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, ipv4.ICMPTypeDestinationUnreachable, icmpMsg.Type)
|
||||
require.Equal(t, 4, icmpMsg.Code)
|
||||
require.Equal(t, uint16(1200), binary.BigEndian.Uint16(data[ipv4.HeaderLen+6:]))
|
||||
require.Equal(t, origBytes, data[ipv4.HeaderLen+8:])
|
||||
}
|
||||
|
||||
func TestICMPTooLargeIPv4Options(t *testing.T) {
|
||||
options := []byte{0x94, 0x04, 0x00, 0x00}
|
||||
orig := ipv4Packet(64, 6, netip.MustParseAddr("192.168.1.1"), netip.MustParseAddr("8.8.8.8"), options, make([]byte, 20))
|
||||
data, err := composeICMPTooLargePacket(orig, 1200)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, orig[:ipv4.HeaderLen+len(options)+8], data[ipv4.HeaderLen+8:])
|
||||
}
|
||||
|
||||
func TestICMPTooLargeIPv6(t *testing.T) {
|
||||
const mtu = 1337
|
||||
src := netip.MustParseAddr("2001:db8::1")
|
||||
dst := netip.MustParseAddr("1:2:3:4::5")
|
||||
orig := []byte{
|
||||
0x60, 0x00, 0x00, 0x00,
|
||||
0x00, 0x00,
|
||||
0x00, 0x2a,
|
||||
}
|
||||
orig = append(orig, src.AsSlice()...)
|
||||
orig = append(orig, dst.AsSlice()...)
|
||||
orig = append(orig, []byte("foobar")...)
|
||||
data, err := composeICMPTooLargePacket(orig, mtu)
|
||||
require.NoError(t, err)
|
||||
|
||||
hdr, err := ipv6.ParseHeader(data)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, 6, hdr.Version)
|
||||
require.Equal(t, ipProtoICMPv6, hdr.NextHeader)
|
||||
require.Equal(t, dst.String(), hdr.Src.String())
|
||||
require.Equal(t, src.String(), hdr.Dst.String())
|
||||
icmpMsg, err := icmp.ParseMessage(ipProtoICMPv6, data[ipv6.HeaderLen:])
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, ipv6.ICMPTypePacketTooBig, icmpMsg.Type)
|
||||
icmpBody, ok := icmpMsg.Body.(*icmp.PacketTooBig)
|
||||
require.True(t, ok)
|
||||
require.Equal(t, mtu, icmpBody.MTU)
|
||||
require.Equal(t, orig, icmpBody.Data)
|
||||
}
|
||||
|
||||
func TestICMPTooLargeMinimumMTU(t *testing.T) {
|
||||
ipv4Orig := ipv4Packet(64, 6, testSrc4, testDst4, nil, make([]byte, 100))
|
||||
ipv6Orig := ipv6Packet(64, 6, testSrc6, testDst6, make([]byte, 1300))
|
||||
for _, tc := range []struct {
|
||||
name string
|
||||
packet []byte
|
||||
mtu int
|
||||
tooSmall bool
|
||||
}{
|
||||
{"IPv4 minimum", ipv4Orig, 68, false},
|
||||
{"IPv4 below minimum", ipv4Orig, 67, true},
|
||||
{"IPv4 negative", ipv4Orig, -1, true},
|
||||
{"IPv6 minimum", ipv6Orig, 1280, false},
|
||||
{"IPv6 below minimum", ipv6Orig, 1279, true},
|
||||
} {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
data, err := composeICMPTooLargePacket(tc.packet, tc.mtu)
|
||||
if tc.tooSmall {
|
||||
require.ErrorIs(t, err, ErrMTUTooSmall)
|
||||
require.Nil(t, data)
|
||||
return
|
||||
}
|
||||
require.NoError(t, err)
|
||||
require.NotEmpty(t, data)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestICMPFailures(t *testing.T) {
|
||||
t.Run("empty packet", func(t *testing.T) {
|
||||
_, err := composeICMPTooLargePacket([]byte{}, 1)
|
||||
require.EqualError(t, err, "connect-ip: empty packet")
|
||||
})
|
||||
|
||||
t.Run("too short IPv4 header", func(t *testing.T) {
|
||||
origHdr := &ipv4.Header{
|
||||
Version: 4,
|
||||
Len: ipv4.HeaderLen,
|
||||
TotalLen: 60,
|
||||
Src: net.IPv4(1, 2, 3, 4),
|
||||
Dst: net.IPv4(5, 6, 7, 8),
|
||||
}
|
||||
data, err := origHdr.Marshal()
|
||||
require.NoError(t, err)
|
||||
_, err = composeICMPTooLargePacket(data[:ipv4.HeaderLen-1], 1)
|
||||
require.EqualError(t, err, "connect-ip: IPv4 packet too short")
|
||||
})
|
||||
|
||||
t.Run("too short IPv6 header", func(t *testing.T) {
|
||||
data := []byte{
|
||||
0x60, 0x00, 0x00, 0x00,
|
||||
0x00, 0x00,
|
||||
0x00, 0x40,
|
||||
}
|
||||
data = append(data, net.ParseIP("2001:db8::1").To16()...)
|
||||
data = append(data, net.ParseIP("2001:db8::2").To16()...)
|
||||
_, err := composeICMPTooLargePacket(data[:ipv6.HeaderLen-1], 1)
|
||||
require.EqualError(t, err, "connect-ip: IPv6 packet too short")
|
||||
})
|
||||
|
||||
t.Run("unknown IP version", func(t *testing.T) {
|
||||
data := []byte{
|
||||
0x30, 0x00, 0x00, 0x00,
|
||||
}
|
||||
_, err := composeICMPTooLargePacket(data, 1)
|
||||
require.EqualError(t, err, "connect-ip: unknown IP version: 3")
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,73 @@
|
||||
/* SPDX-License-Identifier: MIT
|
||||
*
|
||||
* Copyright 2024 Marten Seemann
|
||||
* Adapted from github.com/quic-go/connect-ip-go (commit a0c35fa).
|
||||
*/
|
||||
|
||||
package connectip
|
||||
|
||||
import "net/netip"
|
||||
|
||||
func rangeToPrefixes(start, end netip.Addr) []netip.Prefix {
|
||||
var prefixes []netip.Prefix
|
||||
for current := start; current.Compare(end) <= 0; {
|
||||
prefix := findLargestPrefix(current, end)
|
||||
prefixes = append(prefixes, prefix)
|
||||
|
||||
lastIP := lastIPInPrefix(prefix)
|
||||
if lastIP.Compare(end) >= 0 {
|
||||
break
|
||||
}
|
||||
current = lastIP.Next()
|
||||
}
|
||||
return prefixes
|
||||
}
|
||||
|
||||
func findLargestPrefix(start, end netip.Addr) netip.Prefix {
|
||||
if start == end {
|
||||
return netip.PrefixFrom(start, start.BitLen())
|
||||
}
|
||||
|
||||
var prefixLen int
|
||||
for prefixLen = start.BitLen(); prefixLen > 0; prefixLen-- {
|
||||
prefix := netip.PrefixFrom(start, prefixLen-1)
|
||||
if lastIPInPrefix(prefix).Compare(end) > 0 || !isAligned(start, prefixLen-1) {
|
||||
break
|
||||
}
|
||||
}
|
||||
return netip.PrefixFrom(start, prefixLen)
|
||||
}
|
||||
|
||||
func lastIPInPrefix(prefix netip.Prefix) netip.Addr {
|
||||
addr := prefix.Addr()
|
||||
bits := addr.As16()
|
||||
|
||||
hostBits := addr.BitLen() - prefix.Bits()
|
||||
|
||||
for i := len(bits) - 1; i >= 0 && hostBits > 0; i-- {
|
||||
bitsInThisByte := min(8, hostBits)
|
||||
mask := byte((1 << bitsInThisByte) - 1)
|
||||
bits[i] |= mask
|
||||
hostBits -= bitsInThisByte
|
||||
}
|
||||
|
||||
if addr.Is4() {
|
||||
return netip.AddrFrom4([4]byte(bits[12:16]))
|
||||
}
|
||||
return netip.AddrFrom16(bits)
|
||||
}
|
||||
|
||||
func isAligned(addr netip.Addr, prefixLen int) bool {
|
||||
bits := addr.As16()
|
||||
|
||||
hostBits := addr.BitLen() - prefixLen
|
||||
for i := len(bits) - 1; i >= 0 && hostBits > 0; i-- {
|
||||
bitsInThisByte := min(8, hostBits)
|
||||
mask := byte((1 << bitsInThisByte) - 1)
|
||||
if bits[i]&mask != 0 {
|
||||
return false
|
||||
}
|
||||
hostBits -= bitsInThisByte
|
||||
}
|
||||
return true
|
||||
}
|
||||
@@ -0,0 +1,78 @@
|
||||
/* SPDX-License-Identifier: MIT
|
||||
*
|
||||
* Copyright 2024 Marten Seemann
|
||||
* Adapted from github.com/quic-go/connect-ip-go (commit a0c35fa).
|
||||
*/
|
||||
|
||||
package connectip
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"net/netip"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestIPRanges(t *testing.T) {
|
||||
tests := []struct {
|
||||
start, end netip.Addr
|
||||
want []netip.Prefix
|
||||
}{
|
||||
{
|
||||
start: netip.MustParseAddr("192.168.1.1"),
|
||||
end: netip.MustParseAddr("192.168.1.1"),
|
||||
want: []netip.Prefix{netip.MustParsePrefix("192.168.1.1/32")},
|
||||
},
|
||||
{
|
||||
start: netip.MustParseAddr("192.168.1.0"),
|
||||
end: netip.MustParseAddr("192.168.1.1"),
|
||||
want: []netip.Prefix{netip.MustParsePrefix("192.168.1.0/31")},
|
||||
},
|
||||
{
|
||||
start: netip.MustParseAddr("192.168.1.1"),
|
||||
end: netip.MustParseAddr("192.168.1.2"),
|
||||
want: []netip.Prefix{netip.MustParsePrefix("192.168.1.1/32"), netip.MustParsePrefix("192.168.1.2/32")},
|
||||
},
|
||||
{
|
||||
start: netip.MustParseAddr("192.168.1.0"),
|
||||
end: netip.MustParseAddr("192.168.1.255"),
|
||||
want: []netip.Prefix{netip.MustParsePrefix("192.168.1.0/24")},
|
||||
},
|
||||
{
|
||||
start: netip.MustParseAddr("10.0.0.0"),
|
||||
end: netip.MustParseAddr("10.1.0.255"),
|
||||
want: []netip.Prefix{netip.MustParsePrefix("10.0.0.0/16"), netip.MustParsePrefix("10.1.0.0/24")},
|
||||
},
|
||||
{
|
||||
start: netip.MustParseAddr("2001:0db8:85a3::8a2e:0370:7334"),
|
||||
end: netip.MustParseAddr("2001:0db8:85a3::8a2e:0370:7334"),
|
||||
want: []netip.Prefix{netip.MustParsePrefix("2001:0db8:85a3::8a2e:0370:7334/128")},
|
||||
},
|
||||
{
|
||||
start: netip.MustParseAddr("2001:db8::0"),
|
||||
end: netip.MustParseAddr("2001:db8::ffff:ffff:ffff:ffff"),
|
||||
want: []netip.Prefix{netip.MustParsePrefix("2001:db8::/64")},
|
||||
},
|
||||
{
|
||||
start: netip.MustParseAddr("2001:db8::1"),
|
||||
end: netip.MustParseAddr("2001:db8::2"),
|
||||
want: []netip.Prefix{netip.MustParsePrefix("2001:db8::1/128"), netip.MustParsePrefix("2001:db8::2/128")},
|
||||
},
|
||||
{
|
||||
start: netip.MustParseAddr("2001:db8:1234:5678::"),
|
||||
end: netip.MustParseAddr("2001:db8:1234:5679::"),
|
||||
want: []netip.Prefix{
|
||||
netip.MustParsePrefix("2001:db8:1234:5678::/64"),
|
||||
netip.MustParsePrefix("2001:db8:1234:5679::/128"),
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
for _, test := range tests {
|
||||
t.Run(fmt.Sprintf("%s-%s", test.start, test.end), func(t *testing.T) {
|
||||
prefixes := rangeToPrefixes(test.start, test.end)
|
||||
require.Equal(t, test.want, prefixes)
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,30 @@
|
||||
/* SPDX-License-Identifier: MIT
|
||||
*
|
||||
* Copyright 2024 Marten Seemann
|
||||
* Adapted from github.com/quic-go/connect-ip-go (commit a0c35fa).
|
||||
*/
|
||||
|
||||
package connectip
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"net/http"
|
||||
|
||||
"github.com/apernet/quic-go/http3"
|
||||
"github.com/apernet/quic-go/quicvarint"
|
||||
)
|
||||
|
||||
var contextIDZero = quicvarint.Append([]byte{}, 0)
|
||||
|
||||
type Proxy struct{}
|
||||
|
||||
func (s *Proxy) Proxy(w http.ResponseWriter, _ *ProxyRequest) (*Conn, error) {
|
||||
streamer, ok := w.(http3.HTTPStreamer)
|
||||
if !ok {
|
||||
return nil, errors.New("connect-ip: response writer is not an HTTP/3 stream")
|
||||
}
|
||||
w.Header().Set(http3.CapsuleProtocolHeader, capsuleProtocolHeaderValue)
|
||||
w.WriteHeader(http.StatusOK)
|
||||
|
||||
return newProxiedConn(streamer.HTTPStream()), nil
|
||||
}
|
||||
@@ -0,0 +1,413 @@
|
||||
/* SPDX-License-Identifier: MIT
|
||||
*
|
||||
* Copyright 2024 Marten Seemann
|
||||
* Adapted from github.com/quic-go/connect-ip-go (commit a0c35fa).
|
||||
*/
|
||||
|
||||
package connectip
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/tls"
|
||||
"encoding/binary"
|
||||
"fmt"
|
||||
"net"
|
||||
"net/http"
|
||||
"net/netip"
|
||||
"slices"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/apernet/quic-go"
|
||||
"github.com/apernet/quic-go/http3"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"golang.org/x/net/ipv4"
|
||||
"golang.org/x/net/ipv6"
|
||||
)
|
||||
|
||||
func dialHTTP3(t *testing.T, addr string) *http3.ClientConn {
|
||||
t.Helper()
|
||||
|
||||
ctx, cancel := context.WithTimeout(t.Context(), 5*time.Second)
|
||||
defer cancel()
|
||||
qconn, err := quic.DialAddr(
|
||||
ctx,
|
||||
addr,
|
||||
&tls.Config{ServerName: "localhost", RootCAs: certPool, NextProtos: []string{http3.NextProtoH3}},
|
||||
&quic.Config{EnableDatagrams: true, InitialPacketSize: 1350, DisablePathMTUDiscovery: true},
|
||||
)
|
||||
require.NoError(t, err)
|
||||
t.Cleanup(func() { qconn.CloseWithError(0, "") })
|
||||
return (&http3.Transport{EnableDatagrams: true}).NewClientConn(qconn)
|
||||
}
|
||||
|
||||
func setupConns(t *testing.T) (client, server *Conn) {
|
||||
t.Helper()
|
||||
|
||||
p := &Proxy{}
|
||||
conn, err := net.ListenUDP("udp4", &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1), Port: 0})
|
||||
require.NoError(t, err)
|
||||
t.Cleanup(func() { conn.Close() })
|
||||
|
||||
proxyURL := fmt.Sprintf("https://%s/connect-ip", conn.LocalAddr())
|
||||
connChan := make(chan *Conn, 1)
|
||||
mux := http.NewServeMux()
|
||||
mux.HandleFunc("/connect-ip", func(w http.ResponseWriter, r *http.Request) {
|
||||
assert.Equal(t, "Bearer token", r.Header.Get("Authorization"))
|
||||
mreq, err := ParseProxyRequest(r)
|
||||
if !assert.NoError(t, err) {
|
||||
w.WriteHeader(http.StatusBadRequest)
|
||||
return
|
||||
}
|
||||
|
||||
conn, err := p.Proxy(w, mreq)
|
||||
if assert.NoError(t, err) {
|
||||
connChan <- conn
|
||||
}
|
||||
})
|
||||
s := http3.Server{
|
||||
Handler: mux,
|
||||
Addr: ":0",
|
||||
EnableDatagrams: true,
|
||||
TLSConfig: tlsConf,
|
||||
}
|
||||
go func() { s.Serve(conn) }()
|
||||
t.Cleanup(func() { s.Close() })
|
||||
|
||||
ctx, cancel := context.WithTimeout(t.Context(), 5*time.Second)
|
||||
defer cancel()
|
||||
req, err := NewRequest(ctx, proxyURL)
|
||||
require.NoError(t, err)
|
||||
req.Header().Set("Authorization", "Bearer token")
|
||||
client, rsp, err := NewClientConn(dialHTTP3(t, conn.LocalAddr().String())).Dial(req)
|
||||
require.NoError(t, err)
|
||||
t.Cleanup(func() { client.Close() })
|
||||
require.Equal(t, http.StatusOK, rsp.StatusCode)
|
||||
|
||||
select {
|
||||
case <-time.After(5 * time.Second):
|
||||
t.Fatal("timed out")
|
||||
case server = <-connChan:
|
||||
}
|
||||
t.Cleanup(func() { server.Close() })
|
||||
return client, server
|
||||
}
|
||||
|
||||
func TestAddressAssignment(t *testing.T) {
|
||||
client, server := setupConns(t)
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Millisecond)
|
||||
defer cancel()
|
||||
_, err := server.ReceiveAddressAssignment(ctx)
|
||||
require.ErrorIs(t, err, context.DeadlineExceeded)
|
||||
|
||||
ctx, cancel = context.WithTimeout(context.Background(), time.Second)
|
||||
defer cancel()
|
||||
pref1 := netip.MustParsePrefix("1.1.1.0/24")
|
||||
pref2 := netip.MustParsePrefix("2001:db8::/64")
|
||||
require.NoError(t, client.AssignAddresses([]netip.Prefix{pref1, pref2}))
|
||||
assigned, err := server.ReceiveAddressAssignment(ctx)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, []AssignedAddress{{IPPrefix: pref1}, {IPPrefix: pref2}}, assigned)
|
||||
|
||||
require.NoError(t, client.AssignAddresses([]netip.Prefix{}))
|
||||
assigned, err = server.ReceiveAddressAssignment(ctx)
|
||||
require.NoError(t, err)
|
||||
require.Empty(t, assigned)
|
||||
}
|
||||
|
||||
func TestRejectingAddressRequestKeepsPeerUnrestricted(t *testing.T) {
|
||||
client, server := setupConns(t)
|
||||
ctx, cancel := context.WithTimeout(t.Context(), 5*time.Second)
|
||||
defer cancel()
|
||||
|
||||
clientAddr := netip.MustParsePrefix("192.0.2.2/32")
|
||||
require.NoError(t, server.AssignAddresses([]netip.Prefix{clientAddr}))
|
||||
_, err := client.ReceiveAddressAssignment(ctx)
|
||||
require.NoError(t, err)
|
||||
|
||||
_, err = server.RequestAddresses([]netip.Prefix{netip.MustParsePrefix("0.0.0.0/32")})
|
||||
require.NoError(t, err)
|
||||
req, err := client.ReceiveAddressRequest(ctx)
|
||||
require.NoError(t, err)
|
||||
require.NoError(t, req.Respond([]netip.Prefix{{}}, nil))
|
||||
assigned, err := server.ReceiveAddressAssignment(ctx)
|
||||
require.NoError(t, err)
|
||||
require.Len(t, assigned, 1)
|
||||
require.True(t, assigned[0].Rejected())
|
||||
|
||||
packet := ipv4Packet(64, 17, netip.MustParseAddr("203.0.113.9"), clientAddr.Addr(), nil, []byte("foobar"))
|
||||
_, err = server.WritePacket(slices.Clone(packet))
|
||||
require.NoError(t, err)
|
||||
received := make(chan []byte, 1)
|
||||
go func() {
|
||||
b := make([]byte, 1500)
|
||||
if n, err := client.ReadPacket(b); err == nil {
|
||||
received <- b[:n]
|
||||
}
|
||||
}()
|
||||
select {
|
||||
case b := <-received:
|
||||
require.Equal(t, packet[20:], b[20:])
|
||||
case <-ctx.Done():
|
||||
t.Fatal("packet was not received")
|
||||
}
|
||||
}
|
||||
|
||||
func TestRejectingAddressRequestWithdrawsAssignment(t *testing.T) {
|
||||
client, server := setupConns(t)
|
||||
ctx, cancel := context.WithTimeout(t.Context(), 5*time.Second)
|
||||
defer cancel()
|
||||
|
||||
clientAddr := netip.MustParsePrefix("192.0.2.2/32")
|
||||
dst := netip.MustParseAddr("198.51.100.1")
|
||||
require.NoError(t, server.AssignAddresses([]netip.Prefix{clientAddr}))
|
||||
require.NoError(t, server.AdvertiseRoute([]IPRoute{{StartIP: dst, EndIP: dst}}))
|
||||
_, err := client.ReceiveAddressAssignment(ctx)
|
||||
require.NoError(t, err)
|
||||
|
||||
received := make(chan []byte, 4)
|
||||
go func() {
|
||||
b := make([]byte, 1500)
|
||||
for {
|
||||
n, err := server.ReadPacket(b)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
received <- slices.Clone(b[20:n])
|
||||
}
|
||||
}()
|
||||
send := func(payload string) {
|
||||
_, err := client.WritePacket(ipv4Packet(64, 17, clientAddr.Addr(), dst, nil, []byte(payload)))
|
||||
require.NoError(t, err)
|
||||
}
|
||||
|
||||
send("assigned")
|
||||
select {
|
||||
case b := <-received:
|
||||
require.Equal(t, "assigned", string(b))
|
||||
case <-ctx.Done():
|
||||
t.Fatal("packet from the assigned address was not received")
|
||||
}
|
||||
|
||||
_, err = client.RequestAddresses([]netip.Prefix{netip.MustParsePrefix("0.0.0.0/32")})
|
||||
require.NoError(t, err)
|
||||
req, err := server.ReceiveAddressRequest(ctx)
|
||||
require.NoError(t, err)
|
||||
require.NoError(t, req.Respond([]netip.Prefix{{}}, nil))
|
||||
assigned, err := client.ReceiveAddressAssignment(ctx)
|
||||
require.NoError(t, err)
|
||||
require.Len(t, assigned, 1)
|
||||
require.True(t, assigned[0].Rejected())
|
||||
|
||||
send("withdrawn")
|
||||
select {
|
||||
case b := <-received:
|
||||
t.Fatalf("packet from a withdrawn address was received: %q", b)
|
||||
case <-time.After(200 * time.Millisecond):
|
||||
}
|
||||
}
|
||||
|
||||
func TestRouteAdvertisement(t *testing.T) {
|
||||
client, server := setupConns(t)
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Millisecond)
|
||||
defer cancel()
|
||||
_, err := server.Routes(ctx)
|
||||
require.ErrorIs(t, err, context.DeadlineExceeded)
|
||||
|
||||
ctx, cancel = context.WithTimeout(context.Background(), time.Second)
|
||||
defer cancel()
|
||||
|
||||
require.ErrorContains(t,
|
||||
client.AdvertiseRoute([]IPRoute{
|
||||
{StartIP: netip.MustParseAddr("1.1.1.2"), EndIP: netip.MustParseAddr("1.1.1.1"), IPProtocol: 42},
|
||||
}),
|
||||
"connect-ip: invalid route 0: start IP 1.1.1.2 is greater than end IP 1.1.1.1",
|
||||
)
|
||||
|
||||
require.NoError(t, client.AdvertiseRoute([]IPRoute{
|
||||
{StartIP: netip.MustParseAddr("1.1.1.1"), EndIP: netip.MustParseAddr("2.2.2.2"), IPProtocol: 42},
|
||||
{StartIP: netip.MustParseAddr("2001:db8::1"), EndIP: netip.MustParseAddr("2001:db8::100"), IPProtocol: 24},
|
||||
}))
|
||||
routes, err := server.Routes(ctx)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, []IPRoute{
|
||||
{StartIP: netip.MustParseAddr("1.1.1.1"), EndIP: netip.MustParseAddr("2.2.2.2"), IPProtocol: 42},
|
||||
{StartIP: netip.MustParseAddr("2001:db8::1"), EndIP: netip.MustParseAddr("2001:db8::100"), IPProtocol: 24},
|
||||
}, routes)
|
||||
|
||||
require.NoError(t, client.AdvertiseRoute([]IPRoute{}))
|
||||
routes, err = server.Routes(ctx)
|
||||
require.NoError(t, err)
|
||||
require.Empty(t, routes)
|
||||
}
|
||||
|
||||
func TestTTLs(t *testing.T) {
|
||||
t.Run("IPv4", func(t *testing.T) {
|
||||
client, server := setupConns(t)
|
||||
require.NoError(t, server.AssignAddresses([]netip.Prefix{netip.MustParsePrefix("192.168.1.1/32")}))
|
||||
require.NoError(t, server.AdvertiseRoute([]IPRoute{
|
||||
{StartIP: netip.MustParseAddr("0.0.0.0"), EndIP: netip.MustParseAddr("255.255.255.255")},
|
||||
}))
|
||||
|
||||
src, dst := netip.MustParseAddr("192.168.1.1"), netip.MustParseAddr("8.8.8.8")
|
||||
icmp, err := client.WritePacket(ipv4Packet(1, 0, src, dst, nil, nil))
|
||||
require.NoError(t, err)
|
||||
require.Empty(t, icmp)
|
||||
|
||||
icmp, err = client.WritePacket(ipv4Packet(42, 0, src, dst, nil, nil))
|
||||
require.NoError(t, err)
|
||||
require.Empty(t, icmp)
|
||||
|
||||
receivedPacket := make([]byte, 1500)
|
||||
n, err := server.ReadPacket(receivedPacket)
|
||||
require.NoError(t, err)
|
||||
receivedPacket = receivedPacket[:n]
|
||||
|
||||
receivedHdr, err := ipv4.ParseHeader(receivedPacket)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, uint16(receivedHdr.Checksum), calculateIPv4Checksum(receivedPacket[:ipv4.HeaderLen]))
|
||||
require.Equal(t, 41, receivedHdr.TTL)
|
||||
})
|
||||
|
||||
t.Run("IPv6", func(t *testing.T) {
|
||||
client, server := setupConns(t)
|
||||
require.NoError(t, server.AssignAddresses([]netip.Prefix{netip.MustParsePrefix("2001:db8::1/128")}))
|
||||
require.NoError(t, server.AdvertiseRoute([]IPRoute{
|
||||
{StartIP: netip.MustParseAddr("::"), EndIP: netip.MustParseAddr("ffff:ffff:ffff:ffff:ffff:ffff:ffff:ffff")},
|
||||
}))
|
||||
|
||||
packetHopLimit1 := []byte{
|
||||
0x60, 0x00, 0x00, 0x00,
|
||||
0x00, 0x00,
|
||||
0x00, 0x01,
|
||||
0x20, 0x01, 0x0d, 0xb8, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x01,
|
||||
0x20, 0x01, 0x48, 0x60, 0x48, 0x60, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x88, 0x88,
|
||||
}
|
||||
icmp, err := client.WritePacket(packetHopLimit1)
|
||||
require.NoError(t, err)
|
||||
require.Empty(t, icmp)
|
||||
|
||||
packet := []byte{
|
||||
0x60, 0x00, 0x00, 0x00,
|
||||
0x00, 0x00,
|
||||
0x00, 0x2A,
|
||||
0x20, 0x01, 0x0d, 0xb8, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x01,
|
||||
0x20, 0x01, 0x48, 0x60, 0x48, 0x60, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x88, 0x88,
|
||||
}
|
||||
icmp, err = client.WritePacket(packet)
|
||||
require.NoError(t, err)
|
||||
require.Empty(t, icmp)
|
||||
|
||||
receivedPacket := make([]byte, 1500)
|
||||
n, err := server.ReadPacket(receivedPacket)
|
||||
require.NoError(t, err)
|
||||
receivedPacket = receivedPacket[:n]
|
||||
|
||||
receivedHdr, err := ipv6.ParseHeader(receivedPacket)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, 41, receivedHdr.HopLimit)
|
||||
})
|
||||
}
|
||||
|
||||
func TestMaxPacketSizeOverQUIC(t *testing.T) {
|
||||
client, server := setupConns(t)
|
||||
require.NoError(t, server.AdvertiseRoute([]IPRoute{
|
||||
{StartIP: netip.MustParseAddr("0.0.0.0"), EndIP: netip.MustParseAddr("255.255.255.255")},
|
||||
}))
|
||||
|
||||
size := client.MaxPacketSize()
|
||||
require.Greater(t, size, 1200)
|
||||
require.Less(t, size, 1350)
|
||||
|
||||
icmp, err := client.WritePacket(ipv4Packet(64, 17, testSrc4, testDst4, nil, make([]byte, size-ipv4.HeaderLen)))
|
||||
require.NoError(t, err)
|
||||
require.Nil(t, icmp)
|
||||
type readResult struct {
|
||||
n int
|
||||
err error
|
||||
}
|
||||
received := make(chan readResult, 1)
|
||||
go func() {
|
||||
n, err := server.ReadPacket(make([]byte, 1500))
|
||||
received <- readResult{n, err}
|
||||
}()
|
||||
select {
|
||||
case r := <-received:
|
||||
require.NoError(t, r.err)
|
||||
require.Equal(t, size, r.n)
|
||||
case <-time.After(5 * time.Second):
|
||||
t.Fatal("timeout")
|
||||
}
|
||||
|
||||
icmp, err = client.WritePacket(ipv4Packet(64, 17, testSrc4, testDst4, nil, make([]byte, size+1-ipv4.HeaderLen)))
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, icmp)
|
||||
require.Equal(t, uint16(size), binary.BigEndian.Uint16(icmp[ipv4.HeaderLen+6:]))
|
||||
}
|
||||
|
||||
func TestClosing(t *testing.T) {
|
||||
ipv6Packet := []byte{
|
||||
0x60, 0x00, 0x00, 0x00,
|
||||
0x00, 0x00,
|
||||
0x00, 0x2A,
|
||||
0x20, 0x01, 0x0d, 0xb8, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x01,
|
||||
0x20, 0x01, 0x48, 0x60, 0x48, 0x60, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x88, 0x88,
|
||||
}
|
||||
|
||||
client, server := setupConns(t)
|
||||
routeErrChan := make(chan error, 1)
|
||||
prefixErrChan := make(chan error, 1)
|
||||
go func() {
|
||||
_, err := server.Routes(context.Background())
|
||||
routeErrChan <- err
|
||||
}()
|
||||
go func() {
|
||||
_, err := server.ReceiveAddressAssignment(context.Background())
|
||||
prefixErrChan <- err
|
||||
}()
|
||||
|
||||
require.NoError(t, client.Close())
|
||||
_, err := client.ReceiveAddressAssignment(context.Background())
|
||||
require.ErrorIs(t, err, net.ErrClosed)
|
||||
var closeErr *CloseError
|
||||
require.ErrorAs(t, err, &closeErr)
|
||||
require.False(t, closeErr.Remote)
|
||||
_, err = client.Routes(context.Background())
|
||||
require.ErrorIs(t, err, net.ErrClosed)
|
||||
require.ErrorIs(t,
|
||||
client.AssignAddresses([]netip.Prefix{netip.MustParsePrefix("1.1.1.0/24")}),
|
||||
net.ErrClosed,
|
||||
)
|
||||
require.ErrorIs(t,
|
||||
client.AdvertiseRoute([]IPRoute{
|
||||
{StartIP: netip.MustParseAddr("1.1.1.0"), EndIP: netip.MustParseAddr("1.1.1.1"), IPProtocol: 42},
|
||||
}),
|
||||
net.ErrClosed,
|
||||
)
|
||||
_, err = client.ReadPacket([]byte{0})
|
||||
require.ErrorIs(t, err, net.ErrClosed)
|
||||
_, err = client.WritePacket(ipv6Packet)
|
||||
require.ErrorIs(t, err, net.ErrClosed)
|
||||
|
||||
select {
|
||||
case err := <-routeErrChan:
|
||||
require.ErrorIs(t, err, net.ErrClosed)
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("timeout")
|
||||
}
|
||||
|
||||
select {
|
||||
case err := <-prefixErrChan:
|
||||
require.ErrorIs(t, err, net.ErrClosed)
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("timeout")
|
||||
}
|
||||
|
||||
_, err = server.ReadPacket([]byte{0})
|
||||
require.ErrorIs(t, err, net.ErrClosed)
|
||||
_, err = server.WritePacket(ipv6Packet)
|
||||
require.ErrorIs(t, err, net.ErrClosed)
|
||||
}
|
||||
@@ -0,0 +1,91 @@
|
||||
/* SPDX-License-Identifier: MIT
|
||||
*
|
||||
* Copyright 2024 Marten Seemann
|
||||
* Adapted from github.com/quic-go/connect-ip-go (commit a0c35fa).
|
||||
*/
|
||||
|
||||
package connectip
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"strings"
|
||||
|
||||
"github.com/apernet/quic-go/http3"
|
||||
)
|
||||
|
||||
const requestProtocol = "connect-ip"
|
||||
|
||||
const capsuleProtocolHeaderValue = "?1"
|
||||
|
||||
type Request struct {
|
||||
req *http.Request
|
||||
}
|
||||
|
||||
func NewRequest(ctx context.Context, rawURL string) (*Request, error) {
|
||||
if strings.ContainsAny(rawURL, "{}") {
|
||||
return nil, errors.New("connect-ip: IP flow forwarding not supported: URL contains a URI Template expression")
|
||||
}
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodConnect, rawURL, nil)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("connect-ip: failed to create request: %w", err)
|
||||
}
|
||||
if req.URL.Scheme != "https" || req.URL.Host == "" || !strings.HasPrefix(req.URL.Path, "/") {
|
||||
return nil, fmt.Errorf("connect-ip: invalid proxy URL %q: expected an absolute https URL with a host and a path", rawURL)
|
||||
}
|
||||
req.Proto = requestProtocol
|
||||
req.Host = req.URL.Host
|
||||
req.Header.Set(http3.CapsuleProtocolHeader, capsuleProtocolHeaderValue)
|
||||
return &Request{req: req}, nil
|
||||
}
|
||||
|
||||
func (r *Request) Header() http.Header { return r.req.Header }
|
||||
|
||||
func (r *Request) httpRequest() *http.Request { return r.req }
|
||||
|
||||
type ProxyRequest struct{}
|
||||
|
||||
type ProxyRequestParseError struct {
|
||||
HTTPStatus int
|
||||
Err error
|
||||
}
|
||||
|
||||
func (e *ProxyRequestParseError) Error() string { return e.Err.Error() }
|
||||
func (e *ProxyRequestParseError) Unwrap() error { return e.Err }
|
||||
|
||||
func ParseProxyRequest(r *http.Request) (*ProxyRequest, error) {
|
||||
if r.Method != http.MethodConnect {
|
||||
return nil, &ProxyRequestParseError{
|
||||
HTTPStatus: http.StatusMethodNotAllowed,
|
||||
Err: fmt.Errorf("expected CONNECT request, got %s", r.Method),
|
||||
}
|
||||
}
|
||||
if r.Proto != requestProtocol {
|
||||
return nil, &ProxyRequestParseError{
|
||||
HTTPStatus: http.StatusNotImplemented,
|
||||
Err: fmt.Errorf("unexpected protocol: %s", r.Proto),
|
||||
}
|
||||
}
|
||||
capsuleHeaderValues, ok := r.Header[http3.CapsuleProtocolHeader]
|
||||
if !ok {
|
||||
return nil, &ProxyRequestParseError{
|
||||
HTTPStatus: http.StatusBadRequest,
|
||||
Err: fmt.Errorf("missing Capsule-Protocol header"),
|
||||
}
|
||||
}
|
||||
if !isCapsuleProtocolEnabled(capsuleHeaderValues) {
|
||||
return nil, &ProxyRequestParseError{
|
||||
HTTPStatus: http.StatusBadRequest,
|
||||
Err: fmt.Errorf("invalid capsule header value: %s", capsuleHeaderValues),
|
||||
}
|
||||
}
|
||||
|
||||
return &ProxyRequest{}, nil
|
||||
}
|
||||
|
||||
func isCapsuleProtocolEnabled(values []string) bool {
|
||||
v := strings.Trim(strings.Join(values, ","), " ")
|
||||
return v == capsuleProtocolHeaderValue || strings.HasPrefix(v, capsuleProtocolHeaderValue+";")
|
||||
}
|
||||
@@ -0,0 +1,119 @@
|
||||
/* SPDX-License-Identifier: MIT
|
||||
*
|
||||
* Copyright 2024 Marten Seemann
|
||||
* Adapted from github.com/quic-go/connect-ip-go (commit a0c35fa).
|
||||
*/
|
||||
|
||||
package connectip
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
|
||||
"github.com/apernet/quic-go/http3"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func newRequest(target string) *http.Request {
|
||||
req := httptest.NewRequest(http.MethodGet, target, nil)
|
||||
req.Method = http.MethodConnect
|
||||
req.Proto = requestProtocol
|
||||
req.Header.Add("Capsule-Protocol", capsuleProtocolHeaderValue)
|
||||
return req
|
||||
}
|
||||
|
||||
func TestNewRequest(t *testing.T) {
|
||||
req, err := NewRequest(t.Context(), "https://localhost:1234/masque/ip")
|
||||
require.NoError(t, err)
|
||||
httpReq := req.httpRequest()
|
||||
require.Equal(t, http.MethodConnect, httpReq.Method)
|
||||
require.Equal(t, requestProtocol, httpReq.Proto)
|
||||
require.Equal(t, "localhost:1234", httpReq.Host)
|
||||
require.Equal(t, "?1", req.Header().Get(http3.CapsuleProtocolHeader))
|
||||
|
||||
req.Header().Set("Authorization", "Bearer token")
|
||||
require.Equal(t, "Bearer token", httpReq.Header.Get("Authorization"))
|
||||
}
|
||||
|
||||
func TestNewRequestInvalidURL(t *testing.T) {
|
||||
for _, tc := range []struct {
|
||||
name, url, err string
|
||||
}{
|
||||
{"template with variables", "https://localhost/.well-known/masque/ip/{target}/{ipproto}/", "IP flow forwarding not supported"},
|
||||
{"template with query variables", "https://localhost/masque/ip{?target,ipproto}", "IP flow forwarding not supported"},
|
||||
{"not https", "http://localhost/masque/ip", "expected an absolute https URL"},
|
||||
{"no host", "https:///masque/ip", "expected an absolute https URL"},
|
||||
{"no path", "https://localhost", "expected an absolute https URL"},
|
||||
{"relative", "/masque/ip", "expected an absolute https URL"},
|
||||
{"unparsable", "https://local\x7fhost/", "failed to create request"},
|
||||
} {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
_, err := NewRequest(t.Context(), tc.url)
|
||||
require.ErrorContains(t, err, tc.err)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestProxyRequestParsing(t *testing.T) {
|
||||
t.Run("valid request", func(t *testing.T) {
|
||||
req := newRequest("https://localhost:1234/masque/ip")
|
||||
r, err := ParseProxyRequest(req)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, &ProxyRequest{}, r)
|
||||
})
|
||||
|
||||
t.Run("wrong protocol", func(t *testing.T) {
|
||||
req := newRequest("https://localhost:1234/masque")
|
||||
req.Proto = "not-connect-ip"
|
||||
_, err := ParseProxyRequest(req)
|
||||
require.EqualError(t, err, "unexpected protocol: not-connect-ip")
|
||||
require.Equal(t, http.StatusNotImplemented, err.(*ProxyRequestParseError).HTTPStatus)
|
||||
})
|
||||
|
||||
t.Run("wrong request method", func(t *testing.T) {
|
||||
req := newRequest("https://localhost:1234/masque")
|
||||
req.Method = http.MethodHead
|
||||
_, err := ParseProxyRequest(req)
|
||||
require.EqualError(t, err, "expected CONNECT request, got HEAD")
|
||||
require.Equal(t, http.StatusMethodNotAllowed, err.(*ProxyRequestParseError).HTTPStatus)
|
||||
})
|
||||
|
||||
t.Run("missing Capsule-Protocol header", func(t *testing.T) {
|
||||
req := newRequest("https://localhost:1234/masque")
|
||||
req.Header.Del("Capsule-Protocol")
|
||||
_, err := ParseProxyRequest(req)
|
||||
require.EqualError(t, err, "missing Capsule-Protocol header")
|
||||
require.Equal(t, http.StatusBadRequest, err.(*ProxyRequestParseError).HTTPStatus)
|
||||
})
|
||||
|
||||
for _, tc := range []struct {
|
||||
name string
|
||||
values []string
|
||||
valid bool
|
||||
}{
|
||||
{name: "true", values: []string{"?1"}, valid: true},
|
||||
{name: "surrounding spaces", values: []string{" ?1 "}, valid: true},
|
||||
{name: "parameters", values: []string{"?1;a;b=?0;c=\"x\""}, valid: true},
|
||||
{name: "false", values: []string{"?0"}},
|
||||
{name: "integer", values: []string{"1"}},
|
||||
{name: "empty", values: []string{""}},
|
||||
{name: "not a structured field", values: []string{"🤡"}},
|
||||
{name: "longer token", values: []string{"?10"}},
|
||||
{name: "space before parameters", values: []string{"?1 ;a"}},
|
||||
{name: "list", values: []string{"?1, ?1"}},
|
||||
{name: "multiple field lines", values: []string{"?1", "?1"}},
|
||||
} {
|
||||
t.Run("Capsule-Protocol header: "+tc.name, func(t *testing.T) {
|
||||
req := newRequest("https://localhost:1234/masque")
|
||||
req.Header[http3.CapsuleProtocolHeader] = tc.values
|
||||
_, err := ParseProxyRequest(req)
|
||||
if tc.valid {
|
||||
require.NoError(t, err)
|
||||
return
|
||||
}
|
||||
require.ErrorContains(t, err, "invalid capsule header value")
|
||||
require.Equal(t, http.StatusBadRequest, err.(*ProxyRequestParseError).HTTPStatus)
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,237 @@
|
||||
package masque
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net/netip"
|
||||
"reflect"
|
||||
"runtime"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/apernet/quic-go"
|
||||
"github.com/apernet/quic-go/http3"
|
||||
"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/common/net/cnc"
|
||||
"github.com/xtls/xray-core/common/utils"
|
||||
"github.com/xtls/xray-core/transport/internet"
|
||||
"github.com/xtls/xray-core/transport/internet/finalmask"
|
||||
"github.com/xtls/xray-core/transport/internet/hysteria/congestion"
|
||||
"github.com/xtls/xray-core/transport/internet/hysteria/congestion/bbr"
|
||||
"github.com/xtls/xray-core/transport/internet/masque/connectip"
|
||||
"github.com/xtls/xray-core/transport/internet/stat"
|
||||
"github.com/xtls/xray-core/transport/internet/tls"
|
||||
)
|
||||
|
||||
const (
|
||||
MinPacketSize = 1280
|
||||
initialPacketSize = 1350
|
||||
)
|
||||
|
||||
func Dial(ctx context.Context, dest net.Destination, streamSettings *internet.MemoryStreamConfig) (stat.Connection, error) {
|
||||
tlsConfig := tls.ConfigFromStreamSettings(streamSettings)
|
||||
if tlsConfig == nil {
|
||||
return nil, errors.New("tls config is nil")
|
||||
}
|
||||
config := streamSettings.ProtocolSettings.(*Config)
|
||||
dest.Network = net.Network_UDP
|
||||
|
||||
gotlsConfig := tlsConfig.GetTLSConfig(tls.WithDestination(dest))
|
||||
gotlsConfig.NextProtos = []string{http3.NextProtoH3}
|
||||
|
||||
quicParams := streamSettings.QuicParams
|
||||
if quicParams == nil {
|
||||
quicParams = &internet.QuicParams{
|
||||
BbrProfile: string(bbr.ProfileStandard),
|
||||
}
|
||||
}
|
||||
quicConfig := &quic.Config{
|
||||
InitialStreamReceiveWindow: quicParams.InitStreamReceiveWindow,
|
||||
MaxStreamReceiveWindow: quicParams.MaxStreamReceiveWindow,
|
||||
InitialConnectionReceiveWindow: quicParams.InitConnReceiveWindow,
|
||||
MaxConnectionReceiveWindow: quicParams.MaxConnReceiveWindow,
|
||||
MaxIdleTimeout: time.Duration(quicParams.MaxIdleTimeout) * time.Second,
|
||||
KeepAlivePeriod: time.Duration(quicParams.KeepAlivePeriod) * time.Second,
|
||||
MaxIncomingStreams: -1,
|
||||
InitialPacketSize: initialPacketSize,
|
||||
DisablePathMTUDiscovery: quicParams.DisablePathMtuDiscovery || (runtime.GOOS != "linux" && runtime.GOOS != "windows" && runtime.GOOS != "darwin"),
|
||||
EnableDatagrams: true,
|
||||
DisablePathManager: true,
|
||||
}
|
||||
if quicParams.MaxIdleTimeout == 0 {
|
||||
quicConfig.MaxIdleTimeout = 30 * time.Second
|
||||
}
|
||||
if quicParams.KeepAlivePeriod == 0 {
|
||||
quicConfig.KeepAlivePeriod = net.QuicgoH3KeepAlivePeriod
|
||||
}
|
||||
|
||||
var pktConn net.PacketConn
|
||||
var udpAddr net.Addr
|
||||
if streamSettings.FinalMask != nil {
|
||||
conn, err := streamSettings.FinalMask.DialUDP(ctx, dest)
|
||||
if err != nil {
|
||||
return nil, errors.New("failed to dial to dest").Base(err)
|
||||
}
|
||||
pktConn = conn.(*finalmask.PacketConnWrapper).PacketConn
|
||||
udpAddr = conn.RemoteAddr()
|
||||
} else {
|
||||
conn, err := internet.DialSystem(ctx, dest, 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
|
||||
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))
|
||||
}
|
||||
}
|
||||
|
||||
tr := &quic.Transport{Conn: pktConn, DisableGSO: quicParams.DisableGSO}
|
||||
qconn, err := tr.Dial(ctx, udpAddr, gotlsConfig, quicConfig)
|
||||
if err != nil {
|
||||
tr.Close()
|
||||
pktConn.Close()
|
||||
return nil, err
|
||||
}
|
||||
context.AfterFunc(qconn.Context(), func() { tr.Close(); pktConn.Close() })
|
||||
|
||||
switch quicParams.Congestion {
|
||||
case "reno":
|
||||
case "", "bbr", "brutal":
|
||||
congestion.UseBBR(qconn, bbr.Profile(quicParams.BbrProfile))
|
||||
case "force-brutal":
|
||||
congestion.UseBrutal(qconn, quicParams.BrutalUp, quicParams.BrutalDisableLossCompensation)
|
||||
default:
|
||||
qconn.CloseWithError(quic.ApplicationErrorCode(http3.ErrCodeNoError), "")
|
||||
return nil, errors.New("unknown congestion control: ", quicParams.Congestion)
|
||||
}
|
||||
|
||||
conn, err := establish(ctx, qconn, config, authority(config, gotlsConfig.ServerName, dest.Port))
|
||||
if err != nil {
|
||||
qconn.CloseWithError(quic.ApplicationErrorCode(http3.ErrCodeNoError), "")
|
||||
return nil, err
|
||||
}
|
||||
return conn, nil
|
||||
}
|
||||
|
||||
func establish(ctx context.Context, qconn *quic.Conn, config *Config, host string) (*Conn, error) {
|
||||
stop := context.AfterFunc(ctx, func() {
|
||||
qconn.CloseWithError(quic.ApplicationErrorCode(http3.ErrCodeRequestCanceled), "")
|
||||
})
|
||||
defer stop()
|
||||
|
||||
req, err := connectip.NewRequest(ctx, "https://"+host+config.Path)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
header := req.Header()
|
||||
for k, v := range config.Headers {
|
||||
header.Set(k, v)
|
||||
}
|
||||
switch header.Get("User-Agent") {
|
||||
case "":
|
||||
header["User-Agent"] = nil
|
||||
case "chrome":
|
||||
header.Set("User-Agent", utils.ChromeUA)
|
||||
case "firefox":
|
||||
header.Set("User-Agent", utils.FirefoxUA)
|
||||
case "safari":
|
||||
header.Set("User-Agent", utils.SafariUA)
|
||||
case "edge":
|
||||
header.Set("User-Agent", utils.MSEdgeUA)
|
||||
case "curl":
|
||||
header.Set("User-Agent", utils.CurlUA)
|
||||
case "golang":
|
||||
header.Del("User-Agent")
|
||||
}
|
||||
|
||||
cc := (&http3.Transport{EnableDatagrams: true, DisableCompression: true}).NewClientConn(qconn)
|
||||
ipConn, _, err := connectip.NewClientConn(cc).Dial(req)
|
||||
if err != nil {
|
||||
if ctx.Err() != nil {
|
||||
err = context.Cause(ctx)
|
||||
}
|
||||
return nil, errors.New("CONNECT-IP request failed").Base(err)
|
||||
}
|
||||
|
||||
if n := ipConn.MaxPacketSize(); n < MinPacketSize {
|
||||
ipConn.Close()
|
||||
return nil, errors.New("the tunnel can only carry ", n, "-byte packets, less than ", MinPacketSize)
|
||||
}
|
||||
|
||||
if _, err := ipConn.RequestAddresses([]netip.Prefix{
|
||||
netip.PrefixFrom(netip.IPv4Unspecified(), 32),
|
||||
netip.PrefixFrom(netip.IPv6Unspecified(), 128),
|
||||
}); err != nil {
|
||||
ipConn.Close()
|
||||
return nil, err
|
||||
}
|
||||
var local []netip.Addr
|
||||
for len(local) == 0 {
|
||||
assigned, err := ipConn.ReceiveAddressAssignment(ctx)
|
||||
if err != nil {
|
||||
ipConn.Close()
|
||||
return nil, errors.New("no address assigned").Base(err)
|
||||
}
|
||||
local = localAddrs(assigned)
|
||||
}
|
||||
if !stop() {
|
||||
ipConn.Close()
|
||||
return nil, errors.New("no address assigned").Base(context.Cause(ctx))
|
||||
}
|
||||
|
||||
conn := &Conn{
|
||||
ipConn: ipConn,
|
||||
quicConn: qconn,
|
||||
local: local,
|
||||
}
|
||||
go conn.serveAddressAssignments()
|
||||
go conn.serveAddressRequests()
|
||||
return conn, nil
|
||||
}
|
||||
|
||||
func localAddrs(assigned []connectip.AssignedAddress) []netip.Addr {
|
||||
var local []netip.Addr
|
||||
var has4, has6 bool
|
||||
for _, a := range assigned {
|
||||
if a.Rejected() {
|
||||
continue
|
||||
}
|
||||
addr := a.IPPrefix.Addr()
|
||||
if a.IPPrefix.Bits() != addr.BitLen() {
|
||||
addr = a.IPPrefix.Masked().Addr().Next()
|
||||
}
|
||||
if addr.Is4() && !has4 {
|
||||
has4 = true
|
||||
local = append(local, addr)
|
||||
} else if addr.Is6() && !has6 {
|
||||
has6 = true
|
||||
local = append(local, addr)
|
||||
}
|
||||
}
|
||||
return local
|
||||
}
|
||||
|
||||
func authority(config *Config, serverName string, port net.Port) string {
|
||||
if config.Host != "" {
|
||||
return config.Host
|
||||
}
|
||||
host := strings.TrimSuffix(strings.TrimPrefix(serverName, "["), "]")
|
||||
if port == 443 {
|
||||
if addr, err := netip.ParseAddr(host); err == nil && addr.Is6() {
|
||||
return "[" + host + "]"
|
||||
}
|
||||
return host
|
||||
}
|
||||
return net.JoinHostPort(host, port.String())
|
||||
}
|
||||
|
||||
func init() {
|
||||
common.Must(internet.RegisterTransportDialer(protocolName, Dial))
|
||||
}
|
||||
@@ -0,0 +1,60 @@
|
||||
package masque
|
||||
|
||||
import (
|
||||
"net/netip"
|
||||
"slices"
|
||||
"testing"
|
||||
|
||||
"github.com/xtls/xray-core/common/net"
|
||||
"github.com/xtls/xray-core/transport/internet/masque/connectip"
|
||||
)
|
||||
|
||||
func TestAuthority(t *testing.T) {
|
||||
for _, c := range []struct {
|
||||
host, serverName string
|
||||
port net.Port
|
||||
want string
|
||||
}{
|
||||
{serverName: "example.com", port: 443, want: "example.com"},
|
||||
{serverName: "example.com", port: 8443, want: "example.com:8443"},
|
||||
{serverName: "127.0.0.1", port: 443, want: "127.0.0.1"},
|
||||
{serverName: "[2001:db8::1]", port: 443, want: "[2001:db8::1]"},
|
||||
{serverName: "[2001:db8::1]", port: 8443, want: "[2001:db8::1]:8443"},
|
||||
{serverName: "2001:db8::1", port: 8443, want: "[2001:db8::1]:8443"},
|
||||
{host: "proxy.example", serverName: "example.com", port: 8443, want: "proxy.example"},
|
||||
} {
|
||||
if got := authority(&Config{Host: c.host}, c.serverName, c.port); got != c.want {
|
||||
t.Errorf("authority(%q, %q, %d) = %q, want %q", c.host, c.serverName, c.port, got, c.want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestLocalAddrs(t *testing.T) {
|
||||
assigned := func(prefixes ...string) []connectip.AssignedAddress {
|
||||
var a []connectip.AssignedAddress
|
||||
for _, p := range prefixes {
|
||||
a = append(a, connectip.AssignedAddress{IPPrefix: netip.MustParsePrefix(p)})
|
||||
}
|
||||
return a
|
||||
}
|
||||
addrs := func(s ...string) []netip.Addr {
|
||||
var a []netip.Addr
|
||||
for _, v := range s {
|
||||
a = append(a, netip.MustParseAddr(v))
|
||||
}
|
||||
return a
|
||||
}
|
||||
for _, c := range []struct {
|
||||
assigned []connectip.AssignedAddress
|
||||
want []netip.Addr
|
||||
}{
|
||||
{assigned("192.0.2.2/32", "2001:db8::2/128"), addrs("192.0.2.2", "2001:db8::2")},
|
||||
{assigned("2001:db8::/64", "192.0.2.0/24", "198.51.100.7/32"), addrs("2001:db8::1", "192.0.2.1")},
|
||||
{assigned("0.0.0.0/32", "2001:db8::2/128"), addrs("2001:db8::2")},
|
||||
{assigned("0.0.0.0/32", "::/128"), nil},
|
||||
} {
|
||||
if got := localAddrs(c.assigned); !slices.Equal(got, c.want) {
|
||||
t.Errorf("localAddrs(%v) = %v, want %v", c.assigned, got, c.want)
|
||||
}
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user