mirror of
https://github.com/XTLS/Xray-core.git
synced 2026-09-28 09:58:06 +03:00
Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
e5e85ca9da | ||
|
|
7780db9bbe | ||
|
|
2953d44734 | ||
|
|
7b8ade3ec5 | ||
|
|
5dda894e29 | ||
|
|
47a2c2ffdc | ||
|
|
7a018833ec | ||
|
|
65e853ed84 | ||
|
|
3519dfecbd | ||
|
|
df261e4479 | ||
|
|
61cad5ec8b | ||
|
|
a642a190ed | ||
|
|
60e2a0c502 | ||
|
|
7d3e44fee2 | ||
|
|
9927942aaa | ||
|
|
a308ded2e6 | ||
|
|
7741e9e77e | ||
|
|
d562d8947d | ||
|
|
dbb1ea30ba | ||
|
|
efc9e6da62 | ||
|
|
8267cf953a | ||
|
|
24e6f6d551 | ||
|
|
dcdfc57ccd | ||
|
|
3461c511aa | ||
|
|
c412e77a9b | ||
|
|
ccb69ea5e2 |
@@ -23,7 +23,7 @@ func newFakeDNSSniffer(ctx context.Context) (protocolSnifferWithMetadata, error)
|
|||||||
}
|
}
|
||||||
|
|
||||||
if fakeDNSEngine == nil {
|
if fakeDNSEngine == nil {
|
||||||
errNotInit := errors.New("FakeDNSEngine is not initialized, but such a sniffer is used").AtError()
|
errNotInit := errors.New("FakeDNSEngine is not initialized, but such a sniffer is used")
|
||||||
return protocolSnifferWithMetadata{}, errNotInit
|
return protocolSnifferWithMetadata{}, errNotInit
|
||||||
}
|
}
|
||||||
return protocolSnifferWithMetadata{protocolSniffer: func(ctx context.Context, bytes []byte) (SniffResult, error) {
|
return protocolSnifferWithMetadata{protocolSniffer: func(ctx context.Context, bytes []byte) (SniffResult, error) {
|
||||||
|
|||||||
+1
-1
@@ -28,7 +28,7 @@ func toNetIP(addrs []net.Address) ([]net.IP, error) {
|
|||||||
if addr.Family().IsIP() {
|
if addr.Family().IsIP() {
|
||||||
ips = append(ips, addr.IP())
|
ips = append(ips, addr.IP())
|
||||||
} else {
|
} else {
|
||||||
return nil, errors.New("Failed to convert address", addr, "to Net IP.").AtWarning()
|
return nil, errors.New("Failed to convert address", addr, "to Net IP.")
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
return ips, nil
|
return ips, nil
|
||||||
|
|||||||
@@ -212,6 +212,28 @@ func (s *DNS) IsOwnLink(ctx context.Context) bool {
|
|||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// MayUseSystemResolver reports whether any name server configured here could
|
||||||
|
// still resolve through the system resolver. That is what happens when no name
|
||||||
|
// server is configured at all, and it is also what a name server pointed at
|
||||||
|
// "localhost" does. Callers that are about to redirect the system resolver need
|
||||||
|
// to know, because a resolution path that reaches it would then loop back to
|
||||||
|
// them.
|
||||||
|
//
|
||||||
|
// Any such server is enough: name servers can be selected per domain, so a
|
||||||
|
// single local one makes some query reach the system resolver even when
|
||||||
|
// independent upstreams are configured alongside it.
|
||||||
|
func (s *DNS) MayUseSystemResolver() bool {
|
||||||
|
if len(s.clients) == 0 {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
for _, client := range s.clients {
|
||||||
|
if _, isLocal := client.server.(*LocalNameServer); isLocal {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
// LookupIP implements dns.Client.
|
// LookupIP implements dns.Client.
|
||||||
func (s *DNS) LookupIP(domain string, option dns.IPOption) ([]net.IP, uint32, error) {
|
func (s *DNS) LookupIP(domain string, option dns.IPOption) ([]net.IP, uint32, error) {
|
||||||
// Normalize the FQDN form query
|
// Normalize the FQDN form query
|
||||||
|
|||||||
@@ -0,0 +1,59 @@
|
|||||||
|
package dns
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/xtls/xray-core/common/net"
|
||||||
|
feature_dns "github.com/xtls/xray-core/features/dns"
|
||||||
|
)
|
||||||
|
|
||||||
|
// fakeServer stands in for any name server that is not the system resolver.
|
||||||
|
type fakeServer struct{}
|
||||||
|
|
||||||
|
func (fakeServer) Name() string { return "fake" }
|
||||||
|
func (fakeServer) IsDisableCache() bool { return false }
|
||||||
|
func (fakeServer) QueryIP(context.Context, string, feature_dns.IPOption) ([]net.IP, uint32, error) {
|
||||||
|
return nil, 0, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Callers that are about to redirect the system resolver rely on this to tell
|
||||||
|
// whether any resolution path could still reach the system resolver, so the
|
||||||
|
// mixed shape has to be reported as reachable: a domain-specific rule can
|
||||||
|
// select the system resolver even when an independent upstream also exists.
|
||||||
|
func TestMayUseSystemResolver(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
clients []*Client
|
||||||
|
want bool
|
||||||
|
}{
|
||||||
|
{
|
||||||
|
name: "no clients at all",
|
||||||
|
want: true,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "only the system resolver",
|
||||||
|
clients: []*Client{{server: NewLocalNameServer()}},
|
||||||
|
want: true,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "the system resolver alongside an independent name server",
|
||||||
|
clients: []*Client{{server: fakeServer{}}, {server: NewLocalNameServer()}},
|
||||||
|
want: true,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "only independent name servers",
|
||||||
|
clients: []*Client{{server: fakeServer{}}},
|
||||||
|
want: false,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
server := &DNS{clients: tt.clients}
|
||||||
|
if got := server.MayUseSystemResolver(); got != tt.want {
|
||||||
|
t.Errorf("MayUseSystemResolver() = %v, want %v", got, tt.want)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -188,10 +188,10 @@ func parseResponse(payload []byte) (*IPRecord, error) {
|
|||||||
var parser dnsmessage.Parser
|
var parser dnsmessage.Parser
|
||||||
h, err := parser.Start(payload)
|
h, err := parser.Start(payload)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, errors.New("failed to parse DNS response").Base(err).AtWarning()
|
return nil, errors.New("failed to parse DNS response").Base(err)
|
||||||
}
|
}
|
||||||
if err := parser.SkipAllQuestions(); err != nil {
|
if err := parser.SkipAllQuestions(); err != nil {
|
||||||
return nil, errors.New("failed to skip questions in DNS response").Base(err).AtWarning()
|
return nil, errors.New("failed to skip questions in DNS response").Base(err)
|
||||||
}
|
}
|
||||||
|
|
||||||
now := time.Now()
|
now := time.Now()
|
||||||
|
|||||||
@@ -58,7 +58,7 @@ func NewFakeDNSHolder() (*Holder, error) {
|
|||||||
var err error
|
var err error
|
||||||
|
|
||||||
if fkdns, err = NewFakeDNSHolderConfigOnly(nil); err != nil {
|
if fkdns, err = NewFakeDNSHolderConfigOnly(nil); err != nil {
|
||||||
return nil, errors.New("Unable to create Fake Dns Engine").Base(err).AtError()
|
return nil, errors.New("Unable to create Fake Dns Engine").Base(err)
|
||||||
}
|
}
|
||||||
err = fkdns.initialize(dns.FakeIPv4Pool, 65535)
|
err = fkdns.initialize(dns.FakeIPv4Pool, 65535)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -80,13 +80,13 @@ func (fkdns *Holder) initialize(ipPoolCidr string, lruSize int) error {
|
|||||||
var err error
|
var err error
|
||||||
|
|
||||||
if _, ipRange, err = net.ParseCIDR(ipPoolCidr); err != nil {
|
if _, ipRange, err = net.ParseCIDR(ipPoolCidr); err != nil {
|
||||||
return errors.New("Unable to parse CIDR for Fake DNS IP assignment").Base(err).AtError()
|
return errors.New("Unable to parse CIDR for Fake DNS IP assignment").Base(err)
|
||||||
}
|
}
|
||||||
|
|
||||||
ones, bits := ipRange.Mask.Size()
|
ones, bits := ipRange.Mask.Size()
|
||||||
rooms := bits - ones
|
rooms := bits - ones
|
||||||
if math.Log2(float64(lruSize)) >= float64(rooms) {
|
if math.Log2(float64(lruSize)) >= float64(rooms) {
|
||||||
return errors.New("LRU size is bigger than subnet size").AtError()
|
return errors.New("LRU size is bigger than subnet size")
|
||||||
}
|
}
|
||||||
fkdns.domainToIP = cache.NewLru(lruSize)
|
fkdns.domainToIP = cache.NewLru(lruSize)
|
||||||
fkdns.ipRange = ipRange
|
fkdns.ipRange = ipRange
|
||||||
|
|||||||
@@ -84,7 +84,7 @@ func NewServer(ctx context.Context, dest net.Destination, dispatcher routing.Dis
|
|||||||
if dest.Network == net.Network_UDP { // UDP classic DNS mode
|
if dest.Network == net.Network_UDP { // UDP classic DNS mode
|
||||||
return NewClassicNameServer(dest, dispatcher, disableCache, serveStale, serveExpiredTTL, clientIP), nil
|
return NewClassicNameServer(dest, dispatcher, disableCache, serveStale, serveExpiredTTL, clientIP), nil
|
||||||
}
|
}
|
||||||
return nil, errors.New("No available name server could be created from ", dest).AtWarning()
|
return nil, errors.New("No available name server could be created from ", dest)
|
||||||
}
|
}
|
||||||
|
|
||||||
// NewClient creates a DNS client managing a name server with client IP, domain rules and expected IPs.
|
// NewClient creates a DNS client managing a name server with client IP, domain rules and expected IPs.
|
||||||
@@ -102,7 +102,7 @@ func NewClient(
|
|||||||
// Create a new server for each client for now
|
// Create a new server for each client for now
|
||||||
server, err := NewServer(ctx, ns.Address.AsDestination(), dispatcher, disableCache, serveStale, serveExpiredTTL, clientIP)
|
server, err := NewServer(ctx, ns.Address.AsDestination(), dispatcher, disableCache, serveStale, serveExpiredTTL, clientIP)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return errors.New("failed to create nameserver").Base(err).AtWarning()
|
return errors.New("failed to create nameserver").Base(err)
|
||||||
}
|
}
|
||||||
|
|
||||||
_, isLocalDNS := server.(*LocalNameServer)
|
_, isLocalDNS := server.(*LocalNameServer)
|
||||||
@@ -113,7 +113,7 @@ func NewClient(
|
|||||||
if len(ns.ExpectedIp) > 0 {
|
if len(ns.ExpectedIp) > 0 {
|
||||||
expectedMatcher, err = geodata.IPReg.BuildIPMatcher(ns.ExpectedIp)
|
expectedMatcher, err = geodata.IPReg.BuildIPMatcher(ns.ExpectedIp)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return errors.New("failed to create expected ip matcher").Base(err).AtWarning()
|
return errors.New("failed to create expected ip matcher").Base(err)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -122,7 +122,7 @@ func NewClient(
|
|||||||
if len(ns.UnexpectedIp) > 0 {
|
if len(ns.UnexpectedIp) > 0 {
|
||||||
unexpectedMatcher, err = geodata.IPReg.BuildIPMatcher(ns.UnexpectedIp)
|
unexpectedMatcher, err = geodata.IPReg.BuildIPMatcher(ns.UnexpectedIp)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return errors.New("failed to create unexpected ip matcher").Base(err).AtWarning()
|
return errors.New("failed to create unexpected ip matcher").Base(err)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -27,7 +27,7 @@ func (s *FakeDNSServer) IsDisableCache() bool {
|
|||||||
|
|
||||||
func (f *FakeDNSServer) QueryIP(ctx context.Context, domain string, opt dns.IPOption) ([]net.IP, uint32, error) {
|
func (f *FakeDNSServer) QueryIP(ctx context.Context, domain string, opt dns.IPOption) ([]net.IP, uint32, error) {
|
||||||
if f.fakeDNSEngine == nil {
|
if f.fakeDNSEngine == nil {
|
||||||
return nil, 0, errors.New("Unable to locate a fake DNS Engine").AtError()
|
return nil, 0, errors.New("Unable to locate a fake DNS Engine")
|
||||||
}
|
}
|
||||||
|
|
||||||
var ips []net.Address
|
var ips []net.Address
|
||||||
@@ -39,7 +39,7 @@ func (f *FakeDNSServer) QueryIP(ctx context.Context, domain string, opt dns.IPOp
|
|||||||
|
|
||||||
netIP, err := toNetIP(ips)
|
netIP, err := toNetIP(ips)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, 0, errors.New("Unable to convert IP to net ip").Base(err).AtError()
|
return nil, 0, errors.New("Unable to convert IP to net ip").Base(err)
|
||||||
}
|
}
|
||||||
|
|
||||||
errors.LogInfo(ctx, f.Name(), " got answer: ", domain, " -> ", ips)
|
errors.LogInfo(ctx, f.Name(), " got answer: ", domain, " -> ", ips)
|
||||||
|
|||||||
+6
-2
@@ -89,10 +89,10 @@ func (g *Instance) startInternal() error {
|
|||||||
g.active = true
|
g.active = true
|
||||||
|
|
||||||
if err := g.initAccessLogger(); err != nil {
|
if err := g.initAccessLogger(); err != nil {
|
||||||
return errors.New("failed to initialize access logger").Base(err).AtWarning()
|
return errors.New("failed to initialize access logger").Base(err)
|
||||||
}
|
}
|
||||||
if err := g.initErrorLogger(); err != nil {
|
if err := g.initErrorLogger(); err != nil {
|
||||||
return errors.New("failed to initialize error logger").Base(err).AtWarning()
|
return errors.New("failed to initialize error logger").Base(err)
|
||||||
}
|
}
|
||||||
|
|
||||||
return nil
|
return nil
|
||||||
@@ -141,6 +141,10 @@ func (g *Instance) Handle(msg log.Message) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (g *Instance) Severity() log.Severity {
|
||||||
|
return g.config.ErrorLogLevel
|
||||||
|
}
|
||||||
|
|
||||||
// Close implements common.Closable.Close().
|
// Close implements common.Closable.Close().
|
||||||
func (g *Instance) Close() error {
|
func (g *Instance) Close() error {
|
||||||
errors.LogDebug(context.Background(), "Logger closing")
|
errors.LogDebug(context.Background(), "Logger closing")
|
||||||
|
|||||||
@@ -66,7 +66,7 @@ func NewAlwaysOnInboundHandler(ctx context.Context, tag string, receiverConfig *
|
|||||||
}
|
}
|
||||||
mss, err := internet.ToMemoryStreamConfig(receiverConfig.StreamSettings)
|
mss, err := internet.ToMemoryStreamConfig(receiverConfig.StreamSettings)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, errors.New("failed to parse stream config").Base(err).AtWarning()
|
return nil, errors.New("failed to parse stream config").Base(err)
|
||||||
}
|
}
|
||||||
|
|
||||||
newCtx := session.ContextWithInbound(ctx, &session.Inbound{Tag: tag, Source: src})
|
newCtx := session.ContextWithInbound(ctx, &session.Inbound{Tag: tag, Source: src})
|
||||||
|
|||||||
@@ -165,7 +165,7 @@ func NewHandler(ctx context.Context, config *core.InboundHandlerConfig) (inbound
|
|||||||
|
|
||||||
receiverSettings, ok := rawReceiverSettings.(*proxyman.ReceiverConfig)
|
receiverSettings, ok := rawReceiverSettings.(*proxyman.ReceiverConfig)
|
||||||
if !ok {
|
if !ok {
|
||||||
return nil, errors.New("not a ReceiverConfig").AtError()
|
return nil, errors.New("not a ReceiverConfig")
|
||||||
}
|
}
|
||||||
|
|
||||||
streamSettings := receiverSettings.StreamSettings
|
streamSettings := receiverSettings.StreamSettings
|
||||||
|
|||||||
@@ -142,7 +142,7 @@ func (w *tcpWorker) Start() error {
|
|||||||
go w.callback(conn)
|
go w.callback(conn)
|
||||||
})
|
})
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return errors.New("failed to listen TCP on ", w.port).AtWarning().Base(err)
|
return errors.New("failed to listen TCP on ", w.port).Base(err)
|
||||||
}
|
}
|
||||||
w.hub = hub
|
w.hub = hub
|
||||||
return nil
|
return nil
|
||||||
@@ -528,7 +528,7 @@ func (w *dsWorker) Start() error {
|
|||||||
go w.callback(conn)
|
go w.callback(conn)
|
||||||
})
|
})
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return errors.New("failed to listen Unix Domain Socket on ", w.address).AtWarning().Base(err)
|
return errors.New("failed to listen Unix Domain Socket on ", w.address).Base(err)
|
||||||
}
|
}
|
||||||
w.hub = hub
|
w.hub = hub
|
||||||
return nil
|
return nil
|
||||||
|
|||||||
@@ -87,7 +87,7 @@ func NewHandler(ctx context.Context, config *core.OutboundHandlerConfig) (outbou
|
|||||||
h.senderSettings = s
|
h.senderSettings = s
|
||||||
mss, err := internet.ToMemoryStreamConfig(s.StreamSettings)
|
mss, err := internet.ToMemoryStreamConfig(s.StreamSettings)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, errors.New("failed to parse stream settings").Base(err).AtWarning()
|
return nil, errors.New("failed to parse stream settings").Base(err)
|
||||||
}
|
}
|
||||||
h.streamSettings = mss
|
h.streamSettings = mss
|
||||||
default:
|
default:
|
||||||
@@ -217,7 +217,7 @@ func (h *Handler) Dispatch(ctx context.Context, link *transport.Link) {
|
|||||||
if ob.Target.Network == net.Network_UDP && ob.Target.Port == 443 {
|
if ob.Target.Network == net.Network_UDP && ob.Target.Port == 443 {
|
||||||
switch h.udp443 {
|
switch h.udp443 {
|
||||||
case "reject":
|
case "reject":
|
||||||
test(errors.New("XUDP rejected UDP/443 traffic").AtInfo())
|
test(errors.New("XUDP rejected UDP/443 traffic"))
|
||||||
return
|
return
|
||||||
case "skip":
|
case "skip":
|
||||||
goto out
|
goto out
|
||||||
|
|||||||
@@ -68,13 +68,13 @@ func (p *Portal) HandleConnection(ctx context.Context, link *transport.Link) err
|
|||||||
outbounds := session.OutboundsFromContext(ctx)
|
outbounds := session.OutboundsFromContext(ctx)
|
||||||
ob := outbounds[len(outbounds)-1]
|
ob := outbounds[len(outbounds)-1]
|
||||||
if ob == nil {
|
if ob == nil {
|
||||||
return errors.New("outbound metadata not found").AtError()
|
return errors.New("outbound metadata not found")
|
||||||
}
|
}
|
||||||
|
|
||||||
if isDomain(ob.Target, p.domain) {
|
if isDomain(ob.Target, p.domain) {
|
||||||
muxClient, err := mux.NewClientWorker(*link, mux.ClientStrategy{})
|
muxClient, err := mux.NewClientWorker(*link, mux.ClientStrategy{})
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return errors.New("failed to create mux client worker").Base(err).AtWarning()
|
return errors.New("failed to create mux client worker").Base(err)
|
||||||
}
|
}
|
||||||
|
|
||||||
worker, err := NewPortalWorker(muxClient)
|
worker, err := NewPortalWorker(muxClient)
|
||||||
|
|||||||
@@ -115,7 +115,7 @@ func (rr *RoutingRule) BuildCondition() (Condition, error) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
if conds.Len() == 0 {
|
if conds.Len() == 0 {
|
||||||
return nil, errors.New("this rule has no effective fields").AtWarning()
|
return nil, errors.New("this rule has no effective fields")
|
||||||
}
|
}
|
||||||
|
|
||||||
return conds, nil
|
return conds, nil
|
||||||
@@ -145,7 +145,7 @@ func (br *BalancingRule) Build(ohm outbound.Manager, dispatcher routing.Dispatch
|
|||||||
}
|
}
|
||||||
s, ok := i.(*StrategyLeastLoadConfig)
|
s, ok := i.(*StrategyLeastLoadConfig)
|
||||||
if !ok {
|
if !ok {
|
||||||
return nil, errors.New("not a StrategyLeastLoadConfig").AtError()
|
return nil, errors.New("not a StrategyLeastLoadConfig")
|
||||||
}
|
}
|
||||||
leastLoadStrategy := NewLeastLoadStrategy(s)
|
leastLoadStrategy := NewLeastLoadStrategy(s)
|
||||||
return &Balancer{
|
return &Balancer{
|
||||||
|
|||||||
@@ -5,7 +5,8 @@ import (
|
|||||||
)
|
)
|
||||||
|
|
||||||
type windowsReader struct {
|
type windowsReader struct {
|
||||||
bufs []syscall.WSABuf
|
bufs []syscall.WSABuf
|
||||||
|
ready bool
|
||||||
}
|
}
|
||||||
|
|
||||||
func (r *windowsReader) Init(bs []*Buffer) {
|
func (r *windowsReader) Init(bs []*Buffer) {
|
||||||
@@ -15,6 +16,7 @@ func (r *windowsReader) Init(bs []*Buffer) {
|
|||||||
for _, b := range bs {
|
for _, b := range bs {
|
||||||
r.bufs = append(r.bufs, syscall.WSABuf{Len: uint32(Size), Buf: &b.v[0]})
|
r.bufs = append(r.bufs, syscall.WSABuf{Len: uint32(Size), Buf: &b.v[0]})
|
||||||
}
|
}
|
||||||
|
r.ready = false
|
||||||
}
|
}
|
||||||
|
|
||||||
func (r *windowsReader) Clear() {
|
func (r *windowsReader) Clear() {
|
||||||
@@ -25,6 +27,14 @@ func (r *windowsReader) Clear() {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (r *windowsReader) Read(fd uintptr) int32 {
|
func (r *windowsReader) Read(fd uintptr) int32 {
|
||||||
|
// On the first invocation, we return -1 to indicate "not ready"
|
||||||
|
// to make rawConn.Read wait for readability using the runtime's own mechanism
|
||||||
|
// because syscall.WSARecv() is a blocking call when used with nil OVERLAPPED
|
||||||
|
if !r.ready {
|
||||||
|
r.ready = true
|
||||||
|
return -1
|
||||||
|
}
|
||||||
|
|
||||||
var nBytes uint32
|
var nBytes uint32
|
||||||
var flags uint32
|
var flags uint32
|
||||||
err := syscall.WSARecv(syscall.Handle(fd), &r.bufs[0], uint32(len(r.bufs)), &nBytes, &flags, nil, nil)
|
err := syscall.WSARecv(syscall.Handle(fd), &r.bufs[0], uint32(len(r.bufs)), &nBytes, &flags, nil, nil)
|
||||||
|
|||||||
+13
-65
@@ -18,17 +18,12 @@ type hasInnerError interface {
|
|||||||
Unwrap() error
|
Unwrap() error
|
||||||
}
|
}
|
||||||
|
|
||||||
type hasSeverity interface {
|
|
||||||
Severity() log.Severity
|
|
||||||
}
|
|
||||||
|
|
||||||
// Error is an error object with underlying error.
|
// Error is an error object with underlying error.
|
||||||
type Error struct {
|
type Error struct {
|
||||||
prefix []interface{}
|
prefix []interface{}
|
||||||
message []interface{}
|
message []interface{}
|
||||||
caller string
|
caller string
|
||||||
inner error
|
inner error
|
||||||
severity log.Severity
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// Error implements error.Error().
|
// Error implements error.Error().
|
||||||
@@ -69,46 +64,6 @@ func (err *Error) Base(e error) *Error {
|
|||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
func (err *Error) atSeverity(s log.Severity) *Error {
|
|
||||||
err.severity = s
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
|
|
||||||
func (err *Error) Severity() log.Severity {
|
|
||||||
if err.inner == nil {
|
|
||||||
return err.severity
|
|
||||||
}
|
|
||||||
|
|
||||||
if s, ok := err.inner.(hasSeverity); ok {
|
|
||||||
as := s.Severity()
|
|
||||||
if as < err.severity {
|
|
||||||
return as
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
return err.severity
|
|
||||||
}
|
|
||||||
|
|
||||||
// AtDebug sets the severity to debug.
|
|
||||||
func (err *Error) AtDebug() *Error {
|
|
||||||
return err.atSeverity(log.Severity_Debug)
|
|
||||||
}
|
|
||||||
|
|
||||||
// AtInfo sets the severity to info.
|
|
||||||
func (err *Error) AtInfo() *Error {
|
|
||||||
return err.atSeverity(log.Severity_Info)
|
|
||||||
}
|
|
||||||
|
|
||||||
// AtWarning sets the severity to warning.
|
|
||||||
func (err *Error) AtWarning() *Error {
|
|
||||||
return err.atSeverity(log.Severity_Warning)
|
|
||||||
}
|
|
||||||
|
|
||||||
// AtError sets the severity to error.
|
|
||||||
func (err *Error) AtError() *Error {
|
|
||||||
return err.atSeverity(log.Severity_Error)
|
|
||||||
}
|
|
||||||
|
|
||||||
// String returns the string representation of this error.
|
// String returns the string representation of this error.
|
||||||
func (err *Error) String() string {
|
func (err *Error) String() string {
|
||||||
return err.Error()
|
return err.Error()
|
||||||
@@ -132,9 +87,8 @@ func New(msg ...interface{}) *Error {
|
|||||||
details = details[:i]
|
details = details[:i]
|
||||||
}
|
}
|
||||||
return &Error{
|
return &Error{
|
||||||
message: msg,
|
message: msg,
|
||||||
severity: log.Severity_Info,
|
caller: details,
|
||||||
caller: details,
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -171,6 +125,9 @@ func LogErrorInner(ctx context.Context, inner error, msg ...interface{}) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func doLog(ctx context.Context, inner error, severity log.Severity, msg ...interface{}) {
|
func doLog(ctx context.Context, inner error, severity log.Severity, msg ...interface{}) {
|
||||||
|
if log.GetSeverity() < severity {
|
||||||
|
return
|
||||||
|
}
|
||||||
pc, _, _, _ := runtime.Caller(2)
|
pc, _, _, _ := runtime.Caller(2)
|
||||||
details := runtime.FuncForPC(pc).Name()
|
details := runtime.FuncForPC(pc).Name()
|
||||||
if len(details) >= trim {
|
if len(details) >= trim {
|
||||||
@@ -181,10 +138,9 @@ func doLog(ctx context.Context, inner error, severity log.Severity, msg ...inter
|
|||||||
details = details[:i]
|
details = details[:i]
|
||||||
}
|
}
|
||||||
err := &Error{
|
err := &Error{
|
||||||
message: msg,
|
message: msg,
|
||||||
severity: severity,
|
caller: details,
|
||||||
caller: details,
|
inner: inner,
|
||||||
inner: inner,
|
|
||||||
}
|
}
|
||||||
if ctx != nil && ctx != context.Background() {
|
if ctx != nil && ctx != context.Background() {
|
||||||
id := uint32(c.IDFromContext(ctx))
|
id := uint32(c.IDFromContext(ctx))
|
||||||
@@ -193,7 +149,7 @@ func doLog(ctx context.Context, inner error, severity log.Severity, msg ...inter
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
log.Record(&log.GeneralMessage{
|
log.Record(&log.GeneralMessage{
|
||||||
Severity: GetSeverity(err),
|
Severity: severity,
|
||||||
Content: err,
|
Content: err,
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
@@ -217,11 +173,3 @@ L:
|
|||||||
}
|
}
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
// GetSeverity returns the actual severity of the error, including inner errors.
|
|
||||||
func GetSeverity(err error) log.Severity {
|
|
||||||
if s, ok := err.(hasSeverity); ok {
|
|
||||||
return s.Severity()
|
|
||||||
}
|
|
||||||
return log.Severity_Info
|
|
||||||
}
|
|
||||||
|
|||||||
@@ -7,30 +7,21 @@ import (
|
|||||||
|
|
||||||
"github.com/google/go-cmp/cmp"
|
"github.com/google/go-cmp/cmp"
|
||||||
. "github.com/xtls/xray-core/common/errors"
|
. "github.com/xtls/xray-core/common/errors"
|
||||||
"github.com/xtls/xray-core/common/log"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
func TestError(t *testing.T) {
|
func TestError(t *testing.T) {
|
||||||
err := New("TestError")
|
err := New("TestError")
|
||||||
if v := GetSeverity(err); v != log.Severity_Info {
|
if v := err.Error(); !strings.Contains(v, "TestError") {
|
||||||
t.Error("severity: ", v)
|
t.Error("error: ", v)
|
||||||
}
|
}
|
||||||
|
|
||||||
err = New("TestError2").Base(io.EOF)
|
err = New("TestError2").Base(io.EOF)
|
||||||
if v := GetSeverity(err); v != log.Severity_Info {
|
if v := err.Error(); !strings.Contains(v, "EOF") {
|
||||||
t.Error("severity: ", v)
|
t.Error("error: ", v)
|
||||||
}
|
}
|
||||||
|
|
||||||
err = New("TestError3").Base(io.EOF).AtWarning()
|
err = New("TestError3").Base(io.EOF)
|
||||||
if v := GetSeverity(err); v != log.Severity_Warning {
|
err = New("TestError4").Base(err)
|
||||||
t.Error("severity: ", v)
|
|
||||||
}
|
|
||||||
|
|
||||||
err = New("TestError4").Base(io.EOF).AtWarning()
|
|
||||||
err = New("TestError5").Base(err)
|
|
||||||
if v := GetSeverity(err); v != log.Severity_Warning {
|
|
||||||
t.Error("severity: ", v)
|
|
||||||
}
|
|
||||||
if v := err.Error(); !strings.Contains(v, "EOF") {
|
if v := err.Error(); !strings.Contains(v, "EOF") {
|
||||||
t.Error("error: ", v)
|
t.Error("error: ", v)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1,6 +1,7 @@
|
|||||||
package strmatcher_test
|
package strmatcher_test
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"regexp"
|
||||||
"strconv"
|
"strconv"
|
||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
@@ -72,6 +73,64 @@ func BenchmarkSubstrMatcher(b *testing.B) {
|
|||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func BenchmarkRegexMatcher(b *testing.B) {
|
||||||
|
patterns := []string{ // taken from geosite
|
||||||
|
`(^|\.)91porn\.(best|com|cool|fun|group|party|plus|site|tw|work)$`,
|
||||||
|
`(^|\.)91porn[0-9]{3}\.me$`,
|
||||||
|
`(^|\.)apiproxy-device-prod-nlb-.+\.amazonaws\.com$`,
|
||||||
|
`(^|\.)dualstack\.apiproxy-.+\.amazonaws\.com$`,
|
||||||
|
`(^|\.)aqdk[0-9]{3}\.com$`,
|
||||||
|
`(^|\.)bilibili3(0[1-9]|1[0-2])\.xyz$`,
|
||||||
|
`(^|\.)byyum([3589]|2[235689]|3[34]|4[1-9]|5[1-79]|6[0134679])?\.com$`,
|
||||||
|
`(^|\.)fiftymvapi\..+$`,
|
||||||
|
`(^|\.)gossipfuli[0-9]{3,4}\.xyz$`,
|
||||||
|
`(^|\.)kpkuang\.(bond|fun|info|one|us)$`,
|
||||||
|
`(^|\.)rule34\.(asia|us|world|xxx|xyz)$`,
|
||||||
|
`(^|\.)[a-z][1-9][0-9][a-z]\.com$`,
|
||||||
|
`.+\.awsdns-[0-9][0-9]\.(co\.uk|com|net|org)$`,
|
||||||
|
`.+\.dkr\.ecr\.[^\.]+\.amazonaws\.com$`,
|
||||||
|
`^(.+\.)*zh\.okaapps\.com$`,
|
||||||
|
`^cdn\d-epicgames-\d+\.file\.myqcloud\.com$`,
|
||||||
|
`^chatgpt-async-webps-prod-\S+-\d+\.webpubsub\.azure\.com$`,
|
||||||
|
`^r+[0-9]+(---|\.)sn-(2x3|ni5|j5o)\w{5}\.googlevideo\.com$`,
|
||||||
|
`^speed\.(coe|open)\.ad\.[a-z]{2,6}\.prod\.hosts\.ooklaserver\.net$`,
|
||||||
|
`javdb\d+\.com$`,
|
||||||
|
}
|
||||||
|
domains := []string{
|
||||||
|
"www.google.com", "rr3---sn-4g5edndy.googlevideo.com", "r1---sn-2x3abcde.googlevideo.com", "i.ytimg.com",
|
||||||
|
"graph.facebook.com", "api.twitter.com", "www.baidu.com", "github.com", "objects.githubusercontent.com",
|
||||||
|
"login.microsoftonline.com", "e1234.dscb.akamaiedge.net", "d1a2b3c4d5e6f7.cloudfront.net",
|
||||||
|
"s3.us-east-1.amazonaws.com", "123456789012.dkr.ecr.us-east-1.amazonaws.com", "www.wikipedia.org",
|
||||||
|
"discord.com", "telegram.org", "store.steampowered.com", "www.91porn.com", "ns-1234.awsdns-12.org",
|
||||||
|
}
|
||||||
|
bench := func(b *testing.B, ctor func(pattern string) func(string) bool) {
|
||||||
|
var matchers []func(string) bool
|
||||||
|
for _, p := range patterns {
|
||||||
|
matchers = append(matchers, ctor(p))
|
||||||
|
}
|
||||||
|
b.ResetTimer()
|
||||||
|
for i := 0; i < b.N; i++ {
|
||||||
|
for _, d := range domains {
|
||||||
|
for _, match := range matchers {
|
||||||
|
_ = match(d)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
b.Run("regexp", func(b *testing.B) {
|
||||||
|
bench(b, func(pattern string) func(string) bool {
|
||||||
|
return regexp.MustCompile(pattern).MatchString
|
||||||
|
})
|
||||||
|
})
|
||||||
|
b.Run("prefilter", func(b *testing.B) {
|
||||||
|
bench(b, func(pattern string) func(string) bool {
|
||||||
|
m, err := Regex.New(pattern)
|
||||||
|
common.Must(err)
|
||||||
|
return m.Match
|
||||||
|
})
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
// Utility functions for benchmark
|
// Utility functions for benchmark
|
||||||
|
|
||||||
func benchmarkMatcherType(b *testing.B, t Type, ctor func() MatcherGroup) {
|
func benchmarkMatcherType(b *testing.B, t Type, ctor func() MatcherGroup) {
|
||||||
|
|||||||
@@ -1,6 +1,8 @@
|
|||||||
package strmatcher
|
package strmatcher
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"errors"
|
||||||
|
"math"
|
||||||
"math/bits"
|
"math/bits"
|
||||||
"runtime"
|
"runtime"
|
||||||
"sort"
|
"sort"
|
||||||
@@ -38,19 +40,21 @@ type mphRuleInfo struct {
|
|||||||
// MphMatcherGroup is an implementation of MatcherGroup.
|
// MphMatcherGroup is an implementation of MatcherGroup.
|
||||||
// It implements Rabin-Karp algorithm and minimal perfect hash table for Full and Domain matcher.
|
// It implements Rabin-Karp algorithm and minimal perfect hash table for Full and Domain matcher.
|
||||||
type MphMatcherGroup struct {
|
type MphMatcherGroup struct {
|
||||||
rules []string // RuleIdx -> pattern string, index 0 reserved for failed lookup
|
patterns string // All rule patterns concatenated
|
||||||
values [][]uint32 // RuleIdx -> registered matcher values for the pattern (Full Matcher takes precedence)
|
patternOffs []uint32 // RuleIdx -> patterns[patternOffs[i]:patternOffs[i+1]], index 0 reserved for failed lookup
|
||||||
level0 []uint32 // RollingHash & Mask -> seed for Memhash
|
values []uint32 // All registered matcher values concatenated
|
||||||
level0Mask uint32 // Mask restricting RollingHash to 0 ~ len(level0)
|
valueOffs []uint32 // RuleIdx -> values[valueOffs[i]:valueOffs[i+1]] (Full Matcher takes precedence)
|
||||||
level1 []uint32 // Memhash<seed> & Mask -> stored index for rules
|
level0 []uint32 // RollingHash & Mask -> seed for Memhash
|
||||||
level1Mask uint32 // Mask for restricting Memhash<seed> to 0 ~ len(level1)
|
level0Mask uint32 // Mask restricting RollingHash to 0 ~ len(level0)
|
||||||
ruleInfos *map[string]mphRuleInfo
|
level1 []uint32 // Memhash<seed> & Mask -> stored index for rules
|
||||||
|
level1Mask uint32 // Mask for restricting Memhash<seed> to 0 ~ len(level1)
|
||||||
|
rules []string // RuleIdx -> pattern string, only used for building
|
||||||
|
ruleInfos *map[string]mphRuleInfo
|
||||||
}
|
}
|
||||||
|
|
||||||
func NewMphMatcherGroup() *MphMatcherGroup {
|
func NewMphMatcherGroup() *MphMatcherGroup {
|
||||||
return &MphMatcherGroup{
|
return &MphMatcherGroup{
|
||||||
rules: []string{""},
|
rules: []string{""},
|
||||||
values: [][]uint32{nil},
|
|
||||||
level0: nil,
|
level0: nil,
|
||||||
level0Mask: 0,
|
level0Mask: 0,
|
||||||
level1: nil,
|
level1: nil,
|
||||||
@@ -78,7 +82,6 @@ func (g *MphMatcherGroup) addPattern(suffixHash uint32, suffixPattern string, pa
|
|||||||
if !found {
|
if !found {
|
||||||
info = mphRuleInfo{rollingHash: RollingHash(suffixHash, pattern)}
|
info = mphRuleInfo{rollingHash: RollingHash(suffixHash, pattern)}
|
||||||
g.rules = append(g.rules, fullPattern)
|
g.rules = append(g.rules, fullPattern)
|
||||||
g.values = append(g.values, nil)
|
|
||||||
}
|
}
|
||||||
info.matchers[matcherType] = append(info.matchers[matcherType], value)
|
info.matchers[matcherType] = append(info.matchers[matcherType], value)
|
||||||
(*g.ruleInfos)[fullPattern] = info
|
(*g.ruleInfos)[fullPattern] = info
|
||||||
@@ -94,14 +97,30 @@ func (g *MphMatcherGroup) Build() error {
|
|||||||
g.level1 = make([]uint32, nextPow2(ruleCount))
|
g.level1 = make([]uint32, nextPow2(ruleCount))
|
||||||
g.level1Mask = uint32(len(g.level1) - 1)
|
g.level1Mask = uint32(len(g.level1) - 1)
|
||||||
|
|
||||||
|
// Flatten patterns and values so the built group has no per-rule objects
|
||||||
|
valueCount := 0
|
||||||
|
for _, ruleInfo := range *g.ruleInfos {
|
||||||
|
valueCount += len(ruleInfo.matchers[Full]) + len(ruleInfo.matchers[Domain])
|
||||||
|
}
|
||||||
|
g.patterns = strings.Join(g.rules, "")
|
||||||
|
if uint64(len(g.patterns)) > math.MaxUint32 || uint64(valueCount) > math.MaxUint32 {
|
||||||
|
return errors.New("too many rules for MphMatcherGroup")
|
||||||
|
}
|
||||||
|
g.patternOffs = make([]uint32, len(g.rules)+1)
|
||||||
|
g.values = make([]uint32, 0, valueCount)
|
||||||
|
g.valueOffs = make([]uint32, len(g.rules)+1)
|
||||||
|
|
||||||
// Create buckets based on all rule's rolling hash
|
// Create buckets based on all rule's rolling hash
|
||||||
buckets := make([][]uint32, len(g.level0))
|
buckets := make([][]uint32, len(g.level0))
|
||||||
for ruleIdx := 1; ruleIdx < len(g.rules); ruleIdx++ { // Traverse rules starting from index 1 (0 reserved for failed lookup)
|
for ruleIdx := 1; ruleIdx < len(g.rules); ruleIdx++ { // Traverse rules starting from index 1 (0 reserved for failed lookup)
|
||||||
ruleInfo := (*g.ruleInfos)[g.rules[ruleIdx]]
|
ruleInfo := (*g.ruleInfos)[g.rules[ruleIdx]]
|
||||||
bucketIdx := ruleInfo.rollingHash & g.level0Mask
|
bucketIdx := ruleInfo.rollingHash & g.level0Mask
|
||||||
buckets[bucketIdx] = append(buckets[bucketIdx], uint32(ruleIdx))
|
buckets[bucketIdx] = append(buckets[bucketIdx], uint32(ruleIdx))
|
||||||
g.values[ruleIdx] = append(ruleInfo.matchers[Full], ruleInfo.matchers[Domain]...) // nolint:gocritic
|
g.patternOffs[ruleIdx+1] = g.patternOffs[ruleIdx] + uint32(len(g.rules[ruleIdx]))
|
||||||
|
g.values = append(append(g.values, ruleInfo.matchers[Full]...), ruleInfo.matchers[Domain]...)
|
||||||
|
g.valueOffs[ruleIdx+1] = uint32(len(g.values))
|
||||||
}
|
}
|
||||||
|
g.rules = nil
|
||||||
g.ruleInfos = nil // Set ruleInfos nil to release memory
|
g.ruleInfos = nil // Set ruleInfos nil to release memory
|
||||||
runtime.GC() // peak mem
|
runtime.GC() // peak mem
|
||||||
|
|
||||||
@@ -121,7 +140,7 @@ func (g *MphMatcherGroup) Build() error {
|
|||||||
seed := uint32(0)
|
seed := uint32(0)
|
||||||
for len(hashedBucket) != len(bucket) {
|
for len(hashedBucket) != len(bucket) {
|
||||||
for _, ruleIdx := range bucket {
|
for _, ruleIdx := range bucket {
|
||||||
memHash := MemHash(seed, g.rules[ruleIdx]) & g.level1Mask
|
memHash := MemHash(seed, g.pattern(ruleIdx)) & g.level1Mask
|
||||||
if occupied[memHash] { // Collision occurred with this seed
|
if occupied[memHash] { // Collision occurred with this seed
|
||||||
for _, hash := range hashedBucket { // Revert all values in this hashed bucket
|
for _, hash := range hashedBucket { // Revert all values in this hashed bucket
|
||||||
occupied[hash] = false
|
occupied[hash] = false
|
||||||
@@ -141,12 +160,26 @@ func (g *MphMatcherGroup) Build() error {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (g *MphMatcherGroup) pattern(ruleIdx uint32) string {
|
||||||
|
return g.patterns[g.patternOffs[ruleIdx]:g.patternOffs[ruleIdx+1]]
|
||||||
|
}
|
||||||
|
|
||||||
|
// valuesOf caps the capacity, so appending to a Match result can't overwrite the next rule's values.
|
||||||
|
func (g *MphMatcherGroup) valuesOf(ruleIdx uint32) []uint32 {
|
||||||
|
start, end := g.valueOffs[ruleIdx], g.valueOffs[ruleIdx+1]
|
||||||
|
return g.values[start:end:end]
|
||||||
|
}
|
||||||
|
|
||||||
// Lookup searches for input in minimal perfect hash table and returns its index. 0 indicates not found.
|
// Lookup searches for input in minimal perfect hash table and returns its index. 0 indicates not found.
|
||||||
func (g *MphMatcherGroup) Lookup(rollingHash uint32, input string) uint32 {
|
func (g *MphMatcherGroup) Lookup(rollingHash uint32, input string) uint32 {
|
||||||
i0 := rollingHash & g.level0Mask
|
i0 := rollingHash & g.level0Mask
|
||||||
seed := g.level0[i0]
|
seed := g.level0[i0]
|
||||||
i1 := MemHash(seed, input) & g.level1Mask
|
i1 := MemHash(seed, input) & g.level1Mask
|
||||||
if n := g.level1[i1]; g.rules[n] == input {
|
n := g.level1[i1]
|
||||||
|
// Build only puts valid rule indices in level1, so n+1 < len(patternOffs) and the span is inside patterns.
|
||||||
|
// Skip the bounds checks, they made this hot path measurably slower than indexing a []string
|
||||||
|
offs := (*[2]uint32)(unsafe.Add(unsafe.Pointer(unsafe.SliceData(g.patternOffs)), uintptr(n)*4))
|
||||||
|
if start := offs[0]; int(offs[1]-start) == len(input) && unsafe.String((*byte)(unsafe.Add(unsafe.Pointer(unsafe.StringData(g.patterns)), start)), len(input)) == input {
|
||||||
return n
|
return n
|
||||||
}
|
}
|
||||||
return 0
|
return 0
|
||||||
@@ -160,12 +193,12 @@ func (g *MphMatcherGroup) Match(input string) []uint32 {
|
|||||||
hash = hash*PrimeRK + uint32(input[i])
|
hash = hash*PrimeRK + uint32(input[i])
|
||||||
if input[i] == '.' {
|
if input[i] == '.' {
|
||||||
if mphIdx := g.Lookup(hash, input[i:]); mphIdx != 0 {
|
if mphIdx := g.Lookup(hash, input[i:]); mphIdx != 0 {
|
||||||
matches = append(matches, g.values[mphIdx])
|
matches = append(matches, g.valuesOf(mphIdx))
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
if mphIdx := g.Lookup(hash, input); mphIdx != 0 {
|
if mphIdx := g.Lookup(hash, input); mphIdx != 0 {
|
||||||
matches = append(matches, g.values[mphIdx])
|
matches = append(matches, g.valuesOf(mphIdx))
|
||||||
}
|
}
|
||||||
return CompositeMatchesReverse(matches)
|
return CompositeMatchesReverse(matches)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1,7 +1,9 @@
|
|||||||
package strmatcher_test
|
package strmatcher_test
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"math/rand"
|
||||||
"reflect"
|
"reflect"
|
||||||
|
"slices"
|
||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
"github.com/xtls/xray-core/common"
|
"github.com/xtls/xray-core/common"
|
||||||
@@ -276,3 +278,63 @@ func TestEmptyMphMatcherGroup(t *testing.T) {
|
|||||||
t.Error("Expect [], but ", r)
|
t.Error("Expect [], but ", r)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestMphMatcherGroupRandom(t *testing.T) {
|
||||||
|
inputs := []string{""} // All strings over "ab." up to 7 bytes
|
||||||
|
for i := 0; len(inputs[i]) < 7; i++ {
|
||||||
|
for _, c := range []string{"a", "b", "."} {
|
||||||
|
inputs = append(inputs, inputs[i]+c)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
for seed := int64(0); seed < 300; seed++ {
|
||||||
|
r := rand.New(rand.NewSource(seed))
|
||||||
|
g := NewMphMatcherGroup()
|
||||||
|
full, domain := map[string][]uint32{}, map[string][]uint32{} // Stored pattern -> values
|
||||||
|
for value := uint32(r.Intn(200)); value > 0; value-- {
|
||||||
|
pattern := make([]byte, r.Intn(8))
|
||||||
|
for i := range pattern {
|
||||||
|
pattern[i] = "ab."[r.Intn(3)]
|
||||||
|
}
|
||||||
|
if p := string(pattern); r.Intn(2) == 0 {
|
||||||
|
g.AddFullMatcher(FullMatcher(p), value)
|
||||||
|
full[p] = append(full[p], value)
|
||||||
|
} else {
|
||||||
|
g.AddDomainMatcher(DomainMatcher(p), value)
|
||||||
|
domain[p] = append(domain[p], value)
|
||||||
|
domain["."+p] = append(domain["."+p], value)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
g.Build()
|
||||||
|
for _, input := range inputs {
|
||||||
|
keys := []string{input} // Whole input first, then "." suffixes from longest to shortest
|
||||||
|
for i := range len(input) {
|
||||||
|
if input[i] == '.' {
|
||||||
|
keys = append(keys, input[i:])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
var want []uint32
|
||||||
|
for _, k := range keys {
|
||||||
|
want = append(append(want, full[k]...), domain[k]...)
|
||||||
|
}
|
||||||
|
if m := g.Match(input); !slices.Equal(m, want) {
|
||||||
|
t.Fatalf("seed %d: Match(%q) = %v, want %v", seed, input, m, want)
|
||||||
|
}
|
||||||
|
if m := g.MatchAny(input); m != (len(want) > 0) {
|
||||||
|
t.Fatalf("seed %d: MatchAny(%q) = %v", seed, input, m)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMphMatcherGroupAppend(t *testing.T) {
|
||||||
|
g := NewMphMatcherGroup()
|
||||||
|
g.AddFullMatcher(FullMatcher("a.com"), 1)
|
||||||
|
g.AddFullMatcher(FullMatcher("b.com"), 2)
|
||||||
|
g.Build()
|
||||||
|
if m := append(g.Match("a.com"), 3); !slices.Equal(m, []uint32{1, 3}) {
|
||||||
|
t.Error("expect [1 3], but ", m)
|
||||||
|
}
|
||||||
|
if m := g.Match("b.com"); !slices.Equal(m, []uint32{2}) {
|
||||||
|
t.Error("expect [2], but ", m)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
@@ -3,6 +3,7 @@ package strmatcher
|
|||||||
import (
|
import (
|
||||||
"errors"
|
"errors"
|
||||||
"regexp"
|
"regexp"
|
||||||
|
"regexp/syntax"
|
||||||
"slices"
|
"slices"
|
||||||
"strings"
|
"strings"
|
||||||
"unicode/utf8"
|
"unicode/utf8"
|
||||||
@@ -73,7 +74,43 @@ func (m SubstrMatcher) Match(s string) bool {
|
|||||||
|
|
||||||
// RegexMatcher is an implementation of Matcher.
|
// RegexMatcher is an implementation of Matcher.
|
||||||
type RegexMatcher struct {
|
type RegexMatcher struct {
|
||||||
pattern *regexp.Regexp
|
pattern *regexp.Regexp
|
||||||
|
literals []string // every match contains all of them, longest first
|
||||||
|
}
|
||||||
|
|
||||||
|
func newRegexMatcher(pattern string) (Matcher, error) {
|
||||||
|
regex, err := regexp.Compile(pattern)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
m := &RegexMatcher{pattern: regex}
|
||||||
|
if re, err := syntax.Parse(pattern, syntax.Perl); err == nil { // same flags as regexp.Compile
|
||||||
|
m.literals = requiredLiterals(re, nil)
|
||||||
|
slices.SortStableFunc(m.literals, func(a, b string) int { return len(b) - len(a) })
|
||||||
|
}
|
||||||
|
return m, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// requiredLiterals appends to dst the case-sensitive strings that every match of re contains.
|
||||||
|
func requiredLiterals(re *syntax.Regexp, dst []string) []string {
|
||||||
|
switch re.Op {
|
||||||
|
case syntax.OpLiteral:
|
||||||
|
// regexp matches U+FFFD against invalid UTF-8 bytes, strings.Contains does not
|
||||||
|
if re.Flags&syntax.FoldCase == 0 && !slices.Contains(re.Rune, utf8.RuneError) {
|
||||||
|
dst = append(dst, string(re.Rune))
|
||||||
|
}
|
||||||
|
case syntax.OpCapture, syntax.OpPlus:
|
||||||
|
dst = requiredLiterals(re.Sub[0], dst)
|
||||||
|
case syntax.OpRepeat:
|
||||||
|
if re.Min > 0 {
|
||||||
|
dst = requiredLiterals(re.Sub[0], dst)
|
||||||
|
}
|
||||||
|
case syntax.OpConcat:
|
||||||
|
for _, sub := range re.Sub {
|
||||||
|
dst = requiredLiterals(sub, dst)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return dst
|
||||||
}
|
}
|
||||||
|
|
||||||
func (*RegexMatcher) Type() Type {
|
func (*RegexMatcher) Type() Type {
|
||||||
@@ -89,6 +126,11 @@ func (m *RegexMatcher) String() string {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (m *RegexMatcher) Match(s string) bool {
|
func (m *RegexMatcher) Match(s string) bool {
|
||||||
|
for _, l := range m.literals {
|
||||||
|
if !strings.Contains(s, l) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
}
|
||||||
return m.pattern.MatchString(s)
|
return m.pattern.MatchString(s)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -102,11 +144,7 @@ func (t Type) New(pattern string) (Matcher, error) {
|
|||||||
case Domain:
|
case Domain:
|
||||||
return DomainMatcher(pattern), nil
|
return DomainMatcher(pattern), nil
|
||||||
case Regex: // 1. regex matching is case-sensitive
|
case Regex: // 1. regex matching is case-sensitive
|
||||||
regex, err := regexp.Compile(pattern)
|
return newRegexMatcher(pattern)
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
return &RegexMatcher{pattern: regex}, nil
|
|
||||||
default:
|
default:
|
||||||
return nil, errors.New("unknown matcher type")
|
return nil, errors.New("unknown matcher type")
|
||||||
}
|
}
|
||||||
@@ -135,11 +173,7 @@ func (t Type) NewDomainPattern(pattern string) (Matcher, error) {
|
|||||||
}
|
}
|
||||||
return DomainMatcher(pattern), nil
|
return DomainMatcher(pattern), nil
|
||||||
case Regex: // Regex's charset not in LDH subset
|
case Regex: // Regex's charset not in LDH subset
|
||||||
regex, err := regexp.Compile(pattern)
|
return newRegexMatcher(pattern)
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
return &RegexMatcher{pattern: regex}, nil
|
|
||||||
default:
|
default:
|
||||||
return nil, errors.New("unknown matcher type")
|
return nil, errors.New("unknown matcher type")
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -0,0 +1,60 @@
|
|||||||
|
package strmatcher
|
||||||
|
|
||||||
|
import (
|
||||||
|
"regexp"
|
||||||
|
"slices"
|
||||||
|
"testing"
|
||||||
|
)
|
||||||
|
|
||||||
|
var regexLiteralCases = []struct {
|
||||||
|
pattern string
|
||||||
|
literals []string
|
||||||
|
}{
|
||||||
|
{`(^|\.)91porn\.(best|com)$`, []string{"91porn."}},
|
||||||
|
{`.+\.awsdns-cn-[0-9][0-9]\.(biz|com|net|top)$`, []string{".awsdns-cn-", "."}},
|
||||||
|
{`^r+[0-9]+(---|\.)sn-(2x3|ni5|j5o)\w{5}\.googlevideo\.com$`, []string{".googlevideo.com", "sn-", "r"}},
|
||||||
|
{`(?i)abc`, nil},
|
||||||
|
{`ab(?i:CD)ef`, []string{"ab", "ef"}},
|
||||||
|
{`(abc)?x`, []string{"x"}},
|
||||||
|
{`(abc)*x`, []string{"x"}},
|
||||||
|
{`x{0,3}yy`, []string{"yy"}},
|
||||||
|
{`(ab)+c{2}`, []string{"ab", "c"}},
|
||||||
|
{`abc|abd`, []string{"ab"}},
|
||||||
|
{`\Qa.b\E`, []string{"a.b"}},
|
||||||
|
{`a\x{FFFD}b`, nil},
|
||||||
|
{`^[^.]+$`, nil},
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRegexRequiredLiterals(t *testing.T) {
|
||||||
|
for _, test := range regexLiteralCases {
|
||||||
|
m, err := newRegexMatcher(test.pattern)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if got := m.(*RegexMatcher).literals; !slices.Equal(got, test.literals) {
|
||||||
|
t.Errorf("%s: got %q, want %q", test.pattern, got, test.literals)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func FuzzRegexMatcher(f *testing.F) {
|
||||||
|
inputs := []string{
|
||||||
|
"", "x", "yy", "abd", "ccc", "ABC", "abCDef", "abcdef", "abababcc", "a.b", "a\xffb", "a\uFFFDb",
|
||||||
|
"www.91porn.com", "ns1.awsdns-cn-01.top", "r1---sn-2x3abcde.googlevideo.com",
|
||||||
|
}
|
||||||
|
for _, test := range regexLiteralCases {
|
||||||
|
for _, s := range inputs {
|
||||||
|
f.Add(test.pattern, s)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
f.Fuzz(func(t *testing.T, pattern, s string) {
|
||||||
|
re, err := regexp.Compile(pattern)
|
||||||
|
if err != nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
m, _ := newRegexMatcher(pattern)
|
||||||
|
if got, want := m.Match(s), re.MatchString(s); got != want {
|
||||||
|
t.Errorf("pattern %q, input %q: got %v, want %v", pattern, s, got, want)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
+21
-25
@@ -1,7 +1,7 @@
|
|||||||
package log // import "github.com/xtls/xray-core/common/log"
|
package log // import "github.com/xtls/xray-core/common/log"
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"sync"
|
"sync/atomic"
|
||||||
|
|
||||||
"github.com/xtls/xray-core/common/serial"
|
"github.com/xtls/xray-core/common/serial"
|
||||||
)
|
)
|
||||||
@@ -29,36 +29,32 @@ func (m *GeneralMessage) String() string {
|
|||||||
|
|
||||||
// Record writes a message into log stream.
|
// Record writes a message into log stream.
|
||||||
func Record(msg Message) {
|
func Record(msg Message) {
|
||||||
logHandler.Handle(msg)
|
if h := logHandler.Load(); h != nil {
|
||||||
|
(*h).Handle(msg)
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
var logHandler syncHandler
|
type SeverityLogger interface {
|
||||||
|
Handler
|
||||||
|
Severity() Severity
|
||||||
|
}
|
||||||
|
|
||||||
|
func GetSeverity() Severity {
|
||||||
|
if h := logHandler.Load(); h != nil {
|
||||||
|
if sh, ok := (*h).(SeverityLogger); ok {
|
||||||
|
return sh.Severity()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
// log everything by default
|
||||||
|
return Severity_Debug
|
||||||
|
}
|
||||||
|
|
||||||
|
var logHandler atomic.Pointer[Handler]
|
||||||
|
|
||||||
// RegisterHandler registers a new handler as current log handler. Previous registered handler will be discarded.
|
// RegisterHandler registers a new handler as current log handler. Previous registered handler will be discarded.
|
||||||
func RegisterHandler(handler Handler) {
|
func RegisterHandler(handler Handler) {
|
||||||
if handler == nil {
|
if handler == nil {
|
||||||
panic("Log handler is nil")
|
panic("Log handler is nil")
|
||||||
}
|
}
|
||||||
logHandler.Set(handler)
|
logHandler.Store(&handler)
|
||||||
}
|
|
||||||
|
|
||||||
type syncHandler struct {
|
|
||||||
sync.RWMutex
|
|
||||||
Handler
|
|
||||||
}
|
|
||||||
|
|
||||||
func (h *syncHandler) Handle(msg Message) {
|
|
||||||
h.RLock()
|
|
||||||
defer h.RUnlock()
|
|
||||||
|
|
||||||
if h.Handler != nil {
|
|
||||||
h.Handler.Handle(msg)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func (h *syncHandler) Set(handler Handler) {
|
|
||||||
h.Lock()
|
|
||||||
defer h.Unlock()
|
|
||||||
|
|
||||||
h.Handler = handler
|
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -68,6 +68,10 @@ func (l *serverityLogger) Handle(msg Message) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (l *serverityLogger) Severity() Severity {
|
||||||
|
return l.logLevel
|
||||||
|
}
|
||||||
|
|
||||||
func (l *generalLogger) run() {
|
func (l *generalLogger) run() {
|
||||||
defer l.access.Signal()
|
defer l.access.Signal()
|
||||||
|
|
||||||
|
|||||||
@@ -38,7 +38,7 @@ func (m *ClientManager) Dispatch(ctx context.Context, link *transport.Link) erro
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
return errors.New("unable to find an available mux client").AtWarning()
|
return errors.New("unable to find an available mux client")
|
||||||
}
|
}
|
||||||
|
|
||||||
type WorkerPicker interface {
|
type WorkerPicker interface {
|
||||||
|
|||||||
+1
-1
@@ -117,7 +117,7 @@ func (f *FrameMetadata) Unmarshal(reader io.Reader, readSourceAndLocal bool) err
|
|||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
if metaLen > 512 {
|
if metaLen > 512 {
|
||||||
return errors.New("invalid metalen ", metaLen).AtError()
|
return errors.New("invalid metalen ", metaLen)
|
||||||
}
|
}
|
||||||
|
|
||||||
b := buf.New()
|
b := buf.New()
|
||||||
|
|||||||
@@ -351,7 +351,7 @@ func (w *ServerWorker) handleFrame(ctx context.Context, reader *buf.BufferedRead
|
|||||||
err = w.handleStatusKeep(&meta, reader)
|
err = w.handleStatusKeep(&meta, reader)
|
||||||
default:
|
default:
|
||||||
status := meta.SessionStatus
|
status := meta.SessionStatus
|
||||||
return errors.New("unknown status: ", status).AtError()
|
return errors.New("unknown status: ", status)
|
||||||
}
|
}
|
||||||
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
|||||||
@@ -7,7 +7,7 @@ import (
|
|||||||
|
|
||||||
func (u *User) GetTypedAccount() (Account, error) {
|
func (u *User) GetTypedAccount() (Account, error) {
|
||||||
if u.GetAccount() == nil {
|
if u.GetAccount() == nil {
|
||||||
return nil, errors.New("Account is missing").AtWarning()
|
return nil, errors.New("Account is missing")
|
||||||
}
|
}
|
||||||
|
|
||||||
rawAccount, err := u.Account.GetInstance()
|
rawAccount, err := u.Account.GetInstance()
|
||||||
|
|||||||
@@ -1,53 +0,0 @@
|
|||||||
package singbridge
|
|
||||||
|
|
||||||
import (
|
|
||||||
M "github.com/sagernet/sing/common/metadata"
|
|
||||||
N "github.com/sagernet/sing/common/network"
|
|
||||||
"github.com/xtls/xray-core/common/errors"
|
|
||||||
"github.com/xtls/xray-core/common/net"
|
|
||||||
)
|
|
||||||
|
|
||||||
func ToNetwork(network string) net.Network {
|
|
||||||
switch N.NetworkName(network) {
|
|
||||||
case N.NetworkTCP:
|
|
||||||
return net.Network_TCP
|
|
||||||
case N.NetworkUDP:
|
|
||||||
return net.Network_UDP
|
|
||||||
default:
|
|
||||||
return net.Network_Unknown
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func ToDestination(socksaddr M.Socksaddr, network net.Network) (net.Destination, error) {
|
|
||||||
// IsFqdn() implicitly checks if the domain name is valid
|
|
||||||
if socksaddr.IsFqdn() {
|
|
||||||
return net.Destination{
|
|
||||||
Network: network,
|
|
||||||
Address: net.DomainAddress(socksaddr.Fqdn),
|
|
||||||
Port: net.Port(socksaddr.Port),
|
|
||||||
}, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// IsIP() implicitly checks if the IP address is valid
|
|
||||||
if socksaddr.IsIP() {
|
|
||||||
return net.Destination{
|
|
||||||
Network: network,
|
|
||||||
Address: net.IPAddress(socksaddr.Addr.AsSlice()),
|
|
||||||
Port: net.Port(socksaddr.Port),
|
|
||||||
}, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
return net.Destination{}, errors.New("invalid socks address: ", socksaddr)
|
|
||||||
}
|
|
||||||
|
|
||||||
func ToSocksaddr(destination net.Destination) M.Socksaddr {
|
|
||||||
var addr M.Socksaddr
|
|
||||||
switch destination.Address.Family() {
|
|
||||||
case net.AddressFamilyDomain:
|
|
||||||
addr.Fqdn = destination.Address.Domain()
|
|
||||||
default:
|
|
||||||
addr.Addr = M.AddrFromIP(destination.Address.IP())
|
|
||||||
}
|
|
||||||
addr.Port = uint16(destination.Port)
|
|
||||||
return addr
|
|
||||||
}
|
|
||||||
@@ -1,72 +0,0 @@
|
|||||||
package singbridge
|
|
||||||
|
|
||||||
import (
|
|
||||||
"context"
|
|
||||||
"os"
|
|
||||||
|
|
||||||
M "github.com/sagernet/sing/common/metadata"
|
|
||||||
N "github.com/sagernet/sing/common/network"
|
|
||||||
"github.com/xtls/xray-core/common/net"
|
|
||||||
"github.com/xtls/xray-core/common/net/cnc"
|
|
||||||
"github.com/xtls/xray-core/common/session"
|
|
||||||
"github.com/xtls/xray-core/proxy"
|
|
||||||
"github.com/xtls/xray-core/transport"
|
|
||||||
"github.com/xtls/xray-core/transport/internet"
|
|
||||||
"github.com/xtls/xray-core/transport/pipe"
|
|
||||||
)
|
|
||||||
|
|
||||||
var _ N.Dialer = (*XrayDialer)(nil)
|
|
||||||
|
|
||||||
type XrayDialer struct {
|
|
||||||
internet.Dialer
|
|
||||||
}
|
|
||||||
|
|
||||||
func NewDialer(dialer internet.Dialer) *XrayDialer {
|
|
||||||
return &XrayDialer{dialer}
|
|
||||||
}
|
|
||||||
|
|
||||||
func (d *XrayDialer) DialContext(ctx context.Context, network string, destination M.Socksaddr) (net.Conn, error) {
|
|
||||||
dest, err := ToDestination(destination, ToNetwork(network))
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
return d.Dialer.Dial(ctx, dest)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (d *XrayDialer) ListenPacket(ctx context.Context, destination M.Socksaddr) (net.PacketConn, error) {
|
|
||||||
return nil, os.ErrInvalid
|
|
||||||
}
|
|
||||||
|
|
||||||
type XrayOutboundDialer struct {
|
|
||||||
outbound proxy.Outbound
|
|
||||||
dialer internet.Dialer
|
|
||||||
}
|
|
||||||
|
|
||||||
func NewOutboundDialer(outbound proxy.Outbound, dialer internet.Dialer) *XrayOutboundDialer {
|
|
||||||
return &XrayOutboundDialer{outbound, dialer}
|
|
||||||
}
|
|
||||||
|
|
||||||
func (d *XrayOutboundDialer) DialContext(ctx context.Context, network string, destination M.Socksaddr) (net.Conn, error) {
|
|
||||||
dest, err := ToDestination(destination, ToNetwork(network))
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
outbounds := session.OutboundsFromContext(ctx)
|
|
||||||
if len(outbounds) == 0 {
|
|
||||||
outbounds = []*session.Outbound{{}}
|
|
||||||
ctx = session.ContextWithOutbounds(ctx, outbounds)
|
|
||||||
}
|
|
||||||
ob := outbounds[len(outbounds)-1]
|
|
||||||
ob.Target = dest
|
|
||||||
|
|
||||||
opts := []pipe.Option{pipe.WithSizeLimit(64 * 1024)}
|
|
||||||
uplinkReader, uplinkWriter := pipe.New(opts...)
|
|
||||||
downlinkReader, downlinkWriter := pipe.New(opts...)
|
|
||||||
conn := cnc.NewConnection(cnc.ConnectionInputMulti(downlinkWriter), cnc.ConnectionOutputMulti(uplinkReader))
|
|
||||||
go d.outbound.Process(ctx, &transport.Link{Reader: downlinkReader, Writer: uplinkWriter}, d.dialer)
|
|
||||||
return conn, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (d *XrayOutboundDialer) ListenPacket(ctx context.Context, destination M.Socksaddr) (net.PacketConn, error) {
|
|
||||||
return nil, os.ErrInvalid
|
|
||||||
}
|
|
||||||
@@ -1,10 +0,0 @@
|
|||||||
package singbridge
|
|
||||||
|
|
||||||
import E "github.com/sagernet/sing/common/exceptions"
|
|
||||||
|
|
||||||
func ReturnError(err error) error {
|
|
||||||
if E.IsClosedOrCanceled(err) {
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
@@ -1,58 +0,0 @@
|
|||||||
package singbridge
|
|
||||||
|
|
||||||
import (
|
|
||||||
"context"
|
|
||||||
"io"
|
|
||||||
|
|
||||||
M "github.com/sagernet/sing/common/metadata"
|
|
||||||
N "github.com/sagernet/sing/common/network"
|
|
||||||
"github.com/xtls/xray-core/common/buf"
|
|
||||||
"github.com/xtls/xray-core/common/errors"
|
|
||||||
"github.com/xtls/xray-core/common/net"
|
|
||||||
"github.com/xtls/xray-core/features/routing"
|
|
||||||
"github.com/xtls/xray-core/transport"
|
|
||||||
)
|
|
||||||
|
|
||||||
var (
|
|
||||||
_ N.TCPConnectionHandler = (*Dispatcher)(nil)
|
|
||||||
_ N.UDPConnectionHandler = (*Dispatcher)(nil)
|
|
||||||
)
|
|
||||||
|
|
||||||
type Dispatcher struct {
|
|
||||||
upstream routing.Dispatcher
|
|
||||||
newErrorFunc func(values ...any) *errors.Error
|
|
||||||
}
|
|
||||||
|
|
||||||
func NewDispatcher(dispatcher routing.Dispatcher, newErrorFunc func(values ...any) *errors.Error) *Dispatcher {
|
|
||||||
return &Dispatcher{
|
|
||||||
upstream: dispatcher,
|
|
||||||
newErrorFunc: newErrorFunc,
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func (d *Dispatcher) NewConnection(ctx context.Context, conn net.Conn, metadata M.Metadata) error {
|
|
||||||
dest, err := ToDestination(metadata.Destination, net.Network_TCP)
|
|
||||||
if err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
xConn := NewConn(conn)
|
|
||||||
return d.upstream.DispatchLink(ctx, dest, &transport.Link{
|
|
||||||
Reader: xConn,
|
|
||||||
Writer: xConn,
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
func (d *Dispatcher) NewPacketConnection(ctx context.Context, conn N.PacketConn, metadata M.Metadata) error {
|
|
||||||
dest, err := ToDestination(metadata.Destination, net.Network_UDP)
|
|
||||||
if err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
return d.upstream.DispatchLink(ctx, dest, &transport.Link{
|
|
||||||
Reader: buf.NewPacketReader(conn.(io.Reader)),
|
|
||||||
Writer: buf.NewWriter(conn.(io.Writer)),
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
func (d *Dispatcher) NewError(ctx context.Context, err error) {
|
|
||||||
errors.LogInfo(ctx, err.Error())
|
|
||||||
}
|
|
||||||
@@ -1,70 +0,0 @@
|
|||||||
package singbridge
|
|
||||||
|
|
||||||
import (
|
|
||||||
"context"
|
|
||||||
|
|
||||||
"github.com/sagernet/sing/common/logger"
|
|
||||||
"github.com/xtls/xray-core/common/errors"
|
|
||||||
)
|
|
||||||
|
|
||||||
var _ logger.ContextLogger = (*XrayLogger)(nil)
|
|
||||||
|
|
||||||
type XrayLogger struct {
|
|
||||||
newError func(values ...any) *errors.Error
|
|
||||||
}
|
|
||||||
|
|
||||||
func NewLogger(newErrorFunc func(values ...any) *errors.Error) *XrayLogger {
|
|
||||||
return &XrayLogger{
|
|
||||||
newErrorFunc,
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func (l *XrayLogger) Trace(args ...any) {
|
|
||||||
}
|
|
||||||
|
|
||||||
func (l *XrayLogger) Debug(args ...any) {
|
|
||||||
errors.LogDebug(context.Background(), args...)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (l *XrayLogger) Info(args ...any) {
|
|
||||||
errors.LogInfo(context.Background(), args...)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (l *XrayLogger) Warn(args ...any) {
|
|
||||||
errors.LogWarning(context.Background(), args...)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (l *XrayLogger) Error(args ...any) {
|
|
||||||
errors.LogError(context.Background(), args...)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (l *XrayLogger) Fatal(args ...any) {
|
|
||||||
}
|
|
||||||
|
|
||||||
func (l *XrayLogger) Panic(args ...any) {
|
|
||||||
}
|
|
||||||
|
|
||||||
func (l *XrayLogger) TraceContext(ctx context.Context, args ...any) {
|
|
||||||
}
|
|
||||||
|
|
||||||
func (l *XrayLogger) DebugContext(ctx context.Context, args ...any) {
|
|
||||||
errors.LogDebug(ctx, args...)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (l *XrayLogger) InfoContext(ctx context.Context, args ...any) {
|
|
||||||
errors.LogInfo(ctx, args...)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (l *XrayLogger) WarnContext(ctx context.Context, args ...any) {
|
|
||||||
errors.LogWarning(ctx, args...)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (l *XrayLogger) ErrorContext(ctx context.Context, args ...any) {
|
|
||||||
errors.LogError(ctx, args...)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (l *XrayLogger) FatalContext(ctx context.Context, args ...any) {
|
|
||||||
}
|
|
||||||
|
|
||||||
func (l *XrayLogger) PanicContext(ctx context.Context, args ...any) {
|
|
||||||
}
|
|
||||||
@@ -1,107 +0,0 @@
|
|||||||
package singbridge
|
|
||||||
|
|
||||||
import (
|
|
||||||
"context"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
B "github.com/sagernet/sing/common/buf"
|
|
||||||
"github.com/sagernet/sing/common/bufio"
|
|
||||||
M "github.com/sagernet/sing/common/metadata"
|
|
||||||
"github.com/xtls/xray-core/common"
|
|
||||||
"github.com/xtls/xray-core/common/buf"
|
|
||||||
"github.com/xtls/xray-core/common/net"
|
|
||||||
"github.com/xtls/xray-core/common/signal"
|
|
||||||
"github.com/xtls/xray-core/transport"
|
|
||||||
)
|
|
||||||
|
|
||||||
func CopyPacketConn(ctx context.Context, inboundConn net.Conn, link *transport.Link, destination net.Destination, serverConn net.PacketConn) error {
|
|
||||||
cancel := func() {
|
|
||||||
common.Interrupt(link.Reader)
|
|
||||||
common.Interrupt(serverConn)
|
|
||||||
}
|
|
||||||
conn := &PacketConnWrapper{
|
|
||||||
Reader: link.Reader,
|
|
||||||
Writer: link.Writer,
|
|
||||||
Dest: destination,
|
|
||||||
Conn: inboundConn,
|
|
||||||
T: signal.CancelAfterInactivity(ctx, cancel, 300*time.Second),
|
|
||||||
}
|
|
||||||
return ReturnError(bufio.CopyPacketConn(ctx, conn, bufio.NewPacketConn(serverConn)))
|
|
||||||
}
|
|
||||||
|
|
||||||
type PacketConnWrapper struct {
|
|
||||||
buf.Reader
|
|
||||||
buf.Writer
|
|
||||||
net.Conn
|
|
||||||
Dest net.Destination
|
|
||||||
cached buf.MultiBuffer
|
|
||||||
|
|
||||||
// A simple patch to avoid goroutine leak since sing infra cannot awake read block by write err
|
|
||||||
T *signal.ActivityTimer
|
|
||||||
}
|
|
||||||
|
|
||||||
func (w *PacketConnWrapper) ReadPacket(buffer *B.Buffer) (addr M.Socksaddr, err error) {
|
|
||||||
w.T.Update()
|
|
||||||
defer func() {
|
|
||||||
if err != nil {
|
|
||||||
// uplinkonly
|
|
||||||
w.T.SetTimeout(2 * time.Second)
|
|
||||||
}
|
|
||||||
}()
|
|
||||||
if w.cached != nil {
|
|
||||||
mb, bb := buf.SplitFirst(w.cached)
|
|
||||||
if bb == nil {
|
|
||||||
w.cached = nil
|
|
||||||
} else {
|
|
||||||
buffer.Write(bb.Bytes())
|
|
||||||
w.cached = mb
|
|
||||||
var destination net.Destination
|
|
||||||
if bb.UDP != nil {
|
|
||||||
destination = *bb.UDP
|
|
||||||
} else {
|
|
||||||
destination = w.Dest
|
|
||||||
}
|
|
||||||
bb.Release()
|
|
||||||
return ToSocksaddr(destination), nil
|
|
||||||
}
|
|
||||||
}
|
|
||||||
mb, err := w.ReadMultiBuffer()
|
|
||||||
nb, bb := buf.SplitFirst(mb)
|
|
||||||
if bb == nil {
|
|
||||||
return M.Socksaddr{}, nil
|
|
||||||
} else {
|
|
||||||
buffer.Write(bb.Bytes())
|
|
||||||
w.cached = nb
|
|
||||||
var destination net.Destination
|
|
||||||
if bb.UDP != nil {
|
|
||||||
destination = *bb.UDP
|
|
||||||
} else {
|
|
||||||
destination = w.Dest
|
|
||||||
}
|
|
||||||
bb.Release()
|
|
||||||
return ToSocksaddr(destination), nil
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func (w *PacketConnWrapper) WritePacket(buffer *B.Buffer, destination M.Socksaddr) (err error) {
|
|
||||||
w.T.Update()
|
|
||||||
defer func() {
|
|
||||||
if err != nil {
|
|
||||||
// downlinkonly
|
|
||||||
w.T.SetTimeout(5 * time.Second)
|
|
||||||
}
|
|
||||||
}()
|
|
||||||
endpoint, err := ToDestination(destination, net.Network_UDP)
|
|
||||||
if err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
vBuf := buf.New()
|
|
||||||
vBuf.Write(buffer.Bytes())
|
|
||||||
vBuf.UDP = &endpoint
|
|
||||||
return w.WriteMultiBuffer(buf.MultiBuffer{vBuf})
|
|
||||||
}
|
|
||||||
|
|
||||||
func (w *PacketConnWrapper) Close() error {
|
|
||||||
buf.ReleaseMulti(w.cached)
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
@@ -1,81 +0,0 @@
|
|||||||
package singbridge
|
|
||||||
|
|
||||||
import (
|
|
||||||
"context"
|
|
||||||
"io"
|
|
||||||
"net"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
"github.com/sagernet/sing/common/bufio"
|
|
||||||
"github.com/xtls/xray-core/common"
|
|
||||||
"github.com/xtls/xray-core/common/buf"
|
|
||||||
"github.com/xtls/xray-core/common/signal"
|
|
||||||
"github.com/xtls/xray-core/transport"
|
|
||||||
)
|
|
||||||
|
|
||||||
func CopyConn(ctx context.Context, inboundConn net.Conn, link *transport.Link, serverConn net.Conn) error {
|
|
||||||
conn := &PipeConnWrapper{
|
|
||||||
W: link.Writer,
|
|
||||||
Conn: inboundConn,
|
|
||||||
}
|
|
||||||
if ir, ok := link.Reader.(io.Reader); ok {
|
|
||||||
conn.R = ir
|
|
||||||
} else {
|
|
||||||
conn.R = &buf.BufferedReader{Reader: link.Reader}
|
|
||||||
}
|
|
||||||
cancel := func() {
|
|
||||||
common.Interrupt(link.Reader)
|
|
||||||
common.Interrupt(serverConn)
|
|
||||||
}
|
|
||||||
conn.T = signal.CancelAfterInactivity(ctx, cancel, 300*time.Second)
|
|
||||||
return ReturnError(bufio.CopyConn(ctx, conn, serverConn))
|
|
||||||
}
|
|
||||||
|
|
||||||
type PipeConnWrapper struct {
|
|
||||||
R io.Reader
|
|
||||||
W buf.Writer
|
|
||||||
net.Conn
|
|
||||||
|
|
||||||
// A simple patch to avoid goroutine leak since sing infra cannot awake read block by write err
|
|
||||||
T *signal.ActivityTimer
|
|
||||||
}
|
|
||||||
|
|
||||||
func (w *PipeConnWrapper) Close() error {
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (w *PipeConnWrapper) Read(b []byte) (n int, err error) {
|
|
||||||
w.T.Update()
|
|
||||||
n, err = w.R.Read(b)
|
|
||||||
if err != nil {
|
|
||||||
// uplinkonly
|
|
||||||
w.T.SetTimeout(2 * time.Second)
|
|
||||||
}
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
func (w *PipeConnWrapper) Write(p []byte) (n int, err error) {
|
|
||||||
w.T.Update()
|
|
||||||
n = len(p)
|
|
||||||
var mb buf.MultiBuffer
|
|
||||||
pLen := len(p)
|
|
||||||
for pLen > 0 {
|
|
||||||
buffer := buf.New()
|
|
||||||
if pLen > buf.Size {
|
|
||||||
_, err = buffer.Write(p[:buf.Size])
|
|
||||||
p = p[buf.Size:]
|
|
||||||
} else {
|
|
||||||
buffer.Write(p)
|
|
||||||
}
|
|
||||||
pLen -= int(buffer.Len())
|
|
||||||
mb = append(mb, buffer)
|
|
||||||
}
|
|
||||||
err = w.W.WriteMultiBuffer(mb)
|
|
||||||
if err != nil {
|
|
||||||
n = 0
|
|
||||||
buf.ReleaseMulti(mb)
|
|
||||||
// downlinkonly
|
|
||||||
w.T.SetTimeout(5 * time.Second)
|
|
||||||
}
|
|
||||||
return
|
|
||||||
}
|
|
||||||
@@ -1,66 +0,0 @@
|
|||||||
package singbridge
|
|
||||||
|
|
||||||
import (
|
|
||||||
"time"
|
|
||||||
|
|
||||||
"github.com/sagernet/sing/common"
|
|
||||||
"github.com/sagernet/sing/common/bufio"
|
|
||||||
N "github.com/sagernet/sing/common/network"
|
|
||||||
"github.com/xtls/xray-core/common/buf"
|
|
||||||
"github.com/xtls/xray-core/common/net"
|
|
||||||
)
|
|
||||||
|
|
||||||
var (
|
|
||||||
_ buf.Reader = (*Conn)(nil)
|
|
||||||
_ buf.TimeoutReader = (*Conn)(nil)
|
|
||||||
_ buf.Writer = (*Conn)(nil)
|
|
||||||
)
|
|
||||||
|
|
||||||
type Conn struct {
|
|
||||||
net.Conn
|
|
||||||
writer N.VectorisedWriter
|
|
||||||
}
|
|
||||||
|
|
||||||
func NewConn(conn net.Conn) *Conn {
|
|
||||||
writer, _ := bufio.CreateVectorisedWriter(conn)
|
|
||||||
return &Conn{
|
|
||||||
Conn: conn,
|
|
||||||
writer: writer,
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *Conn) ReadMultiBuffer() (buf.MultiBuffer, error) {
|
|
||||||
buffer, err := buf.ReadBuffer(c.Conn)
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
return buf.MultiBuffer{buffer}, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *Conn) ReadMultiBufferTimeout(duration time.Duration) (buf.MultiBuffer, error) {
|
|
||||||
err := c.SetReadDeadline(time.Now().Add(duration))
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
defer c.SetReadDeadline(time.Time{})
|
|
||||||
return c.ReadMultiBuffer()
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *Conn) WriteMultiBuffer(bufferList buf.MultiBuffer) error {
|
|
||||||
defer buf.ReleaseMulti(bufferList)
|
|
||||||
if c.writer != nil {
|
|
||||||
bytesList := make([][]byte, len(bufferList))
|
|
||||||
for i, buffer := range bufferList {
|
|
||||||
bytesList[i] = buffer.Bytes()
|
|
||||||
}
|
|
||||||
return common.Error(bufio.WriteVectorised(c.writer, bytesList))
|
|
||||||
}
|
|
||||||
// Since this conn is only used by tun, we don't force buffer writes to merge.
|
|
||||||
for _, buffer := range bufferList {
|
|
||||||
_, err := c.Conn.Write(buffer.Bytes())
|
|
||||||
if err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
+2
-2
@@ -16,7 +16,7 @@ var typeCreatorRegistry = make(map[reflect.Type]ConfigCreator)
|
|||||||
func RegisterConfig(config interface{}, configCreator ConfigCreator) error {
|
func RegisterConfig(config interface{}, configCreator ConfigCreator) error {
|
||||||
configType := reflect.TypeOf(config)
|
configType := reflect.TypeOf(config)
|
||||||
if _, found := typeCreatorRegistry[configType]; found {
|
if _, found := typeCreatorRegistry[configType]; found {
|
||||||
return errors.New(configType.Name() + " is already registered").AtError()
|
return errors.New(configType.Name() + " is already registered")
|
||||||
}
|
}
|
||||||
typeCreatorRegistry[configType] = configCreator
|
typeCreatorRegistry[configType] = configCreator
|
||||||
return nil
|
return nil
|
||||||
@@ -27,7 +27,7 @@ func CreateObject(ctx context.Context, config interface{}) (interface{}, error)
|
|||||||
configType := reflect.TypeOf(config)
|
configType := reflect.TypeOf(config)
|
||||||
creator, found := typeCreatorRegistry[configType]
|
creator, found := typeCreatorRegistry[configType]
|
||||||
if !found {
|
if !found {
|
||||||
return nil, errors.New(configType.String() + " is not registered").AtError()
|
return nil, errors.New(configType.String() + " is not registered")
|
||||||
}
|
}
|
||||||
return creator(ctx, config)
|
return creator(ctx, config)
|
||||||
}
|
}
|
||||||
|
|||||||
+4
-4
@@ -125,7 +125,7 @@ func LoadConfig(formatName string, input interface{}) (*Config, error) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
if f == "" {
|
if f == "" {
|
||||||
return nil, errors.New("Failed to get format of ", file).AtWarning()
|
return nil, errors.New("Failed to get format of ", file)
|
||||||
}
|
}
|
||||||
|
|
||||||
if f == "protobuf" {
|
if f == "protobuf" {
|
||||||
@@ -142,7 +142,7 @@ func LoadConfig(formatName string, input interface{}) (*Config, error) {
|
|||||||
if len(v) == 1 {
|
if len(v) == 1 {
|
||||||
return configLoaderByName["protobuf"].Loader(v)
|
return configLoaderByName["protobuf"].Loader(v)
|
||||||
} else {
|
} else {
|
||||||
return nil, errors.New("Only one protobuf config file is allowed").AtWarning()
|
return nil, errors.New("Only one protobuf config file is allowed")
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -152,11 +152,11 @@ func LoadConfig(formatName string, input interface{}) (*Config, error) {
|
|||||||
if f, found := configLoaderByName[formatName]; found {
|
if f, found := configLoaderByName[formatName]; found {
|
||||||
return f.Loader(v)
|
return f.Loader(v)
|
||||||
} else {
|
} else {
|
||||||
return nil, errors.New("Unable to load config in", formatName).AtWarning()
|
return nil, errors.New("Unable to load config in", formatName)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
return nil, errors.New("Unable to load config").AtWarning()
|
return nil, errors.New("Unable to load config")
|
||||||
}
|
}
|
||||||
|
|
||||||
func loadProtobufConfig(data []byte) (*Config, error) {
|
func loadProtobufConfig(data []byte) (*Config, error) {
|
||||||
|
|||||||
@@ -13,7 +13,7 @@ type FakeDNSEngine interface {
|
|||||||
|
|
||||||
var (
|
var (
|
||||||
FakeIPv4Pool = "198.18.0.0/15"
|
FakeIPv4Pool = "198.18.0.0/15"
|
||||||
FakeIPv6Pool = "fc00::/18"
|
FakeIPv6Pool = "2001:2::/48"
|
||||||
)
|
)
|
||||||
|
|
||||||
type FakeDNSEngineRev0 interface {
|
type FakeDNSEngineRev0 interface {
|
||||||
|
|||||||
@@ -18,21 +18,19 @@ require (
|
|||||||
github.com/pires/go-proxyproto v0.15.0
|
github.com/pires/go-proxyproto v0.15.0
|
||||||
github.com/refraction-networking/utls v1.8.3-0.20260301010127-aa6edf4b11af
|
github.com/refraction-networking/utls v1.8.3-0.20260301010127-aa6edf4b11af
|
||||||
github.com/robfig/cron/v3 v3.0.1
|
github.com/robfig/cron/v3 v3.0.1
|
||||||
github.com/sagernet/sing v0.5.1
|
|
||||||
github.com/sagernet/sing-shadowsocks v0.2.7
|
|
||||||
github.com/stretchr/testify v1.12.1
|
github.com/stretchr/testify v1.12.1
|
||||||
github.com/vishvananda/netlink v1.3.1
|
github.com/vishvananda/netlink v1.3.1
|
||||||
github.com/xtls/reality v0.0.0-20260908062103-8cdf7bf9c7f0
|
github.com/xtls/reality v0.0.0-20260908062103-8cdf7bf9c7f0
|
||||||
go4.org/netipx v0.0.0-20231129151722-fdeea329fbba
|
go4.org/netipx v0.0.0-20231129151722-fdeea329fbba
|
||||||
golang.org/x/crypto v0.55.0
|
golang.org/x/crypto v0.57.0
|
||||||
golang.org/x/exp v0.0.0-20240506185415-9bf2ced13842
|
golang.org/x/exp v0.0.0-20240506185415-9bf2ced13842
|
||||||
golang.org/x/net v0.58.0
|
golang.org/x/net v0.59.0
|
||||||
golang.org/x/sync v0.22.0
|
golang.org/x/sync v0.23.0
|
||||||
golang.org/x/sys v0.47.0
|
golang.org/x/sys v0.48.0
|
||||||
golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2
|
golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2
|
||||||
golang.zx2c4.com/wireguard v0.0.0-20250521234502-f333402bd9cb
|
golang.zx2c4.com/wireguard v0.0.0-20250521234502-f333402bd9cb
|
||||||
golang.zx2c4.com/wireguard/windows v1.0.1
|
golang.zx2c4.com/wireguard/windows v1.1.1
|
||||||
google.golang.org/grpc v1.83.2
|
google.golang.org/grpc v1.84.0
|
||||||
google.golang.org/protobuf v1.36.12
|
google.golang.org/protobuf v1.36.12
|
||||||
gvisor.dev/gvisor v0.0.0-20260122175437-89a5d21be8f0
|
gvisor.dev/gvisor v0.0.0-20260122175437-89a5d21be8f0
|
||||||
h12.io/socks v1.0.3
|
h12.io/socks v1.0.3
|
||||||
@@ -57,9 +55,9 @@ require (
|
|||||||
github.com/vishvananda/netns v0.0.5 // indirect
|
github.com/vishvananda/netns v0.0.5 // indirect
|
||||||
github.com/wlynxg/anet v0.0.5 // indirect
|
github.com/wlynxg/anet v0.0.5 // indirect
|
||||||
go.yaml.in/yaml/v3 v3.0.5 // indirect
|
go.yaml.in/yaml/v3 v3.0.5 // indirect
|
||||||
golang.org/x/text v0.41.0 // indirect
|
golang.org/x/text v0.42.0 // indirect
|
||||||
golang.org/x/time v0.14.0 // indirect
|
golang.org/x/time v0.14.0 // indirect
|
||||||
golang.org/x/tools v0.49.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
|
gopkg.in/yaml.v2 v2.4.0 // indirect
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -2,16 +2,10 @@ github.com/andybalholm/brotli v1.0.6 h1:Yf9fFpf49Zrxb9NlQaluyE92/+X7UVHlhMNJN2sx
|
|||||||
github.com/andybalholm/brotli v1.0.6/go.mod h1:fO7iG3H7G2nSZ7m0zPUDn85XEX2GTukHGRSepvi9Eig=
|
github.com/andybalholm/brotli v1.0.6/go.mod h1:fO7iG3H7G2nSZ7m0zPUDn85XEX2GTukHGRSepvi9Eig=
|
||||||
github.com/apernet/quic-go v0.61.1-0.20260806010916-184d081eef3e h1:5mgtR5gwIgBKMiGI1QdXldZZ+SNor06Nbu1wCBulQBg=
|
github.com/apernet/quic-go v0.61.1-0.20260806010916-184d081eef3e h1:5mgtR5gwIgBKMiGI1QdXldZZ+SNor06Nbu1wCBulQBg=
|
||||||
github.com/apernet/quic-go v0.61.1-0.20260806010916-184d081eef3e/go.mod h1:x7qxEvX6MCVtDuBKHj3E+88+BtrbEMuAL5qGUKItjW8=
|
github.com/apernet/quic-go v0.61.1-0.20260806010916-184d081eef3e/go.mod h1:x7qxEvX6MCVtDuBKHj3E+88+BtrbEMuAL5qGUKItjW8=
|
||||||
github.com/cespare/xxhash/v2 v2.3.0 h1:UL815xU9SqsFlibzuggzjXhog7bL6oX9BbNZnL2UFvs=
|
|
||||||
github.com/cespare/xxhash/v2 v2.3.0/go.mod h1:VGX0DQ3Q6kWi7AoAeZDth3/j3BFtOZR5XLFGgcrjCOs=
|
|
||||||
github.com/cloudflare/circl v1.6.5 h1:O64F26HEqNhznd/hrC5KZXVKYuKM2rx4deZDTc4ihQA=
|
github.com/cloudflare/circl v1.6.5 h1:O64F26HEqNhznd/hrC5KZXVKYuKM2rx4deZDTc4ihQA=
|
||||||
github.com/cloudflare/circl v1.6.5/go.mod h1:h5LNyxAc5nTue9DS5jT+48en2PSDYt3zdGnz5OstK6c=
|
github.com/cloudflare/circl v1.6.5/go.mod h1:h5LNyxAc5nTue9DS5jT+48en2PSDYt3zdGnz5OstK6c=
|
||||||
github.com/ghodss/yaml v1.0.1-0.20220118164431-d8423dcdf344 h1:Arcl6UOIS/kgO2nW3A65HN+7CMjSDP/gofXL4CZt1V4=
|
github.com/ghodss/yaml v1.0.1-0.20220118164431-d8423dcdf344 h1:Arcl6UOIS/kgO2nW3A65HN+7CMjSDP/gofXL4CZt1V4=
|
||||||
github.com/ghodss/yaml v1.0.1-0.20220118164431-d8423dcdf344/go.mod h1:GIjDIg/heH5DOkXY3YJ/wNhfHsQHoXGjl8G8amsYQ1I=
|
github.com/ghodss/yaml v1.0.1-0.20220118164431-d8423dcdf344/go.mod h1:GIjDIg/heH5DOkXY3YJ/wNhfHsQHoXGjl8G8amsYQ1I=
|
||||||
github.com/go-logr/logr v1.4.3 h1:CjnDlHq8ikf6E492q6eKboGOC0T8CDaOvkHCIg8idEI=
|
|
||||||
github.com/go-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 h1:HSQxCeh5YZH3EL3W39ixjtyaEhcWSXQHtHnMBzSs474=
|
||||||
github.com/go-quicktest/qt v1.102.0/go.mod h1:p4lGIVX+8Wa6ZPNDvqcxq36XpUDLh42FLetFU7odllI=
|
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=
|
github.com/golang/mock v1.7.0-rc.1 h1:YojYx61/OLFsiv6Rw1Z96LpldJIy31o+UHmwAUMJ6/U=
|
||||||
@@ -76,10 +70,6 @@ github.com/robfig/cron/v3 v3.0.1 h1:WdRxkvbJztn8LMz/QEvLN5sBU+xKpSqwwUO1Pjr4qDs=
|
|||||||
github.com/robfig/cron/v3 v3.0.1/go.mod h1:eQICP3HwyT7UooqI/z+Ov+PtYAWygg1TEWWzGIFLtro=
|
github.com/robfig/cron/v3 v3.0.1/go.mod h1:eQICP3HwyT7UooqI/z+Ov+PtYAWygg1TEWWzGIFLtro=
|
||||||
github.com/rogpeppe/go-internal v1.16.0 h1:O9DK+vNMDVGLr2BeZqmpLeMjiMNkuXfcqntWbZV6S5g=
|
github.com/rogpeppe/go-internal v1.16.0 h1:O9DK+vNMDVGLr2BeZqmpLeMjiMNkuXfcqntWbZV6S5g=
|
||||||
github.com/rogpeppe/go-internal v1.16.0/go.mod h1:DrUVZyrJU+txYW5/1kwtXQSMFio52ZOxX7yM1VHvnxs=
|
github.com/rogpeppe/go-internal v1.16.0/go.mod h1:DrUVZyrJU+txYW5/1kwtXQSMFio52ZOxX7yM1VHvnxs=
|
||||||
github.com/sagernet/sing v0.5.1 h1:mhL/MZVq0TjuvHcpYcFtmSD1BFOxZ/+8ofbNZcg1k1Y=
|
|
||||||
github.com/sagernet/sing v0.5.1/go.mod h1:ARkL0gM13/Iv5VCZmci/NuoOlePoIsW0m7BWfln/Hak=
|
|
||||||
github.com/sagernet/sing-shadowsocks v0.2.7 h1:zaopR1tbHEw5Nk6FAkM05wCslV6ahVegEZaKMv9ipx8=
|
|
||||||
github.com/sagernet/sing-shadowsocks v0.2.7/go.mod h1:0rIKJZBR65Qi0zwdKezt4s57y/Tl1ofkaq6NlkzVuyE=
|
|
||||||
github.com/stretchr/testify v1.12.1 h1:EuwCh5fleGS7H32xRwO3wRGT7DxrDhLAT6FF8MpWDWE=
|
github.com/stretchr/testify v1.12.1 h1:EuwCh5fleGS7H32xRwO3wRGT7DxrDhLAT6FF8MpWDWE=
|
||||||
github.com/stretchr/testify v1.12.1/go.mod h1:MDEgiDPPsNp5cuIrHPPCyornHKgEVbtFUmoNlxoYthg=
|
github.com/stretchr/testify v1.12.1/go.mod h1:MDEgiDPPsNp5cuIrHPPCyornHKgEVbtFUmoNlxoYthg=
|
||||||
github.com/vishvananda/netlink v1.3.1 h1:3AEMt62VKqz90r0tmNhog0r/PpWKmrEShJU0wJW6bV0=
|
github.com/vishvananda/netlink v1.3.1 h1:3AEMt62VKqz90r0tmNhog0r/PpWKmrEShJU0wJW6bV0=
|
||||||
@@ -91,18 +81,6 @@ github.com/wlynxg/anet v0.0.5/go.mod h1:eay5PRQr7fIVAMbTbchTnO9gG65Hg/uYGdc7mguH
|
|||||||
github.com/xtls/reality v0.0.0-20260908062103-8cdf7bf9c7f0 h1:rb+fKQFhz+5I2PPuQsNYxI5mUU840XWYtRF0ZBjvkws=
|
github.com/xtls/reality v0.0.0-20260908062103-8cdf7bf9c7f0 h1:rb+fKQFhz+5I2PPuQsNYxI5mUU840XWYtRF0ZBjvkws=
|
||||||
github.com/xtls/reality v0.0.0-20260908062103-8cdf7bf9c7f0/go.mod h1:DsJblcWDGt76+FVqBVwbwRhxyyNJsGV48gJLch0OOWI=
|
github.com/xtls/reality v0.0.0-20260908062103-8cdf7bf9c7f0/go.mod h1:DsJblcWDGt76+FVqBVwbwRhxyyNJsGV48gJLch0OOWI=
|
||||||
github.com/yuin/goldmark v1.4.1/go.mod h1:mwnBkeHKe2W/ZEtQ+71ViKU8L12m81fl3OWwC1Zlc8k=
|
github.com/yuin/goldmark v1.4.1/go.mod h1:mwnBkeHKe2W/ZEtQ+71ViKU8L12m81fl3OWwC1Zlc8k=
|
||||||
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=
|
|
||||||
go.uber.org/mock v0.5.2 h1:LbtPTcP8A5k9WPXj54PPPbjcI4Y6lhyOZXn+VS7wNko=
|
go.uber.org/mock v0.5.2 h1:LbtPTcP8A5k9WPXj54PPPbjcI4Y6lhyOZXn+VS7wNko=
|
||||||
go.uber.org/mock v0.5.2/go.mod h1:wLlUxC2vVTPTaE3UD51E0BGOAElKrILxhVSDYQLld5o=
|
go.uber.org/mock v0.5.2/go.mod h1:wLlUxC2vVTPTaE3UD51E0BGOAElKrILxhVSDYQLld5o=
|
||||||
go.yaml.in/yaml/v3 v3.0.5 h1:N6y/pJk8buWs9NY5ERU2HSMfm+IuD/OtfdAnq6kESPw=
|
go.yaml.in/yaml/v3 v3.0.5 h1:N6y/pJk8buWs9NY5ERU2HSMfm+IuD/OtfdAnq6kESPw=
|
||||||
@@ -111,8 +89,8 @@ go4.org/netipx v0.0.0-20231129151722-fdeea329fbba h1:0b9z3AuHCjxk0x/opv64kcgZLBs
|
|||||||
go4.org/netipx v0.0.0-20231129151722-fdeea329fbba/go.mod h1:PLyyIXexvUFg3Owu6p/WfdlivPbZJsZdgWZlrGope/Y=
|
go4.org/netipx v0.0.0-20231129151722-fdeea329fbba/go.mod h1:PLyyIXexvUFg3Owu6p/WfdlivPbZJsZdgWZlrGope/Y=
|
||||||
golang.org/x/crypto v0.0.0-20190308221718-c2843e01d9a2/go.mod h1:djNgcEr1/C05ACkg1iLfiJU5Ep61QUkGW8qpdssI0+w=
|
golang.org/x/crypto v0.0.0-20190308221718-c2843e01d9a2/go.mod h1:djNgcEr1/C05ACkg1iLfiJU5Ep61QUkGW8qpdssI0+w=
|
||||||
golang.org/x/crypto v0.0.0-20191011191535-87dc89f01550/go.mod h1:yigFU9vqHzYiE8UmvKecakEJjdnWj3jj499lnFckfCI=
|
golang.org/x/crypto v0.0.0-20191011191535-87dc89f01550/go.mod h1:yigFU9vqHzYiE8UmvKecakEJjdnWj3jj499lnFckfCI=
|
||||||
golang.org/x/crypto v0.55.0 h1:+KWHjbgOaAQ66dh/YlkZKHlz9ZUlq61AFirAR9ntP8M=
|
golang.org/x/crypto v0.57.0 h1:3ZVCjf8Ggz7zneR/EHRVx68Ctf+2pmIMP2UFhh9cC6M=
|
||||||
golang.org/x/crypto v0.55.0/go.mod h1:uq0V9dE/fzQuJtbnL+2EhWOE63vo164FY8xqEnV9xis=
|
golang.org/x/crypto v0.57.0/go.mod h1:Fdz0i5U6CoizGwLda9DttjSk6qlZo25zYNtR+ycvuZA=
|
||||||
golang.org/x/exp v0.0.0-20240506185415-9bf2ced13842 h1:vr/HnozRka3pE4EsMEg1lgkXJkTFJCVUX+S/ZT6wYzM=
|
golang.org/x/exp v0.0.0-20240506185415-9bf2ced13842 h1:vr/HnozRka3pE4EsMEg1lgkXJkTFJCVUX+S/ZT6wYzM=
|
||||||
golang.org/x/exp v0.0.0-20240506185415-9bf2ced13842/go.mod h1:XtvwrStGgqGPLc4cjQfWqZHG1YFdYs6swckp8vpsjnc=
|
golang.org/x/exp v0.0.0-20240506185415-9bf2ced13842/go.mod h1:XtvwrStGgqGPLc4cjQfWqZHG1YFdYs6swckp8vpsjnc=
|
||||||
golang.org/x/lint v0.0.0-20200302205851-738671d3881b/go.mod h1:3xt1FjdF8hUf6vQPIChWIBhFzV8gjjsPE/fR3IyQdNY=
|
golang.org/x/lint v0.0.0-20200302205851-738671d3881b/go.mod h1:3xt1FjdF8hUf6vQPIChWIBhFzV8gjjsPE/fR3IyQdNY=
|
||||||
@@ -121,12 +99,12 @@ golang.org/x/mod v0.5.1/go.mod h1:5OXOZSfqPIIbmVBIIKWRFfZjPR0E5r58TLhUjH0a2Ro=
|
|||||||
golang.org/x/net v0.0.0-20190404232315-eb5bcb51f2a3/go.mod h1:t9HGtf8HONx5eT2rtn7q6eTqICYqUVnKs3thJo3Qplg=
|
golang.org/x/net v0.0.0-20190404232315-eb5bcb51f2a3/go.mod h1:t9HGtf8HONx5eT2rtn7q6eTqICYqUVnKs3thJo3Qplg=
|
||||||
golang.org/x/net v0.0.0-20190620200207-3b0461eec859/go.mod h1:z5CRVTTTmAJ677TzLLGU+0bjPO0LkuOLi4/5GtJWs/s=
|
golang.org/x/net v0.0.0-20190620200207-3b0461eec859/go.mod h1:z5CRVTTTmAJ677TzLLGU+0bjPO0LkuOLi4/5GtJWs/s=
|
||||||
golang.org/x/net v0.0.0-20211015210444-4f30a5c0130f/go.mod h1:9nx3DQGgdP8bBQD5qxJ1jj9UTztislL4KSBs9R2vV5Y=
|
golang.org/x/net v0.0.0-20211015210444-4f30a5c0130f/go.mod h1:9nx3DQGgdP8bBQD5qxJ1jj9UTztislL4KSBs9R2vV5Y=
|
||||||
golang.org/x/net v0.58.0 h1:ynWG7rqYi4ccpTEuPZ2QGWHktVEM9DMCj9yzDE0Q7To=
|
golang.org/x/net v0.59.0 h1:5zfYln+w5XCxwrnMMJPufRgNoXEaGxl0wo5GqPXyues=
|
||||||
golang.org/x/net v0.58.0/go.mod h1:YwCddHnFlT7eLQqVprV19OnhLGtc5xOKgE0RyqgfWAU=
|
golang.org/x/net v0.59.0/go.mod h1:2DA/G1UfVbCpQPeWTmMPGY7Cs2PkBkwu743bVX5PIVg=
|
||||||
golang.org/x/sync v0.0.0-20190423024810-112230192c58/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
golang.org/x/sync v0.0.0-20190423024810-112230192c58/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
||||||
golang.org/x/sync v0.0.0-20210220032951-036812b2e83c/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
golang.org/x/sync v0.0.0-20210220032951-036812b2e83c/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
|
||||||
golang.org/x/sync v0.22.0 h1:SZjpbeLmrCk4xhRSZFNZW5gFUeCeFgjekvI/+gfScek=
|
golang.org/x/sync v0.23.0 h1:KameEIfc1IkluZyXWLn39Wd4tURc6GbCiISGiZm2bQk=
|
||||||
golang.org/x/sync v0.22.0/go.mod h1:9xrNwdLfx4jkKbNva9FpL6vEN7evnE43NNNJQ2LF3+0=
|
golang.org/x/sync v0.23.0/go.mod h1:sUUOizhqBxiL6pEWpqNLUiaJn1ShEbZ6BBqskPbjZm0=
|
||||||
golang.org/x/sys v0.0.0-20190215142949-d0b11bdaac8a/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY=
|
golang.org/x/sys v0.0.0-20190215142949-d0b11bdaac8a/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY=
|
||||||
golang.org/x/sys v0.0.0-20190412213103-97732733099d/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs=
|
golang.org/x/sys v0.0.0-20190412213103-97732733099d/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs=
|
||||||
golang.org/x/sys v0.0.0-20201119102817-f84b799fce68/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs=
|
golang.org/x/sys v0.0.0-20201119102817-f84b799fce68/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs=
|
||||||
@@ -134,14 +112,14 @@ golang.org/x/sys v0.0.0-20210423082822-04245dca01da/go.mod h1:h1NjWce9XRLGQEsW7w
|
|||||||
golang.org/x/sys v0.0.0-20211019181941-9d821ace8654/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
|
golang.org/x/sys v0.0.0-20211019181941-9d821ace8654/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
|
||||||
golang.org/x/sys v0.2.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
|
golang.org/x/sys v0.2.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
|
||||||
golang.org/x/sys v0.10.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
|
golang.org/x/sys v0.10.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
|
||||||
golang.org/x/sys v0.47.0 h1:o7XGOvZQCADBQQ4Y7VNq2dRWQR7JmOUW8Kxx4ZsNgWs=
|
golang.org/x/sys v0.48.0 h1:bbX/i/6MgT9BVLM9RT1thmxL04yeTAhbEz4SyadbXoo=
|
||||||
golang.org/x/sys v0.47.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw=
|
golang.org/x/sys v0.48.0/go.mod h1:hNLxWAXmnKAxqDtdwIYC4bM9oQPEecfsnNMuSxOs3og=
|
||||||
golang.org/x/term v0.0.0-20201126162022-7de9c90e9dd1/go.mod h1:bj7SfCRtBDWHUb9snDiAeCFNEtKQo2Wmx5Cou7ajbmo=
|
golang.org/x/term v0.0.0-20201126162022-7de9c90e9dd1/go.mod h1:bj7SfCRtBDWHUb9snDiAeCFNEtKQo2Wmx5Cou7ajbmo=
|
||||||
golang.org/x/text v0.3.0/go.mod h1:NqM8EUOU14njkJ3fqMW+pc6Ldnwhi/IjpwHt7yyuwOQ=
|
golang.org/x/text v0.3.0/go.mod h1:NqM8EUOU14njkJ3fqMW+pc6Ldnwhi/IjpwHt7yyuwOQ=
|
||||||
golang.org/x/text v0.3.6/go.mod h1:5Zoc/QRtKVWzQhOtBMvqHzDpF6irO9z98xDceosuGiQ=
|
golang.org/x/text v0.3.6/go.mod h1:5Zoc/QRtKVWzQhOtBMvqHzDpF6irO9z98xDceosuGiQ=
|
||||||
golang.org/x/text v0.3.7/go.mod h1:u+2+/6zg+i71rQMx5EYifcz6MCKuco9NR6JIITiCfzQ=
|
golang.org/x/text v0.3.7/go.mod h1:u+2+/6zg+i71rQMx5EYifcz6MCKuco9NR6JIITiCfzQ=
|
||||||
golang.org/x/text v0.41.0 h1:vz/seA0lnX87Othu2f/0L24RcgrXD9/YFTSuGjj3rH8=
|
golang.org/x/text v0.42.0 h1:JbOZXgfeCPU9gacVtYliJqOhD+zhrEqK4LfdpmlUZqI=
|
||||||
golang.org/x/text v0.41.0/go.mod h1:jvf1O8ajNzZqhSrQBPbutR/EB83Cc0CFrezNQIwbb5M=
|
golang.org/x/text v0.42.0/go.mod h1:ojzP1Z+2QtioaF8DTtO8K5q7JWVVYwZKenzujK0Zd0E=
|
||||||
golang.org/x/time v0.14.0 h1:MRx4UaLrDotUKUdCIqzPC48t1Y9hANFKIRpNx+Te8PI=
|
golang.org/x/time v0.14.0 h1:MRx4UaLrDotUKUdCIqzPC48t1Y9hANFKIRpNx+Te8PI=
|
||||||
golang.org/x/time v0.14.0/go.mod h1:eL/Oa2bBBK0TkX57Fyni+NgnyQQN4LitPmob2Hjnqw4=
|
golang.org/x/time v0.14.0/go.mod h1:eL/Oa2bBBK0TkX57Fyni+NgnyQQN4LitPmob2Hjnqw4=
|
||||||
golang.org/x/tools v0.0.0-20180917221912-90fa682c2a6e/go.mod h1:n7NCudcB/nEzxVGmLbDWY5pfWTLqBcC2KZ6jyYvM4mQ=
|
golang.org/x/tools v0.0.0-20180917221912-90fa682c2a6e/go.mod h1:n7NCudcB/nEzxVGmLbDWY5pfWTLqBcC2KZ6jyYvM4mQ=
|
||||||
@@ -157,14 +135,14 @@ golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2 h1:B82qJJgjvYKsXS9jeu
|
|||||||
golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2/go.mod h1:deeaetjYA+DHMHg+sMSMI58GrEteJUUzzw7en6TJQcI=
|
golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2/go.mod h1:deeaetjYA+DHMHg+sMSMI58GrEteJUUzzw7en6TJQcI=
|
||||||
golang.zx2c4.com/wireguard v0.0.0-20250521234502-f333402bd9cb h1:whnFRlWMcXI9d+ZbWg+4sHnLp52d5yiIPUxMBSt4X9A=
|
golang.zx2c4.com/wireguard v0.0.0-20250521234502-f333402bd9cb h1:whnFRlWMcXI9d+ZbWg+4sHnLp52d5yiIPUxMBSt4X9A=
|
||||||
golang.zx2c4.com/wireguard v0.0.0-20250521234502-f333402bd9cb/go.mod h1:rpwXGsirqLqN2L0JDJQlwOboGHmptD5ZD6T2VmcqhTw=
|
golang.zx2c4.com/wireguard v0.0.0-20250521234502-f333402bd9cb/go.mod h1:rpwXGsirqLqN2L0JDJQlwOboGHmptD5ZD6T2VmcqhTw=
|
||||||
golang.zx2c4.com/wireguard/windows v1.0.1 h1:eOxiDVbywPC+ZQqvdCK7x+ZwWXKbYv50TtH8ysFIbw8=
|
golang.zx2c4.com/wireguard/windows v1.1.1 h1:8/H97U1v1PNDNcBsMZgU3KFuND9MQdTsU2NOwmCXArE=
|
||||||
golang.zx2c4.com/wireguard/windows v1.0.1/go.mod h1:+fbT3FFdX4zzYDLwJh5+HPEcNN/3HyNdzhNSVsQM+zs=
|
golang.zx2c4.com/wireguard/windows v1.1.1/go.mod h1:+fbT3FFdX4zzYDLwJh5+HPEcNN/3HyNdzhNSVsQM+zs=
|
||||||
gonum.org/v1/gonum v0.17.0 h1:VbpOemQlsSMrYmn7T2OUvQ4dqxQXU+ouZFQsZOx50z4=
|
gonum.org/v1/gonum v0.17.0 h1:VbpOemQlsSMrYmn7T2OUvQ4dqxQXU+ouZFQsZOx50z4=
|
||||||
gonum.org/v1/gonum v0.17.0/go.mod h1:El3tOrEuMpv2UdMrbNlKEh9vd86bmQ6vqIcDwxEOc1E=
|
gonum.org/v1/gonum v0.17.0/go.mod h1:El3tOrEuMpv2UdMrbNlKEh9vd86bmQ6vqIcDwxEOc1E=
|
||||||
google.golang.org/genproto/googleapis/rpc v0.0.0-20260526163538-3dc84a4a5aaa h1:mZHHdPZl0dbGHCflZgAq/Q468DWVFcU2whhB2KAo8fk=
|
google.golang.org/genproto/googleapis/rpc v0.0.0-20260706201446-f0a921348800 h1:qEHAMpSaUhtD0p3NbEEI83HwNGFxEwaSJ1G9PLnCBZE=
|
||||||
google.golang.org/genproto/googleapis/rpc v0.0.0-20260526163538-3dc84a4a5aaa/go.mod h1:4Hqkh8ycfw05ld/3BWL7rJOSfebL2Q+DVDeRgYgxUU8=
|
google.golang.org/genproto/googleapis/rpc v0.0.0-20260706201446-f0a921348800/go.mod h1:4Hqkh8ycfw05ld/3BWL7rJOSfebL2Q+DVDeRgYgxUU8=
|
||||||
google.golang.org/grpc v1.83.2 h1:EManeRomTObA0BU7I8vXgg/78uE5MJ9M8B39EX2WscU=
|
google.golang.org/grpc v1.84.0 h1:soMyaPJ8pAak5PIQ0DGBUir0XRo2fRoMqhNWMLlLxO0=
|
||||||
google.golang.org/grpc v1.83.2/go.mod h1:YPI1hK3kDked6iHvgX3tR0y+nX/qpMFKhPgFsokw1S8=
|
google.golang.org/grpc v1.84.0/go.mod h1:ljCht0DrxQrXBDRTZp52Qxh3Ffk8CdYm2sj4O2QN2C0=
|
||||||
google.golang.org/protobuf v1.36.12 h1:pJOKDDOyeXErUroCihFAd5LQuwXBSpVnKGrj5o/fwxc=
|
google.golang.org/protobuf v1.36.12 h1:pJOKDDOyeXErUroCihFAd5LQuwXBSpVnKGrj5o/fwxc=
|
||||||
google.golang.org/protobuf v1.36.12/go.mod h1:HTf+CrKn2C3g5S8VImy6tdcUvCska2kB7j23XfzDpco=
|
google.golang.org/protobuf v1.36.12/go.mod h1:HTf+CrKn2C3g5S8VImy6tdcUvCska2kB7j23XfzDpco=
|
||||||
gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0=
|
gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0=
|
||||||
|
|||||||
+2
-2
@@ -97,7 +97,7 @@ func (v *HTTPClientConfig) Build() (proto.Message, error) {
|
|||||||
user.Email = v.Email
|
user.Email = v.Email
|
||||||
} else {
|
} else {
|
||||||
if err := json.Unmarshal(rawUser, user); err != nil {
|
if err := json.Unmarshal(rawUser, user); err != nil {
|
||||||
return nil, errors.New("failed to parse HTTP user").Base(err).AtError()
|
return nil, errors.New("failed to parse HTTP user").Base(err)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
account := new(HTTPAccount)
|
account := new(HTTPAccount)
|
||||||
@@ -106,7 +106,7 @@ func (v *HTTPClientConfig) Build() (proto.Message, error) {
|
|||||||
account.Password = v.Password
|
account.Password = v.Password
|
||||||
} else {
|
} else {
|
||||||
if err := json.Unmarshal(rawUser, account); err != nil {
|
if err := json.Unmarshal(rawUser, account); err != nil {
|
||||||
return nil, errors.New("failed to parse HTTP account").Base(err).AtError()
|
return nil, errors.New("failed to parse HTTP account").Base(err)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
user.Account = serial.ToTypedMessage(account.Build())
|
user.Account = serial.ToTypedMessage(account.Build())
|
||||||
|
|||||||
+1
-1
@@ -18,7 +18,7 @@ func RegisterConfigureFilePostProcessingStage(name string, stage ConfigureFilePo
|
|||||||
func PostProcessConfigureFile(conf *Config) error {
|
func PostProcessConfigureFile(conf *Config) error {
|
||||||
for k, v := range configureFilePostProcessingStages {
|
for k, v := range configureFilePostProcessingStages {
|
||||||
if err := v.Process(conf); err != nil {
|
if err := v.Process(conf); err != nil {
|
||||||
return errors.New("Rejected by Postprocessing Stage ", k).AtError().Base(err)
|
return errors.New("Rejected by Postprocessing Stage ", k).Base(err)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
return nil
|
return nil
|
||||||
|
|||||||
@@ -13,7 +13,7 @@ type ConfigCreatorCache map[string]ConfigCreator
|
|||||||
|
|
||||||
func (v ConfigCreatorCache) RegisterCreator(id string, creator ConfigCreator) error {
|
func (v ConfigCreatorCache) RegisterCreator(id string, creator ConfigCreator) error {
|
||||||
if _, found := v[id]; found {
|
if _, found := v[id]; found {
|
||||||
return errors.New(id, " already registered.").AtError()
|
return errors.New(id, " already registered.")
|
||||||
}
|
}
|
||||||
|
|
||||||
v[id] = creator
|
v[id] = creator
|
||||||
@@ -61,7 +61,7 @@ func (v *JSONConfigLoader) Load(raw []byte) (interface{}, string, error) {
|
|||||||
}
|
}
|
||||||
rawID, found := obj[v.idKey]
|
rawID, found := obj[v.idKey]
|
||||||
if !found {
|
if !found {
|
||||||
return nil, "", errors.New(v.idKey, " not found in JSON context").AtError()
|
return nil, "", errors.New(v.idKey, " not found in JSON context")
|
||||||
}
|
}
|
||||||
var id string
|
var id string
|
||||||
if err := json.Unmarshal(rawID, &id); err != nil {
|
if err := json.Unmarshal(rawID, &id); err != nil {
|
||||||
|
|||||||
@@ -0,0 +1,103 @@
|
|||||||
|
package conf
|
||||||
|
|
||||||
|
import (
|
||||||
|
"net/netip"
|
||||||
|
"strings"
|
||||||
|
|
||||||
|
"github.com/xtls/xray-core/common/errors"
|
||||||
|
"github.com/xtls/xray-core/common/protocol"
|
||||||
|
"github.com/xtls/xray-core/common/serial"
|
||||||
|
"github.com/xtls/xray-core/proxy/masque"
|
||||||
|
"google.golang.org/protobuf/proto"
|
||||||
|
)
|
||||||
|
|
||||||
|
type MasqueClientConfig struct {
|
||||||
|
Address *Address `json:"address"`
|
||||||
|
Port uint16 `json:"port"`
|
||||||
|
RemoteDNS []string `json:"remoteDNS"`
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *MasqueClientConfig) Build() (proto.Message, error) {
|
||||||
|
if c.Address == nil {
|
||||||
|
return nil, errors.New(`MASQUE: "address" is not set`)
|
||||||
|
}
|
||||||
|
if c.Port == 0 {
|
||||||
|
return nil, errors.New(`MASQUE: "port" is not set`)
|
||||||
|
}
|
||||||
|
for _, s := range c.RemoteDNS {
|
||||||
|
if _, err := netip.ParseAddr(s); err != nil {
|
||||||
|
return nil, errors.New(`MASQUE: invalid "remoteDNS" `, s).Base(err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return &masque.ClientConfig{
|
||||||
|
Server: &protocol.ServerEndpoint{
|
||||||
|
Address: c.Address.Build(),
|
||||||
|
Port: uint32(c.Port),
|
||||||
|
},
|
||||||
|
RemoteDns: c.RemoteDNS,
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
type MasqueUserConfig struct {
|
||||||
|
Pass string `json:"pass"`
|
||||||
|
Level uint32 `json:"level"`
|
||||||
|
Email string `json:"email"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type MasqueServerConfig struct {
|
||||||
|
Users []*MasqueUserConfig `json:"users"`
|
||||||
|
Clients []*MasqueUserConfig `json:"clients"`
|
||||||
|
Address []string `json:"address"`
|
||||||
|
MTU uint32 `json:"mtu"`
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *MasqueServerConfig) Build() (proto.Message, error) {
|
||||||
|
if c.Clients != nil {
|
||||||
|
c.Users = c.Clients
|
||||||
|
}
|
||||||
|
config := &masque.ServerConfig{
|
||||||
|
Address: c.Address,
|
||||||
|
Mtu: c.MTU,
|
||||||
|
}
|
||||||
|
emails := make(map[string]bool)
|
||||||
|
for _, user := range c.Users {
|
||||||
|
if user.Email == "" {
|
||||||
|
return nil, errors.New(`MASQUE: "email" is empty`)
|
||||||
|
}
|
||||||
|
if strings.Contains(user.Email, ":") {
|
||||||
|
return nil, errors.New(`MASQUE: invalid "email" `, user.Email)
|
||||||
|
}
|
||||||
|
if user.Pass == "" {
|
||||||
|
return nil, errors.New(`MASQUE: "pass" of `, user.Email, ` is empty`)
|
||||||
|
}
|
||||||
|
email := strings.ToLower(user.Email)
|
||||||
|
if emails[email] {
|
||||||
|
return nil, errors.New(`MASQUE: duplicate "email" `, user.Email)
|
||||||
|
}
|
||||||
|
emails[email] = true
|
||||||
|
config.Users = append(config.Users, &protocol.User{
|
||||||
|
Email: user.Email,
|
||||||
|
Level: user.Level,
|
||||||
|
Account: serial.ToTypedMessage(&masque.Account{Password: user.Pass}),
|
||||||
|
})
|
||||||
|
}
|
||||||
|
if len(c.Address) == 0 {
|
||||||
|
return nil, errors.New(`MASQUE: "address" is not set`)
|
||||||
|
}
|
||||||
|
var v4, v6 bool
|
||||||
|
for _, s := range c.Address {
|
||||||
|
prefix, err := netip.ParsePrefix(s)
|
||||||
|
if err != nil {
|
||||||
|
return nil, errors.New(`MASQUE: invalid "address" `, s).Base(err)
|
||||||
|
}
|
||||||
|
if prefix.Addr().Is4() && v4 || prefix.Addr().Is6() && v6 {
|
||||||
|
return nil, errors.New(`MASQUE: "address" takes at most one IPv4 and one IPv6 prefix`)
|
||||||
|
}
|
||||||
|
v4 = v4 || prefix.Addr().Is4()
|
||||||
|
v6 = v6 || prefix.Addr().Is6()
|
||||||
|
}
|
||||||
|
if c.MTU != 0 && (c.MTU < 1280 || c.MTU > 65535) {
|
||||||
|
return nil, errors.New(`MASQUE: "mtu" must be between 1280 and 65535`)
|
||||||
|
}
|
||||||
|
return config, nil
|
||||||
|
}
|
||||||
@@ -0,0 +1,188 @@
|
|||||||
|
package conf_test
|
||||||
|
|
||||||
|
import (
|
||||||
|
"encoding/json"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/xtls/xray-core/common/protocol"
|
||||||
|
"github.com/xtls/xray-core/common/serial"
|
||||||
|
. "github.com/xtls/xray-core/infra/conf"
|
||||||
|
masqueproxy "github.com/xtls/xray-core/proxy/masque"
|
||||||
|
"github.com/xtls/xray-core/transport/internet/masque"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestMasqueConfig(t *testing.T) {
|
||||||
|
creator := func() Buildable {
|
||||||
|
return new(MasqueConfig)
|
||||||
|
}
|
||||||
|
|
||||||
|
runMultiTestCase(t, []TestCase{
|
||||||
|
{
|
||||||
|
Input: `{}`,
|
||||||
|
Parser: loadJSON(creator),
|
||||||
|
Output: &masque.Config{Path: "/.well-known/masque/ip/*/*/"},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
Input: `{
|
||||||
|
"host": "example.com:8443",
|
||||||
|
"path": "/.well-known/masque/ip/{target}/{ipproto}/",
|
||||||
|
"headers": {"Authorization": "Basic dTpw"}
|
||||||
|
}`,
|
||||||
|
Parser: loadJSON(creator),
|
||||||
|
Output: &masque.Config{
|
||||||
|
Host: "example.com:8443",
|
||||||
|
Path: "/.well-known/masque/ip/*/*/",
|
||||||
|
Headers: map[string]string{"Authorization": "Basic dTpw"},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
Input: `{"path": "/masque/ip{?target,ipproto}"}`,
|
||||||
|
Parser: loadJSON(creator),
|
||||||
|
Output: &masque.Config{Path: "/masque/ip?target=*&ipproto=*"},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
Input: `{"user": "u", "pass": "p:q", "headers": {"X-Token": "a"}}`,
|
||||||
|
Parser: loadJSON(creator),
|
||||||
|
Output: &masque.Config{
|
||||||
|
Path: "/.well-known/masque/ip/*/*/",
|
||||||
|
Headers: map[string]string{"Authorization": "Basic dTpwOnE=", "X-Token": "a"},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
})
|
||||||
|
|
||||||
|
for _, input := range []string{
|
||||||
|
`{"path": "/masque/{target}/{ipproto}/{dns}"}`,
|
||||||
|
`{"path": "masque"}`,
|
||||||
|
`{"host": "example.com/path"}`,
|
||||||
|
`{"headers": {"host": "example.com"}}`,
|
||||||
|
`{"headers": {"Capsule-Protocol": "?0"}}`,
|
||||||
|
`{"headers": {"X Token": "a"}}`,
|
||||||
|
`{"headers": {"X-Token": "a\r\nb"}}`,
|
||||||
|
`{"user": "u:v", "pass": "p"}`,
|
||||||
|
`{"user": "u", "pass": "p", "headers": {"authorization": "Basic dTpw"}}`,
|
||||||
|
} {
|
||||||
|
if _, err := loadJSON(creator)(input); err == nil {
|
||||||
|
t.Errorf("expected an error for %s", input)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMasqueOutboundConfig(t *testing.T) {
|
||||||
|
build := func(s string) error {
|
||||||
|
c := new(OutboundDetourConfig)
|
||||||
|
if err := json.Unmarshal([]byte(s), c); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
_, err := c.Build()
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := build(`{
|
||||||
|
"protocol": "masque",
|
||||||
|
"settings": {"address": "example.com", "port": 443},
|
||||||
|
"streamSettings": {"network": "masque", "security": "tls"},
|
||||||
|
"mux": {"enabled": false, "concurrency": -1}
|
||||||
|
}`); err != nil {
|
||||||
|
t.Error(err)
|
||||||
|
}
|
||||||
|
for _, input := range []string{
|
||||||
|
`{"protocol": "masque", "settings": {"address": "example.com"}, "streamSettings": {"network": "masque", "security": "tls"}}`,
|
||||||
|
`{"protocol": "masque", "settings": {"address": "example.com", "port": 443}, "streamSettings": {"network": "masque", "security": "tls"}, "mux": {"enabled": true}}`,
|
||||||
|
`{"protocol": "masque", "settings": {"address": "example.com", "port": 443}, "streamSettings": {"network": "masque", "security": "tls"}, "mux": {"enabled": true, "concurrency": -1}}`,
|
||||||
|
`{"protocol": "freedom", "streamSettings": {"network": "masque", "security": "tls"}}`,
|
||||||
|
} {
|
||||||
|
if err := build(input); err == nil {
|
||||||
|
t.Errorf("expected an error for %s", input)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMasqueServerConfig(t *testing.T) {
|
||||||
|
creator := func() Buildable {
|
||||||
|
return new(MasqueServerConfig)
|
||||||
|
}
|
||||||
|
|
||||||
|
runMultiTestCase(t, []TestCase{
|
||||||
|
{
|
||||||
|
Input: `{
|
||||||
|
"users": [{"email": "u@example.com", "pass": "p", "level": 1}],
|
||||||
|
"address": ["10.13.0.1/24", "fd13::1/64"],
|
||||||
|
"mtu": 1400
|
||||||
|
}`,
|
||||||
|
Parser: loadJSON(creator),
|
||||||
|
Output: &masqueproxy.ServerConfig{
|
||||||
|
Users: []*protocol.User{{
|
||||||
|
Email: "u@example.com",
|
||||||
|
Level: 1,
|
||||||
|
Account: serial.ToTypedMessage(&masqueproxy.Account{Password: "p"}),
|
||||||
|
}},
|
||||||
|
Address: []string{"10.13.0.1/24", "fd13::1/64"},
|
||||||
|
Mtu: 1400,
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
Input: `{"clients": [{"email": "u", "pass": "p:q"}], "address": ["10.13.0.1/24"]}`,
|
||||||
|
Parser: loadJSON(creator),
|
||||||
|
Output: &masqueproxy.ServerConfig{
|
||||||
|
Users: []*protocol.User{{
|
||||||
|
Email: "u",
|
||||||
|
Account: serial.ToTypedMessage(&masqueproxy.Account{Password: "p:q"}),
|
||||||
|
}},
|
||||||
|
Address: []string{"10.13.0.1/24"},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
Input: `{"address": ["10.13.0.1/24"]}`,
|
||||||
|
Parser: loadJSON(creator),
|
||||||
|
Output: &masqueproxy.ServerConfig{
|
||||||
|
Address: []string{"10.13.0.1/24"},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
})
|
||||||
|
|
||||||
|
for _, input := range []string{
|
||||||
|
`{"users": [{"email": "u:v", "pass": "p"}], "address": ["10.13.0.1/24"]}`,
|
||||||
|
`{"users": [{"email": "", "pass": "p"}], "address": ["10.13.0.1/24"]}`,
|
||||||
|
`{"users": [{"pass": "p"}], "address": ["10.13.0.1/24"]}`,
|
||||||
|
`{"users": [{"email": "u", "pass": ""}], "address": ["10.13.0.1/24"]}`,
|
||||||
|
`{"users": [{"email": "u", "pass": "p"}, {"email": "U", "pass": "q"}], "address": ["10.13.0.1/24"]}`,
|
||||||
|
`{"users": [{"email": "u", "pass": "p"}]}`,
|
||||||
|
`{"users": [{"email": "u", "pass": "p"}], "address": ["10.13.0.1"]}`,
|
||||||
|
`{"users": [{"email": "u", "pass": "p"}], "address": ["10.13.0.1/24", "10.14.0.1/24"]}`,
|
||||||
|
`{"users": [{"email": "u", "pass": "p"}], "address": ["fd13::1/64", "fd14::1/64"]}`,
|
||||||
|
`{"users": [{"email": "u", "pass": "p"}], "address": ["10.13.0.1/24"], "mtu": 1000}`,
|
||||||
|
`{"users": [{"email": "u", "pass": "p"}], "address": ["10.13.0.1/24"], "mtu": 70000}`,
|
||||||
|
} {
|
||||||
|
if _, err := loadJSON(creator)(input); err == nil {
|
||||||
|
t.Errorf("expected an error for %s", input)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMasqueInboundConfig(t *testing.T) {
|
||||||
|
build := func(s string) error {
|
||||||
|
c := new(InboundDetourConfig)
|
||||||
|
if err := json.Unmarshal([]byte(s), c); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
_, err := c.Build()
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := build(`{
|
||||||
|
"protocol": "masque",
|
||||||
|
"port": 443,
|
||||||
|
"settings": {"users": [{"email": "u@example.com", "pass": "p"}], "address": ["10.13.0.1/24"]},
|
||||||
|
"streamSettings": {"network": "masque", "security": "tls"}
|
||||||
|
}`); err != nil {
|
||||||
|
t.Error(err)
|
||||||
|
}
|
||||||
|
if err := build(`{
|
||||||
|
"protocol": "vless",
|
||||||
|
"port": 443,
|
||||||
|
"settings": {"users": [{"id": "27848739-7e62-4138-9fd3-098a63964b6b"}], "decryption": "none"},
|
||||||
|
"streamSettings": {"network": "masque", "security": "tls"}
|
||||||
|
}`); err == nil {
|
||||||
|
t.Error("expected an error for the masque transport on a vless inbound")
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -30,7 +30,7 @@ func MergeConfigFromFiles(files []*core.ConfigSource) (string, error) {
|
|||||||
if j, ok := creflect.MarshalToJson(c, true); ok {
|
if j, ok := creflect.MarshalToJson(c, true); ok {
|
||||||
return j, nil
|
return j, nil
|
||||||
}
|
}
|
||||||
return "", errors.New("marshal to json failed.").AtError()
|
return "", errors.New("marshal to json failed.")
|
||||||
}
|
}
|
||||||
|
|
||||||
func mergeConfigs(files []*core.ConfigSource) (*conf.Config, error) {
|
func mergeConfigs(files []*core.ConfigSource) (*conf.Config, error) {
|
||||||
|
|||||||
+37
-56
@@ -3,8 +3,6 @@ package conf
|
|||||||
import (
|
import (
|
||||||
"strings"
|
"strings"
|
||||||
|
|
||||||
"github.com/sagernet/sing-shadowsocks/shadowaead_2022"
|
|
||||||
C "github.com/sagernet/sing/common"
|
|
||||||
"github.com/xtls/xray-core/common/errors"
|
"github.com/xtls/xray-core/common/errors"
|
||||||
"github.com/xtls/xray-core/common/protocol"
|
"github.com/xtls/xray-core/common/protocol"
|
||||||
"github.com/xtls/xray-core/common/serial"
|
"github.com/xtls/xray-core/common/serial"
|
||||||
@@ -55,7 +53,7 @@ func (v *ShadowsocksServerConfig) Build() (proto.Message, error) {
|
|||||||
v.Users = v.Clients
|
v.Users = v.Clients
|
||||||
}
|
}
|
||||||
|
|
||||||
if C.Contains(shadowaead_2022.List, v.Cipher) {
|
if _, err := shadowsocks_2022.GetCipherMethod(v.Cipher); err == nil {
|
||||||
return buildShadowsocks2022(v)
|
return buildShadowsocks2022(v)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -111,12 +109,14 @@ func (v *ShadowsocksServerConfig) Build() (proto.Message, error) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func buildShadowsocks2022(v *ShadowsocksServerConfig) (proto.Message, error) {
|
func buildShadowsocks2022(v *ShadowsocksServerConfig) (proto.Message, error) {
|
||||||
|
v.Cipher = strings.ToLower(v.Cipher)
|
||||||
if len(v.Users) == 0 {
|
if len(v.Users) == 0 {
|
||||||
config := new(shadowsocks_2022.ServerConfig)
|
config := new(shadowsocks_2022.ServerConfig)
|
||||||
config.Method = v.Cipher
|
config.Method = v.Cipher
|
||||||
config.Key = v.Password
|
config.Key = v.Password
|
||||||
config.Network = v.NetworkList.Build()
|
config.Network = v.NetworkList.Build()
|
||||||
config.Email = v.Email
|
config.Email = v.Email
|
||||||
|
config.Level = int32(v.Level)
|
||||||
return config, nil
|
return config, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -171,6 +171,7 @@ func buildShadowsocks2022(v *ShadowsocksServerConfig) (proto.Message, error) {
|
|||||||
Email: user.Email,
|
Email: user.Email,
|
||||||
Address: user.Address.Build(),
|
Address: user.Address.Build(),
|
||||||
Port: uint32(user.Port),
|
Port: uint32(user.Port),
|
||||||
|
Level: int32(user.Level),
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
return config, nil
|
return config, nil
|
||||||
@@ -214,63 +215,43 @@ func (v *ShadowsocksClientConfig) Build() (proto.Message, error) {
|
|||||||
return nil, errors.New(`Shadowsocks settings: "servers" should have one and only one member. Multiple endpoints in "servers" should use multiple Shadowsocks outbounds and routing balancer instead`)
|
return nil, errors.New(`Shadowsocks settings: "servers" should have one and only one member. Multiple endpoints in "servers" should use multiple Shadowsocks outbounds and routing balancer instead`)
|
||||||
}
|
}
|
||||||
|
|
||||||
if len(v.Servers) == 1 {
|
server := v.Servers[0]
|
||||||
server := v.Servers[0]
|
if server.Address == nil {
|
||||||
if C.Contains(shadowaead_2022.List, server.Cipher) {
|
return nil, errors.New("Shadowsocks server address is not set.")
|
||||||
if server.Address == nil {
|
}
|
||||||
return nil, errors.New("Shadowsocks server address is not set.")
|
if server.Port == 0 {
|
||||||
}
|
return nil, errors.New("Invalid Shadowsocks port.")
|
||||||
if server.Port == 0 {
|
}
|
||||||
return nil, errors.New("Invalid Shadowsocks port.")
|
if server.Password == "" {
|
||||||
}
|
return nil, errors.New("Shadowsocks password is not specified.")
|
||||||
if server.Password == "" {
|
|
||||||
return nil, errors.New("Shadowsocks password is not specified.")
|
|
||||||
}
|
|
||||||
|
|
||||||
config := new(shadowsocks_2022.ClientConfig)
|
|
||||||
config.Address = server.Address.Build()
|
|
||||||
config.Port = uint32(server.Port)
|
|
||||||
config.Method = server.Cipher
|
|
||||||
config.Key = server.Password
|
|
||||||
return config, nil
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
|
if _, err := shadowsocks_2022.GetCipherMethod(server.Cipher); err == nil {
|
||||||
|
config := new(shadowsocks_2022.ClientConfig)
|
||||||
|
config.Address = server.Address.Build()
|
||||||
|
config.Port = uint32(server.Port)
|
||||||
|
config.Method = server.Cipher
|
||||||
|
config.Key = server.Password
|
||||||
|
return config, nil
|
||||||
|
}
|
||||||
config := new(shadowsocks.ClientConfig)
|
config := new(shadowsocks.ClientConfig)
|
||||||
for _, server := range v.Servers {
|
account := &shadowsocks.Account{
|
||||||
if C.Contains(shadowaead_2022.List, server.Cipher) {
|
Password: server.Password,
|
||||||
return nil, errors.New("Shadowsocks 2022 accept no multi servers")
|
|
||||||
}
|
|
||||||
if server.Address == nil {
|
|
||||||
return nil, errors.New("Shadowsocks server address is not set.")
|
|
||||||
}
|
|
||||||
if server.Port == 0 {
|
|
||||||
return nil, errors.New("Invalid Shadowsocks port.")
|
|
||||||
}
|
|
||||||
if server.Password == "" {
|
|
||||||
return nil, errors.New("Shadowsocks password is not specified.")
|
|
||||||
}
|
|
||||||
account := &shadowsocks.Account{
|
|
||||||
Password: server.Password,
|
|
||||||
}
|
|
||||||
account.CipherType = cipherFromString(server.Cipher)
|
|
||||||
if account.CipherType == shadowsocks.CipherType_UNKNOWN {
|
|
||||||
return nil, errors.New("unknown cipher method: ", server.Cipher)
|
|
||||||
}
|
|
||||||
|
|
||||||
ss := &protocol.ServerEndpoint{
|
|
||||||
Address: server.Address.Build(),
|
|
||||||
Port: uint32(server.Port),
|
|
||||||
User: &protocol.User{
|
|
||||||
Level: uint32(server.Level),
|
|
||||||
Email: server.Email,
|
|
||||||
Account: serial.ToTypedMessage(account),
|
|
||||||
},
|
|
||||||
}
|
|
||||||
|
|
||||||
config.Server = ss
|
|
||||||
break
|
|
||||||
}
|
}
|
||||||
|
account.CipherType = cipherFromString(server.Cipher)
|
||||||
|
if account.CipherType == shadowsocks.CipherType_UNKNOWN {
|
||||||
|
return nil, errors.New("unknown cipher method: ", server.Cipher)
|
||||||
|
}
|
||||||
|
ss := &protocol.ServerEndpoint{
|
||||||
|
Address: server.Address.Build(),
|
||||||
|
Port: uint32(server.Port),
|
||||||
|
User: &protocol.User{
|
||||||
|
Level: uint32(server.Level),
|
||||||
|
Email: server.Email,
|
||||||
|
Account: serial.ToTypedMessage(account),
|
||||||
|
},
|
||||||
|
}
|
||||||
|
config.Server = ss
|
||||||
|
|
||||||
return config, nil
|
return config, nil
|
||||||
}
|
}
|
||||||
|
|||||||
+2
-3
@@ -44,7 +44,6 @@ func (v *SocksServerConfig) Build() (proto.Message, error) {
|
|||||||
case AuthMethodUserPass:
|
case AuthMethodUserPass:
|
||||||
config.AuthType = socks.AuthType_PASSWORD
|
config.AuthType = socks.AuthType_PASSWORD
|
||||||
default:
|
default:
|
||||||
// errors.New("unknown socks auth method: ", v.AuthMethod, ". Default to noauth.").AtWarning().WriteToLog()
|
|
||||||
config.AuthType = socks.AuthType_NO_AUTH
|
config.AuthType = socks.AuthType_NO_AUTH
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -115,7 +114,7 @@ func (v *SocksClientConfig) Build() (proto.Message, error) {
|
|||||||
user.Email = v.Email
|
user.Email = v.Email
|
||||||
} else {
|
} else {
|
||||||
if err := json.Unmarshal(rawUser, user); err != nil {
|
if err := json.Unmarshal(rawUser, user); err != nil {
|
||||||
return nil, errors.New("failed to parse Socks user").Base(err).AtError()
|
return nil, errors.New("failed to parse Socks user").Base(err)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
account := new(SocksAccount)
|
account := new(SocksAccount)
|
||||||
@@ -124,7 +123,7 @@ func (v *SocksClientConfig) Build() (proto.Message, error) {
|
|||||||
account.Password = v.Password
|
account.Password = v.Password
|
||||||
} else {
|
} else {
|
||||||
if err := json.Unmarshal(rawUser, account); err != nil {
|
if err := json.Unmarshal(rawUser, account); err != nil {
|
||||||
return nil, errors.New("failed to parse socks account").Base(err).AtError()
|
return nil, errors.New("failed to parse socks account").Base(err)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
user.Account = serial.ToTypedMessage(account.Build())
|
user.Account = serial.ToTypedMessage(account.Build())
|
||||||
|
|||||||
@@ -14,7 +14,6 @@ import (
|
|||||||
googleuuid "github.com/google/uuid"
|
googleuuid "github.com/google/uuid"
|
||||||
"github.com/xtls/xray-core/common/errors"
|
"github.com/xtls/xray-core/common/errors"
|
||||||
"github.com/xtls/xray-core/common/net"
|
"github.com/xtls/xray-core/common/net"
|
||||||
"github.com/xtls/xray-core/transport/internet"
|
|
||||||
"github.com/xtls/xray-core/transport/internet/finalmask/fragment"
|
"github.com/xtls/xray-core/transport/internet/finalmask/fragment"
|
||||||
"github.com/xtls/xray-core/transport/internet/finalmask/header/custom"
|
"github.com/xtls/xray-core/transport/internet/finalmask/header/custom"
|
||||||
"github.com/xtls/xray-core/transport/internet/finalmask/mkcp/aes128gcm"
|
"github.com/xtls/xray-core/transport/internet/finalmask/mkcp/aes128gcm"
|
||||||
@@ -909,22 +908,13 @@ func (c *Realm) Build() (proto.Message, error) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
type UDPHop struct {
|
type UDPHop struct {
|
||||||
Sockopt *SocketConfig `json:"sockopt"`
|
Mode string `json:"mode"`
|
||||||
Mode string `json:"mode"`
|
Interval Int32Range `json:"interval"`
|
||||||
Interval Int32Range `json:"interval"`
|
RemoteIPs []string `json:"remoteIPs"`
|
||||||
RemotePorts PortList `json:"remotePorts"`
|
RemotePorts PortList `json:"remotePorts"`
|
||||||
RemoteIPs []string `json:"remoteIPs"`
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *UDPHop) Build() (proto.Message, error) {
|
func (c *UDPHop) Build() (proto.Message, error) {
|
||||||
var sockopt *internet.SocketConfig
|
|
||||||
if c.Sockopt != nil {
|
|
||||||
var err error
|
|
||||||
sockopt, err = c.Sockopt.Build()
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
}
|
|
||||||
var local, remote, remoteOnce bool
|
var local, remote, remoteOnce bool
|
||||||
for _, mode := range strings.Split(c.Mode, ",") {
|
for _, mode := range strings.Split(c.Mode, ",") {
|
||||||
switch strings.ToLower(mode) {
|
switch strings.ToLower(mode) {
|
||||||
@@ -952,15 +942,21 @@ func (c *UDPHop) Build() (proto.Message, error) {
|
|||||||
}
|
}
|
||||||
return nil, errors.New("invalid ip ", ip)
|
return nil, errors.New("invalid ip ", ip)
|
||||||
}
|
}
|
||||||
|
interval := c.Interval
|
||||||
|
if interval.From == 0 && interval.To == 0 {
|
||||||
|
interval.From, interval.To = 30, 30
|
||||||
|
}
|
||||||
|
if interval.From < 5 {
|
||||||
|
return nil, errors.New("interval must be at least 5")
|
||||||
|
}
|
||||||
return &udphop.Config{
|
return &udphop.Config{
|
||||||
Sockopt: sockopt,
|
|
||||||
Local: local,
|
Local: local,
|
||||||
Remote: remote,
|
Remote: remote,
|
||||||
RemoteOnce: remoteOnce,
|
RemoteOnce: remoteOnce,
|
||||||
IntervalMin: int64(c.Interval.From),
|
IntervalMin: int64(interval.From),
|
||||||
IntervalMax: int64(c.Interval.To),
|
IntervalMax: int64(interval.To),
|
||||||
RemotePorts: c.RemotePorts.Build().Ports(),
|
|
||||||
RemoteIPs: remoteIPs,
|
RemoteIPs: remoteIPs,
|
||||||
|
RemotePorts: c.RemotePorts.Build().Ports(),
|
||||||
}, nil
|
}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -36,6 +36,10 @@ func (p TransportProtocol) Build() (string, error) {
|
|||||||
return "", errors.PrintRemovedFeatureError("QUIC transport (without web service, etc.)", "XHTTP stream-one H3")
|
return "", errors.PrintRemovedFeatureError("QUIC transport (without web service, etc.)", "XHTTP stream-one H3")
|
||||||
case "hysteria":
|
case "hysteria":
|
||||||
return "hysteria", nil
|
return "hysteria", nil
|
||||||
|
case "masque":
|
||||||
|
return "masque", nil
|
||||||
|
case "xdrive":
|
||||||
|
return "xdrive", nil
|
||||||
default:
|
default:
|
||||||
return "", errors.New("Config: unknown transport protocol: ", p)
|
return "", errors.New("Config: unknown transport protocol: ", p)
|
||||||
}
|
}
|
||||||
@@ -59,6 +63,8 @@ type StreamConfig struct {
|
|||||||
WSSettings *WebSocketConfig `json:"wsSettings"`
|
WSSettings *WebSocketConfig `json:"wsSettings"`
|
||||||
HTTPUPGRADESettings *HttpUpgradeConfig `json:"httpupgradeSettings"`
|
HTTPUPGRADESettings *HttpUpgradeConfig `json:"httpupgradeSettings"`
|
||||||
HysteriaSettings *HysteriaConfig `json:"hysteriaSettings"`
|
HysteriaSettings *HysteriaConfig `json:"hysteriaSettings"`
|
||||||
|
MASQUESettings *MasqueConfig `json:"masqueSettings"`
|
||||||
|
XDRIVESettings *XDriveConfig `json:"xdriveSettings"`
|
||||||
SocketSettings *SocketConfig `json:"sockopt"`
|
SocketSettings *SocketConfig `json:"sockopt"`
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -192,6 +198,26 @@ func (c *StreamConfig) Build() (*internet.StreamConfig, error) {
|
|||||||
Settings: serial.ToTypedMessage(hs),
|
Settings: serial.ToTypedMessage(hs),
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
if c.MASQUESettings != nil {
|
||||||
|
ms, err := c.MASQUESettings.Build()
|
||||||
|
if err != nil {
|
||||||
|
return nil, errors.New("Failed to build MASQUE config.").Base(err)
|
||||||
|
}
|
||||||
|
config.TransportSettings = append(config.TransportSettings, &internet.TransportConfig{
|
||||||
|
ProtocolName: "masque",
|
||||||
|
Settings: serial.ToTypedMessage(ms),
|
||||||
|
})
|
||||||
|
}
|
||||||
|
if c.XDRIVESettings != nil {
|
||||||
|
xs, err := c.XDRIVESettings.Build()
|
||||||
|
if err != nil {
|
||||||
|
return nil, errors.New("Failed to build XDRIVE config.").Base(err)
|
||||||
|
}
|
||||||
|
config.TransportSettings = append(config.TransportSettings, &internet.TransportConfig{
|
||||||
|
ProtocolName: "xdrive",
|
||||||
|
Settings: serial.ToTypedMessage(xs),
|
||||||
|
})
|
||||||
|
}
|
||||||
if c.SocketSettings != nil {
|
if c.SocketSettings != nil {
|
||||||
ss, err := c.SocketSettings.Build()
|
ss, err := c.SocketSettings.Build()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
|||||||
@@ -1,7 +1,9 @@
|
|||||||
package conf
|
package conf
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"encoding/base64"
|
||||||
"encoding/json"
|
"encoding/json"
|
||||||
|
"maps"
|
||||||
"math/big"
|
"math/big"
|
||||||
"net/url"
|
"net/url"
|
||||||
"sort"
|
"sort"
|
||||||
@@ -20,9 +22,12 @@ import (
|
|||||||
"github.com/xtls/xray-core/transport/internet/httpupgrade"
|
"github.com/xtls/xray-core/transport/internet/httpupgrade"
|
||||||
"github.com/xtls/xray-core/transport/internet/hysteria"
|
"github.com/xtls/xray-core/transport/internet/hysteria"
|
||||||
"github.com/xtls/xray-core/transport/internet/kcp"
|
"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/splithttp"
|
||||||
"github.com/xtls/xray-core/transport/internet/tcp"
|
"github.com/xtls/xray-core/transport/internet/tcp"
|
||||||
"github.com/xtls/xray-core/transport/internet/websocket"
|
"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"
|
"google.golang.org/protobuf/proto"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -121,7 +126,7 @@ func (v *AuthenticatorRequest) Build() (*http.RequestConfig, error) {
|
|||||||
for _, key := range headerNames {
|
for _, key := range headerNames {
|
||||||
value := v.Headers[key]
|
value := v.Headers[key]
|
||||||
if value == nil {
|
if value == nil {
|
||||||
return nil, errors.New("empty HTTP header value: " + key).AtError()
|
return nil, errors.New("empty HTTP header value: " + key)
|
||||||
}
|
}
|
||||||
config.Header = append(config.Header, &http.Header{
|
config.Header = append(config.Header, &http.Header{
|
||||||
Name: key,
|
Name: key,
|
||||||
@@ -189,7 +194,7 @@ func (v *AuthenticatorResponse) Build() (*http.ResponseConfig, error) {
|
|||||||
for _, key := range headerNames {
|
for _, key := range headerNames {
|
||||||
value := v.Headers[key]
|
value := v.Headers[key]
|
||||||
if value == nil {
|
if value == nil {
|
||||||
return nil, errors.New("empty HTTP header value: " + key).AtError()
|
return nil, errors.New("empty HTTP header value: " + key)
|
||||||
}
|
}
|
||||||
config.Header = append(config.Header, &http.Header{
|
config.Header = append(config.Header, &http.Header{
|
||||||
Name: key,
|
Name: key,
|
||||||
@@ -239,11 +244,11 @@ func (c *TCPConfig) Build() (proto.Message, error) {
|
|||||||
if len(c.HeaderConfig) > 0 {
|
if len(c.HeaderConfig) > 0 {
|
||||||
headerConfig, _, err := tcpHeaderLoader.Load(c.HeaderConfig)
|
headerConfig, _, err := tcpHeaderLoader.Load(c.HeaderConfig)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, errors.New("invalid TCP header config").Base(err).AtError()
|
return nil, errors.New("invalid TCP header config").Base(err)
|
||||||
}
|
}
|
||||||
ts, err := headerConfig.(Buildable).Build()
|
ts, err := headerConfig.(Buildable).Build()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, errors.New("invalid TCP header config").Base(err).AtError()
|
return nil, errors.New("invalid TCP header config").Base(err)
|
||||||
}
|
}
|
||||||
config.HeaderSettings = serial.ToTypedMessage(ts)
|
config.HeaderSettings = serial.ToTypedMessage(ts)
|
||||||
}
|
}
|
||||||
@@ -785,6 +790,63 @@ func (c *HysteriaConfig) Build() (proto.Message, error) {
|
|||||||
return config, nil
|
return config, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
type MasqueConfig struct {
|
||||||
|
Host string `json:"host"`
|
||||||
|
Path string `json:"path"`
|
||||||
|
User string `json:"user"`
|
||||||
|
Pass string `json:"pass"`
|
||||||
|
Headers map[string]string `json:"headers"`
|
||||||
|
}
|
||||||
|
|
||||||
|
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, `"`)
|
||||||
|
case "authorization":
|
||||||
|
if c.User != "" || c.Pass != "" {
|
||||||
|
return nil, errors.New(`"headers" can't contain "`, k, `" when "user" or "pass" is set`)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
headers := c.Headers
|
||||||
|
if c.User != "" || c.Pass != "" {
|
||||||
|
if strings.Contains(c.User, ":") {
|
||||||
|
return nil, errors.New(`invalid "user": `, c.User)
|
||||||
|
}
|
||||||
|
headers = maps.Clone(c.Headers)
|
||||||
|
if headers == nil {
|
||||||
|
headers = make(map[string]string)
|
||||||
|
}
|
||||||
|
headers["Authorization"] = "Basic " + base64.StdEncoding.EncodeToString([]byte(c.User+":"+c.Pass))
|
||||||
|
}
|
||||||
|
return &masque.Config{
|
||||||
|
Host: c.Host,
|
||||||
|
Path: path,
|
||||||
|
Headers: headers,
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
|
|
||||||
func readFileOrString(f string, s []string) ([]byte, error) {
|
func readFileOrString(f string, s []string) ([]byte, error) {
|
||||||
if len(f) > 0 {
|
if len(f) > 0 {
|
||||||
return filesystem.ReadCert(f)
|
return filesystem.ReadCert(f)
|
||||||
@@ -794,3 +856,50 @@ func readFileOrString(f string, s []string) ([]byte, error) {
|
|||||||
}
|
}
|
||||||
return nil, errors.New("both file and bytes are empty.")
|
return nil, errors.New("both file and bytes are empty.")
|
||||||
}
|
}
|
||||||
|
|
||||||
|
type XDriveConfig struct {
|
||||||
|
RemoteFolder string `json:"remoteFolder"`
|
||||||
|
Service string `json:"service"`
|
||||||
|
Secrets []string `json:"secrets"`
|
||||||
|
SegmentBytes uint32 `json:"segmentBytes"`
|
||||||
|
FlushIntervalMs uint32 `json:"flushIntervalMs"`
|
||||||
|
PollIntervalMs uint32 `json:"pollIntervalMs"`
|
||||||
|
MaxPollIntervalMs uint32 `json:"maxPollIntervalMs"`
|
||||||
|
SessionTTLSeconds uint32 `json:"sessionTtlSeconds"`
|
||||||
|
Concurrency uint32 `json:"concurrency"`
|
||||||
|
EagerWindowMs uint32 `json:"eagerWindowMs"`
|
||||||
|
HoleTimeoutMs uint32 `json:"holeTimeoutMs"`
|
||||||
|
Template json.RawMessage `json:"template"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// Build implements Buildable.
|
||||||
|
func (c *XDriveConfig) Build() (proto.Message, error) {
|
||||||
|
switch c.Service {
|
||||||
|
case "local":
|
||||||
|
case "Google Drive":
|
||||||
|
if len(c.Secrets) != 3 {
|
||||||
|
return nil, errors.New("Google Drive needs 3 secrets in order of ClientID, ClientSecret, RefreshToken")
|
||||||
|
}
|
||||||
|
case "template":
|
||||||
|
if len(c.Template) == 0 {
|
||||||
|
return nil, errors.New(`service "template" needs a "template" object`)
|
||||||
|
}
|
||||||
|
default:
|
||||||
|
return nil, errors.New("unsupported service")
|
||||||
|
}
|
||||||
|
config := &xdrive.Config{
|
||||||
|
RemoteFolder: c.RemoteFolder,
|
||||||
|
Service: c.Service,
|
||||||
|
Secrets: c.Secrets,
|
||||||
|
SegmentBytes: c.SegmentBytes,
|
||||||
|
FlushIntervalMs: c.FlushIntervalMs,
|
||||||
|
PollIntervalMs: c.PollIntervalMs,
|
||||||
|
MaxPollIntervalMs: c.MaxPollIntervalMs,
|
||||||
|
SessionTtlSeconds: c.SessionTTLSeconds,
|
||||||
|
Concurrency: c.Concurrency,
|
||||||
|
EagerWindowMs: c.EagerWindowMs,
|
||||||
|
HoleTimeoutMs: c.HoleTimeoutMs,
|
||||||
|
Template: string(c.Template),
|
||||||
|
}
|
||||||
|
return config, nil
|
||||||
|
}
|
||||||
|
|||||||
@@ -291,3 +291,76 @@ func TestHeaderCustomUDPBuildRejectsExprWithoutArgs(t *testing.T) {
|
|||||||
t.Fatalf("expected transform arg rejection, got %v", err)
|
t.Fatalf("expected transform arg rejection, got %v", err)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestXDriveStreamConfig(t *testing.T) {
|
||||||
|
config := new(StreamConfig)
|
||||||
|
if err := json.Unmarshal([]byte(`{
|
||||||
|
"method": "xdrive",
|
||||||
|
"xdriveSettings": {
|
||||||
|
"remoteFolder": "/tmp/xdrive",
|
||||||
|
"service": "local"
|
||||||
|
}
|
||||||
|
}`), config); err != nil {
|
||||||
|
t.Fatalf("Unmarshal: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
built, err := config.Build()
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Build: %v", err)
|
||||||
|
}
|
||||||
|
if built.ProtocolName != "xdrive" {
|
||||||
|
t.Fatalf("ProtocolName is %q, want %q", built.ProtocolName, "xdrive")
|
||||||
|
}
|
||||||
|
if len(built.TransportSettings) != 1 || built.TransportSettings[0].ProtocolName != "xdrive" {
|
||||||
|
t.Fatalf("TransportSettings is %v, want a single xdrive entry", built.TransportSettings)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestXDriveRejectsUnknownService(t *testing.T) {
|
||||||
|
config := new(XDriveConfig)
|
||||||
|
if err := json.Unmarshal([]byte(`{"remoteFolder": "/tmp/xdrive", "service": "Dropbox"}`), config); err != nil {
|
||||||
|
t.Fatalf("Unmarshal: %v", err)
|
||||||
|
}
|
||||||
|
if _, err := config.Build(); err == nil {
|
||||||
|
t.Fatal("Build accepted an unsupported service")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestXDriveTemplateStreamConfig(t *testing.T) {
|
||||||
|
config := new(StreamConfig)
|
||||||
|
if err := json.Unmarshal([]byte(`{
|
||||||
|
"method": "xdrive",
|
||||||
|
"xdriveSettings": {
|
||||||
|
"remoteFolder": "folder",
|
||||||
|
"service": "template",
|
||||||
|
"secrets": ["user", "pass"],
|
||||||
|
"template": {
|
||||||
|
"flatten": true,
|
||||||
|
"auth": {"type": "basic", "username": "{secret0}", "password": "{secret1}"},
|
||||||
|
"put": {"method": "PUT", "url": "https://dav.example/{folder}/{name}"},
|
||||||
|
"get": {"method": "GET", "url": "https://dav.example/{folder}/{name}"},
|
||||||
|
"delete": {"method": "DELETE", "url": "https://dav.example/{folder}/{name}"},
|
||||||
|
"list": {"method": "PROPFIND", "url": "https://dav.example/{folder}/", "namesRegex": "<d:href>/folder/([^<]+)</d:href>"}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}`), config); err != nil {
|
||||||
|
t.Fatalf("Unmarshal: %v", err)
|
||||||
|
}
|
||||||
|
built, err := config.Build()
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Build: %v", err)
|
||||||
|
}
|
||||||
|
if built.ProtocolName != "xdrive" {
|
||||||
|
t.Fatalf("ProtocolName is %q, want xdrive", built.ProtocolName)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestXDriveTemplateNeedsTemplate(t *testing.T) {
|
||||||
|
config := new(XDriveConfig)
|
||||||
|
if err := json.Unmarshal([]byte(`{"remoteFolder": "f", "service": "template"}`), config); err != nil {
|
||||||
|
t.Fatalf("Unmarshal: %v", err)
|
||||||
|
}
|
||||||
|
if _, err := config.Build(); err == nil {
|
||||||
|
t.Fatal("Build accepted a template service without a template")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
@@ -20,6 +20,7 @@ type TunConfig struct {
|
|||||||
UserLevel uint32 `json:"userLevel"`
|
UserLevel uint32 `json:"userLevel"`
|
||||||
AutoSystemRoutingTable []string `json:"autoSystemRoutingTable"`
|
AutoSystemRoutingTable []string `json:"autoSystemRoutingTable"`
|
||||||
AutoOutboundsInterface *string `json:"autoOutboundsInterface"`
|
AutoOutboundsInterface *string `json:"autoOutboundsInterface"`
|
||||||
|
AutoSystemDNS bool `json:"autoSystemDNS"`
|
||||||
}
|
}
|
||||||
|
|
||||||
func (v *TunConfig) Build() (proto.Message, error) {
|
func (v *TunConfig) Build() (proto.Message, error) {
|
||||||
@@ -31,6 +32,7 @@ func (v *TunConfig) Build() (proto.Message, error) {
|
|||||||
DNS: v.DNS,
|
DNS: v.DNS,
|
||||||
UserLevel: v.UserLevel,
|
UserLevel: v.UserLevel,
|
||||||
AutoSystemRoutingTable: v.AutoSystemRoutingTable,
|
AutoSystemRoutingTable: v.AutoSystemRoutingTable,
|
||||||
|
AutoSystemDns: v.AutoSystemDNS,
|
||||||
}
|
}
|
||||||
if v.AutoOutboundsInterface != nil {
|
if v.AutoOutboundsInterface != nil {
|
||||||
config.AutoOutboundsInterface = *v.AutoOutboundsInterface
|
config.AutoOutboundsInterface = *v.AutoOutboundsInterface
|
||||||
|
|||||||
+7
-23
@@ -59,14 +59,13 @@ func (c *WireGuardPeerConfig) Build() (*wireguard.PeerConfig, error) {
|
|||||||
type WireGuardConfig struct {
|
type WireGuardConfig struct {
|
||||||
IsClient bool `json:""`
|
IsClient bool `json:""`
|
||||||
|
|
||||||
NoKernelTun bool `json:"noKernelTun"`
|
NoKernelTun bool `json:"noKernelTun"`
|
||||||
SecretKey string `json:"secretKey"`
|
SecretKey string `json:"secretKey"`
|
||||||
Address []string `json:"address"`
|
Address []string `json:"address"`
|
||||||
Peers []*WireGuardPeerConfig `json:"peers"`
|
Peers []*WireGuardPeerConfig `json:"peers"`
|
||||||
MTU int32 `json:"mtu"`
|
MTU int32 `json:"mtu"`
|
||||||
Reserved []byte `json:"reserved"`
|
Reserved []byte `json:"reserved"`
|
||||||
DomainStrategy string `json:"domainStrategy"`
|
DNS []string `json:"remoteDNS"`
|
||||||
DNS []string `json:"remoteDNS"`
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *WireGuardConfig) Build() (proto.Message, error) {
|
func (c *WireGuardConfig) Build() (proto.Message, error) {
|
||||||
@@ -125,21 +124,6 @@ func (c *WireGuardConfig) Build() (proto.Message, error) {
|
|||||||
}
|
}
|
||||||
config.Reserved = c.Reserved
|
config.Reserved = c.Reserved
|
||||||
|
|
||||||
switch strings.ToLower(c.DomainStrategy) {
|
|
||||||
case "forceip", "":
|
|
||||||
config.DomainStrategy = wireguard.DeviceConfig_FORCE_IP
|
|
||||||
case "forceipv4":
|
|
||||||
config.DomainStrategy = wireguard.DeviceConfig_FORCE_IP4
|
|
||||||
case "forceipv6":
|
|
||||||
config.DomainStrategy = wireguard.DeviceConfig_FORCE_IP6
|
|
||||||
case "forceipv4v6":
|
|
||||||
config.DomainStrategy = wireguard.DeviceConfig_FORCE_IP46
|
|
||||||
case "forceipv6v4":
|
|
||||||
config.DomainStrategy = wireguard.DeviceConfig_FORCE_IP64
|
|
||||||
default:
|
|
||||||
return nil, errors.New("unsupported domain strategy: ", c.DomainStrategy)
|
|
||||||
}
|
|
||||||
|
|
||||||
config.IsClient = c.IsClient
|
config.IsClient = c.IsClient
|
||||||
config.NoKernelTun = c.NoKernelTun
|
config.NoKernelTun = c.NoKernelTun
|
||||||
config.DNS = c.DNS
|
config.DNS = c.DNS
|
||||||
|
|||||||
@@ -16,6 +16,7 @@ import (
|
|||||||
"github.com/xtls/xray-core/common/serial"
|
"github.com/xtls/xray-core/common/serial"
|
||||||
core "github.com/xtls/xray-core/core"
|
core "github.com/xtls/xray-core/core"
|
||||||
"github.com/xtls/xray-core/proxy/freedom"
|
"github.com/xtls/xray-core/proxy/freedom"
|
||||||
|
"github.com/xtls/xray-core/proxy/masque"
|
||||||
"github.com/xtls/xray-core/transport/internet"
|
"github.com/xtls/xray-core/transport/internet"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -32,6 +33,7 @@ var (
|
|||||||
"trojan": func() interface{} { return new(TrojanServerConfig) },
|
"trojan": func() interface{} { return new(TrojanServerConfig) },
|
||||||
"wireguard": func() interface{} { return &WireGuardConfig{IsClient: false} },
|
"wireguard": func() interface{} { return &WireGuardConfig{IsClient: false} },
|
||||||
"hysteria": func() interface{} { return new(HysteriaServerConfig) },
|
"hysteria": func() interface{} { return new(HysteriaServerConfig) },
|
||||||
|
"masque": func() interface{} { return new(MasqueServerConfig) },
|
||||||
"tun": func() interface{} { return new(TunConfig) },
|
"tun": func() interface{} { return new(TunConfig) },
|
||||||
}, "protocol", "settings")
|
}, "protocol", "settings")
|
||||||
|
|
||||||
@@ -48,6 +50,7 @@ var (
|
|||||||
"vmess": func() interface{} { return new(VMessOutboundConfig) },
|
"vmess": func() interface{} { return new(VMessOutboundConfig) },
|
||||||
"trojan": func() interface{} { return new(TrojanClientConfig) },
|
"trojan": func() interface{} { return new(TrojanClientConfig) },
|
||||||
"hysteria": func() interface{} { return new(HysteriaClientConfig) },
|
"hysteria": func() interface{} { return new(HysteriaClientConfig) },
|
||||||
|
"masque": func() interface{} { return new(MasqueClientConfig) },
|
||||||
"dns": func() interface{} { return new(DNSOutboundConfig) },
|
"dns": func() interface{} { return new(DNSOutboundConfig) },
|
||||||
"wireguard": func() interface{} { return &WireGuardConfig{IsClient: true} },
|
"wireguard": func() interface{} { return &WireGuardConfig{IsClient: true} },
|
||||||
}, "protocol", "settings")
|
}, "protocol", "settings")
|
||||||
@@ -203,6 +206,9 @@ func (c *InboundDetourConfig) Build() (*core.InboundHandlerConfig, error) {
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, errors.New("failed to build inbound handler for protocol ", c.Protocol).Base(err)
|
return nil, errors.New("failed to build inbound handler for protocol ", c.Protocol).Base(err)
|
||||||
}
|
}
|
||||||
|
if _, ok := ts.(*masque.ServerConfig); !ok && receiverSettings.StreamSettings != nil && receiverSettings.StreamSettings.ProtocolName == "masque" {
|
||||||
|
return nil, errors.New("the masque transport can only be used by the masque inbound")
|
||||||
|
}
|
||||||
|
|
||||||
return &core.InboundHandlerConfig{
|
return &core.InboundHandlerConfig{
|
||||||
Tag: c.Tag,
|
Tag: c.Tag,
|
||||||
@@ -338,6 +344,14 @@ func (c *OutboundDetourConfig) Build() (*core.OutboundHandlerConfig, error) {
|
|||||||
return nil, err
|
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 fc, ok := ts.(*freedom.Config); ok {
|
||||||
if senderSettings.StreamSettings != nil &&
|
if senderSettings.StreamSettings != nil &&
|
||||||
senderSettings.StreamSettings.SocketSettings != nil &&
|
senderSettings.StreamSettings.SocketSettings != nil &&
|
||||||
|
|||||||
@@ -12,6 +12,7 @@ import (
|
|||||||
"github.com/xtls/xray-core/core"
|
"github.com/xtls/xray-core/core"
|
||||||
"github.com/xtls/xray-core/infra/conf"
|
"github.com/xtls/xray-core/infra/conf"
|
||||||
"github.com/xtls/xray-core/infra/conf/serial"
|
"github.com/xtls/xray-core/infra/conf/serial"
|
||||||
|
"github.com/xtls/xray-core/proxy/masque"
|
||||||
"github.com/xtls/xray-core/proxy/shadowsocks"
|
"github.com/xtls/xray-core/proxy/shadowsocks"
|
||||||
"github.com/xtls/xray-core/proxy/shadowsocks_2022"
|
"github.com/xtls/xray-core/proxy/shadowsocks_2022"
|
||||||
"github.com/xtls/xray-core/proxy/trojan"
|
"github.com/xtls/xray-core/proxy/trojan"
|
||||||
@@ -88,6 +89,8 @@ func extractInboundUsers(inb *core.InboundHandlerConfig) []*protocol.User {
|
|||||||
return ty.Users
|
return ty.Users
|
||||||
case *shadowsocks_2022.MultiUserServerConfig:
|
case *shadowsocks_2022.MultiUserServerConfig:
|
||||||
return ty.Users
|
return ty.Users
|
||||||
|
case *masque.ServerConfig:
|
||||||
|
return ty.Users
|
||||||
default:
|
default:
|
||||||
fmt.Println("unsupported inbound type")
|
fmt.Println("unsupported inbound type")
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -41,6 +41,7 @@ import (
|
|||||||
_ "github.com/xtls/xray-core/proxy/freedom"
|
_ "github.com/xtls/xray-core/proxy/freedom"
|
||||||
_ "github.com/xtls/xray-core/proxy/http"
|
_ "github.com/xtls/xray-core/proxy/http"
|
||||||
_ "github.com/xtls/xray-core/proxy/loopback"
|
_ "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/shadowsocks"
|
||||||
_ "github.com/xtls/xray-core/proxy/socks"
|
_ "github.com/xtls/xray-core/proxy/socks"
|
||||||
_ "github.com/xtls/xray-core/proxy/trojan"
|
_ "github.com/xtls/xray-core/proxy/trojan"
|
||||||
@@ -54,12 +55,14 @@ import (
|
|||||||
_ "github.com/xtls/xray-core/transport/internet/grpc"
|
_ "github.com/xtls/xray-core/transport/internet/grpc"
|
||||||
_ "github.com/xtls/xray-core/transport/internet/httpupgrade"
|
_ "github.com/xtls/xray-core/transport/internet/httpupgrade"
|
||||||
_ "github.com/xtls/xray-core/transport/internet/kcp"
|
_ "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/reality"
|
||||||
_ "github.com/xtls/xray-core/transport/internet/splithttp"
|
_ "github.com/xtls/xray-core/transport/internet/splithttp"
|
||||||
_ "github.com/xtls/xray-core/transport/internet/tcp"
|
_ "github.com/xtls/xray-core/transport/internet/tcp"
|
||||||
_ "github.com/xtls/xray-core/transport/internet/tls"
|
_ "github.com/xtls/xray-core/transport/internet/tls"
|
||||||
_ "github.com/xtls/xray-core/transport/internet/udp"
|
_ "github.com/xtls/xray-core/transport/internet/udp"
|
||||||
_ "github.com/xtls/xray-core/transport/internet/websocket"
|
_ "github.com/xtls/xray-core/transport/internet/websocket"
|
||||||
|
_ "github.com/xtls/xray-core/transport/internet/xdrive"
|
||||||
|
|
||||||
// Transport headers
|
// Transport headers
|
||||||
_ "github.com/xtls/xray-core/transport/internet/headers/http"
|
_ "github.com/xtls/xray-core/transport/internet/headers/http"
|
||||||
|
|||||||
@@ -115,11 +115,7 @@ Start:
|
|||||||
|
|
||||||
request, err := http.ReadRequest(reader)
|
request, err := http.ReadRequest(reader)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
trace := errors.New("failed to read http request").Base(err)
|
return errors.New("failed to read http request").Base(err)
|
||||||
if errors.Cause(err) != io.EOF && !isTimeout(errors.Cause(err)) {
|
|
||||||
trace.AtWarning()
|
|
||||||
}
|
|
||||||
return trace
|
|
||||||
}
|
}
|
||||||
|
|
||||||
if len(s.config.Accounts) > 0 {
|
if len(s.config.Accounts) > 0 {
|
||||||
@@ -147,7 +143,7 @@ Start:
|
|||||||
}
|
}
|
||||||
dest, err := http_proto.ParseHost(host, defaultPort)
|
dest, err := http_proto.ParseHost(host, defaultPort)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return errors.New("malformed proxy host: ", host).AtWarning().Base(err)
|
return errors.New("malformed proxy host: ", host).Base(err)
|
||||||
}
|
}
|
||||||
ctx = log.ContextWithAccessMessage(ctx, &log.AccessMessage{
|
ctx = log.ContextWithAccessMessage(ctx, &log.AccessMessage{
|
||||||
From: conn.RemoteAddr(),
|
From: conn.RemoteAddr(),
|
||||||
@@ -262,7 +258,7 @@ func (s *Server) handlePlainHTTP(ctx context.Context, request *http.Request, wri
|
|||||||
requestWriter := buf.NewBufferedWriter(link.Writer)
|
requestWriter := buf.NewBufferedWriter(link.Writer)
|
||||||
common.Must(requestWriter.SetBuffered(false))
|
common.Must(requestWriter.SetBuffered(false))
|
||||||
if err := request.Write(requestWriter); err != nil {
|
if err := request.Write(requestWriter); err != nil {
|
||||||
return errors.New("failed to write whole request").Base(err).AtWarning()
|
return errors.New("failed to write whole request").Base(err)
|
||||||
}
|
}
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
@@ -299,7 +295,7 @@ func (s *Server) handlePlainHTTP(ctx context.Context, request *http.Request, wri
|
|||||||
response.Header.Set("Proxy-Connection", "close")
|
response.Header.Set("Proxy-Connection", "close")
|
||||||
}
|
}
|
||||||
if err := response.Write(writer); err != nil {
|
if err := response.Write(writer); err != nil {
|
||||||
return errors.New("failed to write response").Base(err).AtWarning()
|
return errors.New("failed to write response").Base(err)
|
||||||
}
|
}
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -62,7 +62,7 @@ func (c *Client) Process(ctx context.Context, link *transport.Link, dialer inter
|
|||||||
|
|
||||||
conn, err := dialer.Dial(hysteria.ContextWithDatagram(ctx, target.Network == net.Network_UDP), c.server.Destination)
|
conn, err := dialer.Dial(hysteria.ContextWithDatagram(ctx, target.Network == net.Network_UDP), c.server.Destination)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return errors.New("failed to find an available destination").AtWarning().Base(err)
|
return errors.New("failed to find an available destination").Base(err)
|
||||||
}
|
}
|
||||||
defer conn.Close()
|
defer conn.Close()
|
||||||
errors.LogInfo(ctx, "tunneling request to ", target, " via ", target.Network, ":", c.server.Destination.NetAddr())
|
errors.LogInfo(ctx, "tunneling request to ", target, " via ", target.Network, ":", c.server.Destination.NetAddr())
|
||||||
@@ -236,14 +236,14 @@ type UDPReader struct {
|
|||||||
|
|
||||||
func (r *UDPReader) ReadFrom(p []byte) (n int, addr *net.Destination, err error) {
|
func (r *UDPReader) ReadFrom(p []byte) (n int, addr *net.Destination, err error) {
|
||||||
for {
|
for {
|
||||||
var buf [hysteria.MaxDatagramFrameSize]byte
|
var packet [1500]byte
|
||||||
|
|
||||||
n, err := r.reader.Read(buf[:])
|
n, err := r.reader.Read(packet[:])
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return 0, nil, err
|
return 0, nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
msg, err := ParseUDPMessage(buf[:n])
|
msg, err := ParseUDPMessage(packet[:n])
|
||||||
if err != nil {
|
if err != nil {
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -40,11 +40,11 @@ func NewServer(ctx context.Context, config *ServerConfig) (*Server, error) {
|
|||||||
for _, user := range config.Users {
|
for _, user := range config.Users {
|
||||||
u, err := user.ToMemoryUser()
|
u, err := user.ToMemoryUser()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, errors.New("failed to get hysteria user").Base(err).AtError()
|
return nil, errors.New("failed to get hysteria user").Base(err)
|
||||||
}
|
}
|
||||||
|
|
||||||
if err := validator.Add(u); err != nil {
|
if err := validator.Add(u); err != nil {
|
||||||
return nil, errors.New("failed to add user").Base(err).AtError()
|
return nil, errors.New("failed to add user").Base(err)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -56,7 +56,7 @@ func (l *Loopback) init(config *Config, dispatcherInstance routing.Dispatcher) e
|
|||||||
if config.Sniffing.GetEnabled() {
|
if config.Sniffing.GetEnabled() {
|
||||||
request, err := proxyman.BuildSniffingRequest(config.Sniffing)
|
request, err := proxyman.BuildSniffingRequest(config.Sniffing)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return errors.New("failed to build loopback sniffing request").Base(err).AtError()
|
return errors.New("failed to build loopback sniffing request").Base(err)
|
||||||
}
|
}
|
||||||
l.sniffingRequest = request
|
l.sniffingRequest = request
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -0,0 +1,108 @@
|
|||||||
|
package masque
|
||||||
|
|
||||||
|
import (
|
||||||
|
"crypto/subtle"
|
||||||
|
"strings"
|
||||||
|
"sync"
|
||||||
|
|
||||||
|
"github.com/xtls/xray-core/common/errors"
|
||||||
|
"github.com/xtls/xray-core/common/protocol"
|
||||||
|
"google.golang.org/protobuf/proto"
|
||||||
|
)
|
||||||
|
|
||||||
|
func (a *Account) AsAccount() (protocol.Account, error) {
|
||||||
|
return &MemoryAccount{Password: a.Password}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
type MemoryAccount struct {
|
||||||
|
Password string
|
||||||
|
}
|
||||||
|
|
||||||
|
func (a *MemoryAccount) Equals(other protocol.Account) bool {
|
||||||
|
b, ok := other.(*MemoryAccount)
|
||||||
|
return ok && a.Password == b.Password
|
||||||
|
}
|
||||||
|
|
||||||
|
func (a *MemoryAccount) ToProto() proto.Message {
|
||||||
|
return &Account{Password: a.Password}
|
||||||
|
}
|
||||||
|
|
||||||
|
type validator struct {
|
||||||
|
mu sync.RWMutex
|
||||||
|
users map[string]*protocol.MemoryUser
|
||||||
|
}
|
||||||
|
|
||||||
|
func newValidator() *validator {
|
||||||
|
return &validator{users: make(map[string]*protocol.MemoryUser)}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (v *validator) add(user *protocol.MemoryUser) error {
|
||||||
|
account, ok := user.Account.(*MemoryAccount)
|
||||||
|
if !ok {
|
||||||
|
return errors.New("not a MASQUE account")
|
||||||
|
}
|
||||||
|
if user.Email == "" || strings.Contains(user.Email, ":") {
|
||||||
|
return errors.New("invalid email ", user.Email)
|
||||||
|
}
|
||||||
|
if account.Password == "" {
|
||||||
|
return errors.New("empty password for ", user.Email)
|
||||||
|
}
|
||||||
|
email := strings.ToLower(user.Email)
|
||||||
|
v.mu.Lock()
|
||||||
|
defer v.mu.Unlock()
|
||||||
|
if _, found := v.users[email]; found {
|
||||||
|
return errors.New("user ", user.Email, " already exists")
|
||||||
|
}
|
||||||
|
v.users[email] = user
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (v *validator) delByEmail(email string) (*protocol.MemoryUser, error) {
|
||||||
|
key := strings.ToLower(email)
|
||||||
|
v.mu.Lock()
|
||||||
|
defer v.mu.Unlock()
|
||||||
|
user, found := v.users[key]
|
||||||
|
if !found {
|
||||||
|
return nil, errors.New("user ", email, " not found")
|
||||||
|
}
|
||||||
|
delete(v.users, key)
|
||||||
|
return user, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (v *validator) contains(user *protocol.MemoryUser) bool {
|
||||||
|
v.mu.RLock()
|
||||||
|
defer v.mu.RUnlock()
|
||||||
|
return v.users[strings.ToLower(user.Email)] == user
|
||||||
|
}
|
||||||
|
|
||||||
|
func (v *validator) get(email, password string) *protocol.MemoryUser {
|
||||||
|
v.mu.RLock()
|
||||||
|
user := v.users[strings.ToLower(email)]
|
||||||
|
v.mu.RUnlock()
|
||||||
|
if user == nil || subtle.ConstantTimeCompare([]byte(user.Account.(*MemoryAccount).Password), []byte(password)) != 1 {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
return user
|
||||||
|
}
|
||||||
|
|
||||||
|
func (v *validator) getByEmail(email string) *protocol.MemoryUser {
|
||||||
|
v.mu.RLock()
|
||||||
|
defer v.mu.RUnlock()
|
||||||
|
return v.users[strings.ToLower(email)]
|
||||||
|
}
|
||||||
|
|
||||||
|
func (v *validator) getAll() []*protocol.MemoryUser {
|
||||||
|
v.mu.RLock()
|
||||||
|
defer v.mu.RUnlock()
|
||||||
|
users := make([]*protocol.MemoryUser, 0, len(v.users))
|
||||||
|
for _, user := range v.users {
|
||||||
|
users = append(users, user)
|
||||||
|
}
|
||||||
|
return users
|
||||||
|
}
|
||||||
|
|
||||||
|
func (v *validator) count() int64 {
|
||||||
|
v.mu.RLock()
|
||||||
|
defer v.mu.RUnlock()
|
||||||
|
return int64(len(v.users))
|
||||||
|
}
|
||||||
@@ -0,0 +1,49 @@
|
|||||||
|
package masque
|
||||||
|
|
||||||
|
import (
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
"github.com/xtls/xray-core/common/protocol"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestValidator(t *testing.T) {
|
||||||
|
v := newValidator()
|
||||||
|
user := &protocol.MemoryUser{Email: "U@example.com", Account: &MemoryAccount{Password: "p"}}
|
||||||
|
require.NoError(t, v.add(user))
|
||||||
|
for _, u := range []*protocol.MemoryUser{
|
||||||
|
{Email: "u@example.com", Account: &MemoryAccount{Password: "other"}},
|
||||||
|
{Account: &MemoryAccount{Password: "p"}},
|
||||||
|
{Email: "a:b", Account: &MemoryAccount{Password: "p"}},
|
||||||
|
{Email: "b@example.com", Account: &MemoryAccount{}},
|
||||||
|
} {
|
||||||
|
require.Error(t, v.add(u), u.Email)
|
||||||
|
}
|
||||||
|
|
||||||
|
require.Equal(t, user, v.get("u@example.com", "p"))
|
||||||
|
require.Equal(t, user, v.get("U@EXAMPLE.COM", "p"))
|
||||||
|
require.Nil(t, v.get("u@example.com", "x"))
|
||||||
|
require.Nil(t, v.get("x@example.com", "p"))
|
||||||
|
require.Nil(t, v.get("", ""))
|
||||||
|
require.Equal(t, user, v.getByEmail("u@example.com"))
|
||||||
|
require.Equal(t, []*protocol.MemoryUser{user}, v.getAll())
|
||||||
|
require.Equal(t, int64(1), v.count())
|
||||||
|
|
||||||
|
require.True(t, v.contains(user))
|
||||||
|
removed, err := v.delByEmail("u@EXAMPLE.com")
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.Equal(t, user, removed)
|
||||||
|
_, err = v.delByEmail("u@example.com")
|
||||||
|
require.Error(t, err)
|
||||||
|
require.False(t, v.contains(user))
|
||||||
|
require.Nil(t, v.get("u@example.com", "p"))
|
||||||
|
require.Zero(t, v.count())
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestAccount(t *testing.T) {
|
||||||
|
account, err := (&Account{Password: "p"}).AsAccount()
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.True(t, account.Equals(&MemoryAccount{Password: "p"}))
|
||||||
|
require.False(t, account.Equals(&MemoryAccount{Password: "x"}))
|
||||||
|
require.Equal(t, &Account{Password: "p"}, account.ToProto())
|
||||||
|
}
|
||||||
@@ -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,250 @@
|
|||||||
|
// 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
|
||||||
|
}
|
||||||
|
|
||||||
|
type Account struct {
|
||||||
|
state protoimpl.MessageState `protogen:"open.v1"`
|
||||||
|
Password string `protobuf:"bytes,1,opt,name=password,proto3" json:"password,omitempty"`
|
||||||
|
unknownFields protoimpl.UnknownFields
|
||||||
|
sizeCache protoimpl.SizeCache
|
||||||
|
}
|
||||||
|
|
||||||
|
func (x *Account) Reset() {
|
||||||
|
*x = Account{}
|
||||||
|
mi := &file_proxy_masque_config_proto_msgTypes[1]
|
||||||
|
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
|
||||||
|
ms.StoreMessageInfo(mi)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (x *Account) String() string {
|
||||||
|
return protoimpl.X.MessageStringOf(x)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (*Account) ProtoMessage() {}
|
||||||
|
|
||||||
|
func (x *Account) ProtoReflect() protoreflect.Message {
|
||||||
|
mi := &file_proxy_masque_config_proto_msgTypes[1]
|
||||||
|
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 Account.ProtoReflect.Descriptor instead.
|
||||||
|
func (*Account) Descriptor() ([]byte, []int) {
|
||||||
|
return file_proxy_masque_config_proto_rawDescGZIP(), []int{1}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (x *Account) GetPassword() string {
|
||||||
|
if x != nil {
|
||||||
|
return x.Password
|
||||||
|
}
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
|
||||||
|
type ServerConfig struct {
|
||||||
|
state protoimpl.MessageState `protogen:"open.v1"`
|
||||||
|
Users []*protocol.User `protobuf:"bytes,1,rep,name=users,proto3" json:"users,omitempty"`
|
||||||
|
Address []string `protobuf:"bytes,2,rep,name=address,proto3" json:"address,omitempty"`
|
||||||
|
Mtu uint32 `protobuf:"varint,3,opt,name=mtu,proto3" json:"mtu,omitempty"`
|
||||||
|
unknownFields protoimpl.UnknownFields
|
||||||
|
sizeCache protoimpl.SizeCache
|
||||||
|
}
|
||||||
|
|
||||||
|
func (x *ServerConfig) Reset() {
|
||||||
|
*x = ServerConfig{}
|
||||||
|
mi := &file_proxy_masque_config_proto_msgTypes[2]
|
||||||
|
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
|
||||||
|
ms.StoreMessageInfo(mi)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (x *ServerConfig) String() string {
|
||||||
|
return protoimpl.X.MessageStringOf(x)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (*ServerConfig) ProtoMessage() {}
|
||||||
|
|
||||||
|
func (x *ServerConfig) ProtoReflect() protoreflect.Message {
|
||||||
|
mi := &file_proxy_masque_config_proto_msgTypes[2]
|
||||||
|
if x != nil {
|
||||||
|
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
|
||||||
|
if ms.LoadMessageInfo() == nil {
|
||||||
|
ms.StoreMessageInfo(mi)
|
||||||
|
}
|
||||||
|
return ms
|
||||||
|
}
|
||||||
|
return mi.MessageOf(x)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Deprecated: Use ServerConfig.ProtoReflect.Descriptor instead.
|
||||||
|
func (*ServerConfig) Descriptor() ([]byte, []int) {
|
||||||
|
return file_proxy_masque_config_proto_rawDescGZIP(), []int{2}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (x *ServerConfig) GetUsers() []*protocol.User {
|
||||||
|
if x != nil {
|
||||||
|
return x.Users
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (x *ServerConfig) GetAddress() []string {
|
||||||
|
if x != nil {
|
||||||
|
return x.Address
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (x *ServerConfig) GetMtu() uint32 {
|
||||||
|
if x != nil {
|
||||||
|
return x.Mtu
|
||||||
|
}
|
||||||
|
return 0
|
||||||
|
}
|
||||||
|
|
||||||
|
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\x1a\x1acommon/protocol/user.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\tremoteDns\"%\n" +
|
||||||
|
"\aAccount\x12\x1a\n" +
|
||||||
|
"\bpassword\x18\x01 \x01(\tR\bpassword\"l\n" +
|
||||||
|
"\fServerConfig\x120\n" +
|
||||||
|
"\x05users\x18\x01 \x03(\v2\x1a.xray.common.protocol.UserR\x05users\x12\x18\n" +
|
||||||
|
"\aaddress\x18\x02 \x03(\tR\aaddress\x12\x10\n" +
|
||||||
|
"\x03mtu\x18\x03 \x01(\rR\x03mtuBU\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, 3)
|
||||||
|
var file_proxy_masque_config_proto_goTypes = []any{
|
||||||
|
(*ClientConfig)(nil), // 0: xray.proxy.masque.ClientConfig
|
||||||
|
(*Account)(nil), // 1: xray.proxy.masque.Account
|
||||||
|
(*ServerConfig)(nil), // 2: xray.proxy.masque.ServerConfig
|
||||||
|
(*protocol.ServerEndpoint)(nil), // 3: xray.common.protocol.ServerEndpoint
|
||||||
|
(*protocol.User)(nil), // 4: xray.common.protocol.User
|
||||||
|
}
|
||||||
|
var file_proxy_masque_config_proto_depIdxs = []int32{
|
||||||
|
3, // 0: xray.proxy.masque.ClientConfig.server:type_name -> xray.common.protocol.ServerEndpoint
|
||||||
|
4, // 1: xray.proxy.masque.ServerConfig.users:type_name -> xray.common.protocol.User
|
||||||
|
2, // [2:2] is the sub-list for method output_type
|
||||||
|
2, // [2:2] is the sub-list for method input_type
|
||||||
|
2, // [2:2] is the sub-list for extension type_name
|
||||||
|
2, // [2:2] is the sub-list for extension extendee
|
||||||
|
0, // [0:2] is the sub-list for field type_name
|
||||||
|
}
|
||||||
|
|
||||||
|
func init() { file_proxy_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: 3,
|
||||||
|
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,25 @@
|
|||||||
|
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";
|
||||||
|
import "common/protocol/user.proto";
|
||||||
|
|
||||||
|
message ClientConfig {
|
||||||
|
xray.common.protocol.ServerEndpoint server = 1;
|
||||||
|
repeated string remote_dns = 2;
|
||||||
|
}
|
||||||
|
|
||||||
|
message Account {
|
||||||
|
string password = 1;
|
||||||
|
}
|
||||||
|
|
||||||
|
message ServerConfig {
|
||||||
|
repeated xray.common.protocol.User users = 1;
|
||||||
|
repeated string address = 2;
|
||||||
|
uint32 mtu = 3;
|
||||||
|
}
|
||||||
@@ -0,0 +1,80 @@
|
|||||||
|
package masque
|
||||||
|
|
||||||
|
import (
|
||||||
|
"net/netip"
|
||||||
|
"sync"
|
||||||
|
|
||||||
|
"github.com/xtls/xray-core/common/errors"
|
||||||
|
)
|
||||||
|
|
||||||
|
type addressPool struct {
|
||||||
|
mu sync.Mutex
|
||||||
|
prefix netip.Prefix
|
||||||
|
server netip.Addr
|
||||||
|
first netip.Addr
|
||||||
|
last netip.Addr
|
||||||
|
next netip.Addr
|
||||||
|
used map[netip.Addr]struct{}
|
||||||
|
}
|
||||||
|
|
||||||
|
func newAddressPool(address netip.Prefix) (*addressPool, error) {
|
||||||
|
server := address.Addr()
|
||||||
|
if server.Is4In6() || server.Zone() != "" {
|
||||||
|
return nil, errors.New("invalid address ", address)
|
||||||
|
}
|
||||||
|
prefix := address.Masked()
|
||||||
|
last := lastAddr(prefix)
|
||||||
|
if server == prefix.Addr() || server.Is4() && server == last {
|
||||||
|
return nil, errors.New("address ", address, " is not a host address")
|
||||||
|
}
|
||||||
|
if server.Is4() {
|
||||||
|
last = last.Prev()
|
||||||
|
}
|
||||||
|
first := prefix.Addr().Next()
|
||||||
|
if first == last {
|
||||||
|
return nil, errors.New("address ", address, " leaves no addresses to assign")
|
||||||
|
}
|
||||||
|
return &addressPool{
|
||||||
|
prefix: prefix,
|
||||||
|
server: server,
|
||||||
|
first: first,
|
||||||
|
last: last,
|
||||||
|
next: first,
|
||||||
|
used: make(map[netip.Addr]struct{}),
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func lastAddr(prefix netip.Prefix) netip.Addr {
|
||||||
|
b := prefix.Addr().AsSlice()
|
||||||
|
for i := prefix.Bits(); i < len(b)*8; i++ {
|
||||||
|
b[i/8] |= 1 << (7 - i%8)
|
||||||
|
}
|
||||||
|
addr, _ := netip.AddrFromSlice(b)
|
||||||
|
return addr
|
||||||
|
}
|
||||||
|
|
||||||
|
func (p *addressPool) allocate() (netip.Addr, bool) {
|
||||||
|
p.mu.Lock()
|
||||||
|
defer p.mu.Unlock()
|
||||||
|
for addr := p.next; ; {
|
||||||
|
next := addr.Next()
|
||||||
|
if addr == p.last {
|
||||||
|
next = p.first
|
||||||
|
}
|
||||||
|
if _, found := p.used[addr]; !found && addr != p.server {
|
||||||
|
p.used[addr] = struct{}{}
|
||||||
|
p.next = next
|
||||||
|
return addr, true
|
||||||
|
}
|
||||||
|
if next == p.next {
|
||||||
|
return netip.Addr{}, false
|
||||||
|
}
|
||||||
|
addr = next
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (p *addressPool) release(addr netip.Addr) {
|
||||||
|
p.mu.Lock()
|
||||||
|
defer p.mu.Unlock()
|
||||||
|
delete(p.used, addr)
|
||||||
|
}
|
||||||
@@ -0,0 +1,59 @@
|
|||||||
|
package masque
|
||||||
|
|
||||||
|
import (
|
||||||
|
"net/netip"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
)
|
||||||
|
|
||||||
|
func allocateAll(p *addressPool) []netip.Addr {
|
||||||
|
var addrs []netip.Addr
|
||||||
|
for {
|
||||||
|
addr, ok := p.allocate()
|
||||||
|
if !ok {
|
||||||
|
return addrs
|
||||||
|
}
|
||||||
|
addrs = append(addrs, addr)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestAddressPool(t *testing.T) {
|
||||||
|
p, err := newAddressPool(netip.MustParsePrefix("10.0.0.1/29"))
|
||||||
|
require.NoError(t, err)
|
||||||
|
var want []netip.Addr
|
||||||
|
for _, s := range []string{"10.0.0.2", "10.0.0.3", "10.0.0.4", "10.0.0.5", "10.0.0.6"} {
|
||||||
|
want = append(want, netip.MustParseAddr(s))
|
||||||
|
}
|
||||||
|
require.Equal(t, want, allocateAll(p))
|
||||||
|
|
||||||
|
p.release(netip.MustParseAddr("10.0.0.4"))
|
||||||
|
addr, ok := p.allocate()
|
||||||
|
require.True(t, ok)
|
||||||
|
require.Equal(t, netip.MustParseAddr("10.0.0.4"), addr)
|
||||||
|
_, ok = p.allocate()
|
||||||
|
require.False(t, ok)
|
||||||
|
|
||||||
|
p, err = newAddressPool(netip.MustParsePrefix("fd00::1/126"))
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.Equal(t, []netip.Addr{netip.MustParseAddr("fd00::2"), netip.MustParseAddr("fd00::3")}, allocateAll(p))
|
||||||
|
|
||||||
|
p, err = newAddressPool(netip.MustParsePrefix("10.0.0.2/30"))
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.Equal(t, []netip.Addr{netip.MustParseAddr("10.0.0.1")}, allocateAll(p))
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestAddressPoolRejects(t *testing.T) {
|
||||||
|
for _, s := range []string{
|
||||||
|
"10.0.0.0/24",
|
||||||
|
"10.0.0.255/24",
|
||||||
|
"10.0.0.1/31",
|
||||||
|
"10.0.0.1/32",
|
||||||
|
"fd00::1/127",
|
||||||
|
"fd00::1/128",
|
||||||
|
"::ffff:10.0.0.1/120",
|
||||||
|
} {
|
||||||
|
_, err := newAddressPool(netip.MustParsePrefix(s))
|
||||||
|
require.Error(t, err, s)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,550 @@
|
|||||||
|
package masque
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
go_errors "errors"
|
||||||
|
"io"
|
||||||
|
stdnet "net"
|
||||||
|
"net/http"
|
||||||
|
"net/netip"
|
||||||
|
"slices"
|
||||||
|
"sync"
|
||||||
|
|
||||||
|
"golang.zx2c4.com/wireguard/tun"
|
||||||
|
|
||||||
|
"github.com/xtls/xray-core/common"
|
||||||
|
"github.com/xtls/xray-core/common/buf"
|
||||||
|
c "github.com/xtls/xray-core/common/ctx"
|
||||||
|
"github.com/xtls/xray-core/common/errors"
|
||||||
|
"github.com/xtls/xray-core/common/log"
|
||||||
|
"github.com/xtls/xray-core/common/net"
|
||||||
|
"github.com/xtls/xray-core/common/protocol"
|
||||||
|
"github.com/xtls/xray-core/common/session"
|
||||||
|
"github.com/xtls/xray-core/core"
|
||||||
|
"github.com/xtls/xray-core/features/routing"
|
||||||
|
"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/masque/connectip"
|
||||||
|
"github.com/xtls/xray-core/transport/internet/stat"
|
||||||
|
"github.com/xtls/xray-core/transport/internet/tls"
|
||||||
|
)
|
||||||
|
|
||||||
|
const (
|
||||||
|
authenticateHeader = `Basic realm="masque", charset="UTF-8"`
|
||||||
|
tunnelQueueSize = 512
|
||||||
|
)
|
||||||
|
|
||||||
|
type Server struct {
|
||||||
|
validator *validator
|
||||||
|
dispatcher routing.Dispatcher
|
||||||
|
ctx context.Context
|
||||||
|
tag string
|
||||||
|
sniffing session.SniffingRequest
|
||||||
|
mtu int
|
||||||
|
|
||||||
|
dev tun.Device
|
||||||
|
pools []*addressPool
|
||||||
|
local []netip.Addr
|
||||||
|
|
||||||
|
mu sync.RWMutex
|
||||||
|
tunnels map[netip.Addr]*serverTunnel
|
||||||
|
closed bool
|
||||||
|
started bool
|
||||||
|
}
|
||||||
|
|
||||||
|
type serverTunnel struct {
|
||||||
|
conn stat.Connection
|
||||||
|
ipConn *connectip.Conn
|
||||||
|
user *protocol.MemoryUser
|
||||||
|
addrs []netip.Addr
|
||||||
|
queue chan *buf.Buffer
|
||||||
|
done chan struct{}
|
||||||
|
|
||||||
|
mu sync.Mutex
|
||||||
|
conns map[net.Conn]struct{}
|
||||||
|
}
|
||||||
|
|
||||||
|
func newServerTunnel(conn stat.Connection, user *protocol.MemoryUser) *serverTunnel {
|
||||||
|
return &serverTunnel{
|
||||||
|
conn: conn,
|
||||||
|
user: user,
|
||||||
|
queue: make(chan *buf.Buffer, tunnelQueueSize),
|
||||||
|
done: make(chan struct{}),
|
||||||
|
conns: make(map[net.Conn]struct{}),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *serverTunnel) send(b *buf.Buffer) bool {
|
||||||
|
select {
|
||||||
|
case <-t.done:
|
||||||
|
return false
|
||||||
|
default:
|
||||||
|
}
|
||||||
|
select {
|
||||||
|
case t.queue <- b:
|
||||||
|
return true
|
||||||
|
default:
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *serverTunnel) track(conn net.Conn) bool {
|
||||||
|
t.mu.Lock()
|
||||||
|
defer t.mu.Unlock()
|
||||||
|
if t.conns == nil {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
t.conns[conn] = struct{}{}
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *serverTunnel) untrack(conn net.Conn) {
|
||||||
|
t.mu.Lock()
|
||||||
|
delete(t.conns, conn)
|
||||||
|
t.mu.Unlock()
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *serverTunnel) close() {
|
||||||
|
t.mu.Lock()
|
||||||
|
conns := t.conns
|
||||||
|
if conns != nil {
|
||||||
|
t.conns = nil
|
||||||
|
close(t.done)
|
||||||
|
}
|
||||||
|
t.mu.Unlock()
|
||||||
|
for conn := range conns {
|
||||||
|
conn.Close()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func NewServer(ctx context.Context, config *ServerConfig) (*Server, error) {
|
||||||
|
v := core.MustFromContext(ctx)
|
||||||
|
|
||||||
|
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"`)
|
||||||
|
}
|
||||||
|
|
||||||
|
users := newValidator()
|
||||||
|
for _, user := range config.Users {
|
||||||
|
u, err := user.ToMemoryUser()
|
||||||
|
if err != nil {
|
||||||
|
return nil, errors.New("failed to get MASQUE user").Base(err)
|
||||||
|
}
|
||||||
|
if err := users.add(u); err != nil {
|
||||||
|
return nil, errors.New("failed to add user").Base(err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
var pools []*addressPool
|
||||||
|
var local []netip.Addr
|
||||||
|
for _, s := range config.Address {
|
||||||
|
prefix, err := netip.ParsePrefix(s)
|
||||||
|
if err != nil {
|
||||||
|
return nil, errors.New("invalid address ", s).Base(err)
|
||||||
|
}
|
||||||
|
if slices.ContainsFunc(local, func(addr netip.Addr) bool { return addr.Is4() == prefix.Addr().Is4() }) {
|
||||||
|
return nil, errors.New("only one address per IP family is supported")
|
||||||
|
}
|
||||||
|
pool, err := newAddressPool(prefix)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
pools = append(pools, pool)
|
||||||
|
local = append(local, prefix.Addr())
|
||||||
|
}
|
||||||
|
if len(pools) == 0 {
|
||||||
|
return nil, errors.New("no address to assign")
|
||||||
|
}
|
||||||
|
|
||||||
|
mtu := int(config.Mtu)
|
||||||
|
if mtu == 0 {
|
||||||
|
mtu = masque.MinPacketSize
|
||||||
|
}
|
||||||
|
dev, _, gstack, err := wireguard.CreateNetTUN(local, nil, mtu, false)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
s := &Server{
|
||||||
|
validator: users,
|
||||||
|
dispatcher: v.GetFeature(routing.DispatcherType()).(routing.Dispatcher),
|
||||||
|
ctx: core.ToBackgroundDetachedContext(ctx),
|
||||||
|
mtu: mtu,
|
||||||
|
dev: dev,
|
||||||
|
pools: pools,
|
||||||
|
local: local,
|
||||||
|
tunnels: make(map[netip.Addr]*serverTunnel),
|
||||||
|
}
|
||||||
|
if inbound := session.InboundFromContext(ctx); inbound != nil {
|
||||||
|
s.tag = inbound.Tag
|
||||||
|
}
|
||||||
|
if content := session.ContentFromContext(ctx); content != nil {
|
||||||
|
s.sniffing = content.SniffingRequest
|
||||||
|
}
|
||||||
|
wireguard.CreateForwarder(gstack, s.handleConnection)
|
||||||
|
return s, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *Server) Start() error {
|
||||||
|
s.mu.Lock()
|
||||||
|
defer s.mu.Unlock()
|
||||||
|
if s.started || s.closed {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
s.started = true
|
||||||
|
go s.readFromStack()
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *Server) Close() error {
|
||||||
|
s.mu.Lock()
|
||||||
|
if s.closed {
|
||||||
|
s.mu.Unlock()
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
s.closed = true
|
||||||
|
var tunnels []*serverTunnel
|
||||||
|
for _, t := range s.tunnels {
|
||||||
|
if !slices.Contains(tunnels, t) {
|
||||||
|
tunnels = append(tunnels, t)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
s.mu.Unlock()
|
||||||
|
for _, t := range tunnels {
|
||||||
|
t.conn.Close()
|
||||||
|
}
|
||||||
|
return s.dev.Close()
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *Server) AddUser(ctx context.Context, user *protocol.MemoryUser) error {
|
||||||
|
return s.validator.add(user)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *Server) RemoveUser(ctx context.Context, email string) error {
|
||||||
|
user, err := s.validator.delByEmail(email)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
s.mu.RLock()
|
||||||
|
var conns []stat.Connection
|
||||||
|
for _, t := range s.tunnels {
|
||||||
|
if t.user == user && !slices.Contains(conns, t.conn) {
|
||||||
|
conns = append(conns, t.conn)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
s.mu.RUnlock()
|
||||||
|
for _, conn := range conns {
|
||||||
|
conn.Close()
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *Server) GetUser(ctx context.Context, email string) *protocol.MemoryUser {
|
||||||
|
return s.validator.getByEmail(email)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *Server) GetUsers(ctx context.Context) []*protocol.MemoryUser {
|
||||||
|
return s.validator.getAll()
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *Server) GetUsersCount(context.Context) int64 {
|
||||||
|
return s.validator.count()
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *Server) Network() []net.Network {
|
||||||
|
return []net.Network{net.Network_TCP}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *Server) Process(ctx context.Context, network net.Network, conn stat.Connection, dispatcher routing.Dispatcher) error {
|
||||||
|
sconn, ok := stat.TryUnwrapStatsConn(conn).(*masque.ServerConn)
|
||||||
|
if !ok {
|
||||||
|
return errors.New("not a MASQUE connection")
|
||||||
|
}
|
||||||
|
inbound := session.InboundFromContext(ctx)
|
||||||
|
inbound.Name = "masque"
|
||||||
|
inbound.CanSpliceCopy = 3
|
||||||
|
|
||||||
|
name, pass, _ := sconn.Request().BasicAuth()
|
||||||
|
user := s.validator.get(name, pass)
|
||||||
|
if user == nil {
|
||||||
|
sconn.Reject(http.StatusUnauthorized, http.Header{"WWW-Authenticate": {authenticateHeader}})
|
||||||
|
log.Record(&log.AccessMessage{
|
||||||
|
From: conn.RemoteAddr(),
|
||||||
|
To: "",
|
||||||
|
Status: log.AccessRejected,
|
||||||
|
Reason: errors.New("invalid credentials"),
|
||||||
|
})
|
||||||
|
return errors.New("MASQUE: authentication failed for ", name)
|
||||||
|
}
|
||||||
|
inbound.User = user
|
||||||
|
|
||||||
|
t := newServerTunnel(conn, user)
|
||||||
|
for _, pool := range s.pools {
|
||||||
|
if addr, ok := pool.allocate(); ok {
|
||||||
|
t.addrs = append(t.addrs, addr)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
defer s.release(t)
|
||||||
|
if len(t.addrs) == 0 {
|
||||||
|
sconn.Reject(http.StatusServiceUnavailable, nil)
|
||||||
|
return errors.New("MASQUE: no address left to assign")
|
||||||
|
}
|
||||||
|
|
||||||
|
ipConn, err := sconn.Accept()
|
||||||
|
if err != nil {
|
||||||
|
return errors.New("MASQUE: failed to accept the tunnel").Base(err)
|
||||||
|
}
|
||||||
|
t.ipConn = ipConn
|
||||||
|
if !s.register(t) {
|
||||||
|
return errors.New("MASQUE: server closed")
|
||||||
|
}
|
||||||
|
if !s.validator.contains(user) {
|
||||||
|
return errors.New("MASQUE: user ", name, " was removed")
|
||||||
|
}
|
||||||
|
go s.writeToTunnel(t)
|
||||||
|
|
||||||
|
prefixes := make([]netip.Prefix, len(t.addrs))
|
||||||
|
for i, addr := range t.addrs {
|
||||||
|
prefixes[i] = netip.PrefixFrom(addr, addr.BitLen())
|
||||||
|
}
|
||||||
|
if err := ipConn.AssignAddresses(prefixes); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if err := ipConn.AdvertiseRoute(fullRoutes(t.addrs)); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
go serveAddressRequests(t)
|
||||||
|
|
||||||
|
ctx = log.ContextWithAccessMessage(ctx, &log.AccessMessage{
|
||||||
|
From: conn.RemoteAddr(),
|
||||||
|
To: "",
|
||||||
|
Status: log.AccessAccepted,
|
||||||
|
Email: user.Email,
|
||||||
|
})
|
||||||
|
errors.LogInfo(ctx, "MASQUE: tunnel from ", inbound.Source, " assigned ", t.addrs)
|
||||||
|
return s.readFromTunnel(t)
|
||||||
|
}
|
||||||
|
|
||||||
|
func fullRoutes(addrs []netip.Addr) []connectip.IPRoute {
|
||||||
|
var routes []connectip.IPRoute
|
||||||
|
if slices.ContainsFunc(addrs, netip.Addr.Is4) {
|
||||||
|
routes = append(routes, connectip.IPRoute{StartIP: netip.IPv4Unspecified(), EndIP: netip.AddrFrom4([4]byte{255, 255, 255, 255})})
|
||||||
|
}
|
||||||
|
if slices.ContainsFunc(addrs, netip.Addr.Is6) {
|
||||||
|
routes = append(routes, connectip.IPRoute{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})})
|
||||||
|
}
|
||||||
|
return routes
|
||||||
|
}
|
||||||
|
|
||||||
|
func serveAddressRequests(t *serverTunnel) {
|
||||||
|
for {
|
||||||
|
req, err := t.ipConn.ReceiveAddressRequest(context.Background())
|
||||||
|
if err != nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
assigned := make([]netip.Prefix, len(req.Prefixes))
|
||||||
|
used := make(map[netip.Addr]bool)
|
||||||
|
for i, requested := range req.Prefixes {
|
||||||
|
for _, addr := range t.addrs {
|
||||||
|
if addr.Is4() == requested.Addr().Is4() && !used[addr] {
|
||||||
|
used[addr] = true
|
||||||
|
assigned[i] = netip.PrefixFrom(addr, addr.BitLen())
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
var additional []netip.Prefix
|
||||||
|
for _, addr := range t.addrs {
|
||||||
|
if !used[addr] {
|
||||||
|
additional = append(additional, netip.PrefixFrom(addr, addr.BitLen()))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if err := req.Respond(assigned, additional); err != nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *Server) register(t *serverTunnel) bool {
|
||||||
|
s.mu.Lock()
|
||||||
|
defer s.mu.Unlock()
|
||||||
|
if s.closed {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
for _, addr := range t.addrs {
|
||||||
|
s.tunnels[addr] = t
|
||||||
|
}
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *Server) release(t *serverTunnel) {
|
||||||
|
s.mu.Lock()
|
||||||
|
for _, addr := range t.addrs {
|
||||||
|
if s.tunnels[addr] == t {
|
||||||
|
delete(s.tunnels, addr)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
s.mu.Unlock()
|
||||||
|
t.close()
|
||||||
|
for _, addr := range t.addrs {
|
||||||
|
for _, pool := range s.pools {
|
||||||
|
if pool.prefix.Contains(addr) {
|
||||||
|
pool.release(addr)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *Server) lookup(addr netip.Addr) *serverTunnel {
|
||||||
|
s.mu.RLock()
|
||||||
|
defer s.mu.RUnlock()
|
||||||
|
return s.tunnels[addr]
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *Server) inPool(addr netip.Addr) bool {
|
||||||
|
return slices.ContainsFunc(s.pools, func(pool *addressPool) bool { return pool.prefix.Contains(addr) })
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *Server) readFromTunnel(t *serverTunnel) error {
|
||||||
|
b := make([]byte, 1<<16)
|
||||||
|
for {
|
||||||
|
n, err := t.conn.Read(b)
|
||||||
|
if err != nil {
|
||||||
|
if go_errors.Is(err, io.ErrShortBuffer) {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if go_errors.Is(err, stdnet.ErrClosed) || go_errors.Is(err, io.EOF) {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
dst, ok := packetDestination(b[:n])
|
||||||
|
if !ok || dst.IsLinkLocalUnicast() || dst.IsMulticast() {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if other := s.lookup(dst); other != nil {
|
||||||
|
if other != t {
|
||||||
|
packet := buf.NewWithSize(int32(n))
|
||||||
|
packet.Write(b[:n])
|
||||||
|
if !other.send(packet) {
|
||||||
|
packet.Release()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if s.inPool(dst) && !slices.Contains(s.local, dst) {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
s.dev.Write([][]byte{b[:n]}, 0)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *Server) readFromStack() {
|
||||||
|
sizes := []int{0}
|
||||||
|
var b *buf.Buffer
|
||||||
|
for {
|
||||||
|
if b == nil {
|
||||||
|
b = buf.NewWithSize(int32(s.mtu))
|
||||||
|
}
|
||||||
|
b.Clear()
|
||||||
|
if _, err := s.dev.Read([][]byte{b.Extend(int32(s.mtu))}, sizes, 0); err != nil {
|
||||||
|
b.Release()
|
||||||
|
return
|
||||||
|
}
|
||||||
|
b.Resize(0, int32(sizes[0]))
|
||||||
|
dst, ok := packetDestination(b.Bytes())
|
||||||
|
if !ok {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if t := s.lookup(dst); t != nil && t.send(b) {
|
||||||
|
b = nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *Server) writeToTunnel(t *serverTunnel) {
|
||||||
|
for {
|
||||||
|
select {
|
||||||
|
case b := <-t.queue:
|
||||||
|
_, err := t.conn.Write(b.Bytes())
|
||||||
|
b.Release()
|
||||||
|
if ptb, ok := go_errors.AsType[*masque.PacketTooBigError](err); ok {
|
||||||
|
s.dev.Write([][]byte{ptb.ICMP}, 0)
|
||||||
|
}
|
||||||
|
case <-t.done:
|
||||||
|
return
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func packetDestination(packet []byte) (netip.Addr, bool) {
|
||||||
|
if len(packet) == 0 {
|
||||||
|
return netip.Addr{}, false
|
||||||
|
}
|
||||||
|
switch packet[0] >> 4 {
|
||||||
|
case 4:
|
||||||
|
if len(packet) >= 20 {
|
||||||
|
return netip.AddrFrom4([4]byte(packet[16:20])), true
|
||||||
|
}
|
||||||
|
case 6:
|
||||||
|
if len(packet) >= 40 {
|
||||||
|
return netip.AddrFrom16([16]byte(packet[24:40])), true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return netip.Addr{}, false
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *Server) handleConnection(conn net.Conn, dest net.Destination) {
|
||||||
|
defer conn.Close()
|
||||||
|
source := net.DestinationFromAddr(conn.RemoteAddr())
|
||||||
|
addr, _ := netip.AddrFromSlice(source.Address.IP())
|
||||||
|
t := s.lookup(addr.Unmap())
|
||||||
|
if t == nil || !t.track(conn) {
|
||||||
|
errors.LogInfo(s.ctx, "MASQUE: no tunnel for ", source, " to ", dest)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
defer t.untrack(conn)
|
||||||
|
|
||||||
|
ctx, cancel := context.WithCancel(s.ctx)
|
||||||
|
defer cancel()
|
||||||
|
ctx = c.ContextWithID(ctx, session.NewID())
|
||||||
|
inbound := session.Inbound{
|
||||||
|
Name: "masque",
|
||||||
|
Tag: s.tag,
|
||||||
|
CanSpliceCopy: 3,
|
||||||
|
Source: source,
|
||||||
|
User: t.user,
|
||||||
|
}
|
||||||
|
ctx = session.ContextWithInbound(ctx, &inbound)
|
||||||
|
ctx = session.ContextWithContent(ctx, &session.Content{
|
||||||
|
SniffingRequest: s.sniffing,
|
||||||
|
})
|
||||||
|
ctx = session.SubContextFromMuxInbound(ctx)
|
||||||
|
ctx = log.ContextWithAccessMessage(ctx, &log.AccessMessage{
|
||||||
|
From: source,
|
||||||
|
To: dest,
|
||||||
|
Status: log.AccessAccepted,
|
||||||
|
Email: t.user.Email,
|
||||||
|
})
|
||||||
|
errors.LogInfo(ctx, "processing from ", source, " to ", dest)
|
||||||
|
|
||||||
|
link := &transport.Link{
|
||||||
|
Reader: &buf.TimeoutWrapperReader{Reader: buf.NewReader(conn)},
|
||||||
|
Writer: buf.NewWriter(conn),
|
||||||
|
}
|
||||||
|
if err := s.dispatcher.DispatchLink(ctx, dest, link); err != nil {
|
||||||
|
errors.LogError(ctx, errors.New("connection closed").Base(err))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func init() {
|
||||||
|
common.Must(common.RegisterConfig((*ServerConfig)(nil), func(ctx context.Context, config interface{}) (interface{}, error) {
|
||||||
|
return NewServer(ctx, config.(*ServerConfig))
|
||||||
|
}))
|
||||||
|
}
|
||||||
@@ -0,0 +1,328 @@
|
|||||||
|
package masque
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"context"
|
||||||
|
"io"
|
||||||
|
"net/netip"
|
||||||
|
"os"
|
||||||
|
"sync"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
"github.com/xtls/xray-core/common/buf"
|
||||||
|
"github.com/xtls/xray-core/common/net"
|
||||||
|
"github.com/xtls/xray-core/common/protocol"
|
||||||
|
"golang.zx2c4.com/wireguard/tun"
|
||||||
|
)
|
||||||
|
|
||||||
|
type fakeTunnelConn struct {
|
||||||
|
mu sync.Mutex
|
||||||
|
reads chan []byte
|
||||||
|
written [][]byte
|
||||||
|
closed bool
|
||||||
|
stall chan struct{}
|
||||||
|
}
|
||||||
|
|
||||||
|
func newFakeTunnelConn() *fakeTunnelConn {
|
||||||
|
return &fakeTunnelConn{reads: make(chan []byte, 16)}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *fakeTunnelConn) Read(b []byte) (int, error) {
|
||||||
|
p, ok := <-c.reads
|
||||||
|
if !ok {
|
||||||
|
return 0, io.EOF
|
||||||
|
}
|
||||||
|
return copy(b, p), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *fakeTunnelConn) Write(b []byte) (int, error) {
|
||||||
|
if c.stall != nil {
|
||||||
|
<-c.stall
|
||||||
|
}
|
||||||
|
c.mu.Lock()
|
||||||
|
defer c.mu.Unlock()
|
||||||
|
c.written = append(c.written, bytes.Clone(b))
|
||||||
|
return len(b), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *fakeTunnelConn) Close() error {
|
||||||
|
c.mu.Lock()
|
||||||
|
defer c.mu.Unlock()
|
||||||
|
if !c.closed {
|
||||||
|
c.closed = true
|
||||||
|
close(c.reads)
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *fakeTunnelConn) packets() [][]byte {
|
||||||
|
c.mu.Lock()
|
||||||
|
defer c.mu.Unlock()
|
||||||
|
return c.written
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *fakeTunnelConn) isClosed() bool {
|
||||||
|
c.mu.Lock()
|
||||||
|
defer c.mu.Unlock()
|
||||||
|
return c.closed
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *fakeTunnelConn) LocalAddr() net.Addr { return &net.TCPAddr{} }
|
||||||
|
func (c *fakeTunnelConn) RemoteAddr() net.Addr { return &net.TCPAddr{} }
|
||||||
|
func (c *fakeTunnelConn) SetDeadline(t time.Time) error { return nil }
|
||||||
|
func (c *fakeTunnelConn) SetReadDeadline(t time.Time) error { return nil }
|
||||||
|
func (c *fakeTunnelConn) SetWriteDeadline(t time.Time) error { return nil }
|
||||||
|
|
||||||
|
type fakeDevice struct {
|
||||||
|
mu sync.Mutex
|
||||||
|
reads chan []byte
|
||||||
|
written [][]byte
|
||||||
|
closed bool
|
||||||
|
}
|
||||||
|
|
||||||
|
func (d *fakeDevice) File() *os.File { return nil }
|
||||||
|
func (d *fakeDevice) MTU() (int, error) { return 1280, nil }
|
||||||
|
func (d *fakeDevice) Name() (string, error) { return "fake", nil }
|
||||||
|
func (d *fakeDevice) Events() <-chan tun.Event { return nil }
|
||||||
|
func (d *fakeDevice) BatchSize() int { return 1 }
|
||||||
|
|
||||||
|
func (d *fakeDevice) Read(bufs [][]byte, sizes []int, offset int) (int, error) {
|
||||||
|
p, ok := <-d.reads
|
||||||
|
if !ok {
|
||||||
|
return 0, os.ErrClosed
|
||||||
|
}
|
||||||
|
sizes[0] = copy(bufs[0][offset:], p)
|
||||||
|
return 1, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (d *fakeDevice) Write(bufs [][]byte, offset int) (int, error) {
|
||||||
|
d.mu.Lock()
|
||||||
|
defer d.mu.Unlock()
|
||||||
|
for _, b := range bufs {
|
||||||
|
d.written = append(d.written, bytes.Clone(b[offset:]))
|
||||||
|
}
|
||||||
|
return len(bufs), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (d *fakeDevice) Close() error {
|
||||||
|
d.mu.Lock()
|
||||||
|
defer d.mu.Unlock()
|
||||||
|
if !d.closed {
|
||||||
|
d.closed = true
|
||||||
|
close(d.reads)
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (d *fakeDevice) packets() [][]byte {
|
||||||
|
d.mu.Lock()
|
||||||
|
defer d.mu.Unlock()
|
||||||
|
return d.written
|
||||||
|
}
|
||||||
|
|
||||||
|
func ipPacket(src, dst string) []byte {
|
||||||
|
s, d := netip.MustParseAddr(src), netip.MustParseAddr(dst)
|
||||||
|
if s.Is4() {
|
||||||
|
b := make([]byte, 20)
|
||||||
|
b[0] = 0x45
|
||||||
|
b[8] = 64
|
||||||
|
copy(b[12:16], s.AsSlice())
|
||||||
|
copy(b[16:20], d.AsSlice())
|
||||||
|
return b
|
||||||
|
}
|
||||||
|
b := make([]byte, 40)
|
||||||
|
b[0] = 0x60
|
||||||
|
b[7] = 64
|
||||||
|
copy(b[8:24], s.AsSlice())
|
||||||
|
copy(b[24:40], d.AsSlice())
|
||||||
|
return b
|
||||||
|
}
|
||||||
|
|
||||||
|
func newTestServer(t *testing.T) (*Server, *fakeDevice) {
|
||||||
|
t.Helper()
|
||||||
|
pool4, err := newAddressPool(netip.MustParsePrefix("10.14.0.1/24"))
|
||||||
|
require.NoError(t, err)
|
||||||
|
pool6, err := newAddressPool(netip.MustParsePrefix("fd14::1/64"))
|
||||||
|
require.NoError(t, err)
|
||||||
|
dev := &fakeDevice{reads: make(chan []byte, tunnelQueueSize*2)}
|
||||||
|
s := &Server{
|
||||||
|
mtu: 1280,
|
||||||
|
dev: dev,
|
||||||
|
pools: []*addressPool{pool4, pool6},
|
||||||
|
local: []netip.Addr{netip.MustParseAddr("10.14.0.1"), netip.MustParseAddr("fd14::1")},
|
||||||
|
tunnels: make(map[netip.Addr]*serverTunnel),
|
||||||
|
}
|
||||||
|
return s, dev
|
||||||
|
}
|
||||||
|
|
||||||
|
func addTunnel(t *testing.T, s *Server) (*serverTunnel, *fakeTunnelConn) {
|
||||||
|
t.Helper()
|
||||||
|
return addUserTunnel(t, s, &protocol.MemoryUser{})
|
||||||
|
}
|
||||||
|
|
||||||
|
func addUserTunnel(t *testing.T, s *Server, user *protocol.MemoryUser) (*serverTunnel, *fakeTunnelConn) {
|
||||||
|
t.Helper()
|
||||||
|
conn := newFakeTunnelConn()
|
||||||
|
tunnel := newServerTunnel(conn, user)
|
||||||
|
for _, pool := range s.pools {
|
||||||
|
addr, ok := pool.allocate()
|
||||||
|
require.True(t, ok)
|
||||||
|
tunnel.addrs = append(tunnel.addrs, addr)
|
||||||
|
}
|
||||||
|
require.True(t, s.register(tunnel))
|
||||||
|
go s.writeToTunnel(tunnel)
|
||||||
|
t.Cleanup(tunnel.close)
|
||||||
|
return tunnel, conn
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestServerRoutesTunnelPackets(t *testing.T) {
|
||||||
|
s, dev := newTestServer(t)
|
||||||
|
a, aConn := addTunnel(t, s)
|
||||||
|
b, bConn := addTunnel(t, s)
|
||||||
|
require.Equal(t, []netip.Addr{netip.MustParseAddr("10.14.0.2"), netip.MustParseAddr("fd14::2")}, a.addrs)
|
||||||
|
require.Equal(t, []netip.Addr{netip.MustParseAddr("10.14.0.3"), netip.MustParseAddr("fd14::3")}, b.addrs)
|
||||||
|
|
||||||
|
toB := ipPacket("10.14.0.2", "10.14.0.3")
|
||||||
|
toB6 := ipPacket("fd14::2", "fd14::3")
|
||||||
|
toServer := ipPacket("10.14.0.2", "10.14.0.1")
|
||||||
|
toInternet := ipPacket("fd14::2", "2001:db8::1")
|
||||||
|
for _, p := range [][]byte{
|
||||||
|
toB,
|
||||||
|
toB6,
|
||||||
|
ipPacket("10.14.0.2", "10.14.0.9"),
|
||||||
|
ipPacket("fd14::2", "fd14::99"),
|
||||||
|
ipPacket("fd14::2", "fe80::1"),
|
||||||
|
ipPacket("fd14::2", "ff02::1"),
|
||||||
|
ipPacket("10.14.0.2", "224.0.0.251"),
|
||||||
|
ipPacket("10.14.0.2", "10.14.0.2"),
|
||||||
|
toServer,
|
||||||
|
toInternet,
|
||||||
|
} {
|
||||||
|
aConn.reads <- p
|
||||||
|
}
|
||||||
|
aConn.Close()
|
||||||
|
require.NoError(t, s.readFromTunnel(a))
|
||||||
|
|
||||||
|
require.Eventually(t, func() bool { return len(bConn.packets()) == 2 }, time.Second, time.Millisecond)
|
||||||
|
require.Equal(t, [][]byte{toB, toB6}, bConn.packets())
|
||||||
|
require.Equal(t, [][]byte{toServer, toInternet}, dev.packets())
|
||||||
|
require.Empty(t, aConn.packets())
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestServerRoutesStackPackets(t *testing.T) {
|
||||||
|
s, dev := newTestServer(t)
|
||||||
|
_, aConn := addTunnel(t, s)
|
||||||
|
_, bConn := addTunnel(t, s)
|
||||||
|
require.NoError(t, s.Start())
|
||||||
|
|
||||||
|
toA := ipPacket("192.0.2.1", "10.14.0.2")
|
||||||
|
toB := ipPacket("2001:db8::1", "fd14::3")
|
||||||
|
dev.reads <- toA
|
||||||
|
dev.reads <- ipPacket("192.0.2.1", "10.14.0.9")
|
||||||
|
dev.reads <- toB
|
||||||
|
require.Eventually(t, func() bool {
|
||||||
|
return len(aConn.packets()) == 1 && len(bConn.packets()) == 1
|
||||||
|
}, time.Second, time.Millisecond)
|
||||||
|
require.Equal(t, [][]byte{toA}, aConn.packets())
|
||||||
|
require.Equal(t, [][]byte{toB}, bConn.packets())
|
||||||
|
|
||||||
|
require.NoError(t, s.Close())
|
||||||
|
require.True(t, aConn.isClosed())
|
||||||
|
require.True(t, bConn.isClosed())
|
||||||
|
require.False(t, s.register(&serverTunnel{}))
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestServerSlowTunnelDoesNotBlockOthers(t *testing.T) {
|
||||||
|
s, dev := newTestServer(t)
|
||||||
|
_, aConn := addTunnel(t, s)
|
||||||
|
_, bConn := addTunnel(t, s)
|
||||||
|
aConn.stall = make(chan struct{})
|
||||||
|
defer close(aConn.stall)
|
||||||
|
require.NoError(t, s.Start())
|
||||||
|
defer s.Close()
|
||||||
|
|
||||||
|
for range tunnelQueueSize + 10 {
|
||||||
|
dev.reads <- ipPacket("192.0.2.1", "10.14.0.2")
|
||||||
|
}
|
||||||
|
toB := ipPacket("192.0.2.1", "10.14.0.3")
|
||||||
|
dev.reads <- toB
|
||||||
|
require.Eventually(t, func() bool { return len(bConn.packets()) == 1 }, time.Second, time.Millisecond)
|
||||||
|
require.Equal(t, [][]byte{toB}, bConn.packets())
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestServerClosesTunnelConnections(t *testing.T) {
|
||||||
|
s, _ := newTestServer(t)
|
||||||
|
a, _ := addTunnel(t, s)
|
||||||
|
conn := newFakeTunnelConn()
|
||||||
|
require.True(t, a.track(conn))
|
||||||
|
other := newFakeTunnelConn()
|
||||||
|
require.True(t, a.track(other))
|
||||||
|
a.untrack(other)
|
||||||
|
|
||||||
|
s.release(a)
|
||||||
|
require.True(t, conn.isClosed())
|
||||||
|
require.False(t, other.isClosed())
|
||||||
|
require.False(t, a.track(newFakeTunnelConn()))
|
||||||
|
require.False(t, a.send(buf.New()))
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestServerReleasesAddresses(t *testing.T) {
|
||||||
|
s, _ := newTestServer(t)
|
||||||
|
a, _ := addTunnel(t, s)
|
||||||
|
s.release(a)
|
||||||
|
require.Nil(t, s.lookup(netip.MustParseAddr("10.14.0.2")))
|
||||||
|
b, _ := addTunnel(t, s)
|
||||||
|
require.Equal(t, []netip.Addr{netip.MustParseAddr("10.14.0.3"), netip.MustParseAddr("fd14::3")}, b.addrs)
|
||||||
|
for range 250 {
|
||||||
|
addTunnel(t, s)
|
||||||
|
}
|
||||||
|
c, _ := addTunnel(t, s)
|
||||||
|
require.Equal(t, netip.MustParseAddr("10.14.0.254"), c.addrs[0])
|
||||||
|
addr, ok := s.pools[0].allocate()
|
||||||
|
require.True(t, ok)
|
||||||
|
require.Equal(t, netip.MustParseAddr("10.14.0.2"), addr)
|
||||||
|
_, ok = s.pools[0].allocate()
|
||||||
|
require.False(t, ok)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestServerRemoveUserClosesTunnels(t *testing.T) {
|
||||||
|
s, _ := newTestServer(t)
|
||||||
|
s.validator = newValidator()
|
||||||
|
alice := &protocol.MemoryUser{Email: "a@example.com", Account: &MemoryAccount{Password: "p"}}
|
||||||
|
bob := &protocol.MemoryUser{Email: "b@example.com", Account: &MemoryAccount{Password: "p"}}
|
||||||
|
require.NoError(t, s.AddUser(context.Background(), alice))
|
||||||
|
require.NoError(t, s.AddUser(context.Background(), bob))
|
||||||
|
_, aConn := addUserTunnel(t, s, alice)
|
||||||
|
_, bConn := addUserTunnel(t, s, bob)
|
||||||
|
|
||||||
|
require.NoError(t, s.RemoveUser(context.Background(), "a@example.com"))
|
||||||
|
require.True(t, aConn.isClosed())
|
||||||
|
require.False(t, bConn.isClosed())
|
||||||
|
require.Error(t, s.RemoveUser(context.Background(), "a@example.com"))
|
||||||
|
require.Nil(t, s.validator.get("a@example.com", "p"))
|
||||||
|
require.Equal(t, bob, s.validator.get("b@example.com", "p"))
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestPacketDestination(t *testing.T) {
|
||||||
|
v4 := make([]byte, 20)
|
||||||
|
v4[0] = 0x45
|
||||||
|
copy(v4[16:20], []byte{192, 0, 2, 1})
|
||||||
|
addr, ok := packetDestination(v4)
|
||||||
|
require.True(t, ok)
|
||||||
|
require.Equal(t, netip.MustParseAddr("192.0.2.1"), addr)
|
||||||
|
|
||||||
|
v6 := make([]byte, 40)
|
||||||
|
v6[0] = 0x60
|
||||||
|
dst := netip.MustParseAddr("2001:db8::1").As16()
|
||||||
|
copy(v6[24:40], dst[:])
|
||||||
|
addr, ok = packetDestination(v6)
|
||||||
|
require.True(t, ok)
|
||||||
|
require.Equal(t, netip.MustParseAddr("2001:db8::1"), addr)
|
||||||
|
|
||||||
|
for _, b := range [][]byte{nil, v4[:19], v6[:39], {0x50}} {
|
||||||
|
_, ok = packetDestination(b)
|
||||||
|
require.False(t, ok)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -277,6 +277,7 @@ func (w *VisionReader) ReadMultiBuffer() (buf.MultiBuffer, error) {
|
|||||||
w.ob.CanSpliceCopy = 1
|
w.ob.CanSpliceCopy = 1
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
SuppressOuterCloseNotify(w.conn)
|
||||||
readerConn, readCounter, _ := UnwrapRawConn(w.conn)
|
readerConn, readCounter, _ := UnwrapRawConn(w.conn)
|
||||||
w.directReadCounter = readCounter
|
w.directReadCounter = readCounter
|
||||||
w.Reader = buf.NewReader(readerConn)
|
w.Reader = buf.NewReader(readerConn)
|
||||||
@@ -340,6 +341,7 @@ func (w *VisionWriter) WriteMultiBuffer(mb buf.MultiBuffer) error {
|
|||||||
// w.ob.CanSpliceCopy = 1
|
// w.ob.CanSpliceCopy = 1
|
||||||
// }
|
// }
|
||||||
}
|
}
|
||||||
|
SuppressOuterCloseNotify(w.conn)
|
||||||
rawConn, _, writerCounter := UnwrapRawConn(w.conn)
|
rawConn, _, writerCounter := UnwrapRawConn(w.conn)
|
||||||
w.Writer = buf.NewWriter(rawConn)
|
w.Writer = buf.NewWriter(rawConn)
|
||||||
w.directWriteCounter = writerCounter
|
w.directWriteCounter = writerCounter
|
||||||
@@ -669,6 +671,19 @@ func XtlsFilterTls(buffer buf.MultiBuffer, trafficState *TrafficState, ctx conte
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
type CloseNotifySuppressor interface {
|
||||||
|
SuppressCloseNotify()
|
||||||
|
}
|
||||||
|
|
||||||
|
// Close our local TLS conn instance might send a incorrect close_notify alert
|
||||||
|
// if the XTLS direct copy mode is enabled and cause TLS BAD_RECORD_MAC on users' browser
|
||||||
|
// Close the underlying connection directly to avoid this issue.
|
||||||
|
func SuppressOuterCloseNotify(conn net.Conn) {
|
||||||
|
if suppressor, ok := stat.TryUnwrapStatsConn(conn).(CloseNotifySuppressor); ok {
|
||||||
|
suppressor.SuppressCloseNotify()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
// UnwrapRawConn support unwrap encryption, stats, mask wrappers, tls, utls, reality, proxyproto, uds-wrapper conn and get raw tcp/uds conn from it
|
// UnwrapRawConn support unwrap encryption, stats, mask wrappers, tls, utls, reality, proxyproto, uds-wrapper conn and get raw tcp/uds conn from it
|
||||||
func UnwrapRawConn(conn net.Conn) (net.Conn, stats.Counter, stats.Counter) {
|
func UnwrapRawConn(conn net.Conn) (net.Conn, stats.Counter, stats.Counter) {
|
||||||
var readCounter, writerCounter stats.Counter
|
var readCounter, writerCounter stats.Counter
|
||||||
|
|||||||
@@ -71,7 +71,7 @@ func (c *Client) Process(ctx context.Context, link *transport.Link, dialer inter
|
|||||||
return nil
|
return nil
|
||||||
})
|
})
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return errors.New("failed to find an available destination").AtWarning().Base(err)
|
return errors.New("failed to find an available destination").Base(err)
|
||||||
}
|
}
|
||||||
errors.LogInfo(ctx, "tunneling request to ", destination, " via ", network, ":", server.Destination.NetAddr())
|
errors.LogInfo(ctx, "tunneling request to ", destination, " via ", network, ":", server.Destination.NetAddr())
|
||||||
|
|
||||||
@@ -124,7 +124,7 @@ func (c *Client) Process(ctx context.Context, link *transport.Link, dialer inter
|
|||||||
}
|
}
|
||||||
|
|
||||||
if err = buf.CopyOnceTimeout(link.Reader, bodyWriter, time.Millisecond*100); err != nil && err != buf.ErrNotTimeoutReader && err != buf.ErrReadTimeout {
|
if err = buf.CopyOnceTimeout(link.Reader, bodyWriter, time.Millisecond*100); err != nil && err != buf.ErrNotTimeoutReader && err != buf.ErrReadTimeout {
|
||||||
return errors.New("failed to write A request payload").Base(err).AtWarning()
|
return errors.New("failed to write A request payload").Base(err)
|
||||||
}
|
}
|
||||||
|
|
||||||
if err := bufferedWriter.SetBuffered(false); err != nil {
|
if err := bufferedWriter.SetBuffered(false); err != nil {
|
||||||
|
|||||||
@@ -98,7 +98,7 @@ func ReadTCPSession(validator *Validator, reader io.Reader) (*protocol.RequestHe
|
|||||||
iv := append([]byte(nil), buffer.BytesTo(ivLen)...)
|
iv := append([]byte(nil), buffer.BytesTo(ivLen)...)
|
||||||
r, err = account.Cipher.NewDecryptionReader(account.Key, iv, reader)
|
r, err = account.Cipher.NewDecryptionReader(account.Key, iv, reader)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, nil, drain.WithError(drainer, reader, errors.New("failed to initialize decoding stream").Base(err).AtError())
|
return nil, nil, drain.WithError(drainer, reader, errors.New("failed to initialize decoding stream").Base(err))
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -146,7 +146,7 @@ func WriteTCPRequest(request *protocol.RequestHeader, writer io.Writer) (buf.Wri
|
|||||||
|
|
||||||
w, err := account.Cipher.NewEncryptionWriter(account.Key, iv, writer)
|
w, err := account.Cipher.NewEncryptionWriter(account.Key, iv, writer)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, errors.New("failed to create encoding stream").Base(err).AtError()
|
return nil, errors.New("failed to create encoding stream").Base(err)
|
||||||
}
|
}
|
||||||
|
|
||||||
header := buf.New()
|
header := buf.New()
|
||||||
|
|||||||
@@ -34,11 +34,11 @@ func NewServer(ctx context.Context, config *ServerConfig) (*Server, error) {
|
|||||||
for _, user := range config.Users {
|
for _, user := range config.Users {
|
||||||
u, err := user.ToMemoryUser()
|
u, err := user.ToMemoryUser()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, errors.New("failed to get shadowsocks user").Base(err).AtError()
|
return nil, errors.New("failed to get shadowsocks user").Base(err)
|
||||||
}
|
}
|
||||||
|
|
||||||
if err := validator.Add(u); err != nil {
|
if err := validator.Add(u); err != nil {
|
||||||
return nil, errors.New("failed to add user").Base(err).AtError()
|
return nil, errors.New("failed to add user").Base(err)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -200,7 +200,7 @@ func (s *Server) handleUDPPayload(ctx context.Context, conn stat.Connection, dis
|
|||||||
func (s *Server) handleConnection(ctx context.Context, conn stat.Connection, dispatcher routing.Dispatcher) error {
|
func (s *Server) handleConnection(ctx context.Context, conn stat.Connection, dispatcher routing.Dispatcher) error {
|
||||||
sessionPolicy := s.policyManager.ForLevel(0)
|
sessionPolicy := s.policyManager.ForLevel(0)
|
||||||
if err := conn.SetReadDeadline(time.Now().Add(sessionPolicy.Timeouts.Handshake)); err != nil {
|
if err := conn.SetReadDeadline(time.Now().Add(sessionPolicy.Timeouts.Handshake)); err != nil {
|
||||||
return errors.New("unable to set read deadline").Base(err).AtWarning()
|
return errors.New("unable to set read deadline").Base(err)
|
||||||
}
|
}
|
||||||
|
|
||||||
bufferedReader := buf.BufferedReader{Reader: buf.NewReader(conn)}
|
bufferedReader := buf.BufferedReader{Reader: buf.NewReader(conn)}
|
||||||
|
|||||||
@@ -0,0 +1,55 @@
|
|||||||
|
package shadowsocks_2022
|
||||||
|
|
||||||
|
import (
|
||||||
|
"crypto/aes"
|
||||||
|
"crypto/cipher"
|
||||||
|
"errors"
|
||||||
|
"strings"
|
||||||
|
|
||||||
|
"golang.org/x/crypto/chacha20poly1305"
|
||||||
|
)
|
||||||
|
|
||||||
|
type CipherMethod struct {
|
||||||
|
Name string
|
||||||
|
KeySaltLength int
|
||||||
|
IsChaCha bool
|
||||||
|
}
|
||||||
|
|
||||||
|
var methods = map[string]*CipherMethod{
|
||||||
|
MethodAES128GCM: {Name: MethodAES128GCM, KeySaltLength: 16, IsChaCha: false},
|
||||||
|
MethodAES256GCM: {Name: MethodAES256GCM, KeySaltLength: 32, IsChaCha: false},
|
||||||
|
MethodChaCha20Poly1305: {Name: MethodChaCha20Poly1305, KeySaltLength: 32, IsChaCha: true},
|
||||||
|
}
|
||||||
|
|
||||||
|
func GetCipherMethod(name string) (*CipherMethod, error) {
|
||||||
|
name = strings.ToLower(name)
|
||||||
|
if m, ok := methods[name]; ok {
|
||||||
|
return m, nil
|
||||||
|
}
|
||||||
|
return nil, errors.New("unknown shadowsocks 2022 method")
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewAEAD creates standard stream AEAD cipher instance (AES-GCM or ChaCha20-Poly1305)
|
||||||
|
func (m *CipherMethod) NewAEAD(key []byte) (cipher.AEAD, error) {
|
||||||
|
if m.IsChaCha {
|
||||||
|
return chacha20poly1305.New(key)
|
||||||
|
}
|
||||||
|
block, err := aes.NewCipher(key)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
return cipher.NewGCM(block)
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewBlock creates standard 16-byte block cipher for AES header encryption/decryption
|
||||||
|
func (m *CipherMethod) NewBlock(key []byte) (cipher.Block, error) {
|
||||||
|
return aes.NewCipher(key)
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewUDPCipher creates AEAD cipher for UDP packets (XChaCha20-Poly1305 with 24-byte nonce)
|
||||||
|
func (m *CipherMethod) NewUDPCipher(key []byte) (cipher.AEAD, error) {
|
||||||
|
if m.IsChaCha {
|
||||||
|
return chacha20poly1305.NewX(key)
|
||||||
|
}
|
||||||
|
return nil, errors.New("shadowsocks-2022: udp separate AEAD cipher only available for chacha20 method")
|
||||||
|
}
|
||||||
@@ -1,6 +1,9 @@
|
|||||||
package shadowsocks_2022
|
package shadowsocks_2022
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"bytes"
|
||||||
|
"encoding/base64"
|
||||||
|
|
||||||
"google.golang.org/protobuf/proto"
|
"google.golang.org/protobuf/proto"
|
||||||
|
|
||||||
"github.com/xtls/xray-core/common/protocol"
|
"github.com/xtls/xray-core/common/protocol"
|
||||||
@@ -8,26 +11,31 @@ import (
|
|||||||
|
|
||||||
// MemoryAccount is an account type converted from Account.
|
// MemoryAccount is an account type converted from Account.
|
||||||
type MemoryAccount struct {
|
type MemoryAccount struct {
|
||||||
Key string
|
Key []byte
|
||||||
}
|
}
|
||||||
|
|
||||||
// AsAccount implements protocol.AsAccount.
|
// AsAccount implements protocol.AsAccount.
|
||||||
func (u *Account) AsAccount() (protocol.Account, error) {
|
func (u *Account) AsAccount() (protocol.Account, error) {
|
||||||
|
keyStr := u.GetKey()
|
||||||
|
raw, err := base64.StdEncoding.DecodeString(keyStr)
|
||||||
|
if err != nil {
|
||||||
|
raw = []byte(keyStr)
|
||||||
|
}
|
||||||
return &MemoryAccount{
|
return &MemoryAccount{
|
||||||
Key: u.GetKey(),
|
Key: raw,
|
||||||
}, nil
|
}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// Equals implements protocol.Account.Equals().
|
// Equals implements protocol.Account.Equals().
|
||||||
func (a *MemoryAccount) Equals(another protocol.Account) bool {
|
func (a *MemoryAccount) Equals(another protocol.Account) bool {
|
||||||
if account, ok := another.(*MemoryAccount); ok {
|
if account, ok := another.(*MemoryAccount); ok {
|
||||||
return a.Key == account.Key
|
return bytes.Equal(a.Key, account.Key)
|
||||||
}
|
}
|
||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
|
|
||||||
func (a *MemoryAccount) ToProto() proto.Message {
|
func (a *MemoryAccount) ToProto() proto.Message {
|
||||||
return &Account{
|
return &Account{
|
||||||
Key: a.Key,
|
Key: base64.StdEncoding.EncodeToString(a.Key),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
+208
-126
@@ -2,17 +2,11 @@ package shadowsocks_2022
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
|
"io"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
shadowsocks "github.com/sagernet/sing-shadowsocks"
|
|
||||||
"github.com/sagernet/sing-shadowsocks/shadowaead_2022"
|
|
||||||
C "github.com/sagernet/sing/common"
|
|
||||||
B "github.com/sagernet/sing/common/buf"
|
|
||||||
"github.com/sagernet/sing/common/bufio"
|
|
||||||
E "github.com/sagernet/sing/common/exceptions"
|
|
||||||
M "github.com/sagernet/sing/common/metadata"
|
|
||||||
N "github.com/sagernet/sing/common/network"
|
|
||||||
"github.com/xtls/xray-core/common"
|
"github.com/xtls/xray-core/common"
|
||||||
|
"github.com/xtls/xray-core/common/antireplay"
|
||||||
"github.com/xtls/xray-core/common/buf"
|
"github.com/xtls/xray-core/common/buf"
|
||||||
"github.com/xtls/xray-core/common/errors"
|
"github.com/xtls/xray-core/common/errors"
|
||||||
"github.com/xtls/xray-core/common/log"
|
"github.com/xtls/xray-core/common/log"
|
||||||
@@ -20,7 +14,10 @@ import (
|
|||||||
"github.com/xtls/xray-core/common/protocol"
|
"github.com/xtls/xray-core/common/protocol"
|
||||||
"github.com/xtls/xray-core/common/session"
|
"github.com/xtls/xray-core/common/session"
|
||||||
"github.com/xtls/xray-core/common/signal"
|
"github.com/xtls/xray-core/common/signal"
|
||||||
"github.com/xtls/xray-core/common/singbridge"
|
"github.com/xtls/xray-core/common/task"
|
||||||
|
"github.com/xtls/xray-core/common/utils"
|
||||||
|
"github.com/xtls/xray-core/core"
|
||||||
|
"github.com/xtls/xray-core/features/policy"
|
||||||
"github.com/xtls/xray-core/features/routing"
|
"github.com/xtls/xray-core/features/routing"
|
||||||
"github.com/xtls/xray-core/transport/internet/stat"
|
"github.com/xtls/xray-core/transport/internet/stat"
|
||||||
)
|
)
|
||||||
@@ -32,10 +29,13 @@ func init() {
|
|||||||
}
|
}
|
||||||
|
|
||||||
type Inbound struct {
|
type Inbound struct {
|
||||||
networks []net.Network
|
networks []net.Network
|
||||||
service shadowsocks.Service
|
method *CipherMethod
|
||||||
email string
|
psk []byte
|
||||||
level int
|
user *protocol.MemoryUser
|
||||||
|
saltFilter *antireplay.ReplayFilter[[32]byte]
|
||||||
|
udpCodec *UDPServerCodec
|
||||||
|
policyManager policy.Manager
|
||||||
}
|
}
|
||||||
|
|
||||||
func NewServer(ctx context.Context, config *ServerConfig) (*Inbound, error) {
|
func NewServer(ctx context.Context, config *ServerConfig) (*Inbound, error) {
|
||||||
@@ -46,20 +46,35 @@ func NewServer(ctx context.Context, config *ServerConfig) (*Inbound, error) {
|
|||||||
net.Network_UDP,
|
net.Network_UDP,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
inbound := &Inbound{
|
|
||||||
networks: networks,
|
method, err := GetCipherMethod(config.Method)
|
||||||
email: config.Email,
|
|
||||||
level: int(config.Level),
|
|
||||||
}
|
|
||||||
if !C.Contains(shadowaead_2022.List, config.Method) {
|
|
||||||
return nil, errors.New("unsupported method ", config.Method)
|
|
||||||
}
|
|
||||||
service, err := shadowaead_2022.NewServiceWithPassword(config.Method, config.Key, 500, inbound, nil)
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, errors.New("create service").Base(err)
|
return nil, errors.New("unsupported method: ", config.Method).Base(err)
|
||||||
}
|
}
|
||||||
inbound.service = service
|
|
||||||
return inbound, nil
|
psk, err := ParseKey(config.Key, method.KeySaltLength)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
udpCodec, err := NewUDPServerCodec(method, psk, 500*time.Second)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
v := core.MustFromContext(ctx)
|
||||||
|
return &Inbound{
|
||||||
|
networks: networks,
|
||||||
|
method: method,
|
||||||
|
psk: psk,
|
||||||
|
saltFilter: antireplay.NewMapFilter[[32]byte](60),
|
||||||
|
user: &protocol.MemoryUser{
|
||||||
|
Email: config.Email,
|
||||||
|
Level: uint32(config.Level),
|
||||||
|
},
|
||||||
|
udpCodec: udpCodec,
|
||||||
|
policyManager: v.GetFeature(policy.ManagerType()).(policy.Manager),
|
||||||
|
}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (i *Inbound) Network() []net.Network {
|
func (i *Inbound) Network() []net.Network {
|
||||||
@@ -70,114 +85,181 @@ func (i *Inbound) Process(ctx context.Context, network net.Network, connection s
|
|||||||
inbound := session.InboundFromContext(ctx)
|
inbound := session.InboundFromContext(ctx)
|
||||||
inbound.Name = "shadowsocks-2022"
|
inbound.Name = "shadowsocks-2022"
|
||||||
inbound.CanSpliceCopy = 3
|
inbound.CanSpliceCopy = 3
|
||||||
|
inbound.User = i.user
|
||||||
var metadata M.Metadata
|
|
||||||
if inbound.Source.IsValid() {
|
|
||||||
metadata.Source = M.ParseSocksaddr(inbound.Source.NetAddr())
|
|
||||||
}
|
|
||||||
|
|
||||||
ctx = session.ContextWithDispatcher(ctx, dispatcher)
|
|
||||||
|
|
||||||
if network == net.Network_TCP {
|
if network == net.Network_TCP {
|
||||||
return singbridge.ReturnError(i.service.NewConnection(ctx, connection, metadata))
|
return i.processTCP(ctx, connection, dispatcher)
|
||||||
} else {
|
}
|
||||||
reader := buf.NewReader(connection)
|
return i.processUDP(ctx, connection, dispatcher)
|
||||||
pc := &natPacketConn{connection}
|
}
|
||||||
for {
|
|
||||||
mb, err := reader.ReadMultiBuffer()
|
func (i *Inbound) processTCP(ctx context.Context, conn net.Conn, dispatcher routing.Dispatcher) error {
|
||||||
|
defer conn.Close()
|
||||||
|
|
||||||
|
sessionPolicy := i.policyManager.ForLevel(0)
|
||||||
|
if err := conn.SetReadDeadline(time.Now().Add(sessionPolicy.Timeouts.Handshake)); err != nil {
|
||||||
|
return errors.New("unable to set read deadline").Base(err)
|
||||||
|
}
|
||||||
|
|
||||||
|
var salt [32]byte
|
||||||
|
saltSlice := salt[:i.method.KeySaltLength]
|
||||||
|
if _, err := io.ReadFull(conn, saltSlice); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
if !i.saltFilter.Check(salt) {
|
||||||
|
return ErrSaltNotUnique
|
||||||
|
}
|
||||||
|
|
||||||
|
sessionKey := DeriveSessionSubKey(i.psk, saltSlice, i.method.KeySaltLength)
|
||||||
|
aead, err := i.method.NewAEAD(sessionKey)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
reader := NewStreamReader(conn, aead)
|
||||||
|
|
||||||
|
reqHeader, err := ReadClientRequestHeader(conn, reader)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
conn.SetReadDeadline(time.Time{})
|
||||||
|
dest := reqHeader.Destination
|
||||||
|
|
||||||
|
writer, err := WriteTCPResponse(conn, i.method, i.psk, saltSlice, nil)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
ctx = log.ContextWithAccessMessage(ctx, &log.AccessMessage{
|
||||||
|
From: conn.RemoteAddr(),
|
||||||
|
To: dest,
|
||||||
|
Status: log.AccessAccepted,
|
||||||
|
Email: i.user.Email,
|
||||||
|
})
|
||||||
|
|
||||||
|
errors.LogInfo(ctx, "tunneling request to ", dest)
|
||||||
|
|
||||||
|
link, err := dispatcher.Dispatch(ctx, dest)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(reqHeader.EarlyData) > 0 {
|
||||||
|
earlyBuf := buf.New()
|
||||||
|
earlyBuf.Write(reqHeader.EarlyData)
|
||||||
|
if err := link.Writer.WriteMultiBuffer(buf.MultiBuffer{earlyBuf}); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
sessionPolicy = i.policyManager.ForLevel(uint32(i.user.Level))
|
||||||
|
ctx, cancel := context.WithCancel(ctx)
|
||||||
|
timer := signal.CancelAfterInactivity(ctx, cancel, sessionPolicy.Timeouts.ConnectionIdle)
|
||||||
|
ctx = policy.ContextWithBufferPolicy(ctx, sessionPolicy.Buffer)
|
||||||
|
|
||||||
|
requestDone := func() error {
|
||||||
|
defer timer.SetTimeout(sessionPolicy.Timeouts.DownlinkOnly)
|
||||||
|
return buf.Copy(reader, link.Writer, buf.UpdateActivity(timer))
|
||||||
|
}
|
||||||
|
|
||||||
|
responseDone := func() error {
|
||||||
|
defer timer.SetTimeout(sessionPolicy.Timeouts.UplinkOnly)
|
||||||
|
return buf.Copy(link.Reader, writer, buf.UpdateActivity(timer))
|
||||||
|
}
|
||||||
|
|
||||||
|
responseDoneAndCloseWriter := task.OnSuccess(responseDone, task.Close(link.Writer))
|
||||||
|
return task.Run(ctx, requestDone, responseDoneAndCloseWriter)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (i *Inbound) processUDP(ctx context.Context, conn stat.Connection, dispatcher routing.Dispatcher) error {
|
||||||
|
udpConns := utils.NewTypedSyncMap[uint64, *udpConnEntry]()
|
||||||
|
defer func() {
|
||||||
|
udpConns.Range(func(key uint64, entry *udpConnEntry) bool {
|
||||||
|
entry.timer.SetTimeout(0)
|
||||||
|
return true
|
||||||
|
})
|
||||||
|
}()
|
||||||
|
|
||||||
|
reader := buf.NewReader(conn)
|
||||||
|
for {
|
||||||
|
mb, err := reader.ReadMultiBuffer()
|
||||||
|
if err != nil {
|
||||||
|
buf.ReleaseMulti(mb)
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, b := range mb {
|
||||||
|
decoded, err := i.udpCodec.DecodePacket(b.Bytes())
|
||||||
if err != nil {
|
if err != nil {
|
||||||
buf.ReleaseMulti(mb)
|
b.Release()
|
||||||
return singbridge.ReturnError(err)
|
continue
|
||||||
}
|
}
|
||||||
for _, buffer := range mb {
|
|
||||||
packet := B.As(buffer.Bytes()).ToOwned()
|
entry, ok := udpConns.Load(decoded.SessionID)
|
||||||
buffer.Release()
|
if !ok {
|
||||||
err = i.service.NewPacket(ctx, pc, packet, metadata)
|
sessCtx, cancel := context.WithCancel(ctx)
|
||||||
|
sessCtx = log.ContextWithAccessMessage(sessCtx, &log.AccessMessage{
|
||||||
|
From: conn.RemoteAddr(),
|
||||||
|
To: decoded.Destination,
|
||||||
|
Status: log.AccessAccepted,
|
||||||
|
Email: i.user.Email,
|
||||||
|
})
|
||||||
|
|
||||||
|
link, err := dispatcher.Dispatch(sessCtx, decoded.Destination)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
packet.Release()
|
cancel()
|
||||||
buf.ReleaseMulti(mb)
|
b.Release()
|
||||||
return err
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
newEntry := &udpConnEntry{
|
||||||
|
link: link,
|
||||||
|
cancel: cancel,
|
||||||
|
}
|
||||||
|
sessionPolicy := i.policyManager.ForLevel(uint32(i.user.Level))
|
||||||
|
newEntry.timer = signal.CancelAfterInactivity(sessCtx, func() {
|
||||||
|
udpConns.Delete(decoded.SessionID)
|
||||||
|
common.Interrupt(link.Reader)
|
||||||
|
common.Interrupt(link.Writer)
|
||||||
|
cancel()
|
||||||
|
}, sessionPolicy.Timeouts.ConnectionIdle)
|
||||||
|
|
||||||
|
actual, loaded := udpConns.LoadOrStore(decoded.SessionID, newEntry)
|
||||||
|
if loaded {
|
||||||
|
// Another goroutine/packet beat us to storing, terminate our redundant link
|
||||||
|
newEntry.timer.SetTimeout(0)
|
||||||
|
entry = actual
|
||||||
|
} else {
|
||||||
|
entry = newEntry
|
||||||
|
go func(sessID uint64, dest net.Destination, cEntry *udpConnEntry) {
|
||||||
|
defer func() {
|
||||||
|
cEntry.timer.SetTimeout(0)
|
||||||
|
}()
|
||||||
|
for {
|
||||||
|
resMb, err := cEntry.link.Reader.ReadMultiBuffer()
|
||||||
|
if err != nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
cEntry.timer.Update()
|
||||||
|
for _, rb := range resMb {
|
||||||
|
encPacket, err := i.udpCodec.EncodeServerPacket(sessID, dest, rb.Bytes())
|
||||||
|
rb.Release()
|
||||||
|
if err != nil {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
_, _ = conn.Write(encPacket)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}(decoded.SessionID, decoded.Destination, entry)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
entry.timer.Update()
|
||||||
|
payloadBuf := buf.New()
|
||||||
|
payloadBuf.Write(decoded.Payload)
|
||||||
|
b.Release()
|
||||||
|
_ = entry.link.Writer.WriteMultiBuffer(buf.MultiBuffer{payloadBuf})
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (i *Inbound) NewConnection(ctx context.Context, conn net.Conn, metadata M.Metadata) error {
|
|
||||||
inbound := session.InboundFromContext(ctx)
|
|
||||||
inbound.User = &protocol.MemoryUser{
|
|
||||||
Email: i.email,
|
|
||||||
Level: uint32(i.level),
|
|
||||||
}
|
|
||||||
ctx = log.ContextWithAccessMessage(ctx, &log.AccessMessage{
|
|
||||||
From: metadata.Source,
|
|
||||||
To: metadata.Destination,
|
|
||||||
Status: log.AccessAccepted,
|
|
||||||
Email: i.email,
|
|
||||||
})
|
|
||||||
errors.LogInfo(ctx, "tunnelling request to tcp:", metadata.Destination)
|
|
||||||
dispatcher := session.DispatcherFromContext(ctx)
|
|
||||||
destination, err := singbridge.ToDestination(metadata.Destination, net.Network_TCP)
|
|
||||||
if err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
link, err := dispatcher.Dispatch(ctx, destination)
|
|
||||||
if err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
return singbridge.CopyConn(ctx, nil, link, conn)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (i *Inbound) NewPacketConnection(ctx context.Context, conn N.PacketConn, metadata M.Metadata) error {
|
|
||||||
inbound := session.InboundFromContext(ctx)
|
|
||||||
inbound.User = &protocol.MemoryUser{
|
|
||||||
Email: i.email,
|
|
||||||
Level: uint32(i.level),
|
|
||||||
}
|
|
||||||
ctx = log.ContextWithAccessMessage(ctx, &log.AccessMessage{
|
|
||||||
From: metadata.Source,
|
|
||||||
To: metadata.Destination,
|
|
||||||
Status: log.AccessAccepted,
|
|
||||||
Email: i.email,
|
|
||||||
})
|
|
||||||
errors.LogInfo(ctx, "tunnelling request to udp:", metadata.Destination)
|
|
||||||
dispatcher := session.DispatcherFromContext(ctx)
|
|
||||||
destination, err := singbridge.ToDestination(metadata.Destination, net.Network_UDP)
|
|
||||||
if err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
link, err := dispatcher.Dispatch(ctx, destination)
|
|
||||||
if err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
outConn := &singbridge.PacketConnWrapper{
|
|
||||||
Reader: link.Reader,
|
|
||||||
Writer: link.Writer,
|
|
||||||
Dest: destination,
|
|
||||||
T: signal.CancelAfterInactivity(ctx, func() {
|
|
||||||
common.Interrupt(link.Reader)
|
|
||||||
}, 300*time.Second),
|
|
||||||
}
|
|
||||||
return bufio.CopyPacketConn(ctx, conn, outConn)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (i *Inbound) NewError(ctx context.Context, err error) {
|
|
||||||
if E.IsClosed(err) {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
errors.LogWarning(ctx, err.Error())
|
|
||||||
}
|
|
||||||
|
|
||||||
type natPacketConn struct {
|
|
||||||
net.Conn
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *natPacketConn) ReadPacket(buffer *B.Buffer) (addr M.Socksaddr, err error) {
|
|
||||||
_, err = buffer.ReadFrom(c)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
func (c *natPacketConn) WritePacket(buffer *B.Buffer, addr M.Socksaddr) error {
|
|
||||||
_, err := buffer.WriteTo(c)
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
|
|||||||
@@ -2,21 +2,17 @@ package shadowsocks_2022
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"encoding/base64"
|
"crypto/cipher"
|
||||||
|
"encoding/binary"
|
||||||
|
"io"
|
||||||
"strconv"
|
"strconv"
|
||||||
"strings"
|
"strings"
|
||||||
"sync"
|
"sync"
|
||||||
|
"sync/atomic"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/sagernet/sing-shadowsocks/shadowaead_2022"
|
|
||||||
C "github.com/sagernet/sing/common"
|
|
||||||
A "github.com/sagernet/sing/common/auth"
|
|
||||||
B "github.com/sagernet/sing/common/buf"
|
|
||||||
"github.com/sagernet/sing/common/bufio"
|
|
||||||
E "github.com/sagernet/sing/common/exceptions"
|
|
||||||
M "github.com/sagernet/sing/common/metadata"
|
|
||||||
N "github.com/sagernet/sing/common/network"
|
|
||||||
"github.com/xtls/xray-core/common"
|
"github.com/xtls/xray-core/common"
|
||||||
|
"github.com/xtls/xray-core/common/antireplay"
|
||||||
"github.com/xtls/xray-core/common/buf"
|
"github.com/xtls/xray-core/common/buf"
|
||||||
"github.com/xtls/xray-core/common/errors"
|
"github.com/xtls/xray-core/common/errors"
|
||||||
"github.com/xtls/xray-core/common/log"
|
"github.com/xtls/xray-core/common/log"
|
||||||
@@ -24,8 +20,11 @@ import (
|
|||||||
"github.com/xtls/xray-core/common/protocol"
|
"github.com/xtls/xray-core/common/protocol"
|
||||||
"github.com/xtls/xray-core/common/session"
|
"github.com/xtls/xray-core/common/session"
|
||||||
"github.com/xtls/xray-core/common/signal"
|
"github.com/xtls/xray-core/common/signal"
|
||||||
"github.com/xtls/xray-core/common/singbridge"
|
"github.com/xtls/xray-core/common/task"
|
||||||
|
"github.com/xtls/xray-core/common/utils"
|
||||||
"github.com/xtls/xray-core/common/uuid"
|
"github.com/xtls/xray-core/common/uuid"
|
||||||
|
"github.com/xtls/xray-core/core"
|
||||||
|
"github.com/xtls/xray-core/features/policy"
|
||||||
"github.com/xtls/xray-core/features/routing"
|
"github.com/xtls/xray-core/features/routing"
|
||||||
"github.com/xtls/xray-core/transport/internet/stat"
|
"github.com/xtls/xray-core/transport/internet/stat"
|
||||||
)
|
)
|
||||||
@@ -38,9 +37,16 @@ func init() {
|
|||||||
|
|
||||||
type MultiUserInbound struct {
|
type MultiUserInbound struct {
|
||||||
sync.Mutex
|
sync.Mutex
|
||||||
networks []net.Network
|
networks []net.Network
|
||||||
users []*protocol.MemoryUser
|
method *CipherMethod
|
||||||
service *shadowaead_2022.MultiService[int]
|
masterPSK []byte
|
||||||
|
usersByHash *utils.TypedSyncMap[[AESBlockSize]byte, *protocol.MemoryUser]
|
||||||
|
usersByEmail *utils.TypedSyncMap[string, *protocol.MemoryUser]
|
||||||
|
userCount atomic.Int64
|
||||||
|
saltFilter *antireplay.ReplayFilter[[32]byte]
|
||||||
|
udpSessions *UDPSessionManager
|
||||||
|
udpMasterCipher cipher.Block
|
||||||
|
policyManager policy.Manager
|
||||||
}
|
}
|
||||||
|
|
||||||
func NewMultiServer(ctx context.Context, config *MultiUserServerConfig) (*MultiUserInbound, error) {
|
func NewMultiServer(ctx context.Context, config *MultiUserServerConfig) (*MultiUserInbound, error) {
|
||||||
@@ -51,138 +57,131 @@ func NewMultiServer(ctx context.Context, config *MultiUserServerConfig) (*MultiU
|
|||||||
net.Network_UDP,
|
net.Network_UDP,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
memUsers := []*protocol.MemoryUser{}
|
|
||||||
for i, user := range config.Users {
|
method, err := GetCipherMethod(config.Method)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
if method.IsChaCha {
|
||||||
|
return nil, errors.New("shadowsocks 2022 multi-user: only aes methods are supported")
|
||||||
|
}
|
||||||
|
|
||||||
|
masterPSK, err := ParseKey(config.Key, method.KeySaltLength)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
masterBlock, err := method.NewBlock(masterPSK)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
v := core.MustFromContext(ctx)
|
||||||
|
i := &MultiUserInbound{
|
||||||
|
networks: networks,
|
||||||
|
method: method,
|
||||||
|
masterPSK: masterPSK,
|
||||||
|
usersByHash: utils.NewTypedSyncMap[[AESBlockSize]byte, *protocol.MemoryUser](),
|
||||||
|
usersByEmail: utils.NewTypedSyncMap[string, *protocol.MemoryUser](),
|
||||||
|
saltFilter: antireplay.NewMapFilter[[32]byte](60),
|
||||||
|
udpSessions: NewUDPSessionManager(500 * time.Second),
|
||||||
|
udpMasterCipher: masterBlock,
|
||||||
|
policyManager: v.GetFeature(policy.ManagerType()).(policy.Manager),
|
||||||
|
}
|
||||||
|
|
||||||
|
for idx, user := range config.Users {
|
||||||
if user.Email == "" {
|
if user.Email == "" {
|
||||||
u := uuid.New()
|
u := uuid.New()
|
||||||
user.Email = "unnamed-user-" + strconv.Itoa(i) + "-" + u.String()
|
user.Email = "unnamed-user-" + strconv.Itoa(idx) + "-" + u.String()
|
||||||
}
|
}
|
||||||
u, err := user.ToMemoryUser()
|
memUser, err := user.ToMemoryUser()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, errors.New("failed to get shadowsocks user").Base(err).AtError()
|
return nil, errors.New("failed to parse shadowsocks user").Base(err)
|
||||||
|
}
|
||||||
|
if err := i.AddUser(ctx, memUser); err != nil {
|
||||||
|
return nil, err
|
||||||
}
|
}
|
||||||
memUsers = append(memUsers, u)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
inbound := &MultiUserInbound{
|
return i, nil
|
||||||
networks: networks,
|
|
||||||
users: memUsers,
|
|
||||||
}
|
|
||||||
if config.Key == "" {
|
|
||||||
return nil, errors.New("missing key")
|
|
||||||
}
|
|
||||||
psk, err := base64.StdEncoding.DecodeString(config.Key)
|
|
||||||
if err != nil {
|
|
||||||
return nil, errors.New("parse config").Base(err)
|
|
||||||
}
|
|
||||||
service, err := shadowaead_2022.NewMultiService[int](config.Method, psk, 500, inbound, nil)
|
|
||||||
if err != nil {
|
|
||||||
return nil, errors.New("create service").Base(err)
|
|
||||||
}
|
|
||||||
err = service.UpdateUsersWithPasswords(
|
|
||||||
C.MapIndexed(memUsers, func(index int, it *protocol.MemoryUser) int { return index }),
|
|
||||||
C.Map(memUsers, func(it *protocol.MemoryUser) string { return it.Account.(*MemoryAccount).Key }),
|
|
||||||
)
|
|
||||||
if err != nil {
|
|
||||||
return nil, errors.New("create service").Base(err)
|
|
||||||
}
|
|
||||||
|
|
||||||
inbound.service = service
|
|
||||||
return inbound, nil
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// AddUser implements proxy.UserManager.AddUser().
|
// AddUser implements proxy.UserManager.AddUser()
|
||||||
func (i *MultiUserInbound) AddUser(ctx context.Context, u *protocol.MemoryUser) error {
|
func (i *MultiUserInbound) AddUser(ctx context.Context, u *protocol.MemoryUser) error {
|
||||||
i.Lock()
|
i.Lock()
|
||||||
defer i.Unlock()
|
defer i.Unlock()
|
||||||
|
|
||||||
|
var emailKey string
|
||||||
if u.Email != "" {
|
if u.Email != "" {
|
||||||
for idx := range i.users {
|
emailKey = strings.ToLower(u.Email)
|
||||||
if i.users[idx].Email == u.Email {
|
if _, exists := i.usersByEmail.Load(emailKey); exists {
|
||||||
return errors.New("User ", u.Email, " already exists.")
|
return errors.New("user ", u.Email, " already exists")
|
||||||
}
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
i.users = append(i.users, u)
|
|
||||||
|
|
||||||
// sync to multi service
|
memAcc, ok := u.Account.(*MemoryAccount)
|
||||||
// Considering implements shadowsocks2022 in xray-core may have better performance.
|
if !ok {
|
||||||
i.service.UpdateUsersWithPasswords(
|
return errors.New("missing or invalid user account")
|
||||||
C.MapIndexed(i.users, func(index int, it *protocol.MemoryUser) int { return index }),
|
}
|
||||||
C.Map(i.users, func(it *protocol.MemoryUser) string { return it.Account.(*MemoryAccount).Key }),
|
|
||||||
)
|
if len(memAcc.Key) != i.method.KeySaltLength {
|
||||||
|
return ErrBadKey
|
||||||
|
}
|
||||||
|
|
||||||
|
pskHash := DeriveUserPSKHash(memAcc.Key)
|
||||||
|
i.usersByHash.Store(pskHash, u)
|
||||||
|
if emailKey != "" {
|
||||||
|
i.usersByEmail.Store(emailKey, u)
|
||||||
|
}
|
||||||
|
i.userCount.Add(1)
|
||||||
|
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// RemoveUser implements proxy.UserManager.RemoveUser().
|
// RemoveUser implements proxy.UserManager.RemoveUser()
|
||||||
func (i *MultiUserInbound) RemoveUser(ctx context.Context, email string) error {
|
func (i *MultiUserInbound) RemoveUser(ctx context.Context, email string) error {
|
||||||
if email == "" {
|
if email == "" {
|
||||||
return errors.New("Email must not be empty.")
|
return errors.New("email must not be empty")
|
||||||
}
|
}
|
||||||
|
|
||||||
i.Lock()
|
i.Lock()
|
||||||
defer i.Unlock()
|
defer i.Unlock()
|
||||||
|
|
||||||
idx := -1
|
emailKey := strings.ToLower(email)
|
||||||
for ii, u := range i.users {
|
u, loaded := i.usersByEmail.LoadAndDelete(emailKey)
|
||||||
if strings.EqualFold(u.Email, email) {
|
if !loaded {
|
||||||
idx = ii
|
return errors.New("user ", email, " not found")
|
||||||
break
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
if idx == -1 {
|
pskHash := DeriveUserPSKHash(u.Account.(*MemoryAccount).Key)
|
||||||
return errors.New("User ", email, " not found.")
|
i.usersByHash.Delete(pskHash)
|
||||||
}
|
i.userCount.Add(-1)
|
||||||
|
|
||||||
ulen := len(i.users)
|
|
||||||
|
|
||||||
i.users[idx] = i.users[ulen-1]
|
|
||||||
i.users[ulen-1] = nil
|
|
||||||
i.users = i.users[:ulen-1]
|
|
||||||
|
|
||||||
// sync to multi service
|
|
||||||
// Considering implements shadowsocks2022 in xray-core may have better performance.
|
|
||||||
i.service.UpdateUsersWithPasswords(
|
|
||||||
C.MapIndexed(i.users, func(index int, it *protocol.MemoryUser) int { return index }),
|
|
||||||
C.Map(i.users, func(it *protocol.MemoryUser) string { return it.Account.(*MemoryAccount).Key }),
|
|
||||||
)
|
|
||||||
|
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// GetUser implements proxy.UserManager.GetUser().
|
// GetUser implements proxy.UserManager.GetUser()
|
||||||
func (i *MultiUserInbound) GetUser(ctx context.Context, email string) *protocol.MemoryUser {
|
func (i *MultiUserInbound) GetUser(ctx context.Context, email string) *protocol.MemoryUser {
|
||||||
if email == "" {
|
if email == "" {
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
u, _ := i.usersByEmail.Load(strings.ToLower(email))
|
||||||
i.Lock()
|
return u
|
||||||
defer i.Unlock()
|
|
||||||
|
|
||||||
for _, u := range i.users {
|
|
||||||
if strings.EqualFold(u.Email, email) {
|
|
||||||
return u
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return nil
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// GetUsers implements proxy.UserManager.GetUsers().
|
// GetUsers implements proxy.UserManager.GetUsers()
|
||||||
func (i *MultiUserInbound) GetUsers(ctx context.Context) []*protocol.MemoryUser {
|
func (i *MultiUserInbound) GetUsers(ctx context.Context) []*protocol.MemoryUser {
|
||||||
i.Lock()
|
var users []*protocol.MemoryUser
|
||||||
defer i.Unlock()
|
i.usersByEmail.Range(func(_ string, user *protocol.MemoryUser) bool {
|
||||||
dst := make([]*protocol.MemoryUser, len(i.users))
|
users = append(users, user)
|
||||||
copy(dst, i.users)
|
return true
|
||||||
return dst
|
})
|
||||||
|
return users
|
||||||
}
|
}
|
||||||
|
|
||||||
// GetUsersCount implements proxy.UserManager.GetUsersCount().
|
// GetUsersCount implements proxy.UserManager.GetUsersCount()
|
||||||
func (i *MultiUserInbound) GetUsersCount(context.Context) int64 {
|
func (i *MultiUserInbound) GetUsersCount(context.Context) int64 {
|
||||||
i.Lock()
|
return i.userCount.Load()
|
||||||
defer i.Unlock()
|
|
||||||
return int64(len(i.users))
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (i *MultiUserInbound) Network() []net.Network {
|
func (i *MultiUserInbound) Network() []net.Network {
|
||||||
@@ -194,97 +193,317 @@ func (i *MultiUserInbound) Process(ctx context.Context, network net.Network, con
|
|||||||
inbound.Name = "shadowsocks-2022-multi"
|
inbound.Name = "shadowsocks-2022-multi"
|
||||||
inbound.CanSpliceCopy = 3
|
inbound.CanSpliceCopy = 3
|
||||||
|
|
||||||
var metadata M.Metadata
|
if network == net.Network_TCP {
|
||||||
if inbound.Source.IsValid() {
|
return i.processTCP(ctx, connection, dispatcher)
|
||||||
metadata.Source = M.ParseSocksaddr(inbound.Source.NetAddr())
|
}
|
||||||
|
return i.processUDP(ctx, connection, dispatcher)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (i *MultiUserInbound) processTCP(ctx context.Context, conn net.Conn, dispatcher routing.Dispatcher) error {
|
||||||
|
defer conn.Close()
|
||||||
|
|
||||||
|
sessionPolicy := i.policyManager.ForLevel(0)
|
||||||
|
if err := conn.SetReadDeadline(time.Now().Add(sessionPolicy.Timeouts.Handshake)); err != nil {
|
||||||
|
return errors.New("unable to set read deadline").Base(err)
|
||||||
}
|
}
|
||||||
|
|
||||||
ctx = session.ContextWithDispatcher(ctx, dispatcher)
|
// 1. Read Request Salt (16 or 32 bytes)
|
||||||
|
var salt [32]byte
|
||||||
|
saltSlice := salt[:i.method.KeySaltLength]
|
||||||
|
if _, err := io.ReadFull(conn, saltSlice); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
if network == net.Network_TCP {
|
if !i.saltFilter.Check(salt) {
|
||||||
return singbridge.ReturnError(i.service.NewConnection(ctx, connection, metadata))
|
return ErrSaltNotUnique
|
||||||
} else {
|
}
|
||||||
reader := buf.NewReader(connection)
|
|
||||||
pc := &natPacketConn{connection}
|
// 2. Read Extended Identity Header (16 bytes)
|
||||||
for {
|
var eih [AESBlockSize]byte
|
||||||
mb, err := reader.ReadMultiBuffer()
|
if _, err := io.ReadFull(conn, eih[:]); err != nil {
|
||||||
if err != nil {
|
return err
|
||||||
buf.ReleaseMulti(mb)
|
}
|
||||||
return singbridge.ReturnError(err)
|
|
||||||
|
// Decrypt EIH with IdentitySubKey derived from masterPSK and salt
|
||||||
|
identitySubkey := DeriveIdentitySubKey(i.masterPSK, saltSlice, i.method.KeySaltLength)
|
||||||
|
block, err := i.method.NewBlock(identitySubkey)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
var decryptedHash [AESBlockSize]byte
|
||||||
|
block.Decrypt(decryptedHash[:], eih[:])
|
||||||
|
|
||||||
|
// Lookup user
|
||||||
|
user, ok := i.usersByHash.Load(decryptedHash)
|
||||||
|
if !ok || user == nil {
|
||||||
|
return ErrInvalidRequest
|
||||||
|
}
|
||||||
|
userPSK := user.Account.(*MemoryAccount).Key
|
||||||
|
|
||||||
|
// 3. Derive Session Subkey using matched user's PSK
|
||||||
|
sessionKey := DeriveSessionSubKey(userPSK, saltSlice, i.method.KeySaltLength)
|
||||||
|
aead, err := i.method.NewAEAD(sessionKey)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
reader := NewStreamReader(conn, aead)
|
||||||
|
|
||||||
|
// 4 & 5. Read Client Request Header
|
||||||
|
reqHeader, err := ReadClientRequestHeader(conn, reader)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
conn.SetReadDeadline(time.Time{})
|
||||||
|
dest := reqHeader.Destination
|
||||||
|
|
||||||
|
// 6. Send Server Response Handshake
|
||||||
|
writer, err := WriteTCPResponse(conn, i.method, userPSK, saltSlice, nil)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
// 7. Dispatch Connection to Xray routing with matched User
|
||||||
|
inbound := session.InboundFromContext(ctx)
|
||||||
|
inbound.User = user
|
||||||
|
|
||||||
|
ctx = log.ContextWithAccessMessage(ctx, &log.AccessMessage{
|
||||||
|
From: conn.RemoteAddr(),
|
||||||
|
To: dest,
|
||||||
|
Status: log.AccessAccepted,
|
||||||
|
Email: user.Email,
|
||||||
|
})
|
||||||
|
|
||||||
|
errors.LogInfo(ctx, "tunneling request to ", dest, " for user ", user.Email)
|
||||||
|
|
||||||
|
link, err := dispatcher.Dispatch(ctx, dest)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(reqHeader.EarlyData) > 0 {
|
||||||
|
earlyBuf := buf.New()
|
||||||
|
earlyBuf.Write(reqHeader.EarlyData)
|
||||||
|
if err := link.Writer.WriteMultiBuffer(buf.MultiBuffer{earlyBuf}); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
sessionPolicy = i.policyManager.ForLevel(user.Level)
|
||||||
|
ctx, cancel := context.WithCancel(ctx)
|
||||||
|
timer := signal.CancelAfterInactivity(ctx, cancel, sessionPolicy.Timeouts.ConnectionIdle)
|
||||||
|
ctx = policy.ContextWithBufferPolicy(ctx, sessionPolicy.Buffer)
|
||||||
|
|
||||||
|
requestDone := func() error {
|
||||||
|
defer timer.SetTimeout(sessionPolicy.Timeouts.DownlinkOnly)
|
||||||
|
return buf.Copy(reader, link.Writer, buf.UpdateActivity(timer))
|
||||||
|
}
|
||||||
|
|
||||||
|
responseDone := func() error {
|
||||||
|
defer timer.SetTimeout(sessionPolicy.Timeouts.UplinkOnly)
|
||||||
|
return buf.Copy(link.Reader, writer, buf.UpdateActivity(timer))
|
||||||
|
}
|
||||||
|
|
||||||
|
responseDoneAndCloseWriter := task.OnSuccess(responseDone, task.Close(link.Writer))
|
||||||
|
return task.Run(ctx, requestDone, responseDoneAndCloseWriter)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (i *MultiUserInbound) processUDP(ctx context.Context, conn stat.Connection, dispatcher routing.Dispatcher) error {
|
||||||
|
udpConns := utils.NewTypedSyncMap[uint64, *udpConnEntry]()
|
||||||
|
defer func() {
|
||||||
|
udpConns.Range(func(key uint64, entry *udpConnEntry) bool {
|
||||||
|
entry.timer.SetTimeout(0)
|
||||||
|
return true
|
||||||
|
})
|
||||||
|
}()
|
||||||
|
|
||||||
|
reader := buf.NewReader(conn)
|
||||||
|
for {
|
||||||
|
mb, err := reader.ReadMultiBuffer()
|
||||||
|
if err != nil {
|
||||||
|
buf.ReleaseMulti(mb)
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, b := range mb {
|
||||||
|
// In multi-user UDP:
|
||||||
|
// Packet header is 16 bytes: Encrypted(SessionID + PacketID)
|
||||||
|
// Followed by 16 bytes EIH
|
||||||
|
packetBytes := b.Bytes()
|
||||||
|
if len(packetBytes) < 32+1+8+2 {
|
||||||
|
b.Release()
|
||||||
|
continue
|
||||||
}
|
}
|
||||||
for _, buffer := range mb {
|
|
||||||
packet := B.As(buffer.Bytes()).ToOwned()
|
var rawHeader [16]byte
|
||||||
buffer.Release()
|
i.udpMasterCipher.Decrypt(rawHeader[:], packetBytes[:16])
|
||||||
err = i.service.NewPacket(ctx, pc, packet, metadata)
|
|
||||||
|
sessionID := binary.BigEndian.Uint64(rawHeader[:8])
|
||||||
|
packetID := binary.BigEndian.Uint64(rawHeader[8:16])
|
||||||
|
|
||||||
|
// Replay protection & session lookup
|
||||||
|
sessionItem := i.udpSessions.GetOrCreate(sessionID)
|
||||||
|
|
||||||
|
sessionItem.Lock()
|
||||||
|
if !sessionItem.Window.Check(packetID) {
|
||||||
|
sessionItem.Unlock()
|
||||||
|
b.Release()
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
var userPSK []byte
|
||||||
|
var currentUser *protocol.MemoryUser
|
||||||
|
if sessionItem.User != nil {
|
||||||
|
currentUser = sessionItem.User
|
||||||
|
userPSK = sessionItem.UserPSK
|
||||||
|
sessionItem.Unlock()
|
||||||
|
} else {
|
||||||
|
sessionItem.Unlock()
|
||||||
|
// Decrypt EIH
|
||||||
|
identitySubkey := DeriveIdentitySubKey(i.masterPSK, rawHeader[:8], i.method.KeySaltLength)
|
||||||
|
idBlock, err := i.method.NewBlock(identitySubkey)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
packet.Release()
|
b.Release()
|
||||||
buf.ReleaseMulti(mb)
|
continue
|
||||||
return err
|
}
|
||||||
|
|
||||||
|
var decryptedHash [16]byte
|
||||||
|
idBlock.Decrypt(decryptedHash[:], packetBytes[16:32])
|
||||||
|
|
||||||
|
user, ok := i.usersByHash.Load(decryptedHash)
|
||||||
|
if !ok || user == nil {
|
||||||
|
b.Release()
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
currentUser = user
|
||||||
|
userPSK = user.Account.(*MemoryAccount).Key
|
||||||
|
|
||||||
|
sessionItem.Lock()
|
||||||
|
sessionItem.User = user
|
||||||
|
sessionItem.UserPSK = userPSK
|
||||||
|
sessionItem.Unlock()
|
||||||
|
}
|
||||||
|
|
||||||
|
// Decrypt Body (with AEAD caching per session)
|
||||||
|
bodyAead := sessionItem.GetRemoteCipher()
|
||||||
|
if bodyAead == nil {
|
||||||
|
bodyKey := DeriveSessionSubKey(userPSK, rawHeader[:8], i.method.KeySaltLength)
|
||||||
|
var err error
|
||||||
|
bodyAead, err = i.method.NewAEAD(bodyKey)
|
||||||
|
if err != nil {
|
||||||
|
b.Release()
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
sessionItem.SetRemoteCipher(bodyAead)
|
||||||
|
}
|
||||||
|
|
||||||
|
bodyNonce := rawHeader[4:16]
|
||||||
|
bodyCipher := packetBytes[32:]
|
||||||
|
bodyPlain, err := bodyAead.Open(nil, bodyNonce, bodyCipher, nil)
|
||||||
|
b.Release()
|
||||||
|
if err != nil || len(bodyPlain) < 1+8+2 {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
sessionItem.Lock()
|
||||||
|
sessionItem.Window.Add(packetID)
|
||||||
|
sessionItem.Unlock()
|
||||||
|
|
||||||
|
if bodyPlain[0] != HeaderTypeClient {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
epoch := binary.BigEndian.Uint64(bodyPlain[1:9])
|
||||||
|
diff := time.Now().Unix() - int64(epoch)
|
||||||
|
if diff < -30 || diff > 30 {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
paddingLen := int(binary.BigEndian.Uint16(bodyPlain[9:11]))
|
||||||
|
offset := 11 + paddingLen
|
||||||
|
if len(bodyPlain) < offset {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
dest, addrLen, err := parseAddressPort(bodyPlain[offset:])
|
||||||
|
if err != nil {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
payload := bodyPlain[offset+addrLen:]
|
||||||
|
|
||||||
|
entry, ok := udpConns.Load(sessionID)
|
||||||
|
if !ok {
|
||||||
|
sessCtx, cancel := context.WithCancel(ctx)
|
||||||
|
inbound := session.InboundFromContext(sessCtx)
|
||||||
|
inbound.User = currentUser
|
||||||
|
|
||||||
|
sessCtx = log.ContextWithAccessMessage(sessCtx, &log.AccessMessage{
|
||||||
|
From: conn.RemoteAddr(),
|
||||||
|
To: dest,
|
||||||
|
Status: log.AccessAccepted,
|
||||||
|
Email: currentUser.Email,
|
||||||
|
})
|
||||||
|
|
||||||
|
link, err := dispatcher.Dispatch(sessCtx, dest)
|
||||||
|
if err != nil {
|
||||||
|
cancel()
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
newEntry := &udpConnEntry{
|
||||||
|
link: link,
|
||||||
|
cancel: cancel,
|
||||||
|
}
|
||||||
|
sessionPolicy := i.policyManager.ForLevel(currentUser.Level)
|
||||||
|
newEntry.timer = signal.CancelAfterInactivity(sessCtx, func() {
|
||||||
|
udpConns.Delete(sessionID)
|
||||||
|
common.Interrupt(link.Reader)
|
||||||
|
common.Interrupt(link.Writer)
|
||||||
|
cancel()
|
||||||
|
}, sessionPolicy.Timeouts.ConnectionIdle)
|
||||||
|
|
||||||
|
actual, loaded := udpConns.LoadOrStore(sessionID, newEntry)
|
||||||
|
if loaded {
|
||||||
|
newEntry.timer.SetTimeout(0)
|
||||||
|
entry = actual
|
||||||
|
} else {
|
||||||
|
entry = newEntry
|
||||||
|
go func(sessID uint64, uPSK []byte, d net.Destination, cEntry *udpConnEntry) {
|
||||||
|
defer func() {
|
||||||
|
cEntry.timer.SetTimeout(0)
|
||||||
|
}()
|
||||||
|
for {
|
||||||
|
resMb, err := cEntry.link.Reader.ReadMultiBuffer()
|
||||||
|
if err != nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
cEntry.timer.Update()
|
||||||
|
for _, rb := range resMb {
|
||||||
|
encPacket, err := i.encodeServerUDPPacket(sessID, uPSK, d, rb.Bytes())
|
||||||
|
rb.Release()
|
||||||
|
if err != nil {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
_, _ = conn.Write(encPacket)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}(sessionID, userPSK, dest, entry)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
entry.timer.Update()
|
||||||
|
pBuf := buf.New()
|
||||||
|
pBuf.Write(payload)
|
||||||
|
_ = entry.link.Writer.WriteMultiBuffer(buf.MultiBuffer{pBuf})
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (i *MultiUserInbound) NewConnection(ctx context.Context, conn net.Conn, metadata M.Metadata) error {
|
func (i *MultiUserInbound) encodeServerUDPPacket(clientSessionID uint64, userPSK []byte, dest net.Destination, payload []byte) ([]byte, error) {
|
||||||
inbound := session.InboundFromContext(ctx)
|
sessionItem := i.udpSessions.GetOrCreate(clientSessionID)
|
||||||
userInt, _ := A.UserFromContext[int](ctx)
|
if err := sessionItem.EnsureServerState(i.method, i.udpMasterCipher, nil, userPSK); err != nil {
|
||||||
user := i.users[userInt]
|
return nil, err
|
||||||
inbound.User = user
|
|
||||||
ctx = log.ContextWithAccessMessage(ctx, &log.AccessMessage{
|
|
||||||
From: metadata.Source,
|
|
||||||
To: metadata.Destination,
|
|
||||||
Status: log.AccessAccepted,
|
|
||||||
Email: user.Email,
|
|
||||||
})
|
|
||||||
errors.LogInfo(ctx, "tunnelling request to tcp:", metadata.Destination)
|
|
||||||
dispatcher := session.DispatcherFromContext(ctx)
|
|
||||||
destination, err := singbridge.ToDestination(metadata.Destination, net.Network_TCP)
|
|
||||||
if err != nil {
|
|
||||||
return err
|
|
||||||
}
|
}
|
||||||
link, err := dispatcher.Dispatch(ctx, destination)
|
return sessionItem.EncodeServerPacket(i.method, clientSessionID, dest, payload)
|
||||||
if err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
return singbridge.CopyConn(ctx, conn, link, conn)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (i *MultiUserInbound) NewPacketConnection(ctx context.Context, conn N.PacketConn, metadata M.Metadata) error {
|
|
||||||
inbound := session.InboundFromContext(ctx)
|
|
||||||
userInt, _ := A.UserFromContext[int](ctx)
|
|
||||||
user := i.users[userInt]
|
|
||||||
inbound.User = user
|
|
||||||
ctx = log.ContextWithAccessMessage(ctx, &log.AccessMessage{
|
|
||||||
From: metadata.Source,
|
|
||||||
To: metadata.Destination,
|
|
||||||
Status: log.AccessAccepted,
|
|
||||||
Email: user.Email,
|
|
||||||
})
|
|
||||||
errors.LogInfo(ctx, "tunnelling request to udp:", metadata.Destination)
|
|
||||||
dispatcher := session.DispatcherFromContext(ctx)
|
|
||||||
destination, err := singbridge.ToDestination(metadata.Destination, net.Network_UDP)
|
|
||||||
if err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
link, err := dispatcher.Dispatch(ctx, destination)
|
|
||||||
if err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
outConn := &singbridge.PacketConnWrapper{
|
|
||||||
Reader: link.Reader,
|
|
||||||
Writer: link.Writer,
|
|
||||||
Dest: destination,
|
|
||||||
T: signal.CancelAfterInactivity(ctx, func() {
|
|
||||||
common.Interrupt(link.Reader)
|
|
||||||
}, 300*time.Second),
|
|
||||||
}
|
|
||||||
return bufio.CopyPacketConn(ctx, conn, outConn)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (i *MultiUserInbound) NewError(ctx context.Context, err error) {
|
|
||||||
if E.IsClosed(err) {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
errors.LogWarning(ctx, err.Error())
|
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -2,18 +2,12 @@ package shadowsocks_2022
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
|
"crypto/cipher"
|
||||||
|
"encoding/binary"
|
||||||
|
"io"
|
||||||
"strconv"
|
"strconv"
|
||||||
"strings"
|
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/sagernet/sing-shadowsocks/shadowaead_2022"
|
|
||||||
C "github.com/sagernet/sing/common"
|
|
||||||
A "github.com/sagernet/sing/common/auth"
|
|
||||||
B "github.com/sagernet/sing/common/buf"
|
|
||||||
"github.com/sagernet/sing/common/bufio"
|
|
||||||
E "github.com/sagernet/sing/common/exceptions"
|
|
||||||
M "github.com/sagernet/sing/common/metadata"
|
|
||||||
N "github.com/sagernet/sing/common/network"
|
|
||||||
"github.com/xtls/xray-core/common"
|
"github.com/xtls/xray-core/common"
|
||||||
"github.com/xtls/xray-core/common/buf"
|
"github.com/xtls/xray-core/common/buf"
|
||||||
"github.com/xtls/xray-core/common/errors"
|
"github.com/xtls/xray-core/common/errors"
|
||||||
@@ -22,8 +16,11 @@ import (
|
|||||||
"github.com/xtls/xray-core/common/protocol"
|
"github.com/xtls/xray-core/common/protocol"
|
||||||
"github.com/xtls/xray-core/common/session"
|
"github.com/xtls/xray-core/common/session"
|
||||||
"github.com/xtls/xray-core/common/signal"
|
"github.com/xtls/xray-core/common/signal"
|
||||||
"github.com/xtls/xray-core/common/singbridge"
|
"github.com/xtls/xray-core/common/task"
|
||||||
|
"github.com/xtls/xray-core/common/utils"
|
||||||
"github.com/xtls/xray-core/common/uuid"
|
"github.com/xtls/xray-core/common/uuid"
|
||||||
|
"github.com/xtls/xray-core/core"
|
||||||
|
"github.com/xtls/xray-core/features/policy"
|
||||||
"github.com/xtls/xray-core/features/routing"
|
"github.com/xtls/xray-core/features/routing"
|
||||||
"github.com/xtls/xray-core/transport/internet/stat"
|
"github.com/xtls/xray-core/transport/internet/stat"
|
||||||
)
|
)
|
||||||
@@ -34,10 +31,22 @@ func init() {
|
|||||||
}))
|
}))
|
||||||
}
|
}
|
||||||
|
|
||||||
|
type relayDest struct {
|
||||||
|
destination net.Destination
|
||||||
|
email string
|
||||||
|
level uint32
|
||||||
|
key []byte
|
||||||
|
blockCipher cipher.Block
|
||||||
|
}
|
||||||
|
|
||||||
type RelayInbound struct {
|
type RelayInbound struct {
|
||||||
networks []net.Network
|
networks []net.Network
|
||||||
destinations []*RelayDestination
|
method *CipherMethod
|
||||||
service *shadowaead_2022.RelayService[int]
|
relayPSK []byte
|
||||||
|
relayBlock cipher.Block
|
||||||
|
destinations map[[AESBlockSize]byte]*relayDest
|
||||||
|
rawDestinations []*RelayDestination
|
||||||
|
policyManager policy.Manager
|
||||||
}
|
}
|
||||||
|
|
||||||
func NewRelayServer(ctx context.Context, config *RelayServerConfig) (*RelayInbound, error) {
|
func NewRelayServer(ctx context.Context, config *RelayServerConfig) (*RelayInbound, error) {
|
||||||
@@ -48,39 +57,63 @@ func NewRelayServer(ctx context.Context, config *RelayServerConfig) (*RelayInbou
|
|||||||
net.Network_UDP,
|
net.Network_UDP,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
inbound := &RelayInbound{
|
|
||||||
networks: networks,
|
method, err := GetCipherMethod(config.Method)
|
||||||
destinations: config.Destinations,
|
|
||||||
}
|
|
||||||
if !C.Contains(shadowaead_2022.List, config.Method) || !strings.Contains(config.Method, "aes") {
|
|
||||||
return nil, errors.New("unsupported method ", config.Method)
|
|
||||||
}
|
|
||||||
service, err := shadowaead_2022.NewRelayServiceWithPassword[int](config.Method, config.Key, 500, inbound)
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, errors.New("create service").Base(err)
|
return nil, err
|
||||||
|
}
|
||||||
|
if method.IsChaCha {
|
||||||
|
return nil, errors.New("shadowsocks 2022 relay: only aes methods are supported")
|
||||||
}
|
}
|
||||||
|
|
||||||
for i, destination := range config.Destinations {
|
relayPSK, err := ParseKey(config.Key, method.KeySaltLength)
|
||||||
if destination.Email == "" {
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
relayBlock, err := method.NewBlock(relayPSK)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
v := core.MustFromContext(ctx)
|
||||||
|
i := &RelayInbound{
|
||||||
|
networks: networks,
|
||||||
|
method: method,
|
||||||
|
relayPSK: relayPSK,
|
||||||
|
relayBlock: relayBlock,
|
||||||
|
destinations: make(map[[AESBlockSize]byte]*relayDest),
|
||||||
|
rawDestinations: config.Destinations,
|
||||||
|
policyManager: v.GetFeature(policy.ManagerType()).(policy.Manager),
|
||||||
|
}
|
||||||
|
|
||||||
|
for idx, d := range config.Destinations {
|
||||||
|
if d.Email == "" {
|
||||||
u := uuid.New()
|
u := uuid.New()
|
||||||
destination.Email = "unnamed-destination-" + strconv.Itoa(i) + "-" + u.String()
|
d.Email = "unnamed-destination-" + strconv.Itoa(idx) + "-" + u.String()
|
||||||
|
}
|
||||||
|
destKey, err := ParseKey(d.Key, method.KeySaltLength)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
destBlock, err := method.NewBlock(destKey)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
hash := DeriveUserPSKHash(destKey)
|
||||||
|
|
||||||
|
i.destinations[hash] = &relayDest{
|
||||||
|
destination: net.TCPDestination(d.Address.AsAddress(), net.Port(d.Port)),
|
||||||
|
email: d.Email,
|
||||||
|
level: uint32(d.Level),
|
||||||
|
key: destKey,
|
||||||
|
blockCipher: destBlock,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
err = service.UpdateUsersWithPasswords(
|
|
||||||
C.MapIndexed(config.Destinations, func(index int, it *RelayDestination) int { return index }),
|
return i, nil
|
||||||
C.Map(config.Destinations, func(it *RelayDestination) string { return it.Key }),
|
|
||||||
C.Map(config.Destinations, func(it *RelayDestination) M.Socksaddr {
|
|
||||||
return singbridge.ToSocksaddr(net.Destination{
|
|
||||||
Address: it.Address.AsAddress(),
|
|
||||||
Port: net.Port(it.Port),
|
|
||||||
})
|
|
||||||
}),
|
|
||||||
)
|
|
||||||
if err != nil {
|
|
||||||
return nil, errors.New("create service").Base(err)
|
|
||||||
}
|
|
||||||
inbound.service = service
|
|
||||||
return inbound, nil
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (i *RelayInbound) Network() []net.Network {
|
func (i *RelayInbound) Network() []net.Network {
|
||||||
@@ -92,103 +125,206 @@ func (i *RelayInbound) Process(ctx context.Context, network net.Network, connect
|
|||||||
inbound.Name = "shadowsocks-2022-relay"
|
inbound.Name = "shadowsocks-2022-relay"
|
||||||
inbound.CanSpliceCopy = 3
|
inbound.CanSpliceCopy = 3
|
||||||
|
|
||||||
var metadata M.Metadata
|
if network == net.Network_TCP {
|
||||||
if inbound.Source.IsValid() {
|
return i.processTCP(ctx, connection, dispatcher)
|
||||||
metadata.Source = M.ParseSocksaddr(inbound.Source.NetAddr())
|
}
|
||||||
|
return i.processUDP(ctx, connection, dispatcher)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (i *RelayInbound) processTCP(ctx context.Context, conn net.Conn, dispatcher routing.Dispatcher) error {
|
||||||
|
defer conn.Close()
|
||||||
|
|
||||||
|
sessionPolicy := i.policyManager.ForLevel(0)
|
||||||
|
if err := conn.SetReadDeadline(time.Now().Add(sessionPolicy.Timeouts.Handshake)); err != nil {
|
||||||
|
return errors.New("unable to set read deadline").Base(err)
|
||||||
}
|
}
|
||||||
|
|
||||||
ctx = session.ContextWithDispatcher(ctx, dispatcher)
|
// Read Salt + Outer EIH
|
||||||
|
needed := i.method.KeySaltLength + AESBlockSize
|
||||||
|
var headerBuf [48]byte
|
||||||
|
headerSlice := headerBuf[:needed]
|
||||||
|
if _, err := io.ReadFull(conn, headerSlice); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
if network == net.Network_TCP {
|
salt := headerSlice[:i.method.KeySaltLength]
|
||||||
return singbridge.ReturnError(i.service.NewConnection(ctx, connection, metadata))
|
eih := headerSlice[i.method.KeySaltLength:]
|
||||||
} else {
|
|
||||||
reader := buf.NewReader(connection)
|
identitySubkey := DeriveIdentitySubKey(i.relayPSK, salt, i.method.KeySaltLength)
|
||||||
pc := &natPacketConn{connection}
|
block, err := i.method.NewBlock(identitySubkey)
|
||||||
for {
|
if err != nil {
|
||||||
mb, err := reader.ReadMultiBuffer()
|
return err
|
||||||
if err != nil {
|
}
|
||||||
buf.ReleaseMulti(mb)
|
|
||||||
return singbridge.ReturnError(err)
|
var decryptedHash [AESBlockSize]byte
|
||||||
|
block.Decrypt(decryptedHash[:], eih)
|
||||||
|
|
||||||
|
targetDest, ok := i.destinations[decryptedHash]
|
||||||
|
if !ok {
|
||||||
|
return ErrInvalidRequest
|
||||||
|
}
|
||||||
|
conn.SetReadDeadline(time.Time{})
|
||||||
|
|
||||||
|
inbound := session.InboundFromContext(ctx)
|
||||||
|
inbound.User = &protocol.MemoryUser{
|
||||||
|
Email: targetDest.email,
|
||||||
|
Level: targetDest.level,
|
||||||
|
}
|
||||||
|
|
||||||
|
ctx = log.ContextWithAccessMessage(ctx, &log.AccessMessage{
|
||||||
|
From: conn.RemoteAddr(),
|
||||||
|
To: targetDest.destination,
|
||||||
|
Status: log.AccessAccepted,
|
||||||
|
Email: targetDest.email,
|
||||||
|
})
|
||||||
|
|
||||||
|
errors.LogInfo(ctx, "relaying connection to ", targetDest.destination)
|
||||||
|
|
||||||
|
link, err := dispatcher.Dispatch(ctx, targetDest.destination)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
// Unwrap outer EIH: send client salt to next hop, stripping this hop's EIH
|
||||||
|
saltBuf := buf.New()
|
||||||
|
saltBuf.Write(salt)
|
||||||
|
if err := link.Writer.WriteMultiBuffer(buf.MultiBuffer{saltBuf}); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
sessionPolicy = i.policyManager.ForLevel(targetDest.level)
|
||||||
|
ctx, cancel := context.WithCancel(ctx)
|
||||||
|
timer := signal.CancelAfterInactivity(ctx, cancel, sessionPolicy.Timeouts.ConnectionIdle)
|
||||||
|
ctx = policy.ContextWithBufferPolicy(ctx, sessionPolicy.Buffer)
|
||||||
|
|
||||||
|
requestDone := func() error {
|
||||||
|
defer timer.SetTimeout(sessionPolicy.Timeouts.DownlinkOnly)
|
||||||
|
return buf.Copy(buf.NewReader(conn), link.Writer, buf.UpdateActivity(timer))
|
||||||
|
}
|
||||||
|
|
||||||
|
responseDone := func() error {
|
||||||
|
defer timer.SetTimeout(sessionPolicy.Timeouts.UplinkOnly)
|
||||||
|
return buf.Copy(link.Reader, buf.NewWriter(conn), buf.UpdateActivity(timer))
|
||||||
|
}
|
||||||
|
|
||||||
|
responseDoneAndCloseWriter := task.OnSuccess(responseDone, task.Close(link.Writer))
|
||||||
|
return task.Run(ctx, requestDone, responseDoneAndCloseWriter)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (i *RelayInbound) processUDP(ctx context.Context, conn stat.Connection, dispatcher routing.Dispatcher) error {
|
||||||
|
udpConns := utils.NewTypedSyncMap[uint64, *udpConnEntry]()
|
||||||
|
defer func() {
|
||||||
|
udpConns.Range(func(key uint64, entry *udpConnEntry) bool {
|
||||||
|
entry.timer.SetTimeout(0)
|
||||||
|
return true
|
||||||
|
})
|
||||||
|
}()
|
||||||
|
|
||||||
|
reader := buf.NewReader(conn)
|
||||||
|
for {
|
||||||
|
mb, err := reader.ReadMultiBuffer()
|
||||||
|
if err != nil {
|
||||||
|
buf.ReleaseMulti(mb)
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, b := range mb {
|
||||||
|
data := b.Bytes()
|
||||||
|
if len(data) < 2*AESBlockSize {
|
||||||
|
b.Release()
|
||||||
|
continue
|
||||||
}
|
}
|
||||||
for _, buffer := range mb {
|
|
||||||
packet := B.As(buffer.Bytes()).ToOwned()
|
var packetHeader [AESBlockSize]byte
|
||||||
buffer.Release()
|
i.relayBlock.Decrypt(packetHeader[:], data[:AESBlockSize])
|
||||||
err = i.service.NewPacket(ctx, pc, packet, metadata)
|
|
||||||
|
var eiHeader [AESBlockSize]byte
|
||||||
|
i.relayBlock.Decrypt(eiHeader[:], data[AESBlockSize:2*AESBlockSize])
|
||||||
|
for idx := 0; idx < AESBlockSize; idx++ {
|
||||||
|
eiHeader[idx] ^= packetHeader[idx]
|
||||||
|
}
|
||||||
|
|
||||||
|
targetDest, ok := i.destinations[eiHeader]
|
||||||
|
if !ok {
|
||||||
|
b.Release()
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
// Extract sessionID from raw packetHeader for session-level link caching before re-encrypting
|
||||||
|
sessionID := binary.BigEndian.Uint64(packetHeader[:8])
|
||||||
|
|
||||||
|
// Re-encrypt packetHeader with next hop block cipher
|
||||||
|
targetDest.blockCipher.Encrypt(packetHeader[:], packetHeader[:])
|
||||||
|
|
||||||
|
// Strip outer EIH: replace second block with re-encrypted packetHeader and advance
|
||||||
|
copy(data[AESBlockSize:2*AESBlockSize], packetHeader[:])
|
||||||
|
b.Advance(int32(AESBlockSize))
|
||||||
|
|
||||||
|
dest := targetDest.destination
|
||||||
|
dest.Network = net.Network_UDP
|
||||||
|
|
||||||
|
entry, ok := udpConns.Load(sessionID)
|
||||||
|
if !ok {
|
||||||
|
sessCtx, cancel := context.WithCancel(ctx)
|
||||||
|
inbound := session.InboundFromContext(sessCtx)
|
||||||
|
inbound.User = &protocol.MemoryUser{
|
||||||
|
Email: targetDest.email,
|
||||||
|
Level: targetDest.level,
|
||||||
|
}
|
||||||
|
|
||||||
|
sessCtx = log.ContextWithAccessMessage(sessCtx, &log.AccessMessage{
|
||||||
|
From: conn.RemoteAddr(),
|
||||||
|
To: dest,
|
||||||
|
Status: log.AccessAccepted,
|
||||||
|
Email: targetDest.email,
|
||||||
|
})
|
||||||
|
|
||||||
|
link, err := dispatcher.Dispatch(sessCtx, dest)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
packet.Release()
|
cancel()
|
||||||
buf.ReleaseMulti(mb)
|
b.Release()
|
||||||
return err
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
newEntry := &udpConnEntry{
|
||||||
|
link: link,
|
||||||
|
cancel: cancel,
|
||||||
|
}
|
||||||
|
sessionPolicy := i.policyManager.ForLevel(targetDest.level)
|
||||||
|
newEntry.timer = signal.CancelAfterInactivity(sessCtx, func() {
|
||||||
|
udpConns.Delete(sessionID)
|
||||||
|
common.Interrupt(link.Reader)
|
||||||
|
common.Interrupt(link.Writer)
|
||||||
|
cancel()
|
||||||
|
}, sessionPolicy.Timeouts.ConnectionIdle)
|
||||||
|
|
||||||
|
actual, loaded := udpConns.LoadOrStore(sessionID, newEntry)
|
||||||
|
if loaded {
|
||||||
|
newEntry.timer.SetTimeout(0)
|
||||||
|
entry = actual
|
||||||
|
} else {
|
||||||
|
entry = newEntry
|
||||||
|
go func(cEntry *udpConnEntry) {
|
||||||
|
defer func() {
|
||||||
|
cEntry.timer.SetTimeout(0)
|
||||||
|
}()
|
||||||
|
for {
|
||||||
|
resMb, err := cEntry.link.Reader.ReadMultiBuffer()
|
||||||
|
if err != nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
cEntry.timer.Update()
|
||||||
|
for _, rb := range resMb {
|
||||||
|
_, _ = conn.Write(rb.Bytes())
|
||||||
|
rb.Release()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}(entry)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
entry.timer.Update()
|
||||||
|
_ = entry.link.Writer.WriteMultiBuffer(buf.MultiBuffer{b})
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (i *RelayInbound) NewConnection(ctx context.Context, conn net.Conn, metadata M.Metadata) error {
|
|
||||||
inbound := session.InboundFromContext(ctx)
|
|
||||||
userInt, _ := A.UserFromContext[int](ctx)
|
|
||||||
user := i.destinations[userInt]
|
|
||||||
inbound.User = &protocol.MemoryUser{
|
|
||||||
Email: user.Email,
|
|
||||||
Level: uint32(user.Level),
|
|
||||||
}
|
|
||||||
ctx = log.ContextWithAccessMessage(ctx, &log.AccessMessage{
|
|
||||||
From: metadata.Source,
|
|
||||||
To: metadata.Destination,
|
|
||||||
Status: log.AccessAccepted,
|
|
||||||
Email: user.Email,
|
|
||||||
})
|
|
||||||
errors.LogInfo(ctx, "tunnelling request to tcp:", metadata.Destination)
|
|
||||||
dispatcher := session.DispatcherFromContext(ctx)
|
|
||||||
destination, err := singbridge.ToDestination(metadata.Destination, net.Network_TCP)
|
|
||||||
if err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
link, err := dispatcher.Dispatch(ctx, destination)
|
|
||||||
if err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
return singbridge.CopyConn(ctx, nil, link, conn)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (i *RelayInbound) NewPacketConnection(ctx context.Context, conn N.PacketConn, metadata M.Metadata) error {
|
|
||||||
inbound := session.InboundFromContext(ctx)
|
|
||||||
userInt, _ := A.UserFromContext[int](ctx)
|
|
||||||
user := i.destinations[userInt]
|
|
||||||
inbound.User = &protocol.MemoryUser{
|
|
||||||
Email: user.Email,
|
|
||||||
Level: uint32(user.Level),
|
|
||||||
}
|
|
||||||
ctx = log.ContextWithAccessMessage(ctx, &log.AccessMessage{
|
|
||||||
From: metadata.Source,
|
|
||||||
To: metadata.Destination,
|
|
||||||
Status: log.AccessAccepted,
|
|
||||||
Email: user.Email,
|
|
||||||
})
|
|
||||||
errors.LogInfo(ctx, "tunnelling request to udp:", metadata.Destination)
|
|
||||||
dispatcher := session.DispatcherFromContext(ctx)
|
|
||||||
destination, err := singbridge.ToDestination(metadata.Destination, net.Network_UDP)
|
|
||||||
if err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
link, err := dispatcher.Dispatch(ctx, destination)
|
|
||||||
if err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
outConn := &singbridge.PacketConnWrapper{
|
|
||||||
Reader: link.Reader,
|
|
||||||
Writer: link.Writer,
|
|
||||||
Dest: destination,
|
|
||||||
T: signal.CancelAfterInactivity(ctx, func() {
|
|
||||||
common.Interrupt(link.Reader)
|
|
||||||
}, 300*time.Second),
|
|
||||||
}
|
|
||||||
return bufio.CopyPacketConn(ctx, conn, outConn)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (i *RelayInbound) NewError(ctx context.Context, err error) {
|
|
||||||
if E.IsClosed(err) {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
errors.LogWarning(ctx, err.Error())
|
|
||||||
}
|
|
||||||
|
|||||||
@@ -0,0 +1,63 @@
|
|||||||
|
package shadowsocks_2022
|
||||||
|
|
||||||
|
import (
|
||||||
|
"encoding/base64"
|
||||||
|
"strings"
|
||||||
|
|
||||||
|
"lukechampine.com/blake3"
|
||||||
|
)
|
||||||
|
|
||||||
|
const (
|
||||||
|
ContextSessionSubKey = "shadowsocks 2022 session subkey"
|
||||||
|
ContextIdentitySubKey = "shadowsocks 2022 identity subkey"
|
||||||
|
)
|
||||||
|
|
||||||
|
// ParseKey decodes a base64 or raw PSK key string and validates its length
|
||||||
|
func ParseKey(key string, keyLength int) ([]byte, error) {
|
||||||
|
raw, err := base64.StdEncoding.DecodeString(key)
|
||||||
|
if err != nil {
|
||||||
|
raw = []byte(key)
|
||||||
|
}
|
||||||
|
if len(raw) != keyLength {
|
||||||
|
return nil, ErrBadKey
|
||||||
|
}
|
||||||
|
return raw, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func ParsePSKList(password string, keyLength int) ([][]byte, error) {
|
||||||
|
parts := strings.Split(password, ":")
|
||||||
|
pskList := make([][]byte, len(parts))
|
||||||
|
for i, part := range parts {
|
||||||
|
norm, err := ParseKey(part, keyLength)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
pskList[i] = norm
|
||||||
|
}
|
||||||
|
return pskList, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func deriveSubKey(ctx string, psk, salt []byte, keyLength int) []byte {
|
||||||
|
var keyMaterial [64]byte
|
||||||
|
kmLen := len(psk) + len(salt)
|
||||||
|
copy(keyMaterial[:], psk)
|
||||||
|
copy(keyMaterial[len(psk):], salt)
|
||||||
|
out := make([]byte, keyLength)
|
||||||
|
blake3.DeriveKey(out, ctx, keyMaterial[:kmLen])
|
||||||
|
return out
|
||||||
|
}
|
||||||
|
|
||||||
|
func DeriveSessionSubKey(psk, salt []byte, keyLength int) []byte {
|
||||||
|
return deriveSubKey(ContextSessionSubKey, psk, salt, keyLength)
|
||||||
|
}
|
||||||
|
|
||||||
|
func DeriveIdentitySubKey(psk, salt []byte, keyLength int) []byte {
|
||||||
|
return deriveSubKey(ContextIdentitySubKey, psk, salt, keyLength)
|
||||||
|
}
|
||||||
|
|
||||||
|
func DeriveUserPSKHash(userPSK []byte) [AESBlockSize]byte {
|
||||||
|
h := blake3.Sum512(userPSK)
|
||||||
|
var out [AESBlockSize]byte
|
||||||
|
copy(out[:], h[:AESBlockSize])
|
||||||
|
return out
|
||||||
|
}
|
||||||
@@ -2,21 +2,20 @@ package shadowsocks_2022
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
|
"crypto/rand"
|
||||||
|
"io"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
shadowsocks "github.com/sagernet/sing-shadowsocks"
|
|
||||||
"github.com/sagernet/sing-shadowsocks/shadowaead_2022"
|
|
||||||
C "github.com/sagernet/sing/common"
|
|
||||||
B "github.com/sagernet/sing/common/buf"
|
|
||||||
"github.com/sagernet/sing/common/bufio"
|
|
||||||
N "github.com/sagernet/sing/common/network"
|
|
||||||
"github.com/xtls/xray-core/common"
|
"github.com/xtls/xray-core/common"
|
||||||
"github.com/xtls/xray-core/common/buf"
|
"github.com/xtls/xray-core/common/buf"
|
||||||
"github.com/xtls/xray-core/common/errors"
|
"github.com/xtls/xray-core/common/errors"
|
||||||
"github.com/xtls/xray-core/common/net"
|
"github.com/xtls/xray-core/common/net"
|
||||||
|
"github.com/xtls/xray-core/common/retry"
|
||||||
"github.com/xtls/xray-core/common/session"
|
"github.com/xtls/xray-core/common/session"
|
||||||
"github.com/xtls/xray-core/common/signal"
|
"github.com/xtls/xray-core/common/signal"
|
||||||
"github.com/xtls/xray-core/common/singbridge"
|
"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/transport"
|
"github.com/xtls/xray-core/transport"
|
||||||
"github.com/xtls/xray-core/transport/internet"
|
"github.com/xtls/xray-core/transport/internet"
|
||||||
)
|
)
|
||||||
@@ -28,42 +27,47 @@ func init() {
|
|||||||
}
|
}
|
||||||
|
|
||||||
type Outbound struct {
|
type Outbound struct {
|
||||||
ctx context.Context
|
server net.Destination
|
||||||
server net.Destination
|
method *CipherMethod
|
||||||
method shadowsocks.Method
|
pskList [][]byte
|
||||||
|
finalPSK []byte
|
||||||
|
udpCodec *UDPPacketCodec
|
||||||
|
policyManager policy.Manager
|
||||||
}
|
}
|
||||||
|
|
||||||
func NewClient(ctx context.Context, config *ClientConfig) (*Outbound, error) {
|
func NewClient(ctx context.Context, config *ClientConfig) (*Outbound, error) {
|
||||||
o := &Outbound{
|
method, err := GetCipherMethod(config.Method)
|
||||||
ctx: ctx,
|
if err != nil {
|
||||||
|
return nil, errors.New("unsupported method: ", config.Method).Base(err)
|
||||||
|
}
|
||||||
|
|
||||||
|
pskList, err := ParsePSKList(config.Key, method.KeySaltLength)
|
||||||
|
if err != nil {
|
||||||
|
return nil, errors.New("invalid key: ", config.Key).Base(err)
|
||||||
|
}
|
||||||
|
|
||||||
|
finalPSK := pskList[len(pskList)-1]
|
||||||
|
udpCodec, err := NewUDPPacketCodec(method, finalPSK)
|
||||||
|
if err != nil {
|
||||||
|
return nil, errors.New("failed to create udp packet codec").Base(err)
|
||||||
|
}
|
||||||
|
|
||||||
|
v := core.MustFromContext(ctx)
|
||||||
|
return &Outbound{
|
||||||
server: net.Destination{
|
server: net.Destination{
|
||||||
Address: config.Address.AsAddress(),
|
Address: config.Address.AsAddress(),
|
||||||
Port: net.Port(config.Port),
|
Port: net.Port(config.Port),
|
||||||
Network: net.Network_TCP,
|
Network: net.Network_TCP,
|
||||||
},
|
},
|
||||||
}
|
method: method,
|
||||||
if C.Contains(shadowaead_2022.List, config.Method) {
|
pskList: pskList,
|
||||||
if config.Key == "" {
|
finalPSK: finalPSK,
|
||||||
return nil, errors.New("missing psk")
|
udpCodec: udpCodec,
|
||||||
}
|
policyManager: v.GetFeature(policy.ManagerType()).(policy.Manager),
|
||||||
method, err := shadowaead_2022.NewWithPassword(config.Method, config.Key, nil)
|
}, nil
|
||||||
if err != nil {
|
|
||||||
return nil, errors.New("create method").Base(err)
|
|
||||||
}
|
|
||||||
o.method = method
|
|
||||||
} else {
|
|
||||||
return nil, errors.New("unknown method ", config.Method)
|
|
||||||
}
|
|
||||||
return o, nil
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (o *Outbound) Process(ctx context.Context, link *transport.Link, dialer internet.Dialer) error {
|
func (o *Outbound) Process(ctx context.Context, link *transport.Link, dialer internet.Dialer) error {
|
||||||
var inboundConn net.Conn
|
|
||||||
inbound := session.InboundFromContext(ctx)
|
|
||||||
if inbound != nil {
|
|
||||||
inboundConn = inbound.Conn
|
|
||||||
}
|
|
||||||
|
|
||||||
outbounds := session.OutboundsFromContext(ctx)
|
outbounds := session.OutboundsFromContext(ctx)
|
||||||
ob := outbounds[len(outbounds)-1]
|
ob := outbounds[len(outbounds)-1]
|
||||||
if !ob.Target.IsValid() {
|
if !ob.Target.IsValid() {
|
||||||
@@ -78,70 +82,123 @@ func (o *Outbound) Process(ctx context.Context, link *transport.Link, dialer int
|
|||||||
|
|
||||||
serverDestination := o.server
|
serverDestination := o.server
|
||||||
serverDestination.Network = network
|
serverDestination.Network = network
|
||||||
connection, err := dialer.Dial(ctx, serverDestination)
|
|
||||||
if err != nil {
|
|
||||||
return errors.New("failed to connect to server").Base(err)
|
|
||||||
}
|
|
||||||
defer connection.Close()
|
|
||||||
|
|
||||||
|
var conn net.Conn
|
||||||
|
if err := retry.ExponentialBackoff(5, 100).On(func() error {
|
||||||
|
rawConn, err := dialer.Dial(ctx, serverDestination)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
conn = rawConn
|
||||||
|
return nil
|
||||||
|
}); err != nil {
|
||||||
|
return errors.New("failed to find an available destination").Base(err)
|
||||||
|
}
|
||||||
|
defer conn.Close()
|
||||||
|
|
||||||
|
var newCtx context.Context
|
||||||
|
var newCancel context.CancelFunc
|
||||||
if session.TimeoutOnlyFromContext(ctx) {
|
if session.TimeoutOnlyFromContext(ctx) {
|
||||||
ctx, _ = context.WithCancel(context.Background())
|
newCtx, newCancel = context.WithCancel(context.Background())
|
||||||
|
}
|
||||||
|
|
||||||
|
sessionPolicy := o.policyManager.ForLevel(0)
|
||||||
|
ctx, cancel := context.WithCancel(ctx)
|
||||||
|
timer := signal.CancelAfterInactivity(ctx, func() {
|
||||||
|
cancel()
|
||||||
|
if newCancel != nil {
|
||||||
|
newCancel()
|
||||||
|
}
|
||||||
|
}, sessionPolicy.Timeouts.ConnectionIdle)
|
||||||
|
|
||||||
|
ctx = policy.ContextWithBufferPolicy(ctx, sessionPolicy.Buffer)
|
||||||
|
|
||||||
|
if newCtx != nil {
|
||||||
|
ctx = newCtx
|
||||||
}
|
}
|
||||||
|
|
||||||
if network == net.Network_TCP {
|
if network == net.Network_TCP {
|
||||||
serverConn := o.method.DialEarlyConn(connection, singbridge.ToSocksaddr(destination))
|
var clientSalt [32]byte
|
||||||
var handshake bool
|
clientSaltSlice := clientSalt[:o.method.KeySaltLength]
|
||||||
if timeoutReader, isTimeoutReader := link.Reader.(buf.TimeoutReader); isTimeoutReader {
|
if _, err := io.ReadFull(rand.Reader, clientSaltSlice); err != nil {
|
||||||
mb, err := timeoutReader.ReadMultiBufferTimeout(time.Millisecond * 100)
|
return errors.New("failed to generate client salt").Base(err)
|
||||||
if err != nil && err != buf.ErrNotTimeoutReader && err != buf.ErrReadTimeout {
|
|
||||||
return errors.New("read payload").Base(err)
|
|
||||||
}
|
|
||||||
payload := B.New()
|
|
||||||
for {
|
|
||||||
payload.Reset()
|
|
||||||
nb, n := buf.SplitBytes(mb, payload.FreeBytes())
|
|
||||||
if n > 0 {
|
|
||||||
payload.Truncate(n)
|
|
||||||
_, err = serverConn.Write(payload.Bytes())
|
|
||||||
if err != nil {
|
|
||||||
payload.Release()
|
|
||||||
return errors.New("write payload").Base(err)
|
|
||||||
}
|
|
||||||
handshake = true
|
|
||||||
}
|
|
||||||
if nb.IsEmpty() {
|
|
||||||
break
|
|
||||||
}
|
|
||||||
mb = nb
|
|
||||||
}
|
|
||||||
payload.Release()
|
|
||||||
}
|
|
||||||
if !handshake {
|
|
||||||
_, err = serverConn.Write(nil)
|
|
||||||
if err != nil {
|
|
||||||
return errors.New("client handshake").Base(err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return singbridge.CopyConn(ctx, inboundConn, link, serverConn)
|
|
||||||
} else {
|
|
||||||
var packetConn N.PacketConn
|
|
||||||
if pc, isPacketConn := inboundConn.(N.PacketConn); isPacketConn {
|
|
||||||
packetConn = pc
|
|
||||||
} else if nc, isNetPacket := inboundConn.(net.PacketConn); isNetPacket {
|
|
||||||
packetConn = bufio.NewPacketConn(nc)
|
|
||||||
} else {
|
|
||||||
packetConn = &singbridge.PacketConnWrapper{
|
|
||||||
Reader: link.Reader,
|
|
||||||
Writer: link.Writer,
|
|
||||||
Conn: inboundConn,
|
|
||||||
Dest: destination,
|
|
||||||
T: signal.CancelAfterInactivity(ctx, func() {
|
|
||||||
common.Interrupt(link.Reader)
|
|
||||||
}, 300*time.Second),
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
serverConn := o.method.DialPacketConn(connection)
|
requestDone := func() error {
|
||||||
return singbridge.ReturnError(bufio.CopyPacketConn(ctx, packetConn, serverConn))
|
defer timer.SetTimeout(sessionPolicy.Timeouts.DownlinkOnly)
|
||||||
|
bufferedWriter := buf.NewBufferedWriter(buf.NewWriter(conn))
|
||||||
|
bodyWriter, err := WriteTCPRequest(bufferedWriter, o.method, o.pskList, destination, clientSaltSlice, nil)
|
||||||
|
if err != nil {
|
||||||
|
return errors.New("failed to write request").Base(err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if err = buf.CopyOnceTimeout(link.Reader, bodyWriter, time.Millisecond*100); err != nil && err != buf.ErrNotTimeoutReader && err != buf.ErrReadTimeout {
|
||||||
|
return errors.New("failed to write A request payload").Base(err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := bufferedWriter.SetBuffered(false); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
return buf.Copy(link.Reader, bodyWriter, buf.UpdateActivity(timer))
|
||||||
|
}
|
||||||
|
|
||||||
|
responseDone := func() error {
|
||||||
|
defer timer.SetTimeout(sessionPolicy.Timeouts.UplinkOnly)
|
||||||
|
|
||||||
|
responseReader, err := ReadTCPResponse(conn, o.method, o.finalPSK, clientSaltSlice)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
return buf.Copy(responseReader, link.Writer, buf.UpdateActivity(timer))
|
||||||
|
}
|
||||||
|
|
||||||
|
responseDoneAndCloseWriter := task.OnSuccess(responseDone, task.Close(link.Writer))
|
||||||
|
if err := task.Run(ctx, requestDone, responseDoneAndCloseWriter); err != nil {
|
||||||
|
return errors.New("connection ends").Base(err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
if network == net.Network_UDP {
|
||||||
|
requestDone := func() error {
|
||||||
|
defer timer.SetTimeout(sessionPolicy.Timeouts.DownlinkOnly)
|
||||||
|
|
||||||
|
writer := &UDPWriter{
|
||||||
|
Writer: conn,
|
||||||
|
Destination: destination,
|
||||||
|
Codec: o.udpCodec,
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := buf.Copy(link.Reader, writer, buf.UpdateActivity(timer)); err != nil {
|
||||||
|
return errors.New("failed to transport all UDP request").Base(err)
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
responseDone := func() error {
|
||||||
|
defer timer.SetTimeout(sessionPolicy.Timeouts.UplinkOnly)
|
||||||
|
|
||||||
|
reader := &UDPReader{
|
||||||
|
Reader: conn,
|
||||||
|
Codec: o.udpCodec,
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := buf.Copy(reader, link.Writer, buf.UpdateActivity(timer)); err != nil {
|
||||||
|
return errors.New("failed to transport all UDP response").Base(err)
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
responseDoneAndCloseWriter := task.OnSuccess(responseDone, task.Close(link.Writer))
|
||||||
|
if err := task.Run(ctx, requestDone, responseDoneAndCloseWriter); err != nil {
|
||||||
|
return errors.New("connection ends").Base(err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
return errors.New("unsupported network: ", network)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -0,0 +1,512 @@
|
|||||||
|
package shadowsocks_2022
|
||||||
|
|
||||||
|
import (
|
||||||
|
"crypto/cipher"
|
||||||
|
"crypto/rand"
|
||||||
|
"encoding/binary"
|
||||||
|
"io"
|
||||||
|
"math"
|
||||||
|
mrand "math/rand/v2"
|
||||||
|
"sync/atomic"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/xtls/xray-core/common/buf"
|
||||||
|
"github.com/xtls/xray-core/common/errors"
|
||||||
|
"github.com/xtls/xray-core/common/net"
|
||||||
|
)
|
||||||
|
|
||||||
|
type UDPCodec struct {
|
||||||
|
method *CipherMethod
|
||||||
|
psk []byte
|
||||||
|
blockCipher cipher.Block
|
||||||
|
chachaCipher cipher.AEAD
|
||||||
|
clientBodyCipher cipher.AEAD
|
||||||
|
clientSessionID uint64
|
||||||
|
nextPacketID atomic.Uint64
|
||||||
|
sessions *UDPSessionManager
|
||||||
|
}
|
||||||
|
|
||||||
|
type (
|
||||||
|
UDPPacketCodec = UDPCodec
|
||||||
|
UDPServerCodec = UDPCodec
|
||||||
|
)
|
||||||
|
|
||||||
|
func newUDPCodec(method *CipherMethod, psk []byte) (*UDPCodec, error) {
|
||||||
|
c := &UDPCodec{
|
||||||
|
method: method,
|
||||||
|
psk: psk,
|
||||||
|
}
|
||||||
|
var err error
|
||||||
|
if method.IsChaCha {
|
||||||
|
c.chachaCipher, err = method.NewUDPCipher(psk)
|
||||||
|
} else {
|
||||||
|
c.blockCipher, err = method.NewBlock(psk)
|
||||||
|
}
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
return c, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func NewUDPPacketCodec(method *CipherMethod, psk []byte) (*UDPCodec, error) {
|
||||||
|
c, err := newUDPCodec(method, psk)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
var sessID [8]byte
|
||||||
|
if _, err := io.ReadFull(rand.Reader, sessID[:]); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
c.clientSessionID = binary.BigEndian.Uint64(sessID[:])
|
||||||
|
|
||||||
|
if !method.IsChaCha {
|
||||||
|
clientBodyKey := DeriveSessionSubKey(psk, sessID[:], method.KeySaltLength)
|
||||||
|
c.clientBodyCipher, err = method.NewAEAD(clientBodyKey)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return c, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func NewUDPServerCodec(method *CipherMethod, psk []byte, sessionTimeout time.Duration) (*UDPCodec, error) {
|
||||||
|
c, err := newUDPCodec(method, psk)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
c.sessions = NewUDPSessionManager(sessionTimeout)
|
||||||
|
return c, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *UDPCodec) EncodeClientPacket(dest net.Destination, payload []byte) (*buf.Buffer, error) {
|
||||||
|
packetID := c.nextPacketID.Add(1)
|
||||||
|
sessID := c.clientSessionID
|
||||||
|
|
||||||
|
// Padding determination (e.g. DNS port 53 disguise)
|
||||||
|
var paddingLen int
|
||||||
|
if dest.Port == 53 && len(payload) < MaxPaddingLength {
|
||||||
|
paddingLen = mrand.IntN(MaxPaddingLength-len(payload)) + 1
|
||||||
|
}
|
||||||
|
|
||||||
|
addrPortLen := AddrPortLength(dest)
|
||||||
|
|
||||||
|
if c.method.IsChaCha {
|
||||||
|
// ChaCha20 mode: 24-byte nonce + plaintext header (27B) + padding + dest + payload + AEAD tag (16B)
|
||||||
|
totalLen := PacketNonceSize + 27 + paddingLen + addrPortLen + len(payload) + AEADTagSize
|
||||||
|
if totalLen > buf.Size {
|
||||||
|
return nil, ErrPacketTooLarge
|
||||||
|
}
|
||||||
|
|
||||||
|
outBuf := buf.New()
|
||||||
|
|
||||||
|
var nonce [PacketNonceSize]byte
|
||||||
|
if _, err := io.ReadFull(rand.Reader, nonce[:]); err != nil {
|
||||||
|
outBuf.Release()
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
outBuf.Write(nonce[:])
|
||||||
|
|
||||||
|
var hdr [16 + 1 + 8 + 2]byte
|
||||||
|
binary.BigEndian.PutUint64(hdr[0:8], sessID)
|
||||||
|
binary.BigEndian.PutUint64(hdr[8:16], packetID)
|
||||||
|
hdr[16] = HeaderTypeClient
|
||||||
|
binary.BigEndian.PutUint64(hdr[17:25], uint64(time.Now().Unix()))
|
||||||
|
binary.BigEndian.PutUint16(hdr[25:27], uint16(paddingLen))
|
||||||
|
outBuf.Write(hdr[:])
|
||||||
|
if paddingLen > 0 {
|
||||||
|
outBuf.Write(zeroPadding[:paddingLen])
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := WriteAddressPort(outBuf, dest); err != nil {
|
||||||
|
outBuf.Release()
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
outBuf.Write(payload)
|
||||||
|
|
||||||
|
plainBytes := outBuf.Bytes()[PacketNonceSize:]
|
||||||
|
outBuf.Extend(int32(c.chachaCipher.Overhead()))
|
||||||
|
c.chachaCipher.Seal(plainBytes[:0], nonce[:], plainBytes, nil)
|
||||||
|
return outBuf, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// AES mode:
|
||||||
|
// 16B Encrypted Header + (11B header + padding + dest + payload + 16B AEAD tag)
|
||||||
|
totalLen := 16 + 11 + paddingLen + addrPortLen + len(payload) + AEADTagSize
|
||||||
|
if totalLen > buf.Size {
|
||||||
|
return nil, ErrPacketTooLarge
|
||||||
|
}
|
||||||
|
|
||||||
|
outBuf := buf.New()
|
||||||
|
|
||||||
|
var rawHeader [16]byte
|
||||||
|
binary.BigEndian.PutUint64(rawHeader[:8], sessID)
|
||||||
|
binary.BigEndian.PutUint64(rawHeader[8:16], packetID)
|
||||||
|
|
||||||
|
var encryptedHeader [16]byte
|
||||||
|
c.blockCipher.Encrypt(encryptedHeader[:], rawHeader[:])
|
||||||
|
outBuf.Write(encryptedHeader[:])
|
||||||
|
|
||||||
|
bodyAead := c.clientBodyCipher
|
||||||
|
|
||||||
|
var hdr [1 + 8 + 2]byte
|
||||||
|
hdr[0] = HeaderTypeClient
|
||||||
|
binary.BigEndian.PutUint64(hdr[1:9], uint64(time.Now().Unix()))
|
||||||
|
binary.BigEndian.PutUint16(hdr[9:11], uint16(paddingLen))
|
||||||
|
outBuf.Write(hdr[:])
|
||||||
|
if paddingLen > 0 {
|
||||||
|
outBuf.Write(zeroPadding[:paddingLen])
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := WriteAddressPort(outBuf, dest); err != nil {
|
||||||
|
outBuf.Release()
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
outBuf.Write(payload)
|
||||||
|
|
||||||
|
plainBytes := outBuf.Bytes()[16:]
|
||||||
|
bodyNonce := rawHeader[4:16]
|
||||||
|
outBuf.Extend(int32(bodyAead.Overhead()))
|
||||||
|
bodyAead.Seal(plainBytes[:0], bodyNonce, plainBytes, nil)
|
||||||
|
return outBuf, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
type DecodedUDPPacket struct {
|
||||||
|
SessionID uint64
|
||||||
|
PacketID uint64
|
||||||
|
HeaderType byte
|
||||||
|
Timestamp uint64
|
||||||
|
Destination net.Destination
|
||||||
|
Payload []byte
|
||||||
|
}
|
||||||
|
|
||||||
|
func parseAddressPort(data []byte) (net.Destination, int, error) {
|
||||||
|
if len(data) < 1 {
|
||||||
|
return net.Destination{}, 0, ErrPacketTooShort
|
||||||
|
}
|
||||||
|
switch data[0] {
|
||||||
|
case 1: // IPv4
|
||||||
|
if len(data) < 1+4+2 {
|
||||||
|
return net.Destination{}, 0, ErrPacketTooShort
|
||||||
|
}
|
||||||
|
ip := net.IPAddress(data[1:5])
|
||||||
|
port := binary.BigEndian.Uint16(data[5:7])
|
||||||
|
return net.UDPDestination(ip, net.Port(port)), 7, nil
|
||||||
|
case 4: // IPv6
|
||||||
|
if len(data) < 1+16+2 {
|
||||||
|
return net.Destination{}, 0, ErrPacketTooShort
|
||||||
|
}
|
||||||
|
ip := net.IPAddress(data[1:17])
|
||||||
|
port := binary.BigEndian.Uint16(data[17:19])
|
||||||
|
return net.UDPDestination(ip, net.Port(port)), 19, nil
|
||||||
|
case 3: // Domain
|
||||||
|
if len(data) < 2 {
|
||||||
|
return net.Destination{}, 0, ErrPacketTooShort
|
||||||
|
}
|
||||||
|
domainLen := int(data[1])
|
||||||
|
if len(data) < 2+domainLen+2 {
|
||||||
|
return net.Destination{}, 0, ErrPacketTooShort
|
||||||
|
}
|
||||||
|
domain := string(data[2 : 2+domainLen])
|
||||||
|
port := binary.BigEndian.Uint16(data[2+domainLen : 2+domainLen+2])
|
||||||
|
return net.UDPDestination(net.DomainAddress(domain), net.Port(port)), 2 + domainLen + 2, nil
|
||||||
|
default:
|
||||||
|
return net.Destination{}, 0, errors.New("unknown address type")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func parsePlainUDPPacket(sessionID, packetID uint64, bodyPlain []byte) (DecodedUDPPacket, error) {
|
||||||
|
if len(bodyPlain) < 1+8+2 {
|
||||||
|
return DecodedUDPPacket{}, ErrPacketTooShort
|
||||||
|
}
|
||||||
|
|
||||||
|
headerType := bodyPlain[0]
|
||||||
|
epoch := binary.BigEndian.Uint64(bodyPlain[1:9])
|
||||||
|
diff := int(math.Abs(float64(time.Now().Unix() - int64(epoch))))
|
||||||
|
if diff > 30 {
|
||||||
|
return DecodedUDPPacket{}, ErrBadTimestamp
|
||||||
|
}
|
||||||
|
|
||||||
|
offset := 9
|
||||||
|
if headerType == HeaderTypeServer {
|
||||||
|
if len(bodyPlain) < offset+8+2 {
|
||||||
|
return DecodedUDPPacket{}, ErrPacketTooShort
|
||||||
|
}
|
||||||
|
offset += 8 // skip clientSessionID
|
||||||
|
}
|
||||||
|
|
||||||
|
paddingLen := int(binary.BigEndian.Uint16(bodyPlain[offset : offset+2]))
|
||||||
|
offset += 2
|
||||||
|
|
||||||
|
if len(bodyPlain) < offset+paddingLen {
|
||||||
|
return DecodedUDPPacket{}, ErrNoPadding
|
||||||
|
}
|
||||||
|
offset += paddingLen
|
||||||
|
|
||||||
|
dest, addrLen, err := parseAddressPort(bodyPlain[offset:])
|
||||||
|
if err != nil {
|
||||||
|
return DecodedUDPPacket{}, err
|
||||||
|
}
|
||||||
|
payload := bodyPlain[offset+addrLen:]
|
||||||
|
|
||||||
|
return DecodedUDPPacket{
|
||||||
|
SessionID: sessionID,
|
||||||
|
PacketID: packetID,
|
||||||
|
HeaderType: headerType,
|
||||||
|
Timestamp: epoch,
|
||||||
|
Destination: dest,
|
||||||
|
Payload: payload,
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *UDPCodec) DecodePacket(data []byte) (DecodedUDPPacket, error) {
|
||||||
|
if len(data) < PacketMinimalHeaderSize {
|
||||||
|
return DecodedUDPPacket{}, ErrPacketTooShort
|
||||||
|
}
|
||||||
|
|
||||||
|
if c.method.IsChaCha {
|
||||||
|
if len(data) < PacketNonceSize+AEADTagSize {
|
||||||
|
return DecodedUDPPacket{}, ErrPacketTooShort
|
||||||
|
}
|
||||||
|
nonce := data[:PacketNonceSize]
|
||||||
|
ciphertext := data[PacketNonceSize:]
|
||||||
|
plain, err := c.chachaCipher.Open(ciphertext[:0], nonce, ciphertext, nil)
|
||||||
|
if err != nil {
|
||||||
|
return DecodedUDPPacket{}, errors.New("failed to decrypt chacha udp packet").Base(err)
|
||||||
|
}
|
||||||
|
if len(plain) < 16+1+8+2 {
|
||||||
|
return DecodedUDPPacket{}, ErrPacketTooShort
|
||||||
|
}
|
||||||
|
|
||||||
|
sessionID := binary.BigEndian.Uint64(plain[:8])
|
||||||
|
packetID := binary.BigEndian.Uint64(plain[8:16])
|
||||||
|
|
||||||
|
if c.sessions != nil {
|
||||||
|
sessionItem := c.sessions.GetOrCreate(sessionID)
|
||||||
|
sessionItem.Lock()
|
||||||
|
if !sessionItem.Window.CheckAndAdd(packetID) {
|
||||||
|
sessionItem.Unlock()
|
||||||
|
return DecodedUDPPacket{}, ErrPacketIdNotUnique
|
||||||
|
}
|
||||||
|
sessionItem.Unlock()
|
||||||
|
}
|
||||||
|
|
||||||
|
return parsePlainUDPPacket(sessionID, packetID, plain[16:])
|
||||||
|
}
|
||||||
|
|
||||||
|
// AES mode
|
||||||
|
var rawHeader [16]byte
|
||||||
|
c.blockCipher.Decrypt(rawHeader[:], data[:16])
|
||||||
|
sessionID := binary.BigEndian.Uint64(rawHeader[:8])
|
||||||
|
packetID := binary.BigEndian.Uint64(rawHeader[8:16])
|
||||||
|
|
||||||
|
var bodyAead cipher.AEAD
|
||||||
|
var sessionItem *ServerUDPSession
|
||||||
|
|
||||||
|
if c.sessions != nil {
|
||||||
|
sessionItem = c.sessions.GetOrCreate(sessionID)
|
||||||
|
sessionItem.Lock()
|
||||||
|
if !sessionItem.Window.Check(packetID) {
|
||||||
|
sessionItem.Unlock()
|
||||||
|
return DecodedUDPPacket{}, ErrPacketIdNotUnique
|
||||||
|
}
|
||||||
|
sessionItem.Unlock()
|
||||||
|
|
||||||
|
bodyAead = sessionItem.GetRemoteCipher()
|
||||||
|
if bodyAead == nil {
|
||||||
|
bodyKey := DeriveSessionSubKey(c.psk, rawHeader[:8], c.method.KeySaltLength)
|
||||||
|
var err error
|
||||||
|
bodyAead, err = c.method.NewAEAD(bodyKey)
|
||||||
|
if err != nil {
|
||||||
|
return DecodedUDPPacket{}, err
|
||||||
|
}
|
||||||
|
sessionItem.SetRemoteCipher(bodyAead)
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
bodyKey := DeriveSessionSubKey(c.psk, rawHeader[:8], c.method.KeySaltLength)
|
||||||
|
var err error
|
||||||
|
bodyAead, err = c.method.NewAEAD(bodyKey)
|
||||||
|
if err != nil {
|
||||||
|
return DecodedUDPPacket{}, err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
bodyNonce := rawHeader[4:16]
|
||||||
|
bodyCipher := data[16:]
|
||||||
|
bodyPlain, err := bodyAead.Open(bodyCipher[:0], bodyNonce, bodyCipher, nil)
|
||||||
|
if err != nil {
|
||||||
|
return DecodedUDPPacket{}, errors.New("failed to decrypt aes udp body").Base(err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if sessionItem != nil {
|
||||||
|
sessionItem.Lock()
|
||||||
|
sessionItem.Window.Add(packetID)
|
||||||
|
sessionItem.Unlock()
|
||||||
|
}
|
||||||
|
|
||||||
|
return parsePlainUDPPacket(sessionID, packetID, bodyPlain)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *ServerUDPSession) EnsureServerState(method *CipherMethod, headerBlock cipher.Block, chachaCipher cipher.AEAD, psk []byte) error {
|
||||||
|
s.Lock()
|
||||||
|
defer s.Unlock()
|
||||||
|
if s.ServerSessionID != 0 {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
var sidBuf [8]byte
|
||||||
|
for {
|
||||||
|
if _, err := io.ReadFull(rand.Reader, sidBuf[:]); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
s.ServerSessionID = binary.BigEndian.Uint64(sidBuf[:])
|
||||||
|
if s.ServerSessionID != 0 {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if method.IsChaCha {
|
||||||
|
s.ServerChaCha = chachaCipher
|
||||||
|
} else {
|
||||||
|
s.ServerBlockCipher = headerBlock
|
||||||
|
bodyKey := DeriveSessionSubKey(psk, sidBuf[:], method.KeySaltLength)
|
||||||
|
bodyAead, err := method.NewAEAD(bodyKey)
|
||||||
|
if err != nil {
|
||||||
|
s.ServerSessionID = 0
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
s.ServerCipher = bodyAead
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *ServerUDPSession) EncodeServerPacket(method *CipherMethod, clientSessionID uint64, dest net.Destination, payload []byte) ([]byte, error) {
|
||||||
|
serverSessionID := s.ServerSessionID
|
||||||
|
serverPacketID := s.ServerPacketID.Add(1)
|
||||||
|
|
||||||
|
if method.IsChaCha {
|
||||||
|
var nonce [PacketNonceSize]byte
|
||||||
|
if _, err := io.ReadFull(rand.Reader, nonce[:]); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
plainBuf := buf.New()
|
||||||
|
defer plainBuf.Release()
|
||||||
|
|
||||||
|
var hdr [16 + 1 + 8 + 8 + 2]byte
|
||||||
|
binary.BigEndian.PutUint64(hdr[0:8], serverSessionID)
|
||||||
|
binary.BigEndian.PutUint64(hdr[8:16], serverPacketID)
|
||||||
|
hdr[16] = HeaderTypeServer
|
||||||
|
binary.BigEndian.PutUint64(hdr[17:25], uint64(time.Now().Unix()))
|
||||||
|
binary.BigEndian.PutUint64(hdr[25:33], clientSessionID)
|
||||||
|
binary.BigEndian.PutUint16(hdr[33:35], 0)
|
||||||
|
plainBuf.Write(hdr[:])
|
||||||
|
|
||||||
|
if err := WriteAddressPort(plainBuf, dest); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
plainBuf.Write(payload)
|
||||||
|
|
||||||
|
sealed := s.ServerChaCha.Seal(nil, nonce[:], plainBuf.Bytes(), nil)
|
||||||
|
res := make([]byte, PacketNonceSize+len(sealed))
|
||||||
|
copy(res[:PacketNonceSize], nonce[:])
|
||||||
|
copy(res[PacketNonceSize:], sealed)
|
||||||
|
return res, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// AES mode
|
||||||
|
var rawHeader [16]byte
|
||||||
|
binary.BigEndian.PutUint64(rawHeader[:8], serverSessionID)
|
||||||
|
binary.BigEndian.PutUint64(rawHeader[8:16], serverPacketID)
|
||||||
|
|
||||||
|
var encryptedHeader [16]byte
|
||||||
|
s.ServerBlockCipher.Encrypt(encryptedHeader[:], rawHeader[:])
|
||||||
|
|
||||||
|
bodyBuf := buf.New()
|
||||||
|
defer bodyBuf.Release()
|
||||||
|
|
||||||
|
var hdr [1 + 8 + 8 + 2]byte
|
||||||
|
hdr[0] = HeaderTypeServer
|
||||||
|
binary.BigEndian.PutUint64(hdr[1:9], uint64(time.Now().Unix()))
|
||||||
|
binary.BigEndian.PutUint64(hdr[9:17], clientSessionID)
|
||||||
|
binary.BigEndian.PutUint16(hdr[17:19], 0)
|
||||||
|
bodyBuf.Write(hdr[:])
|
||||||
|
|
||||||
|
if err := WriteAddressPort(bodyBuf, dest); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
bodyBuf.Write(payload)
|
||||||
|
|
||||||
|
bodyNonce := rawHeader[4:16]
|
||||||
|
sealedBody := s.ServerCipher.Seal(nil, bodyNonce, bodyBuf.Bytes(), nil)
|
||||||
|
|
||||||
|
res := make([]byte, 16+len(sealedBody))
|
||||||
|
copy(res[:16], encryptedHeader[:])
|
||||||
|
copy(res[16:], sealedBody)
|
||||||
|
return res, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *UDPCodec) EncodeServerPacket(clientSessionID uint64, dest net.Destination, payload []byte) ([]byte, error) {
|
||||||
|
sessionItem := c.sessions.GetOrCreate(clientSessionID)
|
||||||
|
if err := sessionItem.EnsureServerState(c.method, c.blockCipher, c.chachaCipher, c.psk); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
return sessionItem.EncodeServerPacket(c.method, clientSessionID, dest, payload)
|
||||||
|
}
|
||||||
|
|
||||||
|
type UDPWriter struct {
|
||||||
|
Writer io.Writer
|
||||||
|
Destination net.Destination
|
||||||
|
Codec *UDPPacketCodec
|
||||||
|
}
|
||||||
|
|
||||||
|
func (w *UDPWriter) WriteMultiBuffer(mb buf.MultiBuffer) error {
|
||||||
|
for {
|
||||||
|
mb2, b := buf.SplitFirst(mb)
|
||||||
|
mb = mb2
|
||||||
|
if b == nil {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
dest := w.Destination
|
||||||
|
if b.UDP != nil {
|
||||||
|
dest = *b.UDP
|
||||||
|
}
|
||||||
|
pktBuf, err := w.Codec.EncodeClientPacket(dest, b.Bytes())
|
||||||
|
b.Release()
|
||||||
|
if err != nil {
|
||||||
|
buf.ReleaseMulti(mb)
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
_, writeErr := w.Writer.Write(pktBuf.Bytes())
|
||||||
|
pktBuf.Release()
|
||||||
|
if writeErr != nil {
|
||||||
|
buf.ReleaseMulti(mb)
|
||||||
|
return writeErr
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
type UDPReader struct {
|
||||||
|
Reader io.Reader
|
||||||
|
Codec *UDPPacketCodec
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *UDPReader) ReadMultiBuffer() (buf.MultiBuffer, error) {
|
||||||
|
for {
|
||||||
|
buffer := buf.New()
|
||||||
|
_, err := buffer.ReadFrom(r.Reader)
|
||||||
|
if err != nil {
|
||||||
|
buffer.Release()
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
decoded, err := r.Codec.DecodePacket(buffer.Bytes())
|
||||||
|
if err != nil {
|
||||||
|
buffer.Release()
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
buffer.Clear()
|
||||||
|
buffer.Write(decoded.Payload)
|
||||||
|
dest := decoded.Destination
|
||||||
|
buffer.UDP = &dest
|
||||||
|
return buf.MultiBuffer{buffer}, nil
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,271 @@
|
|||||||
|
package shadowsocks_2022_test
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"encoding/base64"
|
||||||
|
"encoding/binary"
|
||||||
|
"errors"
|
||||||
|
gonet "net"
|
||||||
|
"sync"
|
||||||
|
"sync/atomic"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/xtls/xray-core/common"
|
||||||
|
"github.com/xtls/xray-core/common/buf"
|
||||||
|
"github.com/xtls/xray-core/common/net"
|
||||||
|
"github.com/xtls/xray-core/features/routing"
|
||||||
|
. "github.com/xtls/xray-core/proxy/shadowsocks_2022"
|
||||||
|
"github.com/xtls/xray-core/transport"
|
||||||
|
"lukechampine.com/blake3"
|
||||||
|
)
|
||||||
|
|
||||||
|
// encodeRelayClientUDPPacket encodes a Shadowsocks-2022 UDP packet with 1 layer of EIH (Relay)
|
||||||
|
func encodeRelayClientUDPPacket(relayKey, destKey []byte, sessionID, packetID uint64, dest net.Destination, payload []byte) ([]byte, error) {
|
||||||
|
method, err := GetCipherMethod(MethodAES128GCM)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
relayBlock, err := method.NewBlock(relayKey)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
// 1. Plain packet header: sessionID (8B) + packetID (8B)
|
||||||
|
var rawHeader [16]byte
|
||||||
|
binary.BigEndian.PutUint64(rawHeader[:8], sessionID)
|
||||||
|
binary.BigEndian.PutUint64(rawHeader[8:16], packetID)
|
||||||
|
|
||||||
|
// Encrypt packetHeader under relayKey
|
||||||
|
var encPacketHeader [16]byte
|
||||||
|
relayBlock.Encrypt(encPacketHeader[:], rawHeader[:])
|
||||||
|
|
||||||
|
// 2. EI Header: blake3(destKey)[:16] ^ rawHeader
|
||||||
|
var destHash [16]byte
|
||||||
|
hash512 := blake3.Sum512(destKey)
|
||||||
|
copy(destHash[:], hash512[:16])
|
||||||
|
|
||||||
|
var eiHeader [16]byte
|
||||||
|
for i := 0; i < 16; i++ {
|
||||||
|
eiHeader[i] = destHash[i] ^ rawHeader[i]
|
||||||
|
}
|
||||||
|
var encEIHeader [16]byte
|
||||||
|
relayBlock.Encrypt(encEIHeader[:], eiHeader[:])
|
||||||
|
|
||||||
|
// 3. Payload under destination server's AEAD
|
||||||
|
bodyKey := DeriveSessionSubKey(destKey, rawHeader[:8], 16)
|
||||||
|
bodyAead, err := method.NewAEAD(bodyKey)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
bodyNonce := rawHeader[4:16]
|
||||||
|
|
||||||
|
outBuf := buf.New()
|
||||||
|
defer outBuf.Release()
|
||||||
|
|
||||||
|
// VarHeader: client type (1) + timestamp (8) + paddingLen (2) + padding + dest + payload
|
||||||
|
var hdr [1 + 8 + 2]byte
|
||||||
|
hdr[0] = HeaderTypeClient
|
||||||
|
binary.BigEndian.PutUint64(hdr[1:9], uint64(time.Now().Unix()))
|
||||||
|
binary.BigEndian.PutUint16(hdr[9:11], 0)
|
||||||
|
outBuf.Write(hdr[:])
|
||||||
|
|
||||||
|
if err := WriteAddressPort(outBuf, dest); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
outBuf.Write(payload)
|
||||||
|
|
||||||
|
plainBytes := outBuf.Bytes()
|
||||||
|
outBuf.Extend(int32(bodyAead.Overhead()))
|
||||||
|
bodyAead.Seal(plainBytes[:0], bodyNonce, plainBytes, nil)
|
||||||
|
|
||||||
|
// Full packet: encPacketHeader (16B) + encEIHeader (16B) + sealedBody
|
||||||
|
packet := make([]byte, 0, 32+outBuf.Len())
|
||||||
|
packet = append(packet, encPacketHeader[:]...)
|
||||||
|
packet = append(packet, encEIHeader[:]...)
|
||||||
|
packet = append(packet, outBuf.Bytes()...)
|
||||||
|
return packet, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRelayUDPSessionStabilityAndDispatch(t *testing.T) {
|
||||||
|
relayKey := []byte("0123456789abcdef")
|
||||||
|
destKey := []byte("fedcba9876543210")
|
||||||
|
relayKeyB64 := base64.StdEncoding.EncodeToString(relayKey)
|
||||||
|
destKeyB64 := base64.StdEncoding.EncodeToString(destKey)
|
||||||
|
|
||||||
|
config := &RelayServerConfig{
|
||||||
|
Method: MethodAES128GCM,
|
||||||
|
Key: relayKeyB64,
|
||||||
|
Destinations: []*RelayDestination{
|
||||||
|
{
|
||||||
|
Key: destKeyB64,
|
||||||
|
Address: &net.IPOrDomain{Address: &net.IPOrDomain_Ip{Ip: []byte{127, 0, 0, 1}}},
|
||||||
|
Port: 8388,
|
||||||
|
Email: "dest@example.com",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
inbound, err := NewRelayServer(newTestContext(), config)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("failed to create RelayServer: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
sessionID := uint64(0x1122334455667788)
|
||||||
|
dest := net.UDPDestination(net.LocalHostIP, 8388)
|
||||||
|
|
||||||
|
pkt1, err := encodeRelayClientUDPPacket(relayKey, destKey, sessionID, 1, dest, []byte("xray packet 1"))
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("failed to encode pkt1: %v", err)
|
||||||
|
}
|
||||||
|
pkt2, err := encodeRelayClientUDPPacket(relayKey, destKey, sessionID, 2, dest, []byte("xray packet 2"))
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("failed to encode pkt2: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
var dispatchCount atomic.Int32
|
||||||
|
var receivedPackets [][]byte
|
||||||
|
var mu sync.Mutex
|
||||||
|
|
||||||
|
disp := &dummyDispatcher{
|
||||||
|
onDispatch: func(ctx context.Context, d net.Destination) (*transport.Link, error) {
|
||||||
|
dispatchCount.Add(1)
|
||||||
|
linkR, linkW := gonet.Pipe()
|
||||||
|
t.Cleanup(func() {
|
||||||
|
linkW.Close()
|
||||||
|
linkR.Close()
|
||||||
|
})
|
||||||
|
link := &transport.Link{
|
||||||
|
Reader: buf.NewReader(linkR),
|
||||||
|
Writer: &customWriter{
|
||||||
|
write: func(mb buf.MultiBuffer) error {
|
||||||
|
mu.Lock()
|
||||||
|
defer mu.Unlock()
|
||||||
|
for _, b := range mb {
|
||||||
|
cpy := make([]byte, b.Len())
|
||||||
|
copy(cpy, b.Bytes())
|
||||||
|
receivedPackets = append(receivedPackets, cpy)
|
||||||
|
b.Release()
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
return link, nil
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
clientConn, serverConn := gonet.Pipe()
|
||||||
|
defer clientConn.Close()
|
||||||
|
defer serverConn.Close()
|
||||||
|
|
||||||
|
inboundConn := &dummyStatConn{Conn: serverConn}
|
||||||
|
ctx, cancel := context.WithCancel(newTestContext())
|
||||||
|
defer cancel()
|
||||||
|
|
||||||
|
go func() {
|
||||||
|
_ = inbound.Process(ctx, net.Network_UDP, inboundConn, disp)
|
||||||
|
}()
|
||||||
|
|
||||||
|
// Send Packet 1
|
||||||
|
_, err = clientConn.Write(pkt1)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("write pkt1 failed: %v", err)
|
||||||
|
}
|
||||||
|
time.Sleep(50 * time.Millisecond)
|
||||||
|
|
||||||
|
// Send Packet 2 (same sessionID, packetID=2)
|
||||||
|
_, err = clientConn.Write(pkt2)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("write pkt2 failed: %v", err)
|
||||||
|
}
|
||||||
|
time.Sleep(50 * time.Millisecond)
|
||||||
|
|
||||||
|
// Check dispatch count: For the SAME UDP session, Dispatch MUST be called exactly ONCE!
|
||||||
|
if count := dispatchCount.Load(); count != 1 {
|
||||||
|
t.Fatalf("CRITICAL BUG CONFIRMED: expected dispatchCount = 1 for same session, got %d (sessionID was corrupted by Encrypt!)", count)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Verify downstream destination can decode both packets
|
||||||
|
method, err := GetCipherMethod(MethodAES128GCM)
|
||||||
|
common.Must(err)
|
||||||
|
destCodec, err := NewUDPServerCodec(method, destKey, 300*time.Second)
|
||||||
|
common.Must(err)
|
||||||
|
|
||||||
|
mu.Lock()
|
||||||
|
pkts := receivedPackets
|
||||||
|
mu.Unlock()
|
||||||
|
|
||||||
|
if len(pkts) != 2 {
|
||||||
|
t.Fatalf("expected 2 received packets at destination, got %d", len(pkts))
|
||||||
|
}
|
||||||
|
|
||||||
|
dec1, err := destCodec.DecodePacket(pkts[0])
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("dest failed to decode packet 1: %v", err)
|
||||||
|
}
|
||||||
|
if dec1.SessionID != sessionID || dec1.PacketID != 1 || string(dec1.Payload) != "xray packet 1" {
|
||||||
|
t.Fatalf("dec1 mismatch: sess=%x, pktID=%d, payload=%s", dec1.SessionID, dec1.PacketID, string(dec1.Payload))
|
||||||
|
}
|
||||||
|
|
||||||
|
dec2, err := destCodec.DecodePacket(pkts[1])
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("dest failed to decode packet 2: %v", err)
|
||||||
|
}
|
||||||
|
if dec2.SessionID != sessionID || dec2.PacketID != 2 || string(dec2.Payload) != "xray packet 2" {
|
||||||
|
t.Fatalf("dec2 mismatch: sess=%x, pktID=%d, payload=%s", dec2.SessionID, dec2.PacketID, string(dec2.Payload))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
type customWriter struct {
|
||||||
|
write func(mb buf.MultiBuffer) error
|
||||||
|
}
|
||||||
|
|
||||||
|
func (w *customWriter) WriteMultiBuffer(mb buf.MultiBuffer) error {
|
||||||
|
return w.write(mb)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (w *customWriter) Close() error {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (w *customWriter) Interrupt() {}
|
||||||
|
|
||||||
|
type dummyDispatcher struct {
|
||||||
|
onDispatch func(ctx context.Context, dest net.Destination) (*transport.Link, error)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (d *dummyDispatcher) Dispatch(ctx context.Context, dest net.Destination) (*transport.Link, error) {
|
||||||
|
if d.onDispatch != nil {
|
||||||
|
return d.onDispatch(ctx, dest)
|
||||||
|
}
|
||||||
|
return nil, errors.New("not handled")
|
||||||
|
}
|
||||||
|
|
||||||
|
func (d *dummyDispatcher) DispatchLink(ctx context.Context, dest net.Destination, link *transport.Link) error {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (d *dummyDispatcher) Start() error { return nil }
|
||||||
|
func (d *dummyDispatcher) Close() error { return nil }
|
||||||
|
func (d *dummyDispatcher) Type() interface{} { return routing.DispatcherType() }
|
||||||
|
|
||||||
|
type dummyStatConn struct {
|
||||||
|
gonet.Conn
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *dummyStatConn) ReadMultiBuffer() (buf.MultiBuffer, error) {
|
||||||
|
b := buf.New()
|
||||||
|
_, err := b.ReadFrom(c.Conn)
|
||||||
|
return buf.MultiBuffer{b}, err
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *dummyStatConn) WriteMultiBuffer(mb buf.MultiBuffer) error {
|
||||||
|
defer buf.ReleaseMulti(mb)
|
||||||
|
for _, b := range mb {
|
||||||
|
if _, err := c.Conn.Write(b.Bytes()); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
@@ -0,0 +1,158 @@
|
|||||||
|
package shadowsocks_2022
|
||||||
|
|
||||||
|
import (
|
||||||
|
"crypto/cipher"
|
||||||
|
"sync"
|
||||||
|
"sync/atomic"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/xtls/xray-core/common/protocol"
|
||||||
|
"github.com/xtls/xray-core/common/utils"
|
||||||
|
)
|
||||||
|
|
||||||
|
const (
|
||||||
|
swBlockBitLog = 6 // 1<<6 == 64 bits
|
||||||
|
swBlockBits = 1 << swBlockBitLog // 64
|
||||||
|
swRingBlocks = 1 << 7 // 128
|
||||||
|
swBlockMask = swRingBlocks - 1 // 127
|
||||||
|
swBitMask = swBlockBits - 1 // 63
|
||||||
|
swSize = (swRingBlocks - 1) * swBlockBits // 8128
|
||||||
|
)
|
||||||
|
|
||||||
|
type SlidingWindow struct {
|
||||||
|
last uint64
|
||||||
|
ring [swRingBlocks]uint64
|
||||||
|
}
|
||||||
|
|
||||||
|
func (f *SlidingWindow) Reset() {
|
||||||
|
*f = SlidingWindow{}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (f *SlidingWindow) Check(counter uint64) bool {
|
||||||
|
switch {
|
||||||
|
case counter > f.last:
|
||||||
|
return true
|
||||||
|
case f.last-counter > swSize:
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
blockIndex := (counter >> swBlockBitLog) & swBlockMask
|
||||||
|
bitIndex := counter & swBitMask
|
||||||
|
return (f.ring[blockIndex]>>bitIndex)&1 == 0
|
||||||
|
}
|
||||||
|
|
||||||
|
func (f *SlidingWindow) Add(counter uint64) {
|
||||||
|
blockIndex := counter >> swBlockBitLog
|
||||||
|
|
||||||
|
if counter > f.last {
|
||||||
|
lastBlockIndex := f.last >> swBlockBitLog
|
||||||
|
diff := int(blockIndex - lastBlockIndex)
|
||||||
|
if diff > swRingBlocks {
|
||||||
|
diff = swRingBlocks
|
||||||
|
}
|
||||||
|
|
||||||
|
for i := 0; i < diff; i++ {
|
||||||
|
lastBlockIndex = (lastBlockIndex + 1) & swBlockMask
|
||||||
|
f.ring[lastBlockIndex] = 0
|
||||||
|
}
|
||||||
|
|
||||||
|
f.last = counter
|
||||||
|
}
|
||||||
|
|
||||||
|
blockIndex &= swBlockMask
|
||||||
|
bitIndex := counter & swBitMask
|
||||||
|
f.ring[blockIndex] |= 1 << bitIndex
|
||||||
|
}
|
||||||
|
|
||||||
|
func (f *SlidingWindow) CheckAndAdd(counter uint64) bool {
|
||||||
|
if !f.Check(counter) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
f.Add(counter)
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
type ServerUDPSession struct {
|
||||||
|
sync.Mutex
|
||||||
|
SessionID uint64
|
||||||
|
RemoteCipher atomic.Pointer[cipher.AEAD]
|
||||||
|
Window SlidingWindow
|
||||||
|
User *protocol.MemoryUser
|
||||||
|
UserPSK []byte
|
||||||
|
LastActive atomic.Int64 // Unix timestamp in seconds
|
||||||
|
|
||||||
|
ServerSessionID uint64
|
||||||
|
ServerPacketID atomic.Uint64
|
||||||
|
ServerCipher cipher.AEAD
|
||||||
|
ServerBlockCipher cipher.Block
|
||||||
|
ServerChaCha cipher.AEAD
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *ServerUDPSession) GetRemoteCipher() cipher.AEAD {
|
||||||
|
ptr := s.RemoteCipher.Load()
|
||||||
|
if ptr == nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
return *ptr
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *ServerUDPSession) SetRemoteCipher(c cipher.AEAD) {
|
||||||
|
s.RemoteCipher.Store(&c)
|
||||||
|
}
|
||||||
|
|
||||||
|
type UDPSessionManager struct {
|
||||||
|
sessions *utils.TypedSyncMap[uint64, *ServerUDPSession]
|
||||||
|
timeout time.Duration
|
||||||
|
lastClean atomic.Int64 // Unix timestamp in seconds
|
||||||
|
}
|
||||||
|
|
||||||
|
func NewUDPSessionManager(timeout time.Duration) *UDPSessionManager {
|
||||||
|
return &UDPSessionManager{
|
||||||
|
sessions: utils.NewTypedSyncMap[uint64, *ServerUDPSession](),
|
||||||
|
timeout: timeout,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *UDPSessionManager) GetOrCreate(sessionID uint64) *ServerUDPSession {
|
||||||
|
now := time.Now().Unix()
|
||||||
|
if s, ok := m.sessions.Load(sessionID); ok {
|
||||||
|
s.LastActive.Store(now)
|
||||||
|
return s
|
||||||
|
}
|
||||||
|
|
||||||
|
s := &ServerUDPSession{
|
||||||
|
SessionID: sessionID,
|
||||||
|
}
|
||||||
|
s.LastActive.Store(now)
|
||||||
|
|
||||||
|
actual, loaded := m.sessions.LoadOrStore(sessionID, s)
|
||||||
|
if loaded {
|
||||||
|
actual.LastActive.Store(now)
|
||||||
|
return actual
|
||||||
|
}
|
||||||
|
|
||||||
|
// Trigger cleanup if at least 30 seconds have passed since last cleanup
|
||||||
|
last := m.lastClean.Load()
|
||||||
|
if now-last > 30 && m.lastClean.CompareAndSwap(last, now) {
|
||||||
|
go m.cleanup(now)
|
||||||
|
}
|
||||||
|
|
||||||
|
return s
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *UDPSessionManager) cleanup(now int64) {
|
||||||
|
timeoutSec := int64(m.timeout.Seconds())
|
||||||
|
if timeoutSec <= 0 {
|
||||||
|
timeoutSec = 60
|
||||||
|
}
|
||||||
|
m.sessions.Range(func(k uint64, v *ServerUDPSession) bool {
|
||||||
|
if now-v.LastActive.Load() > timeoutSec {
|
||||||
|
m.sessions.Delete(k)
|
||||||
|
}
|
||||||
|
return true
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *UDPSessionManager) Delete(sessionID uint64) {
|
||||||
|
m.sessions.Delete(sessionID)
|
||||||
|
}
|
||||||
@@ -1 +1,50 @@
|
|||||||
package shadowsocks_2022
|
package shadowsocks_2022
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"sync"
|
||||||
|
|
||||||
|
"github.com/xtls/xray-core/common/errors"
|
||||||
|
"github.com/xtls/xray-core/common/signal"
|
||||||
|
"github.com/xtls/xray-core/transport"
|
||||||
|
)
|
||||||
|
|
||||||
|
type udpConnEntry struct {
|
||||||
|
sync.Mutex
|
||||||
|
link *transport.Link
|
||||||
|
timer *signal.ActivityTimer
|
||||||
|
cancel context.CancelFunc
|
||||||
|
}
|
||||||
|
|
||||||
|
const (
|
||||||
|
HeaderTypeClient = 0
|
||||||
|
HeaderTypeServer = 1
|
||||||
|
MaxPaddingLength = 900
|
||||||
|
PacketNonceSize = 24
|
||||||
|
MaxPacketSize = 65535
|
||||||
|
RequestHeaderFixedChunkLength = 1 + 8 + 2 // Type (1B) + Timestamp (8B) + VarHeaderLen (2B)
|
||||||
|
PacketMinimalHeaderSize = 30
|
||||||
|
StreamNonceSize = 12
|
||||||
|
AESBlockSize = 16
|
||||||
|
AEADTagSize = 16
|
||||||
|
)
|
||||||
|
|
||||||
|
var zeroPadding [MaxPaddingLength]byte
|
||||||
|
|
||||||
|
const (
|
||||||
|
MethodAES128GCM = "2022-blake3-aes-128-gcm"
|
||||||
|
MethodAES256GCM = "2022-blake3-aes-256-gcm"
|
||||||
|
MethodChaCha20Poly1305 = "2022-blake3-chacha20-poly1305"
|
||||||
|
)
|
||||||
|
|
||||||
|
var (
|
||||||
|
ErrBadKey = errors.New("bad key")
|
||||||
|
ErrBadHeaderType = errors.New("bad header type")
|
||||||
|
ErrBadTimestamp = errors.New("bad timestamp")
|
||||||
|
ErrSaltNotUnique = errors.New("salt not unique")
|
||||||
|
ErrPacketIdNotUnique = errors.New("packet id not unique")
|
||||||
|
ErrPacketTooShort = errors.New("packet too short")
|
||||||
|
ErrPacketTooLarge = errors.New("packet too large")
|
||||||
|
ErrNoPadding = errors.New("bad request: missing payload or padding")
|
||||||
|
ErrInvalidRequest = errors.New("invalid request")
|
||||||
|
)
|
||||||
|
|||||||
@@ -0,0 +1,362 @@
|
|||||||
|
package shadowsocks_2022_test
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"context"
|
||||||
|
"crypto/rand"
|
||||||
|
"encoding/base64"
|
||||||
|
"encoding/binary"
|
||||||
|
"io"
|
||||||
|
gonet "net"
|
||||||
|
"sync"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/google/go-cmp/cmp"
|
||||||
|
"github.com/xtls/xray-core/common"
|
||||||
|
"github.com/xtls/xray-core/common/buf"
|
||||||
|
"github.com/xtls/xray-core/common/net"
|
||||||
|
"github.com/xtls/xray-core/common/protocol"
|
||||||
|
"github.com/xtls/xray-core/common/serial"
|
||||||
|
"github.com/xtls/xray-core/common/session"
|
||||||
|
"github.com/xtls/xray-core/core"
|
||||||
|
. "github.com/xtls/xray-core/proxy/shadowsocks_2022"
|
||||||
|
)
|
||||||
|
|
||||||
|
func newTestContext() context.Context {
|
||||||
|
v, err := core.New(&core.Config{})
|
||||||
|
common.Must(err)
|
||||||
|
ctx := context.WithValue(context.Background(), core.XrayKey(1), v)
|
||||||
|
ctx = session.ContextWithInbound(ctx, &session.Inbound{})
|
||||||
|
return ctx
|
||||||
|
}
|
||||||
|
|
||||||
|
func generateRandomKey(size int) string {
|
||||||
|
b := make([]byte, size)
|
||||||
|
_, _ = rand.Read(b)
|
||||||
|
return base64.StdEncoding.EncodeToString(b)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestKDF(t *testing.T) {
|
||||||
|
// Test ParseKey
|
||||||
|
if _, err := ParseKey("", 16); err != ErrBadKey {
|
||||||
|
t.Fatalf("expected ErrBadKey for empty key, got %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
shortKey := base64.StdEncoding.EncodeToString([]byte("short"))
|
||||||
|
if _, err := ParseKey(shortKey, 16); err != ErrBadKey {
|
||||||
|
t.Fatalf("expected ErrBadKey for short key, got %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
exactKey := []byte("0123456789abcdef")
|
||||||
|
exactKeyB64 := base64.StdEncoding.EncodeToString(exactKey)
|
||||||
|
normExact, err := ParseKey(exactKeyB64, 16)
|
||||||
|
if err != nil || !bytes.Equal(normExact, exactKey) {
|
||||||
|
t.Fatalf("unexpected parsed exact key: %v, err: %v", normExact, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
longKey := base64.StdEncoding.EncodeToString([]byte("0123456789abcdef_longer_key_for_testing"))
|
||||||
|
if _, err := ParseKey(longKey, 16); err != ErrBadKey {
|
||||||
|
t.Fatalf("expected ErrBadKey for long key, got %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Test Session Subkey determinism
|
||||||
|
salt := []byte("random_salt_1234")
|
||||||
|
k1 := DeriveSessionSubKey(normExact, salt, 16)
|
||||||
|
k2 := DeriveSessionSubKey(normExact, salt, 16)
|
||||||
|
if !bytes.Equal(k1, k2) {
|
||||||
|
t.Fatal("DeriveSessionSubKey should be deterministic")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Identity subkey must differ from session subkey with same inputs
|
||||||
|
idKey := DeriveIdentitySubKey(normExact, salt, 16)
|
||||||
|
if bytes.Equal(k1, idKey) {
|
||||||
|
t.Fatal("DeriveIdentitySubKey must differ from DeriveSessionSubKey")
|
||||||
|
}
|
||||||
|
|
||||||
|
// User PSK hash
|
||||||
|
h1 := DeriveUserPSKHash(normExact)
|
||||||
|
h2 := DeriveUserPSKHash(normExact)
|
||||||
|
if h1 != h2 {
|
||||||
|
t.Fatal("DeriveUserPSKHash should be deterministic")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSlidingWindow(t *testing.T) {
|
||||||
|
var window SlidingWindow
|
||||||
|
if !window.Check(1) {
|
||||||
|
t.Fatal("packet 1 should be accepted")
|
||||||
|
}
|
||||||
|
window.Add(1)
|
||||||
|
|
||||||
|
if window.Check(1) {
|
||||||
|
t.Fatal("duplicate packet 1 should be rejected")
|
||||||
|
}
|
||||||
|
|
||||||
|
if !window.Check(100) {
|
||||||
|
t.Fatal("packet 100 should be accepted")
|
||||||
|
}
|
||||||
|
window.Add(100)
|
||||||
|
|
||||||
|
if window.Check(100) {
|
||||||
|
t.Fatal("duplicate packet 100 should be rejected")
|
||||||
|
}
|
||||||
|
|
||||||
|
if !window.Check(50) {
|
||||||
|
t.Fatal("out-of-order packet 50 within window should be accepted")
|
||||||
|
}
|
||||||
|
window.Add(50)
|
||||||
|
if window.Check(50) {
|
||||||
|
t.Fatal("duplicate packet 50 should be rejected")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Check packet far behind window (> 8128)
|
||||||
|
window.Add(10000)
|
||||||
|
if window.Check(1) {
|
||||||
|
t.Fatal("packet 1 should be rejected as behind window")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestTCPStream(t *testing.T) {
|
||||||
|
methods := []struct {
|
||||||
|
name string
|
||||||
|
keySize int
|
||||||
|
}{
|
||||||
|
{MethodAES128GCM, 16},
|
||||||
|
{MethodAES256GCM, 32},
|
||||||
|
{MethodChaCha20Poly1305, 32},
|
||||||
|
}
|
||||||
|
|
||||||
|
dest := net.TCPDestination(net.LocalHostIP, net.Port(8080))
|
||||||
|
testPayload := []byte("Hello, Shadowsocks 2022 Native Implementation!")
|
||||||
|
|
||||||
|
for _, m := range methods {
|
||||||
|
t.Run(m.name, func(t *testing.T) {
|
||||||
|
rawKey := make([]byte, m.keySize)
|
||||||
|
_, _ = rand.Read(rawKey)
|
||||||
|
method, err := GetCipherMethod(m.name)
|
||||||
|
common.Must(err)
|
||||||
|
|
||||||
|
clientConn, serverConn := gonet.Pipe()
|
||||||
|
defer clientConn.Close()
|
||||||
|
defer serverConn.Close()
|
||||||
|
|
||||||
|
var wg sync.WaitGroup
|
||||||
|
wg.Add(2)
|
||||||
|
|
||||||
|
var receivedDest net.Destination
|
||||||
|
var receivedPayload []byte
|
||||||
|
|
||||||
|
// Server goroutine
|
||||||
|
go func() {
|
||||||
|
defer wg.Done()
|
||||||
|
salt := make([]byte, method.KeySaltLength)
|
||||||
|
_, err := io.ReadFull(serverConn, salt)
|
||||||
|
common.Must(err)
|
||||||
|
|
||||||
|
sessionKey := DeriveSessionSubKey(rawKey, salt, method.KeySaltLength)
|
||||||
|
aead, err := method.NewAEAD(sessionKey)
|
||||||
|
common.Must(err)
|
||||||
|
|
||||||
|
reader := NewStreamReader(serverConn, aead)
|
||||||
|
|
||||||
|
// Read fixed chunk (11 + 16 bytes)
|
||||||
|
var fixedBuf [RequestHeaderFixedChunkLength + AEADTagSize]byte
|
||||||
|
_, err = io.ReadFull(serverConn, fixedBuf[:])
|
||||||
|
common.Must(err)
|
||||||
|
|
||||||
|
plainFixed, err := aead.Open(fixedBuf[:0], reader.Nonce(), fixedBuf[:], nil)
|
||||||
|
common.Must(err)
|
||||||
|
IncreaseNonce(reader.Nonce())
|
||||||
|
if plainFixed[0] != HeaderTypeClient {
|
||||||
|
t.Errorf("expected client header type, got %d", plainFixed[0])
|
||||||
|
}
|
||||||
|
|
||||||
|
// Read variable chunk
|
||||||
|
varLen := int(plainFixed[9])<<8 | int(plainFixed[10])
|
||||||
|
varBuf := make([]byte, varLen+AEADTagSize)
|
||||||
|
_, err = io.ReadFull(serverConn, varBuf)
|
||||||
|
common.Must(err)
|
||||||
|
|
||||||
|
plainVar, err := aead.Open(varBuf[:0], reader.Nonce(), varBuf, nil)
|
||||||
|
common.Must(err)
|
||||||
|
IncreaseNonce(reader.Nonce())
|
||||||
|
|
||||||
|
vBuf := buf.New()
|
||||||
|
vBuf.Write(plainVar)
|
||||||
|
receivedDest, err = ReadAddressPort(vBuf)
|
||||||
|
common.Must(err)
|
||||||
|
|
||||||
|
// Skip padding
|
||||||
|
var padBytes [2]byte
|
||||||
|
_, _ = vBuf.Read(padBytes[:])
|
||||||
|
padLen := int(padBytes[0])<<8 | int(padBytes[1])
|
||||||
|
vBuf.Advance(int32(padLen))
|
||||||
|
|
||||||
|
receivedPayload = make([]byte, vBuf.Len())
|
||||||
|
copy(receivedPayload, vBuf.Bytes())
|
||||||
|
vBuf.Release()
|
||||||
|
|
||||||
|
// Server sends response handshake
|
||||||
|
serverSalt := make([]byte, method.KeySaltLength)
|
||||||
|
_, _ = rand.Read(serverSalt)
|
||||||
|
respKey := DeriveSessionSubKey(rawKey, serverSalt, method.KeySaltLength)
|
||||||
|
respAead, err := method.NewAEAD(respKey)
|
||||||
|
writer := NewStreamWriter(serverConn, respAead)
|
||||||
|
_, _ = serverConn.Write(serverSalt)
|
||||||
|
|
||||||
|
fixedResp := make([]byte, 1+8+method.KeySaltLength+2)
|
||||||
|
fixedResp[0] = HeaderTypeServer
|
||||||
|
binary.BigEndian.PutUint64(fixedResp[1:9], uint64(time.Now().Unix()))
|
||||||
|
copy(fixedResp[9:9+method.KeySaltLength], salt)
|
||||||
|
binary.BigEndian.PutUint16(fixedResp[9+method.KeySaltLength:11+method.KeySaltLength], 0)
|
||||||
|
|
||||||
|
fixedChunk := respAead.Seal(nil, writer.Nonce(), fixedResp, nil)
|
||||||
|
IncreaseNonce(writer.Nonce())
|
||||||
|
_, _ = serverConn.Write(fixedChunk)
|
||||||
|
|
||||||
|
// Echo stream data
|
||||||
|
mb, err := reader.ReadMultiBuffer()
|
||||||
|
common.Must(err)
|
||||||
|
_ = writer.WriteMultiBuffer(mb)
|
||||||
|
}()
|
||||||
|
|
||||||
|
// Client goroutine
|
||||||
|
go func() {
|
||||||
|
defer wg.Done()
|
||||||
|
clientSalt, writer, err := ClientHandshake(clientConn, method, [][]byte{rawKey}, dest, testPayload)
|
||||||
|
common.Must(err)
|
||||||
|
|
||||||
|
reader, _, err := ClientVerifyServerResponse(clientConn, method, rawKey, clientSalt)
|
||||||
|
common.Must(err)
|
||||||
|
|
||||||
|
// Send additional stream data
|
||||||
|
streamData := []byte("stream chunk test")
|
||||||
|
_ = writer.WriteChunk(streamData)
|
||||||
|
|
||||||
|
mb, err := reader.ReadMultiBuffer()
|
||||||
|
common.Must(err)
|
||||||
|
if !bytes.Equal(mb[0].Bytes(), streamData) {
|
||||||
|
t.Errorf("echoed stream data mismatch: got %s, want %s", mb[0].Bytes(), streamData)
|
||||||
|
}
|
||||||
|
buf.ReleaseMulti(mb)
|
||||||
|
}()
|
||||||
|
|
||||||
|
wg.Wait()
|
||||||
|
|
||||||
|
if receivedDest.NetAddr() != dest.NetAddr() {
|
||||||
|
t.Errorf("destination mismatch: got %s, want %s", receivedDest.NetAddr(), dest.NetAddr())
|
||||||
|
}
|
||||||
|
if diff := cmp.Diff(receivedPayload, testPayload); diff != "" {
|
||||||
|
t.Errorf("payload mismatch: %s", diff)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestUDPCodec(t *testing.T) {
|
||||||
|
methods := []string{
|
||||||
|
MethodAES128GCM,
|
||||||
|
MethodAES256GCM,
|
||||||
|
MethodChaCha20Poly1305,
|
||||||
|
}
|
||||||
|
|
||||||
|
dest := net.UDPDestination(net.LocalHostIP, net.Port(53))
|
||||||
|
payload := []byte("DNS query payload")
|
||||||
|
|
||||||
|
for _, methodName := range methods {
|
||||||
|
t.Run(methodName, func(t *testing.T) {
|
||||||
|
method, err := GetCipherMethod(methodName)
|
||||||
|
common.Must(err)
|
||||||
|
|
||||||
|
psk := make([]byte, method.KeySaltLength)
|
||||||
|
_, _ = rand.Read(psk)
|
||||||
|
|
||||||
|
clientCodec, err := NewUDPPacketCodec(method, psk)
|
||||||
|
common.Must(err)
|
||||||
|
serverCodec, err := NewUDPServerCodec(method, psk, time.Minute)
|
||||||
|
common.Must(err)
|
||||||
|
|
||||||
|
pktBuf, err := clientCodec.EncodeClientPacket(dest, payload)
|
||||||
|
common.Must(err)
|
||||||
|
defer pktBuf.Release()
|
||||||
|
|
||||||
|
rawCopy := make([]byte, pktBuf.Len())
|
||||||
|
copy(rawCopy, pktBuf.Bytes())
|
||||||
|
|
||||||
|
decoded, err := serverCodec.DecodePacket(pktBuf.Bytes())
|
||||||
|
common.Must(err)
|
||||||
|
|
||||||
|
if decoded.HeaderType != HeaderTypeClient {
|
||||||
|
t.Errorf("expected header type %d, got %d", HeaderTypeClient, decoded.HeaderType)
|
||||||
|
}
|
||||||
|
if decoded.Destination.Port != dest.Port {
|
||||||
|
t.Errorf("port mismatch: got %d, want %d", decoded.Destination.Port, dest.Port)
|
||||||
|
}
|
||||||
|
if !bytes.Equal(decoded.Payload, payload) {
|
||||||
|
t.Errorf("payload mismatch: got %s, want %s", decoded.Payload, payload)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Replay same packet wire bytes should fail with ErrPacketIdNotUnique
|
||||||
|
_, err = serverCodec.DecodePacket(rawCopy)
|
||||||
|
if err != ErrPacketIdNotUnique {
|
||||||
|
t.Fatalf("expected ErrPacketIdNotUnique on replay, got: %v", err)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMultiUserManager(t *testing.T) {
|
||||||
|
masterKey := generateRandomKey(16)
|
||||||
|
userKey1 := generateRandomKey(16)
|
||||||
|
userKey2 := generateRandomKey(16)
|
||||||
|
|
||||||
|
config := &MultiUserServerConfig{
|
||||||
|
Method: MethodAES128GCM,
|
||||||
|
Key: masterKey,
|
||||||
|
Users: []*protocol.User{
|
||||||
|
{
|
||||||
|
Email: "user1@example.com",
|
||||||
|
Account: serial.ToTypedMessage(&Account{Key: userKey1}),
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
inbound, err := NewMultiServer(newTestContext(), config)
|
||||||
|
common.Must(err)
|
||||||
|
|
||||||
|
if inbound.GetUsersCount(context.Background()) != 1 {
|
||||||
|
t.Fatalf("expected 1 user, got %d", inbound.GetUsersCount(context.Background()))
|
||||||
|
}
|
||||||
|
|
||||||
|
u1 := inbound.GetUser(context.Background(), "user1@example.com")
|
||||||
|
if u1 == nil || u1.Email != "user1@example.com" {
|
||||||
|
t.Fatal("user1 not found")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Add User 2
|
||||||
|
rawKey2, _ := base64.StdEncoding.DecodeString(userKey2)
|
||||||
|
u2 := &protocol.MemoryUser{
|
||||||
|
Email: "user2@example.com",
|
||||||
|
Account: &MemoryAccount{
|
||||||
|
Key: rawKey2,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
err = inbound.AddUser(context.Background(), u2)
|
||||||
|
common.Must(err)
|
||||||
|
|
||||||
|
if inbound.GetUsersCount(context.Background()) != 2 {
|
||||||
|
t.Fatalf("expected 2 users, got %d", inbound.GetUsersCount(context.Background()))
|
||||||
|
}
|
||||||
|
|
||||||
|
// Remove User 1
|
||||||
|
err = inbound.RemoveUser(context.Background(), "user1@example.com")
|
||||||
|
common.Must(err)
|
||||||
|
|
||||||
|
if inbound.GetUsersCount(context.Background()) != 1 {
|
||||||
|
t.Fatalf("expected 1 user, got %d", inbound.GetUsersCount(context.Background()))
|
||||||
|
}
|
||||||
|
if inbound.GetUser(context.Background(), "user1@example.com") != nil {
|
||||||
|
t.Fatal("user1 should have been removed")
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,529 @@
|
|||||||
|
package shadowsocks_2022
|
||||||
|
|
||||||
|
import (
|
||||||
|
"crypto/cipher"
|
||||||
|
"crypto/rand"
|
||||||
|
"encoding/binary"
|
||||||
|
"io"
|
||||||
|
"math"
|
||||||
|
mrand "math/rand/v2"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"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"
|
||||||
|
)
|
||||||
|
|
||||||
|
var addrParser = protocol.NewAddressParser(
|
||||||
|
protocol.AddressFamilyByte(0x01, net.AddressFamilyIPv4),
|
||||||
|
protocol.AddressFamilyByte(0x04, net.AddressFamilyIPv6),
|
||||||
|
protocol.AddressFamilyByte(0x03, net.AddressFamilyDomain),
|
||||||
|
protocol.WithAddressTypeParser(func(b byte) byte {
|
||||||
|
return b & 0x0F
|
||||||
|
}),
|
||||||
|
)
|
||||||
|
|
||||||
|
func IncreaseNonce(nonce []byte) {
|
||||||
|
for i := range nonce {
|
||||||
|
nonce[i]++
|
||||||
|
if nonce[i] != 0 {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// WriteAddressPort writes a destination address and port in SOCKS5 format
|
||||||
|
func WriteAddressPort(w io.Writer, dest net.Destination) error {
|
||||||
|
return addrParser.WriteAddressPort(w, dest.Address, dest.Port)
|
||||||
|
}
|
||||||
|
|
||||||
|
// ReadAddressPort reads a destination address and port in SOCKS5 format
|
||||||
|
func ReadAddressPort(r io.Reader) (net.Destination, error) {
|
||||||
|
addr, port, err := addrParser.ReadAddressPort(nil, r)
|
||||||
|
if err != nil {
|
||||||
|
return net.Destination{}, err
|
||||||
|
}
|
||||||
|
return net.TCPDestination(addr, port), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// AddrPortLength returns the serialized length of a destination in SOCKS5 format
|
||||||
|
func AddrPortLength(dest net.Destination) int {
|
||||||
|
switch dest.Address.Family() {
|
||||||
|
case net.AddressFamilyIPv4:
|
||||||
|
return 1 + 4 + 2
|
||||||
|
case net.AddressFamilyDomain:
|
||||||
|
return 1 + 1 + len(dest.Address.Domain()) + 2
|
||||||
|
case net.AddressFamilyIPv6:
|
||||||
|
return 1 + 16 + 2
|
||||||
|
default:
|
||||||
|
return 0
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
type StreamWriter struct {
|
||||||
|
writer io.Writer
|
||||||
|
cipher cipher.AEAD
|
||||||
|
nonce [StreamNonceSize]byte
|
||||||
|
lenBuf [2]byte
|
||||||
|
buf []byte
|
||||||
|
}
|
||||||
|
|
||||||
|
func NewStreamWriter(w io.Writer, c cipher.AEAD) *StreamWriter {
|
||||||
|
return &StreamWriter{
|
||||||
|
writer: w,
|
||||||
|
cipher: c,
|
||||||
|
buf: make([]byte, 0, MaxPacketSize+2+2*AEADTagSize),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (w *StreamWriter) Nonce() []byte {
|
||||||
|
return w.nonce[:]
|
||||||
|
}
|
||||||
|
|
||||||
|
func (w *StreamWriter) WriteChunk(payload []byte) error {
|
||||||
|
payloadLen := len(payload)
|
||||||
|
if payloadLen == 0 {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
if payloadLen > MaxPacketSize {
|
||||||
|
return errors.New("payload exceeds MaxPacketSize")
|
||||||
|
}
|
||||||
|
|
||||||
|
binary.BigEndian.PutUint16(w.lenBuf[:], uint16(payloadLen))
|
||||||
|
w.buf = w.cipher.Seal(w.buf[:0], w.nonce[:], w.lenBuf[:], nil)
|
||||||
|
IncreaseNonce(w.nonce[:])
|
||||||
|
|
||||||
|
w.buf = w.cipher.Seal(w.buf, w.nonce[:], payload, nil)
|
||||||
|
IncreaseNonce(w.nonce[:])
|
||||||
|
|
||||||
|
_, err := w.writer.Write(w.buf)
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
func (w *StreamWriter) Write(p []byte) (int, error) {
|
||||||
|
n := len(p)
|
||||||
|
for len(p) > 0 {
|
||||||
|
chunkSize := len(p)
|
||||||
|
if chunkSize > MaxPacketSize {
|
||||||
|
chunkSize = MaxPacketSize
|
||||||
|
}
|
||||||
|
if err := w.WriteChunk(p[:chunkSize]); err != nil {
|
||||||
|
return 0, err
|
||||||
|
}
|
||||||
|
p = p[chunkSize:]
|
||||||
|
}
|
||||||
|
return n, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (w *StreamWriter) WriteMultiBuffer(mb buf.MultiBuffer) error {
|
||||||
|
defer buf.ReleaseMulti(mb)
|
||||||
|
for _, b := range mb {
|
||||||
|
if err := w.WriteChunk(b.Bytes()); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
type StreamReader struct {
|
||||||
|
reader io.Reader
|
||||||
|
cipher cipher.AEAD
|
||||||
|
nonce [StreamNonceSize]byte
|
||||||
|
lenBuf [2 + AEADTagSize]byte
|
||||||
|
buffer []byte
|
||||||
|
cached int
|
||||||
|
offset int
|
||||||
|
}
|
||||||
|
|
||||||
|
func NewStreamReader(r io.Reader, c cipher.AEAD) *StreamReader {
|
||||||
|
return &StreamReader{
|
||||||
|
reader: r,
|
||||||
|
cipher: c,
|
||||||
|
buffer: make([]byte, MaxPacketSize+AEADTagSize),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *StreamReader) Nonce() []byte {
|
||||||
|
return r.nonce[:]
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *StreamReader) Read(p []byte) (int, error) {
|
||||||
|
if r.cached > 0 {
|
||||||
|
n := copy(p, r.buffer[r.offset:r.offset+r.cached])
|
||||||
|
r.cached -= n
|
||||||
|
r.offset += n
|
||||||
|
return n, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Read 2-byte length + AEAD tag (18 bytes)
|
||||||
|
if _, err := io.ReadFull(r.reader, r.lenBuf[:]); err != nil {
|
||||||
|
return 0, err
|
||||||
|
}
|
||||||
|
|
||||||
|
decryptedLen, err := r.cipher.Open(r.lenBuf[:0], r.nonce[:], r.lenBuf[:], nil)
|
||||||
|
if err != nil {
|
||||||
|
return 0, errors.New("failed to decrypt chunk length").Base(err)
|
||||||
|
}
|
||||||
|
IncreaseNonce(r.nonce[:])
|
||||||
|
|
||||||
|
payloadLen := int(binary.BigEndian.Uint16(decryptedLen))
|
||||||
|
if payloadLen == 0 {
|
||||||
|
return 0, ErrInvalidRequest
|
||||||
|
}
|
||||||
|
|
||||||
|
chunkEnd := payloadLen + AEADTagSize
|
||||||
|
if _, err := io.ReadFull(r.reader, r.buffer[:chunkEnd]); err != nil {
|
||||||
|
return 0, err
|
||||||
|
}
|
||||||
|
|
||||||
|
decryptedPayload, err := r.cipher.Open(r.buffer[:0], r.nonce[:], r.buffer[:chunkEnd], nil)
|
||||||
|
if err != nil {
|
||||||
|
return 0, errors.New("failed to decrypt chunk payload").Base(err)
|
||||||
|
}
|
||||||
|
IncreaseNonce(r.nonce[:])
|
||||||
|
|
||||||
|
r.cached = len(decryptedPayload)
|
||||||
|
r.offset = 0
|
||||||
|
|
||||||
|
n := copy(p, r.buffer[r.offset:r.offset+r.cached])
|
||||||
|
r.cached -= n
|
||||||
|
r.offset += n
|
||||||
|
return n, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *StreamReader) ReadMultiBuffer() (buf.MultiBuffer, error) {
|
||||||
|
if r.cached > 0 {
|
||||||
|
b := buf.New()
|
||||||
|
b.Write(r.buffer[r.offset : r.offset+r.cached])
|
||||||
|
r.cached = 0
|
||||||
|
r.offset = 0
|
||||||
|
return buf.MultiBuffer{b}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
if _, err := io.ReadFull(r.reader, r.lenBuf[:]); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
decryptedLen, err := r.cipher.Open(r.lenBuf[:0], r.nonce[:], r.lenBuf[:], nil)
|
||||||
|
if err != nil {
|
||||||
|
return nil, errors.New("failed to decrypt chunk length").Base(err)
|
||||||
|
}
|
||||||
|
IncreaseNonce(r.nonce[:])
|
||||||
|
|
||||||
|
payloadLen := int(binary.BigEndian.Uint16(decryptedLen))
|
||||||
|
if payloadLen == 0 {
|
||||||
|
return nil, ErrInvalidRequest
|
||||||
|
}
|
||||||
|
|
||||||
|
chunkEnd := payloadLen + AEADTagSize
|
||||||
|
if _, err := io.ReadFull(r.reader, r.buffer[:chunkEnd]); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
decryptedPayload, err := r.cipher.Open(r.buffer[:0], r.nonce[:], r.buffer[:chunkEnd], nil)
|
||||||
|
if err != nil {
|
||||||
|
return nil, errors.New("failed to decrypt chunk payload").Base(err)
|
||||||
|
}
|
||||||
|
IncreaseNonce(r.nonce[:])
|
||||||
|
|
||||||
|
b := buf.New()
|
||||||
|
b.Write(decryptedPayload)
|
||||||
|
return buf.MultiBuffer{b}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
type ClientRequestHeader struct {
|
||||||
|
Destination net.Destination
|
||||||
|
EarlyData []byte
|
||||||
|
}
|
||||||
|
|
||||||
|
func ReadClientRequestHeader(conn io.Reader, reader *StreamReader) (*ClientRequestHeader, error) {
|
||||||
|
var fixedBuf [RequestHeaderFixedChunkLength + AEADTagSize]byte
|
||||||
|
if _, err := io.ReadFull(conn, fixedBuf[:]); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
plainFixed, err := reader.cipher.Open(fixedBuf[:0], reader.Nonce(), fixedBuf[:], nil)
|
||||||
|
if err != nil {
|
||||||
|
return nil, errors.New("failed to decrypt client request header").Base(err)
|
||||||
|
}
|
||||||
|
IncreaseNonce(reader.Nonce())
|
||||||
|
|
||||||
|
if plainFixed[0] != HeaderTypeClient {
|
||||||
|
return nil, ErrBadHeaderType
|
||||||
|
}
|
||||||
|
|
||||||
|
epoch := binary.BigEndian.Uint64(plainFixed[1:9])
|
||||||
|
diff := int(math.Abs(float64(time.Now().Unix() - int64(epoch))))
|
||||||
|
if diff > 30 {
|
||||||
|
return nil, ErrBadTimestamp
|
||||||
|
}
|
||||||
|
|
||||||
|
varHeaderLen := int(binary.BigEndian.Uint16(plainFixed[9:11]))
|
||||||
|
if varHeaderLen == 0 {
|
||||||
|
return nil, ErrInvalidRequest
|
||||||
|
}
|
||||||
|
|
||||||
|
var stackVarChunk [512]byte
|
||||||
|
var varChunkCipher []byte
|
||||||
|
needed := varHeaderLen + AEADTagSize
|
||||||
|
if needed <= len(stackVarChunk) {
|
||||||
|
varChunkCipher = stackVarChunk[:needed]
|
||||||
|
} else {
|
||||||
|
varChunkCipher = make([]byte, needed)
|
||||||
|
}
|
||||||
|
if _, err := io.ReadFull(conn, varChunkCipher); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
plainVar, err := reader.cipher.Open(varChunkCipher[:0], reader.Nonce(), varChunkCipher, nil)
|
||||||
|
if err != nil {
|
||||||
|
return nil, errors.New("failed to decrypt variable request header").Base(err)
|
||||||
|
}
|
||||||
|
IncreaseNonce(reader.Nonce())
|
||||||
|
|
||||||
|
b := buf.New()
|
||||||
|
b.Write(plainVar)
|
||||||
|
defer b.Release()
|
||||||
|
|
||||||
|
dest, err := ReadAddressPort(b)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
var padLenBytes [2]byte
|
||||||
|
if _, err := b.Read(padLenBytes[:]); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
paddingLen := int(binary.BigEndian.Uint16(padLenBytes[:]))
|
||||||
|
if int(b.Len()) < paddingLen {
|
||||||
|
return nil, ErrNoPadding
|
||||||
|
}
|
||||||
|
if paddingLen > 0 {
|
||||||
|
b.Advance(int32(paddingLen))
|
||||||
|
}
|
||||||
|
|
||||||
|
var earlyData []byte
|
||||||
|
if b.Len() > 0 {
|
||||||
|
earlyData = make([]byte, b.Len())
|
||||||
|
copy(earlyData, b.Bytes())
|
||||||
|
}
|
||||||
|
|
||||||
|
return &ClientRequestHeader{
|
||||||
|
Destination: dest,
|
||||||
|
EarlyData: earlyData,
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// ClientHandshake writes the full client request header to w
|
||||||
|
func ClientHandshake(w io.Writer, method *CipherMethod, pskList [][]byte, dest net.Destination, payload []byte) ([]byte, *StreamWriter, error) {
|
||||||
|
salt := make([]byte, method.KeySaltLength)
|
||||||
|
if _, err := io.ReadFull(rand.Reader, salt); err != nil {
|
||||||
|
return nil, nil, err
|
||||||
|
}
|
||||||
|
writer, err := WriteTCPRequest(w, method, pskList, dest, salt, payload)
|
||||||
|
if err != nil {
|
||||||
|
return nil, nil, err
|
||||||
|
}
|
||||||
|
return salt, writer.(*StreamWriter), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// ClientVerifyServerResponse reads and verifies the server's handshake response
|
||||||
|
func ClientVerifyServerResponse(r io.Reader, method *CipherMethod, psk []byte, clientSalt []byte) (*StreamReader, []byte, error) {
|
||||||
|
reader, err := ReadTCPResponse(r, method, psk, clientSalt)
|
||||||
|
if err != nil {
|
||||||
|
return nil, nil, err
|
||||||
|
}
|
||||||
|
sr := reader.(*StreamReader)
|
||||||
|
var initialPayload []byte
|
||||||
|
if sr.cached > 0 {
|
||||||
|
initialPayload = make([]byte, sr.cached)
|
||||||
|
copy(initialPayload, sr.buffer[sr.offset:sr.offset+sr.cached])
|
||||||
|
}
|
||||||
|
return sr, initialPayload, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// WriteTCPRequest writes the Shadowsocks 2022 request header into w and returns a body writer.
|
||||||
|
func WriteTCPRequest(w io.Writer, method *CipherMethod, pskList [][]byte, dest net.Destination, clientSalt []byte, payload []byte) (buf.Writer, error) {
|
||||||
|
finalPSK := pskList[len(pskList)-1]
|
||||||
|
sessionKey := DeriveSessionSubKey(finalPSK, clientSalt, method.KeySaltLength)
|
||||||
|
aead, err := method.NewAEAD(sessionKey)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
writer := NewStreamWriter(w, aead)
|
||||||
|
|
||||||
|
handshakeBuf := buf.New()
|
||||||
|
defer handshakeBuf.Release()
|
||||||
|
|
||||||
|
handshakeBuf.Write(clientSalt)
|
||||||
|
|
||||||
|
for i, currPSK := range pskList[:len(pskList)-1] {
|
||||||
|
identitySubkey := DeriveIdentitySubKey(currPSK, clientSalt, method.KeySaltLength)
|
||||||
|
block, err := method.NewBlock(identitySubkey)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
nextPSK := pskList[i+1]
|
||||||
|
pskHash := DeriveUserPSKHash(nextPSK)
|
||||||
|
var encryptedEIH [AESBlockSize]byte
|
||||||
|
block.Encrypt(encryptedEIH[:], pskHash[:])
|
||||||
|
handshakeBuf.Write(encryptedEIH[:])
|
||||||
|
}
|
||||||
|
|
||||||
|
payloadLen := len(payload)
|
||||||
|
var paddingLen int
|
||||||
|
if payloadLen < MaxPaddingLength {
|
||||||
|
paddingLen = mrand.IntN(MaxPaddingLength-payloadLen) + 1
|
||||||
|
}
|
||||||
|
addrPortLen := AddrPortLength(dest)
|
||||||
|
varHeaderLen := addrPortLen + 2 + paddingLen + payloadLen
|
||||||
|
|
||||||
|
var fixedHeaderPlaintext [RequestHeaderFixedChunkLength]byte
|
||||||
|
fixedHeaderPlaintext[0] = HeaderTypeClient
|
||||||
|
binary.BigEndian.PutUint64(fixedHeaderPlaintext[1:9], uint64(time.Now().Unix()))
|
||||||
|
binary.BigEndian.PutUint16(fixedHeaderPlaintext[9:11], uint16(varHeaderLen))
|
||||||
|
|
||||||
|
fixedChunk := writer.cipher.Seal(nil, writer.nonce[:], fixedHeaderPlaintext[:], nil)
|
||||||
|
IncreaseNonce(writer.nonce[:])
|
||||||
|
handshakeBuf.Write(fixedChunk)
|
||||||
|
|
||||||
|
varHeaderBuf := buf.New()
|
||||||
|
defer varHeaderBuf.Release()
|
||||||
|
|
||||||
|
if err := WriteAddressPort(varHeaderBuf, dest); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
var padLenBytes [2]byte
|
||||||
|
binary.BigEndian.PutUint16(padLenBytes[:], uint16(paddingLen))
|
||||||
|
varHeaderBuf.Write(padLenBytes[:])
|
||||||
|
|
||||||
|
if paddingLen > 0 {
|
||||||
|
varHeaderBuf.Write(zeroPadding[:paddingLen])
|
||||||
|
}
|
||||||
|
|
||||||
|
if payloadLen > 0 {
|
||||||
|
varHeaderBuf.Write(payload)
|
||||||
|
}
|
||||||
|
|
||||||
|
varChunk := writer.cipher.Seal(nil, writer.nonce[:], varHeaderBuf.Bytes(), nil)
|
||||||
|
IncreaseNonce(writer.nonce[:])
|
||||||
|
handshakeBuf.Write(varChunk)
|
||||||
|
|
||||||
|
if _, err := w.Write(handshakeBuf.Bytes()); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
return writer, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// ReadTCPResponse reads and verifies the server's handshake response and returns a reader for the stream.
|
||||||
|
func ReadTCPResponse(r io.Reader, method *CipherMethod, psk []byte, clientSalt []byte) (buf.Reader, error) {
|
||||||
|
var serverSalt [32]byte
|
||||||
|
serverSaltSlice := serverSalt[:method.KeySaltLength]
|
||||||
|
if _, err := io.ReadFull(r, serverSaltSlice); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
sessionKey := DeriveSessionSubKey(psk, serverSaltSlice, method.KeySaltLength)
|
||||||
|
aead, err := method.NewAEAD(sessionKey)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
reader := NewStreamReader(r, aead)
|
||||||
|
|
||||||
|
fixedPlainLen := 1 + 8 + method.KeySaltLength + 2
|
||||||
|
chunkCipherLen := fixedPlainLen + AEADTagSize
|
||||||
|
var chunkBuf [64]byte
|
||||||
|
chunkSlice := chunkBuf[:chunkCipherLen]
|
||||||
|
if _, err := io.ReadFull(r, chunkSlice); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
decryptedFixed, err := reader.cipher.Open(chunkSlice[:0], reader.nonce[:], chunkSlice, nil)
|
||||||
|
if err != nil {
|
||||||
|
return nil, errors.New("failed to decrypt server response header").Base(err)
|
||||||
|
}
|
||||||
|
IncreaseNonce(reader.nonce[:])
|
||||||
|
|
||||||
|
if decryptedFixed[0] != HeaderTypeServer {
|
||||||
|
return nil, ErrBadHeaderType
|
||||||
|
}
|
||||||
|
|
||||||
|
serverEpoch := binary.BigEndian.Uint64(decryptedFixed[1:9])
|
||||||
|
diff := int(math.Abs(float64(time.Now().Unix() - int64(serverEpoch))))
|
||||||
|
if diff > 30 {
|
||||||
|
return nil, ErrBadTimestamp
|
||||||
|
}
|
||||||
|
|
||||||
|
echoedSalt := decryptedFixed[9 : 9+method.KeySaltLength]
|
||||||
|
for i := 0; i < method.KeySaltLength; i++ {
|
||||||
|
if echoedSalt[i] != clientSalt[i] {
|
||||||
|
return nil, errors.New("bad request salt")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
initialPayloadLen := int(binary.BigEndian.Uint16(decryptedFixed[9+method.KeySaltLength : 11+method.KeySaltLength]))
|
||||||
|
if initialPayloadLen > 0 {
|
||||||
|
initialCipherLen := initialPayloadLen + AEADTagSize
|
||||||
|
if _, err := io.ReadFull(r, reader.buffer[:initialCipherLen]); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
decryptedInitial, err := reader.cipher.Open(reader.buffer[:0], reader.nonce[:], reader.buffer[:initialCipherLen], nil)
|
||||||
|
if err != nil {
|
||||||
|
return nil, errors.New("failed to decrypt initial response payload").Base(err)
|
||||||
|
}
|
||||||
|
IncreaseNonce(reader.nonce[:])
|
||||||
|
reader.cached = len(decryptedInitial)
|
||||||
|
reader.offset = 0
|
||||||
|
}
|
||||||
|
|
||||||
|
return reader, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// WriteTCPResponse writes the server handshake response and returns a body writer for server stream.
|
||||||
|
func WriteTCPResponse(w io.Writer, method *CipherMethod, psk []byte, clientSalt []byte, initialPayload []byte) (buf.Writer, error) {
|
||||||
|
var serverSalt [32]byte
|
||||||
|
serverSaltSlice := serverSalt[:method.KeySaltLength]
|
||||||
|
if _, err := io.ReadFull(rand.Reader, serverSaltSlice); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
respKey := DeriveSessionSubKey(psk, serverSaltSlice, method.KeySaltLength)
|
||||||
|
respAead, err := method.NewAEAD(respKey)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
writer := NewStreamWriter(w, respAead)
|
||||||
|
|
||||||
|
respBuf := buf.New()
|
||||||
|
defer respBuf.Release()
|
||||||
|
|
||||||
|
respBuf.Write(serverSaltSlice)
|
||||||
|
|
||||||
|
var fixedRespPlain [1 + 8 + 32 + 2]byte
|
||||||
|
fixedRespSlice := fixedRespPlain[:1+8+method.KeySaltLength+2]
|
||||||
|
fixedRespSlice[0] = HeaderTypeServer
|
||||||
|
binary.BigEndian.PutUint64(fixedRespSlice[1:9], uint64(time.Now().Unix()))
|
||||||
|
copy(fixedRespSlice[9:9+method.KeySaltLength], clientSalt)
|
||||||
|
binary.BigEndian.PutUint16(fixedRespSlice[9+method.KeySaltLength:11+method.KeySaltLength], uint16(len(initialPayload)))
|
||||||
|
|
||||||
|
fixedRespChunk := writer.cipher.Seal(nil, writer.nonce[:], fixedRespSlice, nil)
|
||||||
|
IncreaseNonce(writer.nonce[:])
|
||||||
|
respBuf.Write(fixedRespChunk)
|
||||||
|
|
||||||
|
if len(initialPayload) > 0 {
|
||||||
|
initialChunk := writer.cipher.Seal(nil, writer.nonce[:], initialPayload, nil)
|
||||||
|
IncreaseNonce(writer.nonce[:])
|
||||||
|
respBuf.Write(initialChunk)
|
||||||
|
}
|
||||||
|
|
||||||
|
if _, err := w.Write(respBuf.Bytes()); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
return writer, nil
|
||||||
|
}
|
||||||
@@ -105,7 +105,7 @@ func (c *Client) Process(ctx context.Context, link *transport.Link, dialer inter
|
|||||||
}
|
}
|
||||||
udpRequest, err := ClientHandshake(request, conn, conn)
|
udpRequest, err := ClientHandshake(request, conn, conn)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return errors.New("failed to establish connection to server").AtWarning().Base(err)
|
return errors.New("failed to establish connection to server").Base(err)
|
||||||
}
|
}
|
||||||
if udpRequest != nil {
|
if udpRequest != nil {
|
||||||
if udpRequest.Address == net.AnyIP || udpRequest.Address == net.AnyIPv6 {
|
if udpRequest.Address == net.AnyIP || udpRequest.Address == net.AnyIPv6 {
|
||||||
|
|||||||
@@ -458,10 +458,10 @@ func ClientHandshake(request *protocol.RequestHeader, reader io.Reader, writer i
|
|||||||
}
|
}
|
||||||
|
|
||||||
if b.Byte(0) != socks5Version {
|
if b.Byte(0) != socks5Version {
|
||||||
return nil, errors.New("unexpected server version: ", b.Byte(0)).AtWarning()
|
return nil, errors.New("unexpected server version: ", b.Byte(0))
|
||||||
}
|
}
|
||||||
if b.Byte(1) != authByte {
|
if b.Byte(1) != authByte {
|
||||||
return nil, errors.New("auth method not supported.").AtWarning()
|
return nil, errors.New("auth method not supported.")
|
||||||
}
|
}
|
||||||
|
|
||||||
if authByte == authPassword {
|
if authByte == authPassword {
|
||||||
|
|||||||
@@ -69,7 +69,7 @@ func (c *Client) Process(ctx context.Context, link *transport.Link, dialer inter
|
|||||||
return nil
|
return nil
|
||||||
})
|
})
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return errors.New("failed to find an available destination").AtWarning().Base(err)
|
return errors.New("failed to find an available destination").Base(err)
|
||||||
}
|
}
|
||||||
errors.LogInfo(ctx, "tunneling request to ", destination, " via ", server.Destination.NetAddr())
|
errors.LogInfo(ctx, "tunneling request to ", destination, " via ", server.Destination.NetAddr())
|
||||||
|
|
||||||
@@ -116,21 +116,21 @@ func (c *Client) Process(ctx context.Context, link *transport.Link, dialer inter
|
|||||||
|
|
||||||
// write some request payload to buffer
|
// write some request payload to buffer
|
||||||
if err = buf.CopyOnceTimeout(link.Reader, bodyWriter, time.Millisecond*100); err != nil && err != buf.ErrNotTimeoutReader && err != buf.ErrReadTimeout {
|
if err = buf.CopyOnceTimeout(link.Reader, bodyWriter, time.Millisecond*100); err != nil && err != buf.ErrNotTimeoutReader && err != buf.ErrReadTimeout {
|
||||||
return errors.New("failed to write A request payload").Base(err).AtWarning()
|
return errors.New("failed to write A request payload").Base(err)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Flush; bufferWriter.WriteMultiBuffer now is bufferWriter.writer.WriteMultiBuffer
|
// Flush; bufferWriter.WriteMultiBuffer now is bufferWriter.writer.WriteMultiBuffer
|
||||||
if err = bufferWriter.SetBuffered(false); err != nil {
|
if err = bufferWriter.SetBuffered(false); err != nil {
|
||||||
return errors.New("failed to flush payload").Base(err).AtWarning()
|
return errors.New("failed to flush payload").Base(err)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Send header if not sent yet
|
// Send header if not sent yet
|
||||||
if _, err = connWriter.Write([]byte{}); err != nil {
|
if _, err = connWriter.Write([]byte{}); err != nil {
|
||||||
return err.(*errors.Error).AtWarning()
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
if err = buf.Copy(link.Reader, bodyWriter, buf.UpdateActivity(timer)); err != nil {
|
if err = buf.Copy(link.Reader, bodyWriter, buf.UpdateActivity(timer)); err != nil {
|
||||||
return errors.New("failed to transfer request payload").Base(err).AtInfo()
|
return errors.New("failed to transfer request payload").Base(err)
|
||||||
}
|
}
|
||||||
|
|
||||||
return nil
|
return nil
|
||||||
|
|||||||
+12
-12
@@ -47,11 +47,11 @@ func NewServer(ctx context.Context, config *ServerConfig) (*Server, error) {
|
|||||||
for _, user := range config.Users {
|
for _, user := range config.Users {
|
||||||
u, err := user.ToMemoryUser()
|
u, err := user.ToMemoryUser()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, errors.New("failed to get trojan user").Base(err).AtError()
|
return nil, errors.New("failed to get trojan user").Base(err)
|
||||||
}
|
}
|
||||||
|
|
||||||
if err := validator.Add(u); err != nil {
|
if err := validator.Add(u); err != nil {
|
||||||
return nil, errors.New("failed to add user").Base(err).AtError()
|
return nil, errors.New("failed to add user").Base(err)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -151,7 +151,7 @@ func (s *Server) Process(ctx context.Context, network net.Network, conn stat.Con
|
|||||||
|
|
||||||
sessionPolicy := s.policyManager.ForLevel(0)
|
sessionPolicy := s.policyManager.ForLevel(0)
|
||||||
if err := conn.SetReadDeadline(time.Now().Add(sessionPolicy.Timeouts.Handshake)); err != nil {
|
if err := conn.SetReadDeadline(time.Now().Add(sessionPolicy.Timeouts.Handshake)); err != nil {
|
||||||
return errors.New("unable to set read deadline").Base(err).AtWarning()
|
return errors.New("unable to set read deadline").Base(err)
|
||||||
}
|
}
|
||||||
|
|
||||||
first := buf.FromBytes(make([]byte, buf.Size))
|
first := buf.FromBytes(make([]byte, buf.Size))
|
||||||
@@ -219,7 +219,7 @@ func (s *Server) Process(ctx context.Context, network net.Network, conn stat.Con
|
|||||||
|
|
||||||
destination := clientReader.Target
|
destination := clientReader.Target
|
||||||
if err := conn.SetReadDeadline(time.Time{}); err != nil {
|
if err := conn.SetReadDeadline(time.Time{}); err != nil {
|
||||||
return errors.New("unable to set read deadline").Base(err).AtWarning()
|
return errors.New("unable to set read deadline").Base(err)
|
||||||
}
|
}
|
||||||
|
|
||||||
inbound := session.InboundFromContext(ctx)
|
inbound := session.InboundFromContext(ctx)
|
||||||
@@ -402,7 +402,7 @@ func (s *Server) fallback(ctx context.Context, err error, sessionPolicy policy.S
|
|||||||
}
|
}
|
||||||
apfb := napfb[name]
|
apfb := napfb[name]
|
||||||
if apfb == nil {
|
if apfb == nil {
|
||||||
return errors.New(`failed to find the default "name" config`).AtWarning()
|
return errors.New(`failed to find the default "name" config`)
|
||||||
}
|
}
|
||||||
|
|
||||||
if apfb[alpn] == nil {
|
if apfb[alpn] == nil {
|
||||||
@@ -410,7 +410,7 @@ func (s *Server) fallback(ctx context.Context, err error, sessionPolicy policy.S
|
|||||||
}
|
}
|
||||||
pfb := apfb[alpn]
|
pfb := apfb[alpn]
|
||||||
if pfb == nil {
|
if pfb == nil {
|
||||||
return errors.New(`failed to find the default "alpn" config`).AtWarning()
|
return errors.New(`failed to find the default "alpn" config`)
|
||||||
}
|
}
|
||||||
|
|
||||||
path := ""
|
path := ""
|
||||||
@@ -444,7 +444,7 @@ func (s *Server) fallback(ctx context.Context, err error, sessionPolicy policy.S
|
|||||||
}
|
}
|
||||||
fb := pfb[path]
|
fb := pfb[path]
|
||||||
if fb == nil {
|
if fb == nil {
|
||||||
return errors.New(`failed to find the default "path" config`).AtWarning()
|
return errors.New(`failed to find the default "path" config`)
|
||||||
}
|
}
|
||||||
|
|
||||||
ctx, cancel := context.WithCancel(ctx)
|
ctx, cancel := context.WithCancel(ctx)
|
||||||
@@ -460,7 +460,7 @@ func (s *Server) fallback(ctx context.Context, err error, sessionPolicy policy.S
|
|||||||
}
|
}
|
||||||
return nil
|
return nil
|
||||||
}); err != nil {
|
}); err != nil {
|
||||||
return errors.New("failed to dial to " + fb.Dest).Base(err).AtWarning()
|
return errors.New("failed to dial to " + fb.Dest).Base(err)
|
||||||
}
|
}
|
||||||
defer conn.Close()
|
defer conn.Close()
|
||||||
|
|
||||||
@@ -520,11 +520,11 @@ func (s *Server) fallback(ctx context.Context, err error, sessionPolicy policy.S
|
|||||||
common.Must2(pro.Write([]byte{byte(p1 >> 8), byte(p1), byte(p2 >> 8), byte(p2)}))
|
common.Must2(pro.Write([]byte{byte(p1 >> 8), byte(p1), byte(p2 >> 8), byte(p2)}))
|
||||||
}
|
}
|
||||||
if err := serverWriter.WriteMultiBuffer(buf.MultiBuffer{pro}); err != nil {
|
if err := serverWriter.WriteMultiBuffer(buf.MultiBuffer{pro}); err != nil {
|
||||||
return errors.New("failed to set PROXY protocol v", fb.Xver).Base(err).AtWarning()
|
return errors.New("failed to set PROXY protocol v", fb.Xver).Base(err)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
if err := buf.Copy(reader, serverWriter, buf.UpdateActivity(timer)); err != nil {
|
if err := buf.Copy(reader, serverWriter, buf.UpdateActivity(timer)); err != nil {
|
||||||
return errors.New("failed to fallback request payload").Base(err).AtInfo()
|
return errors.New("failed to fallback request payload").Base(err)
|
||||||
}
|
}
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
@@ -534,7 +534,7 @@ func (s *Server) fallback(ctx context.Context, err error, sessionPolicy policy.S
|
|||||||
getResponse := func() error {
|
getResponse := func() error {
|
||||||
defer timer.SetTimeout(sessionPolicy.Timeouts.UplinkOnly)
|
defer timer.SetTimeout(sessionPolicy.Timeouts.UplinkOnly)
|
||||||
if err := buf.Copy(serverReader, writer, buf.UpdateActivity(timer)); err != nil {
|
if err := buf.Copy(serverReader, writer, buf.UpdateActivity(timer)); err != nil {
|
||||||
return errors.New("failed to deliver response payload").Base(err).AtInfo()
|
return errors.New("failed to deliver response payload").Base(err)
|
||||||
}
|
}
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
@@ -542,7 +542,7 @@ func (s *Server) fallback(ctx context.Context, err error, sessionPolicy policy.S
|
|||||||
if err := task.Run(ctx, task.OnSuccess(postRequest, task.Close(serverWriter)), task.OnSuccess(getResponse, task.Close(writer))); err != nil {
|
if err := task.Run(ctx, task.OnSuccess(postRequest, task.Close(serverWriter)), task.OnSuccess(getResponse, task.Close(writer))); err != nil {
|
||||||
common.Must(common.Interrupt(serverReader))
|
common.Must(common.Interrupt(serverReader))
|
||||||
common.Must(common.Interrupt(serverWriter))
|
common.Must(common.Interrupt(serverWriter))
|
||||||
return errors.New("fallback ends").Base(err).AtInfo()
|
return errors.New("fallback ends").Base(err)
|
||||||
}
|
}
|
||||||
|
|
||||||
return nil
|
return nil
|
||||||
|
|||||||
+44
-1
@@ -17,11 +17,54 @@ Plainly enabling it in the config probably will result nothing, or lock your rou
|
|||||||
By default, enabling the feature will only bring the tun interface up. \
|
By default, enabling the feature will only bring the tun interface up. \
|
||||||
When configured explicitly, Windows and Linux can apply interface addresses from `gateway`, while macOS uses the first IPv4 prefix from `gateway` to configure the utun point-to-point address. \
|
When configured explicitly, Windows and Linux can apply interface addresses from `gateway`, while macOS uses the first IPv4 prefix from `gateway` to configure the utun point-to-point address. \
|
||||||
Windows, Linux and macOS can also apply system routes from `autoSystemRoutingTable`.
|
Windows, Linux and macOS can also apply system routes from `autoSystemRoutingTable`.
|
||||||
Linux and macOS do not configure system DNS from the `dns` field; system DNS remains managed by the OS or distribution-specific network services. \
|
macOS does not configure system DNS from the `dns` field, and neither does Linux by default; system DNS remains managed by the OS or distribution-specific network services. \
|
||||||
For more advanced routing policies or rules, OS level configuration can still manage the named interface (e.g. xray0) when it appears.
|
For more advanced routing policies or rules, OS level configuration can still manage the named interface (e.g. xray0) when it appears.
|
||||||
This keeps complex system level routing and rules in a single place of responsibility - the OS itself. \
|
This keeps complex system level routing and rules in a single place of responsibility - the OS itself. \
|
||||||
Examples of how to achieve this on a simple Linux system (Ubuntu with systemd-networkd) can be found at the end of this README.
|
Examples of how to achieve this on a simple Linux system (Ubuntu with systemd-networkd) can be found at the end of this README.
|
||||||
|
|
||||||
|
### SYSTEM DNS ON LINUX (`autoSystemDNS`)
|
||||||
|
|
||||||
|
On Linux, setting `autoSystemDNS` to `true` lets the inbound point the system resolver at the tun interface, so name lookups resolve through Xray instead of going out over the physical link. It is off by default, and it is Linux-only.
|
||||||
|
|
||||||
|
It uses `resolvectl`, which means it applies only when all of these hold:
|
||||||
|
|
||||||
|
- the system runs systemd and `resolvectl` is on `PATH`
|
||||||
|
- `systemd-resolved` is enabled and actually managing DNS (installed but not running has no effect)
|
||||||
|
- systemd-resolved is version 240 or newer, where `default-route` exists
|
||||||
|
- no `dns` upstream resolves through the system resolver, directly or through its own bootstrap (see below)
|
||||||
|
|
||||||
|
The address handed over is the first IPv4 `gateway` incremented by one (e.g. `192.168.100.1/30` -> `192.168.100.2`). It is not taken from `dns`: handing `1.1.1.1` to `resolvectl dns` would make systemd-resolved query that server directly over the physical link, which is the leak this option exists to close.
|
||||||
|
|
||||||
|
Because that address has to actually answer, the takeover is checked before it happens. A query from the interface address to that address is routed through the configured rules, and host-wide DNS is only changed when the result is a DNS-capable outbound. Otherwise the option does nothing and DNS is left to the OS. In practice this means you also need a routing rule sending the interface's port 53 to a `dns` outbound, for example:
|
||||||
|
|
||||||
|
```json
|
||||||
|
"routing": {
|
||||||
|
"rules": [
|
||||||
|
{ "type": "field", "inboundTag": ["tun"], "port": 53, "outboundTag": "dns" }
|
||||||
|
]
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
The check is a preflight, not a proof for arbitrary rules. It sends its query from the interface address and from a representative ephemeral source port, so a rule that matches on the source port cannot be predicted ahead of time: if the interface's port 53 reaches the `dns` outbound only from some source ports, the takeover is accepted and queries from the other ports fail. Supported configurations are those where the DNS path does not depend on the source port, that is, where the interface's port 53 reaches a `dns` outbound whatever its source.
|
||||||
|
|
||||||
|
It is also a check for the dependencies it knows about, not a proof that no indirect one exists. A hostname-based upstream that bootstraps through system DNS is the case in point: `https+local://dns.google/dns-query` resolves its own hostname with `DialSystem`, so once the takeover is in place that bootstrap goes `resolved -> TUN -> DNS outbound -> bootstrap -> resolved` and the query times out. The preflight does not see it, because the dependency sits in the upstream's bootstrap rather than in the clients it inspects. Upstream resolution, bootstrap included, therefore has to stay independent of the resolver path being redirected; configuring the address instead of the hostname, or resolving the hostname beforehand, avoids it.
|
||||||
|
|
||||||
|
The upstream requirement in the list above matters as much as the routing rule. With no name servers configured, Core resolves through a client that forwards to the system resolver; pointing the system resolver at the TUN would then close a loop through the DNS outbound, `resolved -> TUN -> DNS outbound -> system resolver -> resolved`, and resolution stops. The takeover is refused in that case.
|
||||||
|
|
||||||
|
The same applies to a name server pointed at `localhost`, and to a `dns` section that is present but lists no name servers. One such upstream is enough to refuse the takeover even when independent upstreams are configured alongside it: name servers are selected per domain, so a domain-specific rule can still choose the local one, and the loop then affects whichever domains reach it. The check is deliberately broader than the loop it observed, because the alternative would be to drop a name server the user configured.
|
||||||
|
|
||||||
|
Where it does not apply, DNS is left alone and the leak described in XTLS/Xray-core#6454 remains:
|
||||||
|
|
||||||
|
| Environment | Behaviour |
|
||||||
|
|---|---|
|
||||||
|
| systemd distribution with systemd-resolved enabled | applies |
|
||||||
|
| Alpine, Void, Devuan, OpenRC-based, OpenWrt | no `resolvectl`, skipped |
|
||||||
|
| DNS managed by dnsmasq / unbound / BIND / static `resolv.conf` | unreachable by `resolvectl`, skipped |
|
||||||
|
| Containers without a systemd-resolved daemon | skipped |
|
||||||
|
| systemd older than 240 | `default-route` unavailable, skipped |
|
||||||
|
|
||||||
|
On `Close()` the setting is reverted. It is **not** reverted if the process is killed with `SIGKILL`, since a process cannot handle that signal; run `resolvectl revert <iface>` to clean up by hand. An application that brings its own DNS endpoint is unaffected either way — this only covers the system resolver.
|
||||||
|
|
||||||
Due to this inbound not actually being a proxy, the configuration ignore required listen and port options, and never listen on any port. \
|
Due to this inbound not actually being a proxy, the configuration ignore required listen and port options, and never listen on any port. \
|
||||||
Here is simple Xray config snippet to enable the inbound:
|
Here is simple Xray config snippet to enable the inbound:
|
||||||
```
|
```
|
||||||
|
|||||||
+11
-2
@@ -32,6 +32,7 @@ type Config struct {
|
|||||||
AutoSystemRoutingTable []string `protobuf:"bytes,6,rep,name=auto_system_routing_table,json=autoSystemRoutingTable,proto3" json:"auto_system_routing_table,omitempty"`
|
AutoSystemRoutingTable []string `protobuf:"bytes,6,rep,name=auto_system_routing_table,json=autoSystemRoutingTable,proto3" json:"auto_system_routing_table,omitempty"`
|
||||||
AutoOutboundsInterface string `protobuf:"bytes,7,opt,name=auto_outbounds_interface,json=autoOutboundsInterface,proto3" json:"auto_outbounds_interface,omitempty"`
|
AutoOutboundsInterface string `protobuf:"bytes,7,opt,name=auto_outbounds_interface,json=autoOutboundsInterface,proto3" json:"auto_outbounds_interface,omitempty"`
|
||||||
Desc string `protobuf:"bytes,8,opt,name=desc,proto3" json:"desc,omitempty"`
|
Desc string `protobuf:"bytes,8,opt,name=desc,proto3" json:"desc,omitempty"`
|
||||||
|
AutoSystemDns bool `protobuf:"varint,9,opt,name=auto_system_dns,json=autoSystemDns,proto3" json:"auto_system_dns,omitempty"`
|
||||||
unknownFields protoimpl.UnknownFields
|
unknownFields protoimpl.UnknownFields
|
||||||
sizeCache protoimpl.SizeCache
|
sizeCache protoimpl.SizeCache
|
||||||
}
|
}
|
||||||
@@ -122,11 +123,18 @@ func (x *Config) GetDesc() string {
|
|||||||
return ""
|
return ""
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (x *Config) GetAutoSystemDns() bool {
|
||||||
|
if x != nil {
|
||||||
|
return x.AutoSystemDns
|
||||||
|
}
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
var File_proxy_tun_config_proto protoreflect.FileDescriptor
|
var File_proxy_tun_config_proto protoreflect.FileDescriptor
|
||||||
|
|
||||||
const file_proxy_tun_config_proto_rawDesc = "" +
|
const file_proxy_tun_config_proto_rawDesc = "" +
|
||||||
"\n" +
|
"\n" +
|
||||||
"\x16proxy/tun/config.proto\x12\x0exray.proxy.tun\"\x82\x02\n" +
|
"\x16proxy/tun/config.proto\x12\x0exray.proxy.tun\"\xaa\x02\n" +
|
||||||
"\x06Config\x12\x12\n" +
|
"\x06Config\x12\x12\n" +
|
||||||
"\x04name\x18\x01 \x01(\tR\x04name\x12\x10\n" +
|
"\x04name\x18\x01 \x01(\tR\x04name\x12\x10\n" +
|
||||||
"\x03MTU\x18\x02 \x01(\rR\x03MTU\x12\x18\n" +
|
"\x03MTU\x18\x02 \x01(\rR\x03MTU\x12\x18\n" +
|
||||||
@@ -136,7 +144,8 @@ const file_proxy_tun_config_proto_rawDesc = "" +
|
|||||||
"user_level\x18\x05 \x01(\rR\tuserLevel\x129\n" +
|
"user_level\x18\x05 \x01(\rR\tuserLevel\x129\n" +
|
||||||
"\x19auto_system_routing_table\x18\x06 \x03(\tR\x16autoSystemRoutingTable\x128\n" +
|
"\x19auto_system_routing_table\x18\x06 \x03(\tR\x16autoSystemRoutingTable\x128\n" +
|
||||||
"\x18auto_outbounds_interface\x18\a \x01(\tR\x16autoOutboundsInterface\x12\x12\n" +
|
"\x18auto_outbounds_interface\x18\a \x01(\tR\x16autoOutboundsInterface\x12\x12\n" +
|
||||||
"\x04desc\x18\b \x01(\tR\x04descBL\n" +
|
"\x04desc\x18\b \x01(\tR\x04desc\x12&\n" +
|
||||||
|
"\x0fauto_system_dns\x18\t \x01(\bR\rautoSystemDnsBL\n" +
|
||||||
"\x12com.xray.proxy.tunP\x01Z#github.com/xtls/xray-core/proxy/tun\xaa\x02\x0eXray.Proxy.Tunb\x06proto3"
|
"\x12com.xray.proxy.tunP\x01Z#github.com/xtls/xray-core/proxy/tun\xaa\x02\x0eXray.Proxy.Tunb\x06proto3"
|
||||||
|
|
||||||
var (
|
var (
|
||||||
|
|||||||
@@ -15,4 +15,5 @@ message Config {
|
|||||||
repeated string auto_system_routing_table = 6;
|
repeated string auto_system_routing_table = 6;
|
||||||
string auto_outbounds_interface = 7;
|
string auto_outbounds_interface = 7;
|
||||||
string desc = 8;
|
string desc = 8;
|
||||||
|
bool auto_system_dns = 9;
|
||||||
}
|
}
|
||||||
|
|||||||
+43
-4
@@ -37,6 +37,25 @@ type Handler struct {
|
|||||||
downlinkCounter stats.Counter
|
downlinkCounter stats.Counter
|
||||||
}
|
}
|
||||||
|
|
||||||
|
type tunUDPStatsWriter struct {
|
||||||
|
writer buf.Writer
|
||||||
|
counter stats.Counter
|
||||||
|
}
|
||||||
|
|
||||||
|
func (w *tunUDPStatsWriter) WriteMultiBuffer(mb buf.MultiBuffer) error {
|
||||||
|
for len(mb) > 0 {
|
||||||
|
remaining, packet := buf.SplitFirst(mb)
|
||||||
|
packetSize := packet.Len()
|
||||||
|
if err := w.writer.WriteMultiBuffer(buf.MultiBuffer{packet}); err != nil {
|
||||||
|
buf.ReleaseMulti(remaining)
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
w.counter.Add(int64(packetSize))
|
||||||
|
mb = remaining
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
// ConnectionHandler interface with the only method that stack is going to push new connections to
|
// ConnectionHandler interface with the only method that stack is going to push new connections to
|
||||||
type ConnectionHandler interface {
|
type ConnectionHandler interface {
|
||||||
HandleConnection(conn net.Conn, destination net.Destination)
|
HandleConnection(conn net.Conn, destination net.Destination)
|
||||||
@@ -104,7 +123,7 @@ func (t *Handler) Start() error {
|
|||||||
iface := updater.Get()
|
iface := updater.Get()
|
||||||
if iface == nil {
|
if iface == nil {
|
||||||
errors.LogInfo(context.Background(), "[tun] falied to set interface > iface == nil")
|
errors.LogInfo(context.Background(), "[tun] falied to set interface > iface == nil")
|
||||||
return nil
|
return errors.New("iface not found")
|
||||||
}
|
}
|
||||||
return c.Control(func(fd uintptr) {
|
return c.Control(func(fd uintptr) {
|
||||||
addrPort, _ := netip.ParseAddrPort(address)
|
addrPort, _ := netip.ParseAddrPort(address)
|
||||||
@@ -146,6 +165,16 @@ func (t *Handler) Start() error {
|
|||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Platform-specific system DNS takeover, where the platform implements it.
|
||||||
|
// Non-fatal: a failure leaves DNS management with the OS.
|
||||||
|
if c, ok := tunInterface.(interface {
|
||||||
|
ConfigureSystemDNS(context.Context, string) error
|
||||||
|
}); ok {
|
||||||
|
if err := c.ConfigureSystemDNS(t.ctx, t.tag); err != nil {
|
||||||
|
errors.LogInfoInner(t.ctx, err, "[tun] system DNS not configured")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
t.stack = tunStack
|
t.stack = tunStack
|
||||||
t.tun = tunInterface
|
t.tun = tunInterface
|
||||||
|
|
||||||
@@ -171,7 +200,8 @@ func (t *Handler) HandleConnection(conn net.Conn, destination net.Destination) {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
source := net.DestinationFromAddr(remote)
|
source := net.DestinationFromAddr(remote)
|
||||||
if t.uplinkCounter != nil || t.downlinkCounter != nil {
|
isUDP := destination.Network == net.Network_UDP
|
||||||
|
if !isUDP && (t.uplinkCounter != nil || t.downlinkCounter != nil) {
|
||||||
conn = &stat.CounterConnection{
|
conn = &stat.CounterConnection{
|
||||||
Connection: conn,
|
Connection: conn,
|
||||||
ReadCounter: t.uplinkCounter,
|
ReadCounter: t.uplinkCounter,
|
||||||
@@ -203,9 +233,18 @@ func (t *Handler) HandleConnection(conn net.Conn, destination net.Destination) {
|
|||||||
})
|
})
|
||||||
errors.LogInfo(ctx, "processing from ", source, " to ", destination)
|
errors.LogInfo(ctx, "processing from ", source, " to ", destination)
|
||||||
|
|
||||||
|
reader := &buf.TimeoutWrapperReader{Reader: buf.NewReader(conn)}
|
||||||
|
writer := buf.NewWriter(conn)
|
||||||
|
if isUDP {
|
||||||
|
reader.Counter = t.uplinkCounter
|
||||||
|
if t.downlinkCounter != nil {
|
||||||
|
writer = &tunUDPStatsWriter{writer: writer, counter: t.downlinkCounter}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
link := &transport.Link{
|
link := &transport.Link{
|
||||||
Reader: &buf.TimeoutWrapperReader{Reader: buf.NewReader(conn)},
|
Reader: reader,
|
||||||
Writer: buf.NewWriter(conn),
|
Writer: writer,
|
||||||
}
|
}
|
||||||
if err := t.dispatcher.DispatchLink(ctx, destination, link); err != nil {
|
if err := t.dispatcher.DispatchLink(ctx, destination, link); err != nil {
|
||||||
errors.LogError(ctx, errors.New("connection closed").Base(err))
|
errors.LogError(ctx, errors.New("connection closed").Base(err))
|
||||||
|
|||||||
@@ -6,12 +6,24 @@ import (
|
|||||||
"context"
|
"context"
|
||||||
"net"
|
"net"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
|
"os/exec"
|
||||||
"strconv"
|
"strconv"
|
||||||
"sync"
|
"sync"
|
||||||
|
|
||||||
"github.com/vishvananda/netlink"
|
"github.com/vishvananda/netlink"
|
||||||
|
appdns "github.com/xtls/xray-core/app/dns"
|
||||||
"github.com/xtls/xray-core/common/errors"
|
"github.com/xtls/xray-core/common/errors"
|
||||||
|
xnet "github.com/xtls/xray-core/common/net"
|
||||||
"github.com/xtls/xray-core/common/platform"
|
"github.com/xtls/xray-core/common/platform"
|
||||||
|
"github.com/xtls/xray-core/common/serial"
|
||||||
|
"github.com/xtls/xray-core/common/session"
|
||||||
|
"github.com/xtls/xray-core/core"
|
||||||
|
feature_dns "github.com/xtls/xray-core/features/dns"
|
||||||
|
"github.com/xtls/xray-core/features/dns/localdns"
|
||||||
|
"github.com/xtls/xray-core/features/outbound"
|
||||||
|
"github.com/xtls/xray-core/features/routing"
|
||||||
|
routingsession "github.com/xtls/xray-core/features/routing/session"
|
||||||
|
"github.com/xtls/xray-core/proxy/dns"
|
||||||
"golang.org/x/sys/unix"
|
"golang.org/x/sys/unix"
|
||||||
"gvisor.dev/gvisor/pkg/tcpip/link/fdbased"
|
"gvisor.dev/gvisor/pkg/tcpip/link/fdbased"
|
||||||
"gvisor.dev/gvisor/pkg/tcpip/stack"
|
"gvisor.dev/gvisor/pkg/tcpip/stack"
|
||||||
@@ -30,6 +42,238 @@ type LinuxTun struct {
|
|||||||
systemRoutes []netlink.Route
|
systemRoutes []netlink.Route
|
||||||
routeMonitorStop chan struct{}
|
routeMonitorStop chan struct{}
|
||||||
routeMonitorOnce sync.Once
|
routeMonitorOnce sync.Once
|
||||||
|
|
||||||
|
systemDNSSet bool
|
||||||
|
systemDNSDirty bool
|
||||||
|
}
|
||||||
|
|
||||||
|
// resolvectlRunner runs a resolvectl command. Overridable for tests.
|
||||||
|
var resolvectlRunner = func(name string, args ...string) ([]byte, error) {
|
||||||
|
return exec.Command(name, args...).CombinedOutput()
|
||||||
|
}
|
||||||
|
|
||||||
|
// systemDNSAddrs derives the addresses used for the system DNS takeover from the
|
||||||
|
// first IPv4 gateway: the gateway address itself is what a query from this
|
||||||
|
// interface appears to come from, and the next address is what the resolver is
|
||||||
|
// pointed at. The latter belongs to the TUN and is answered inside Xray;
|
||||||
|
// handing the configured public resolvers to resolvectl instead would leave the
|
||||||
|
// system querying them directly over the physical link, defeating the point of
|
||||||
|
// the TUN.
|
||||||
|
func systemDNSAddrs(gateway []string) (source, dns netip.Addr, ok bool) {
|
||||||
|
for _, address := range gateway {
|
||||||
|
prefix, err := netip.ParsePrefix(address)
|
||||||
|
if err != nil {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
addr := prefix.Addr()
|
||||||
|
if !addr.Is4() {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
return addr, addr.Next(), true
|
||||||
|
}
|
||||||
|
return netip.Addr{}, netip.Addr{}, false
|
||||||
|
}
|
||||||
|
|
||||||
|
func buildResolvectlArgs(action, iface string, extra ...string) []string {
|
||||||
|
args := make([]string, 0, 2+len(extra))
|
||||||
|
args = append(args, action, iface)
|
||||||
|
args = append(args, extra...)
|
||||||
|
return args
|
||||||
|
}
|
||||||
|
|
||||||
|
func runResolvectl(action, iface string, extra ...string) error {
|
||||||
|
args := buildResolvectlArgs(action, iface, extra...)
|
||||||
|
if _, err := resolvectlRunner("resolvectl", args...); err != nil {
|
||||||
|
return errors.New("resolvectl ", action, " failed").Base(err)
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// ifaceName returns the TUN interface name, or empty when the link is not
|
||||||
|
// available. Callers must treat empty as "nothing to configure".
|
||||||
|
func (t *LinuxTun) ifaceName() string {
|
||||||
|
if t.tunLink == nil {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
attrs := t.tunLink.Attrs()
|
||||||
|
if attrs == nil {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
return attrs.Name
|
||||||
|
}
|
||||||
|
|
||||||
|
// probeSourcePort is a representative client port for the routing probe. A real
|
||||||
|
// query arrives from an ephemeral port that cannot be known in advance, so this
|
||||||
|
// only matters for a rule that matches on a source port.
|
||||||
|
const probeSourcePort = 49152
|
||||||
|
|
||||||
|
// verifyDNSRouting reports whether a DNS query to address would actually be
|
||||||
|
// handled. Redirecting the system resolver at an address nothing answers would
|
||||||
|
// break name resolution outright, so the takeover only proceeds when routing
|
||||||
|
// hands such a query to a DNS-capable outbound.
|
||||||
|
//
|
||||||
|
// Overridable for tests.
|
||||||
|
var verifyDNSRouting = func(ctx context.Context, inboundTag, source, address string) error {
|
||||||
|
ip, err := netip.ParseAddr(address)
|
||||||
|
if err != nil || !ip.Is4() {
|
||||||
|
return errors.New("invalid DNS address ", address).Base(err)
|
||||||
|
}
|
||||||
|
src, err := netip.ParseAddr(source)
|
||||||
|
if err != nil || !src.Is4() {
|
||||||
|
return errors.New("invalid source address ", source).Base(err)
|
||||||
|
}
|
||||||
|
|
||||||
|
instance := core.MustFromContext(ctx)
|
||||||
|
|
||||||
|
// Any resolution path that could still reach the system resolver has to be
|
||||||
|
// refused, because pointing the system resolver at the TUN would close a
|
||||||
|
// loop through the DNS outbound. With no `dns` section Core installs such a
|
||||||
|
// client; with a `dns` section that has no name servers app/dns falls back
|
||||||
|
// to one; and a name server pointed at "localhost" is one even when
|
||||||
|
// independent upstreams are configured alongside it, because name servers
|
||||||
|
// are selected per domain.
|
||||||
|
switch dnsFeature := instance.GetFeature(feature_dns.ClientType()).(type) {
|
||||||
|
case *localdns.Client:
|
||||||
|
return errors.New("DNS feature is the system resolver, takeover would loop")
|
||||||
|
case *appdns.DNS:
|
||||||
|
if dnsFeature.MayUseSystemResolver() {
|
||||||
|
return errors.New("DNS configuration may resolve through the system resolver, takeover would loop")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
router, ok := instance.GetFeature(routing.RouterType()).(routing.Router)
|
||||||
|
if !ok {
|
||||||
|
return errors.New("router feature unavailable")
|
||||||
|
}
|
||||||
|
|
||||||
|
// A real query from this interface carries a source address, and rules may
|
||||||
|
// match on it, so the probe has to carry one too.
|
||||||
|
queryCtx := session.ContextWithInbound(ctx, &session.Inbound{
|
||||||
|
Name: "tun",
|
||||||
|
Tag: inboundTag,
|
||||||
|
Source: xnet.UDPDestination(xnet.IPAddress(src.AsSlice()), probeSourcePort),
|
||||||
|
})
|
||||||
|
queryCtx = session.ContextWithOutbounds(queryCtx, []*session.Outbound{{
|
||||||
|
Target: xnet.UDPDestination(xnet.IPAddress(ip.AsSlice()), 53),
|
||||||
|
}})
|
||||||
|
|
||||||
|
route, err := router.PickRoute(routingsession.AsRoutingContext(queryCtx))
|
||||||
|
if err != nil {
|
||||||
|
return errors.New("no route for ", address, ":53").Base(err)
|
||||||
|
}
|
||||||
|
|
||||||
|
manager, ok := instance.GetFeature(outbound.ManagerType()).(outbound.Manager)
|
||||||
|
if !ok {
|
||||||
|
return errors.New("outbound manager unavailable")
|
||||||
|
}
|
||||||
|
|
||||||
|
handler := manager.GetHandler(route.GetOutboundTag())
|
||||||
|
if handler == nil {
|
||||||
|
return errors.New("outbound ", route.GetOutboundTag(), " does not exist")
|
||||||
|
}
|
||||||
|
if settings := handler.ProxySettings(); settings == nil || settings.Type != serial.GetMessageType(&dns.Config{}) {
|
||||||
|
return errors.New("outbound ", route.GetOutboundTag(), " does not handle DNS")
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// ConfigureSystemDNS points systemd-resolved at this interface so name lookups
|
||||||
|
// resolve through Xray instead of leaking to the physical link.
|
||||||
|
//
|
||||||
|
// It acts only when the config opts in, and it verifies the data path first:
|
||||||
|
// unless a query to the advertised address would actually be handled, host-wide
|
||||||
|
// resolution is left to the OS, which is the documented default. Errors are
|
||||||
|
// returned to the caller, which treats them as non-fatal.
|
||||||
|
func (t *LinuxTun) ConfigureSystemDNS(ctx context.Context, inboundTag string) error {
|
||||||
|
if !t.options.AutoSystemDns {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
if t.systemDNSSet {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// A previous revert may have failed. Retry before applying anything, so a
|
||||||
|
// dirty resolver does not silently outlive the attempt to clean it up.
|
||||||
|
if t.systemDNSDirty {
|
||||||
|
if err := t.revertSystemDNS(); err != nil {
|
||||||
|
return errors.New("previous system DNS revert still failing").Base(err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
source, address, ok := systemDNSAddrs(t.options.Gateway)
|
||||||
|
if !ok {
|
||||||
|
return errors.New("no IPv4 gateway, cannot derive a system DNS address")
|
||||||
|
}
|
||||||
|
|
||||||
|
iface := t.ifaceName()
|
||||||
|
if iface == "" {
|
||||||
|
return errors.New("interface not available")
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := verifyDNSRouting(ctx, inboundTag, source.String(), address.String()); err != nil {
|
||||||
|
return errors.New("no DNS path at ", address.String(), ":53").Base(err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Applied as a sequence with rollback: a half-configured resolver would be
|
||||||
|
// worse than none at all.
|
||||||
|
if err := runResolvectl("dns", iface, address.String()); err != nil {
|
||||||
|
return errors.New("resolvectl dns failed").Base(err)
|
||||||
|
}
|
||||||
|
if err := runResolvectl("domain", iface, "~."); err != nil {
|
||||||
|
return t.rollbackSystemDNS(iface, errors.New("resolvectl domain failed").Base(err))
|
||||||
|
}
|
||||||
|
if err := runResolvectl("default-route", iface, "true"); err != nil {
|
||||||
|
return t.rollbackSystemDNS(iface, errors.New("resolvectl default-route failed").Base(err))
|
||||||
|
}
|
||||||
|
|
||||||
|
t.systemDNSSet = true
|
||||||
|
errors.LogInfo(ctx, "[tun] system DNS set to ", address.String(), " on ", iface)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// rollbackSystemDNS undoes a partially applied takeover. A failed revert is
|
||||||
|
// recorded so the next attempt retries it, and is reported rather than
|
||||||
|
// swallowed.
|
||||||
|
func (t *LinuxTun) rollbackSystemDNS(iface string, cause error) error {
|
||||||
|
if err := runResolvectl("revert", iface); err != nil {
|
||||||
|
t.systemDNSDirty = true
|
||||||
|
// Combine, because Base overwrites: reporting only the cause would hide
|
||||||
|
// the revert failure, and reporting only the revert failure would hide
|
||||||
|
// why the revert was attempted.
|
||||||
|
return errors.New("revert failed, per-link DNS settings may remain").Base(errors.Combine(err, cause))
|
||||||
|
}
|
||||||
|
return cause
|
||||||
|
}
|
||||||
|
|
||||||
|
// revertSystemDNS issues the revert and keeps the dirty flag in step with the
|
||||||
|
// outcome.
|
||||||
|
func (t *LinuxTun) revertSystemDNS() error {
|
||||||
|
err := runResolvectl("revert", t.ifaceName())
|
||||||
|
t.systemDNSDirty = err != nil
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
t.systemDNSSet = false
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// unsetSystemDNS hands DNS back to the OS. Only meaningful when
|
||||||
|
// ConfigureSystemDNS applied something, or a previous revert failed.
|
||||||
|
func (t *LinuxTun) unsetSystemDNS() {
|
||||||
|
if !t.systemDNSSet && !t.systemDNSDirty {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
if t.ifaceName() == "" {
|
||||||
|
// The link is gone, and its per-link settings went with it.
|
||||||
|
t.systemDNSSet = false
|
||||||
|
t.systemDNSDirty = false
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := t.revertSystemDNS(); err != nil {
|
||||||
|
errors.LogInfoInner(context.Background(), err, "[tun] failed to revert system DNS; per-link settings may remain until revert succeeds")
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// LinuxTun implements Tun
|
// LinuxTun implements Tun
|
||||||
@@ -200,6 +444,7 @@ func (t *LinuxTun) Close() error {
|
|||||||
}
|
}
|
||||||
})
|
})
|
||||||
|
|
||||||
|
t.unsetSystemDNS()
|
||||||
_ = t.unsetSystemRoutes()
|
_ = t.unsetSystemRoutes()
|
||||||
_ = t.unsetInterfaceAddresses()
|
_ = t.unsetInterfaceAddresses()
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,193 @@
|
|||||||
|
//go:build linux && !android
|
||||||
|
|
||||||
|
package tun
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"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/inbound"
|
||||||
|
_ "github.com/xtls/xray-core/app/proxyman/outbound"
|
||||||
|
"github.com/xtls/xray-core/app/router"
|
||||||
|
"github.com/xtls/xray-core/common/geodata"
|
||||||
|
"github.com/xtls/xray-core/common/net"
|
||||||
|
"github.com/xtls/xray-core/common/serial"
|
||||||
|
"github.com/xtls/xray-core/core"
|
||||||
|
"github.com/xtls/xray-core/proxy/blackhole"
|
||||||
|
proxydns "github.com/xtls/xray-core/proxy/dns"
|
||||||
|
"github.com/xtls/xray-core/proxy/freedom"
|
||||||
|
)
|
||||||
|
|
||||||
|
const (
|
||||||
|
routeTestInboundTag = "tun"
|
||||||
|
routeTestSource = "192.168.100.1"
|
||||||
|
routeTestDNSAddress = "192.168.100.2"
|
||||||
|
)
|
||||||
|
|
||||||
|
// port53Rule sends DNS queries arriving from the interface to the dns outbound.
|
||||||
|
func port53Rule() *router.RoutingRule {
|
||||||
|
return &router.RoutingRule{
|
||||||
|
InboundTag: []string{routeTestInboundTag},
|
||||||
|
PortList: &net.PortList{Range: []*net.PortRange{net.SinglePortRange(53)}},
|
||||||
|
TargetTag: &router.RoutingRule_Tag{Tag: "dns"},
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// sourceBlockRule diverts traffic from one address, which is the shape of a rule
|
||||||
|
// that only matches because the real request carries a source.
|
||||||
|
func sourceBlockRule(ip []byte) *router.RoutingRule {
|
||||||
|
return &router.RoutingRule{
|
||||||
|
SourceIp: []*geodata.IPRule{{
|
||||||
|
Value: &geodata.IPRule_Custom{
|
||||||
|
Custom: &geodata.CIDRRule{
|
||||||
|
Cidr: &geodata.CIDR{Ip: ip, Prefix: 32},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}},
|
||||||
|
TargetTag: &router.RoutingRule_Tag{Tag: "block"},
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// newRouteTestContext builds a real but unstarted instance: no TUN device, no
|
||||||
|
// running resolver. The instance is placed in the context through the key core
|
||||||
|
// exports for tests.
|
||||||
|
func newRouteTestContext(t *testing.T, withDNSApp bool, nameServers []*appdns.NameServer, rules []*router.RoutingRule) context.Context {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
apps := []*serial.TypedMessage{
|
||||||
|
serial.ToTypedMessage(&dispatcher.Config{}),
|
||||||
|
serial.ToTypedMessage(&proxyman.InboundConfig{}),
|
||||||
|
serial.ToTypedMessage(&proxyman.OutboundConfig{}),
|
||||||
|
serial.ToTypedMessage(&router.Config{Rule: rules}),
|
||||||
|
}
|
||||||
|
if withDNSApp {
|
||||||
|
apps = append(apps, serial.ToTypedMessage(&appdns.Config{NameServer: nameServers}))
|
||||||
|
}
|
||||||
|
|
||||||
|
instance, err := core.New(&core.Config{
|
||||||
|
App: apps,
|
||||||
|
Outbound: []*core.OutboundHandlerConfig{
|
||||||
|
{Tag: "direct", ProxySettings: serial.ToTypedMessage(&freedom.Config{})},
|
||||||
|
{Tag: "dns", ProxySettings: serial.ToTypedMessage(&proxydns.Config{})},
|
||||||
|
{Tag: "block", ProxySettings: serial.ToTypedMessage(&blackhole.Config{})},
|
||||||
|
},
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("core.New: %v", err)
|
||||||
|
}
|
||||||
|
t.Cleanup(func() { _ = instance.Close() })
|
||||||
|
|
||||||
|
return context.WithValue(context.Background(), core.XrayKey(1), instance)
|
||||||
|
}
|
||||||
|
|
||||||
|
func udpNameServer(ip []byte) []*appdns.NameServer {
|
||||||
|
return []*appdns.NameServer{{
|
||||||
|
Address: &net.Endpoint{
|
||||||
|
Network: net.Network_UDP,
|
||||||
|
Address: &net.IPOrDomain{Address: &net.IPOrDomain_Ip{Ip: ip}},
|
||||||
|
Port: 53,
|
||||||
|
},
|
||||||
|
}}
|
||||||
|
}
|
||||||
|
|
||||||
|
// localNameServer is a name server pointed at "localhost", which app/dns
|
||||||
|
// resolves through the system resolver.
|
||||||
|
func localNameServer() *appdns.NameServer {
|
||||||
|
return &appdns.NameServer{
|
||||||
|
Address: &net.Endpoint{
|
||||||
|
Network: net.Network_UDP,
|
||||||
|
Address: &net.IPOrDomain{Address: &net.IPOrDomain_Domain{Domain: "localhost"}},
|
||||||
|
Port: 53,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// These drive the real feature lookup and the real router. verifyDNSRouting is
|
||||||
|
// the same function ConfigureSystemDNS calls, so a false positive here is a
|
||||||
|
// false positive in the takeover decision itself, which is what assertions on
|
||||||
|
// the resolvectl arguments could never catch.
|
||||||
|
func TestVerifyDNSRoutingDecisions(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
withDNSApp bool
|
||||||
|
nameServers []*appdns.NameServer
|
||||||
|
rules []*router.RoutingRule
|
||||||
|
wantErr string
|
||||||
|
}{
|
||||||
|
{
|
||||||
|
name: "independent upstream reaches the dns outbound",
|
||||||
|
withDNSApp: true,
|
||||||
|
nameServers: udpNameServer([]byte{9, 9, 9, 9}),
|
||||||
|
rules: []*router.RoutingRule{port53Rule()},
|
||||||
|
wantErr: "",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "no dns section falls back to the system resolver",
|
||||||
|
rules: []*router.RoutingRule{port53Rule()},
|
||||||
|
wantErr: "system resolver",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "dns section without name servers falls back too",
|
||||||
|
withDNSApp: true,
|
||||||
|
rules: []*router.RoutingRule{port53Rule()},
|
||||||
|
wantErr: "system resolver",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
// An independent upstream is not enough on its own: name servers are
|
||||||
|
// selected per domain, so a local one can still be the one chosen.
|
||||||
|
// The refusal is deliberately domain-agnostic for that reason.
|
||||||
|
name: "a local name server alongside an independent one",
|
||||||
|
withDNSApp: true,
|
||||||
|
nameServers: append(udpNameServer([]byte{9, 9, 9, 9}), localNameServer()),
|
||||||
|
rules: []*router.RoutingRule{port53Rule()},
|
||||||
|
wantErr: "system resolver",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "a rule on the interface address diverts the real query",
|
||||||
|
withDNSApp: true, nameServers: udpNameServer([]byte{9, 9, 9, 9}),
|
||||||
|
rules: []*router.RoutingRule{
|
||||||
|
sourceBlockRule([]byte{192, 168, 100, 1}),
|
||||||
|
port53Rule(),
|
||||||
|
},
|
||||||
|
wantErr: "does not handle DNS",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "a rule on another address does not match it",
|
||||||
|
withDNSApp: true, nameServers: udpNameServer([]byte{9, 9, 9, 9}),
|
||||||
|
rules: []*router.RoutingRule{
|
||||||
|
sourceBlockRule([]byte{10, 0, 0, 1}),
|
||||||
|
port53Rule(),
|
||||||
|
},
|
||||||
|
wantErr: "",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "no rule matches the query",
|
||||||
|
withDNSApp: true, nameServers: udpNameServer([]byte{9, 9, 9, 9}),
|
||||||
|
wantErr: "no route",
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
ctx := newRouteTestContext(t, tt.withDNSApp, tt.nameServers, tt.rules)
|
||||||
|
err := verifyDNSRouting(ctx, routeTestInboundTag, routeTestSource, routeTestDNSAddress)
|
||||||
|
|
||||||
|
if tt.wantErr == "" {
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("expected the takeover to be accepted, got: %v", err)
|
||||||
|
}
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if err == nil {
|
||||||
|
t.Fatalf("expected the takeover to be refused with %q, got nil", tt.wantErr)
|
||||||
|
}
|
||||||
|
if !strings.Contains(err.Error(), tt.wantErr) {
|
||||||
|
t.Errorf("error = %q, want it to contain %q", err.Error(), tt.wantErr)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,436 @@
|
|||||||
|
//go:build linux && !android
|
||||||
|
|
||||||
|
package tun
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"errors"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/vishvananda/netlink"
|
||||||
|
)
|
||||||
|
|
||||||
|
// testLink returns a minimal netlink.Link whose Attrs().Name is name, so the
|
||||||
|
// DNS helpers can be exercised without a real TUN device.
|
||||||
|
func testLink(name string) netlink.Link {
|
||||||
|
return &netlink.Dummy{LinkAttrs: netlink.LinkAttrs{Name: name}}
|
||||||
|
}
|
||||||
|
|
||||||
|
type probeCall struct {
|
||||||
|
inboundTag string
|
||||||
|
source string
|
||||||
|
address string
|
||||||
|
}
|
||||||
|
|
||||||
|
// stubDNSRouting replaces the routing probe for the duration of a test and
|
||||||
|
// records how it was called, so tests can assert the probe is representative.
|
||||||
|
func stubDNSRouting(t *testing.T, err error) *[]probeCall {
|
||||||
|
t.Helper()
|
||||||
|
original := verifyDNSRouting
|
||||||
|
calls := []probeCall{}
|
||||||
|
verifyDNSRouting = func(_ context.Context, inboundTag, source, address string) error {
|
||||||
|
calls = append(calls, probeCall{inboundTag, source, address})
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
t.Cleanup(func() { verifyDNSRouting = original })
|
||||||
|
return &calls
|
||||||
|
}
|
||||||
|
|
||||||
|
// recorder installs a resolvectl stub for the duration of a test and returns the
|
||||||
|
// captured invocations. An empty failOn succeeds every call; otherwise the named
|
||||||
|
// subcommand fails.
|
||||||
|
func recorder(t *testing.T, failOn string) *[][]string {
|
||||||
|
t.Helper()
|
||||||
|
original := resolvectlRunner
|
||||||
|
calls := [][]string{}
|
||||||
|
resolvectlRunner = func(name string, args ...string) ([]byte, error) {
|
||||||
|
calls = append(calls, append([]string{name}, args...))
|
||||||
|
if failOn != "" && len(args) > 0 && args[0] == failOn {
|
||||||
|
return nil, errors.New("boom")
|
||||||
|
}
|
||||||
|
return nil, nil
|
||||||
|
}
|
||||||
|
t.Cleanup(func() { resolvectlRunner = original })
|
||||||
|
return &calls
|
||||||
|
}
|
||||||
|
|
||||||
|
func optedInTun() *LinuxTun {
|
||||||
|
return &LinuxTun{
|
||||||
|
options: &Config{
|
||||||
|
Name: "xray_tun",
|
||||||
|
Gateway: []string{"192.168.100.1/30"},
|
||||||
|
AutoSystemDns: true,
|
||||||
|
},
|
||||||
|
tunLink: testLink("xray_tun"),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func joined(calls [][]string) string {
|
||||||
|
parts := make([]string, 0, len(calls))
|
||||||
|
for _, call := range calls {
|
||||||
|
parts = append(parts, strings.Join(call, " "))
|
||||||
|
}
|
||||||
|
return strings.Join(parts, " | ")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestConfigureSystemDNSDisabledByDefault(t *testing.T) {
|
||||||
|
probes := stubDNSRouting(t, nil)
|
||||||
|
calls := recorder(t, "")
|
||||||
|
|
||||||
|
t1 := optedInTun()
|
||||||
|
t1.options.AutoSystemDns = false
|
||||||
|
|
||||||
|
if err := t1.ConfigureSystemDNS(context.Background(), "tun"); err != nil {
|
||||||
|
t.Fatalf("unexpected error: %v", err)
|
||||||
|
}
|
||||||
|
if len(*probes) != 0 {
|
||||||
|
t.Errorf("routing probe must not run when disabled, got %d calls", len(*probes))
|
||||||
|
}
|
||||||
|
if len(*calls) != 0 {
|
||||||
|
t.Errorf("resolvectl must not run when disabled, got %v", *calls)
|
||||||
|
}
|
||||||
|
if t1.systemDNSSet {
|
||||||
|
t.Error("systemDNSSet should stay false when disabled")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestConfigureSystemDNSNoGateway(t *testing.T) {
|
||||||
|
probes := stubDNSRouting(t, nil)
|
||||||
|
calls := recorder(t, "")
|
||||||
|
|
||||||
|
t1 := optedInTun()
|
||||||
|
t1.options.Gateway = nil
|
||||||
|
|
||||||
|
if err := t1.ConfigureSystemDNS(context.Background(), "tun"); err == nil {
|
||||||
|
t.Fatal("expected an error when no IPv4 gateway is configured")
|
||||||
|
}
|
||||||
|
if len(*probes) != 0 {
|
||||||
|
t.Errorf("routing probe must not run without a gateway, got %d calls", len(*probes))
|
||||||
|
}
|
||||||
|
if len(*calls) != 0 {
|
||||||
|
t.Errorf("resolvectl must not run without a gateway, got %v", *calls)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// This is the case the reviewer flagged: without a routed DNS path, pointing the
|
||||||
|
// system resolver at the derived address would break resolution outright.
|
||||||
|
func TestConfigureSystemDNSLeavesOSDNSWhenNoRoute(t *testing.T) {
|
||||||
|
probes := stubDNSRouting(t, errors.New("no route"))
|
||||||
|
calls := recorder(t, "")
|
||||||
|
|
||||||
|
t1 := optedInTun()
|
||||||
|
|
||||||
|
if err := t1.ConfigureSystemDNS(context.Background(), "tun"); err == nil {
|
||||||
|
t.Fatal("expected an error when the DNS path is unverified")
|
||||||
|
}
|
||||||
|
if len(*probes) != 1 {
|
||||||
|
t.Errorf("routing probe should run once, got %d", len(*probes))
|
||||||
|
}
|
||||||
|
if len(*calls) != 0 {
|
||||||
|
t.Errorf("system DNS must be left untouched, got %v", *calls)
|
||||||
|
}
|
||||||
|
if t1.systemDNSSet {
|
||||||
|
t.Error("systemDNSSet should stay false when the path is unverified")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// A real query from the interface carries a source address, and rules may match
|
||||||
|
// on it, so the probe must not be source-less.
|
||||||
|
func TestConfigureSystemDNSProbeCarriesSource(t *testing.T) {
|
||||||
|
probes := stubDNSRouting(t, nil)
|
||||||
|
recorder(t, "")
|
||||||
|
|
||||||
|
t1 := optedInTun()
|
||||||
|
if err := t1.ConfigureSystemDNS(context.Background(), "tun"); err != nil {
|
||||||
|
t.Fatalf("unexpected error: %v", err)
|
||||||
|
}
|
||||||
|
if len(*probes) != 1 {
|
||||||
|
t.Fatalf("expected one probe call, got %d", len(*probes))
|
||||||
|
}
|
||||||
|
got := (*probes)[0]
|
||||||
|
if got.source != "192.168.100.1" {
|
||||||
|
t.Errorf("probe source = %q, want the interface address %q", got.source, "192.168.100.1")
|
||||||
|
}
|
||||||
|
if got.address != "192.168.100.2" {
|
||||||
|
t.Errorf("probe address = %q, want %q", got.address, "192.168.100.2")
|
||||||
|
}
|
||||||
|
if got.inboundTag != "tun" {
|
||||||
|
t.Errorf("probe inbound tag = %q, want %q", got.inboundTag, "tun")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestConfigureSystemDNSAppliesResolvectl(t *testing.T) {
|
||||||
|
stubDNSRouting(t, nil)
|
||||||
|
calls := recorder(t, "")
|
||||||
|
|
||||||
|
t1 := optedInTun()
|
||||||
|
|
||||||
|
if err := t1.ConfigureSystemDNS(context.Background(), "tun"); err != nil {
|
||||||
|
t.Fatalf("unexpected error: %v", err)
|
||||||
|
}
|
||||||
|
if !t1.systemDNSSet {
|
||||||
|
t.Fatal("systemDNSSet should be true after a successful takeover")
|
||||||
|
}
|
||||||
|
|
||||||
|
want := "resolvectl dns xray_tun 192.168.100.2 | " +
|
||||||
|
"resolvectl domain xray_tun ~. | " +
|
||||||
|
"resolvectl default-route xray_tun true"
|
||||||
|
if got := joined(*calls); got != want {
|
||||||
|
t.Errorf("resolvectl calls = %q, want %q", got, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestConfigureSystemDNSIdempotent(t *testing.T) {
|
||||||
|
stubDNSRouting(t, nil)
|
||||||
|
calls := recorder(t, "")
|
||||||
|
|
||||||
|
t1 := optedInTun()
|
||||||
|
|
||||||
|
if err := t1.ConfigureSystemDNS(context.Background(), "tun"); err != nil {
|
||||||
|
t.Fatalf("unexpected error: %v", err)
|
||||||
|
}
|
||||||
|
before := len(*calls)
|
||||||
|
|
||||||
|
if err := t1.ConfigureSystemDNS(context.Background(), "tun"); err != nil {
|
||||||
|
t.Fatalf("unexpected error: %v", err)
|
||||||
|
}
|
||||||
|
if len(*calls) != before {
|
||||||
|
t.Errorf("second call must be a no-op, calls went %d -> %d", before, len(*calls))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// A half-applied resolver is worse than none, so a failure mid-sequence reverts.
|
||||||
|
func TestConfigureSystemDNSRollsBackOnPartialFailure(t *testing.T) {
|
||||||
|
stubDNSRouting(t, nil)
|
||||||
|
calls := recorder(t, "domain")
|
||||||
|
|
||||||
|
t1 := optedInTun()
|
||||||
|
|
||||||
|
if err := t1.ConfigureSystemDNS(context.Background(), "tun"); err == nil {
|
||||||
|
t.Fatal("expected an error when a resolvectl step fails")
|
||||||
|
}
|
||||||
|
if t1.systemDNSSet {
|
||||||
|
t.Error("systemDNSSet should stay false after a failed takeover")
|
||||||
|
}
|
||||||
|
if t1.systemDNSDirty {
|
||||||
|
t.Error("a successful revert should not leave the resolver dirty")
|
||||||
|
}
|
||||||
|
if !strings.Contains(joined(*calls), "resolvectl revert xray_tun") {
|
||||||
|
t.Errorf("expected a revert after partial failure, got %q", joined(*calls))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// If the revert itself fails the settings may still be installed, so the state
|
||||||
|
// has to be remembered rather than silently dropped.
|
||||||
|
func TestConfigureSystemDNSRollbackFailureKeepsDirty(t *testing.T) {
|
||||||
|
stubDNSRouting(t, nil)
|
||||||
|
calls := recorder(t, "revert")
|
||||||
|
|
||||||
|
t1 := optedInTun()
|
||||||
|
t1.options.Gateway = []string{"192.168.100.1/30"}
|
||||||
|
// Make only the rollback path fail: "dns" succeeds, "domain" fails, "revert" fails.
|
||||||
|
*calls = nil
|
||||||
|
|
||||||
|
original := resolvectlRunner
|
||||||
|
defer func() { resolvectlRunner = original }()
|
||||||
|
resolvectlRunner = func(name string, args ...string) ([]byte, error) {
|
||||||
|
*calls = append(*calls, append([]string{name}, args...))
|
||||||
|
if len(args) > 0 && (args[0] == "domain" || args[0] == "revert") {
|
||||||
|
return nil, errors.New("boom")
|
||||||
|
}
|
||||||
|
return nil, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := t1.ConfigureSystemDNS(context.Background(), "tun"); err == nil {
|
||||||
|
t.Fatal("expected an error when domain fails")
|
||||||
|
}
|
||||||
|
if !t1.systemDNSDirty {
|
||||||
|
t.Error("a failed revert must leave the resolver marked dirty")
|
||||||
|
}
|
||||||
|
if t1.systemDNSSet {
|
||||||
|
t.Error("systemDNSSet must stay false when the takeover did not complete")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// A dirty resolver is retried before anything new is applied.
|
||||||
|
func TestConfigureSystemDNSRetriesDirtyBeforeApplying(t *testing.T) {
|
||||||
|
stubDNSRouting(t, nil)
|
||||||
|
calls := recorder(t, "")
|
||||||
|
|
||||||
|
t1 := optedInTun()
|
||||||
|
t1.systemDNSDirty = true
|
||||||
|
|
||||||
|
if err := t1.ConfigureSystemDNS(context.Background(), "tun"); err != nil {
|
||||||
|
t.Fatalf("unexpected error: %v", err)
|
||||||
|
}
|
||||||
|
got := joined(*calls)
|
||||||
|
if !strings.HasPrefix(got, "resolvectl revert xray_tun") {
|
||||||
|
t.Errorf("expected the stale revert first, got %q", got)
|
||||||
|
}
|
||||||
|
if t1.systemDNSDirty {
|
||||||
|
t.Error("a successful retry should clear the dirty flag")
|
||||||
|
}
|
||||||
|
if !t1.systemDNSSet {
|
||||||
|
t.Error("the takeover should proceed once the retry succeeds")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestUnsetSystemDNSReverts(t *testing.T) {
|
||||||
|
stubDNSRouting(t, nil)
|
||||||
|
calls := recorder(t, "")
|
||||||
|
|
||||||
|
t1 := optedInTun()
|
||||||
|
if err := t1.ConfigureSystemDNS(context.Background(), "tun"); err != nil {
|
||||||
|
t.Fatalf("setup failed: %v", err)
|
||||||
|
}
|
||||||
|
*calls = nil
|
||||||
|
|
||||||
|
t1.unsetSystemDNS()
|
||||||
|
if t1.systemDNSSet {
|
||||||
|
t.Error("systemDNSSet should be false after unset")
|
||||||
|
}
|
||||||
|
if got := joined(*calls); got != "resolvectl revert xray_tun" {
|
||||||
|
t.Errorf("unset calls = %q, want %q", got, "resolvectl revert xray_tun")
|
||||||
|
}
|
||||||
|
|
||||||
|
t1.unsetSystemDNS()
|
||||||
|
if len(*calls) != 1 {
|
||||||
|
t.Errorf("unsetSystemDNS must be idempotent, got %q", joined(*calls))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestUnsetSystemDNSKeepsDirtyWhenRevertFails(t *testing.T) {
|
||||||
|
stubDNSRouting(t, nil)
|
||||||
|
calls := recorder(t, "revert")
|
||||||
|
|
||||||
|
t1 := optedInTun()
|
||||||
|
t1.systemDNSSet = true
|
||||||
|
|
||||||
|
t1.unsetSystemDNS()
|
||||||
|
if !t1.systemDNSDirty {
|
||||||
|
t.Error("a failed revert during unset must be remembered")
|
||||||
|
}
|
||||||
|
if got := joined(*calls); !strings.Contains(got, "resolvectl revert xray_tun") {
|
||||||
|
t.Errorf("expected a revert attempt, got %q", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSystemDNSAddrs(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
gateway []string
|
||||||
|
wantSource string
|
||||||
|
wantDNS string
|
||||||
|
wantOK bool
|
||||||
|
}{
|
||||||
|
{
|
||||||
|
name: "ipv4 /30",
|
||||||
|
gateway: []string{"192.168.100.1/30"},
|
||||||
|
wantSource: "192.168.100.1",
|
||||||
|
wantDNS: "192.168.100.2",
|
||||||
|
wantOK: true,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "ipv4 /16",
|
||||||
|
gateway: []string{"10.0.0.1/16"},
|
||||||
|
wantSource: "10.0.0.1",
|
||||||
|
wantDNS: "10.0.0.2",
|
||||||
|
wantOK: true,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "first ipv4 wins",
|
||||||
|
gateway: []string{"fc00::1/64", "172.18.0.1/30"},
|
||||||
|
wantSource: "172.18.0.1",
|
||||||
|
wantDNS: "172.18.0.2",
|
||||||
|
wantOK: true,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "no gateway",
|
||||||
|
gateway: nil,
|
||||||
|
wantOK: false,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "ipv6 only",
|
||||||
|
gateway: []string{"fc00::1/64"},
|
||||||
|
wantOK: false,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
source, dnsAddr, ok := systemDNSAddrs(tt.gateway)
|
||||||
|
if ok != tt.wantOK {
|
||||||
|
t.Fatalf("ok = %v, want %v", ok, tt.wantOK)
|
||||||
|
}
|
||||||
|
if !tt.wantOK {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if source.String() != tt.wantSource {
|
||||||
|
t.Errorf("source = %q, want %q", source.String(), tt.wantSource)
|
||||||
|
}
|
||||||
|
if dnsAddr.String() != tt.wantDNS {
|
||||||
|
t.Errorf("dns = %q, want %q", dnsAddr.String(), tt.wantDNS)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestBuildResolvectlArgs(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
action string
|
||||||
|
iface string
|
||||||
|
extra []string
|
||||||
|
want []string
|
||||||
|
}{
|
||||||
|
{
|
||||||
|
name: "revert",
|
||||||
|
action: "revert",
|
||||||
|
iface: "xray_tun",
|
||||||
|
want: []string{"revert", "xray_tun"},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "dns single",
|
||||||
|
action: "dns",
|
||||||
|
iface: "xray_tun",
|
||||||
|
extra: []string{"192.168.100.2"},
|
||||||
|
want: []string{"dns", "xray_tun", "192.168.100.2"},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "dns multiple",
|
||||||
|
action: "dns",
|
||||||
|
iface: "xray_tun",
|
||||||
|
extra: []string{"192.168.100.2", "fc00::2"},
|
||||||
|
want: []string{"dns", "xray_tun", "192.168.100.2", "fc00::2"},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "domain wildcard",
|
||||||
|
action: "domain",
|
||||||
|
iface: "xray_tun",
|
||||||
|
extra: []string{"~."},
|
||||||
|
want: []string{"domain", "xray_tun", "~."},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "default-route",
|
||||||
|
action: "default-route",
|
||||||
|
iface: "xray_tun",
|
||||||
|
extra: []string{"true"},
|
||||||
|
want: []string{"default-route", "xray_tun", "true"},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
got := buildResolvectlArgs(tt.action, tt.iface, tt.extra...)
|
||||||
|
if len(got) != len(tt.want) {
|
||||||
|
t.Fatalf("args = %v, want %v", got, tt.want)
|
||||||
|
}
|
||||||
|
for i := range got {
|
||||||
|
if got[i] != tt.want[i] {
|
||||||
|
t.Errorf("args[%d] = %q, want %q (full: %v)", i, got[i], tt.want[i], got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user