mirror of
https://github.com/XTLS/Xray-core.git
synced 2026-10-02 05:46:39 +00:00
Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
b26a91de4f | ||
|
|
1f304916bd | ||
|
|
0086362663 | ||
|
|
e51b3c3621 | ||
|
|
6243d2a26e | ||
|
|
35e616d3b9 | ||
|
|
08cb6e6bca | ||
|
|
48ad0300ea | ||
|
|
0fc379203f | ||
|
|
fc8f8a451d | ||
|
|
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)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -82,19 +82,10 @@ func (f *MphDomainMatcherFactory) BuildMatcher(rules []*DomainRule) (DomainMatch
|
|||||||
}
|
}
|
||||||
g.Add(m, uint32(i))
|
g.Add(m, uint32(i))
|
||||||
case *DomainRule_Geosite:
|
case *DomainRule_Geosite:
|
||||||
domains, err := loadSiteWithAttrs(v.Geosite.File, v.Geosite.Code, v.Geosite.Attrs)
|
err := loadSiteMatchers(v.Geosite, func(m strmatcher.Matcher) { g.Add(m, uint32(i)) })
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
for j, d := range domains {
|
|
||||||
domains[j] = nil // peak mem
|
|
||||||
m, err := parseDomain(d)
|
|
||||||
if err != nil {
|
|
||||||
errors.LogError(context.Background(), "ignore invalid geosite entry in ", v.Geosite.File, ":", v.Geosite.Code, " at index ", j, ", ", err)
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
g.Add(m, uint32(i))
|
|
||||||
}
|
|
||||||
default:
|
default:
|
||||||
panic("unknown domain rule type")
|
panic("unknown domain rule type")
|
||||||
}
|
}
|
||||||
@@ -108,12 +99,12 @@ func (f *MphDomainMatcherFactory) BuildMatcher(rules []*DomainRule) (DomainMatch
|
|||||||
return g, nil
|
return g, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
type CompactDomainMatcherFactory struct {
|
type CompactMphDomainMatcherFactory struct {
|
||||||
sync.Mutex
|
sync.Mutex
|
||||||
shared *utils.WeakCacheMap[string, strmatcher.LinearAnyMatcher]
|
shared *utils.WeakCacheMap[string, strmatcher.MphValueMatcher]
|
||||||
}
|
}
|
||||||
|
|
||||||
func (f *CompactDomainMatcherFactory) getOrCreateFrom(rule *GeoSiteRule) (strmatcher.MatcherSet, error) {
|
func (f *CompactMphDomainMatcherFactory) getOrCreateFrom(rule *GeoSiteRule) (*strmatcher.MphValueMatcher, error) {
|
||||||
key := rule.File + ":" + rule.Code + "@" + rule.Attrs
|
key := rule.File + ":" + rule.Code + "@" + rule.Attrs
|
||||||
|
|
||||||
f.Lock()
|
f.Lock()
|
||||||
@@ -125,33 +116,23 @@ func (f *CompactDomainMatcherFactory) getOrCreateFrom(rule *GeoSiteRule) (strmat
|
|||||||
}
|
}
|
||||||
errors.LogDebug(context.Background(), "geodata geosite matcher cache MISS ", key)
|
errors.LogDebug(context.Background(), "geodata geosite matcher cache MISS ", key)
|
||||||
|
|
||||||
s := strmatcher.NewLinearAnyMatcher()
|
s := strmatcher.NewMphValueMatcher()
|
||||||
domains, err := loadSiteWithAttrs(rule.File, rule.Code, rule.Attrs)
|
if err := loadSiteMatchers(rule, func(m strmatcher.Matcher) { s.Add(m, 0) }); err != nil {
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
for i, d := range domains {
|
if err := s.Build(); err != nil {
|
||||||
domains[i] = nil // peak mem
|
return nil, err
|
||||||
m, err := parseDomain(d)
|
|
||||||
if err != nil {
|
|
||||||
errors.LogError(context.Background(), "ignore invalid geosite entry in ", rule.File, ":", rule.Code, " at index ", i, ", ", err)
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
s.Add(m)
|
|
||||||
}
|
}
|
||||||
f.shared.Store(key, s)
|
f.shared.Store(key, s)
|
||||||
return s, err
|
return s, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// BuildMatcher implements DomainMatcherFactory.
|
// BuildMatcher implements DomainMatcherFactory.
|
||||||
func (f *CompactDomainMatcherFactory) BuildMatcher(rules []*DomainRule) (DomainMatcher, error) {
|
func (f *CompactMphDomainMatcherFactory) BuildMatcher(rules []*DomainRule) (DomainMatcher, error) {
|
||||||
if len(rules) == 0 {
|
if len(rules) == 0 {
|
||||||
return nil, errors.New("empty domain rule list")
|
return nil, errors.New("empty domain rule list")
|
||||||
}
|
}
|
||||||
compact := &CompactDomainMatcher{
|
compact := new(CompactMphDomainMatcher)
|
||||||
matchers: make([]strmatcher.MatcherSet, 0, len(rules)),
|
|
||||||
values: make([]uint32, 0, len(rules)),
|
|
||||||
}
|
|
||||||
for i, r := range rules {
|
for i, r := range rules {
|
||||||
switch v := r.Value.(type) {
|
switch v := r.Value.(type) {
|
||||||
case *DomainRule_Custom:
|
case *DomainRule_Custom:
|
||||||
@@ -168,8 +149,7 @@ func (f *CompactDomainMatcherFactory) BuildMatcher(rules []*DomainRule) (DomainM
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
compact.matchers = append(compact.matchers, m)
|
compact.combiner.Add(m, uint32(i))
|
||||||
compact.values = append(compact.values, uint32(i))
|
|
||||||
default:
|
default:
|
||||||
panic("unknown domain rule type")
|
panic("unknown domain rule type")
|
||||||
}
|
}
|
||||||
@@ -177,37 +157,40 @@ func (f *CompactDomainMatcherFactory) BuildMatcher(rules []*DomainRule) (DomainM
|
|||||||
return compact, nil
|
return compact, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
type CompactDomainMatcher struct {
|
type CompactMphDomainMatcher struct {
|
||||||
custom strmatcher.ValueMatcher
|
custom strmatcher.ValueMatcher
|
||||||
matchers []strmatcher.MatcherSet
|
combiner strmatcher.MphValueMatcherCombiner
|
||||||
values []uint32
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// Match implements DomainMatcher.
|
// Match implements DomainMatcher.
|
||||||
func (c *CompactDomainMatcher) Match(input string) []uint32 {
|
func (c *CompactMphDomainMatcher) Match(input string) []uint32 {
|
||||||
var result []uint32
|
result := c.combiner.Match(input)
|
||||||
if c.custom != nil {
|
if c.custom != nil {
|
||||||
result = append(result, c.custom.Match(input)...)
|
result = append(c.custom.Match(input), result...)
|
||||||
}
|
|
||||||
for i, m := range c.matchers {
|
|
||||||
if m.MatchAny(input) {
|
|
||||||
result = append(result, c.values[i])
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
return result
|
return result
|
||||||
}
|
}
|
||||||
|
|
||||||
// MatchAny implements DomainMatcher.
|
// MatchAny implements DomainMatcher.
|
||||||
func (c *CompactDomainMatcher) MatchAny(input string) bool {
|
func (c *CompactMphDomainMatcher) MatchAny(input string) bool {
|
||||||
if c.custom != nil && c.custom.MatchAny(input) {
|
if c.custom != nil && c.custom.MatchAny(input) {
|
||||||
return true
|
return true
|
||||||
}
|
}
|
||||||
for _, m := range c.matchers {
|
return c.combiner.MatchAny(input)
|
||||||
if m.MatchAny(input) {
|
}
|
||||||
return true
|
|
||||||
|
// loadSiteMatchers calls add with a matcher for every domain of the geosite rule and logs the invalid ones.
|
||||||
|
func loadSiteMatchers(rule *GeoSiteRule, add func(strmatcher.Matcher)) error {
|
||||||
|
i := 0
|
||||||
|
return loadSite(rule.File, rule.Code, rule.Attrs, func(t Domain_Type, value []byte) {
|
||||||
|
m, err := parseDomain(&Domain{Type: t, Value: string(value)})
|
||||||
|
if err != nil {
|
||||||
|
errors.LogError(context.Background(), "ignore invalid geosite entry in ", rule.File, ":", rule.Code, " at index ", i, ", ", err)
|
||||||
|
} else {
|
||||||
|
add(m)
|
||||||
}
|
}
|
||||||
}
|
i++
|
||||||
return false
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
func parseDomain(d *Domain) (strmatcher.Matcher, error) {
|
func parseDomain(d *Domain) (strmatcher.Matcher, error) {
|
||||||
@@ -231,7 +214,7 @@ func parseDomain(d *Domain) (strmatcher.Matcher, error) {
|
|||||||
func newDomainMatcherFactory() DomainMatcherFactory {
|
func newDomainMatcherFactory() DomainMatcherFactory {
|
||||||
switch runtime.GOOS {
|
switch runtime.GOOS {
|
||||||
case "ios", "android":
|
case "ios", "android":
|
||||||
return &CompactDomainMatcherFactory{shared: utils.NewWeakCacheMap[string, strmatcher.LinearAnyMatcher]()}
|
return &CompactMphDomainMatcherFactory{shared: utils.NewWeakCacheMap[string, strmatcher.MphValueMatcher]()}
|
||||||
default:
|
default:
|
||||||
return &MphDomainMatcherFactory{shared: utils.NewWeakCacheMap[string, strmatcher.MphValueMatcher]()}
|
return &MphDomainMatcherFactory{shared: utils.NewWeakCacheMap[string, strmatcher.MphValueMatcher]()}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -4,6 +4,7 @@ import (
|
|||||||
"path/filepath"
|
"path/filepath"
|
||||||
"reflect"
|
"reflect"
|
||||||
"slices"
|
"slices"
|
||||||
|
"sync"
|
||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
"github.com/xtls/xray-core/common/geodata/strmatcher"
|
"github.com/xtls/xray-core/common/geodata/strmatcher"
|
||||||
@@ -11,7 +12,7 @@ import (
|
|||||||
)
|
)
|
||||||
|
|
||||||
func TestCompactDomainMatcher_PreservesCustomRuleIndices(t *testing.T) {
|
func TestCompactDomainMatcher_PreservesCustomRuleIndices(t *testing.T) {
|
||||||
factory := &CompactDomainMatcherFactory{shared: utils.NewWeakCacheMap[string, strmatcher.LinearAnyMatcher]()}
|
factory := &CompactMphDomainMatcherFactory{shared: utils.NewWeakCacheMap[string, strmatcher.MphValueMatcher]()}
|
||||||
matcher, err := factory.BuildMatcher([]*DomainRule{
|
matcher, err := factory.BuildMatcher([]*DomainRule{
|
||||||
{Value: &DomainRule_Custom{Custom: &Domain{Type: Domain_Full, Value: "example.com"}}},
|
{Value: &DomainRule_Custom{Custom: &Domain{Type: Domain_Full, Value: "example.com"}}},
|
||||||
{Value: &DomainRule_Custom{Custom: &Domain{Type: Domain_Domain, Value: "example.com"}}},
|
{Value: &DomainRule_Custom{Custom: &Domain{Type: Domain_Domain, Value: "example.com"}}},
|
||||||
@@ -32,7 +33,7 @@ func TestCompactDomainMatcher_PreservesCustomRuleIndices(t *testing.T) {
|
|||||||
func TestCompactDomainMatcher_PreservesMixedRuleIndices(t *testing.T) {
|
func TestCompactDomainMatcher_PreservesMixedRuleIndices(t *testing.T) {
|
||||||
t.Setenv("xray.location.asset", filepath.Join("..", "..", "resources"))
|
t.Setenv("xray.location.asset", filepath.Join("..", "..", "resources"))
|
||||||
|
|
||||||
factory := &CompactDomainMatcherFactory{shared: utils.NewWeakCacheMap[string, strmatcher.LinearAnyMatcher]()}
|
factory := &CompactMphDomainMatcherFactory{shared: utils.NewWeakCacheMap[string, strmatcher.MphValueMatcher]()}
|
||||||
matcher, err := factory.BuildMatcher([]*DomainRule{
|
matcher, err := factory.BuildMatcher([]*DomainRule{
|
||||||
{Value: &DomainRule_Geosite{Geosite: &GeoSiteRule{File: DefaultGeoSiteDat, Code: "CN"}}},
|
{Value: &DomainRule_Geosite{Geosite: &GeoSiteRule{File: DefaultGeoSiteDat, Code: "CN"}}},
|
||||||
{Value: &DomainRule_Custom{Custom: &Domain{Type: Domain_Full, Value: "163.com"}}},
|
{Value: &DomainRule_Custom{Custom: &Domain{Type: Domain_Full, Value: "163.com"}}},
|
||||||
@@ -72,3 +73,76 @@ func TestMphDomainMatcher_MatchReturnsDetachedSlice(t *testing.T) {
|
|||||||
t.Fatalf("Match() after caller mutation = %v, want %v", gotAgain, []uint32{0, 1})
|
t.Fatalf("Match() after caller mutation = %v, want %v", gotAgain, []uint32{0, 1})
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// DNS sorts every Match result in place, so a matcher must never hand out a
|
||||||
|
// slice it keeps, also when only its keyword or regex part matches.
|
||||||
|
func TestDomainMatcher_MatchResultsCanBeSortedConcurrently(t *testing.T) {
|
||||||
|
t.Setenv("xray.location.asset", filepath.Join("..", "..", "resources"))
|
||||||
|
|
||||||
|
rules := []*DomainRule{
|
||||||
|
{Value: &DomainRule_Custom{Custom: &Domain{Type: Domain_Full, Value: "example.com"}}},
|
||||||
|
{Value: &DomainRule_Custom{Custom: &Domain{Type: Domain_Domain, Value: "example.com"}}},
|
||||||
|
{Value: &DomainRule_Custom{Custom: &Domain{Type: Domain_Substr, Value: "exam"}}},
|
||||||
|
{Value: &DomainRule_Custom{Custom: &Domain{Type: Domain_Regex, Value: `^ex.*\.org$`}}},
|
||||||
|
{Value: &DomainRule_Custom{Custom: &Domain{Type: Domain_Substr, Value: "exam"}}},
|
||||||
|
{Value: &DomainRule_Geosite{Geosite: &GeoSiteRule{File: DefaultGeoSiteDat, Code: "CN"}}},
|
||||||
|
{Value: &DomainRule_Custom{Custom: &Domain{Type: Domain_Full, Value: "only.full.test"}}},
|
||||||
|
}
|
||||||
|
cases := []struct {
|
||||||
|
input string
|
||||||
|
want []uint32
|
||||||
|
}{
|
||||||
|
{"example.com", []uint32{0, 1, 2, 4}},
|
||||||
|
{"www.example.com", []uint32{1, 2, 4}},
|
||||||
|
{"exam.net", []uint32{2, 4}}, // keyword part only
|
||||||
|
{"example.org", []uint32{2, 3, 4}},
|
||||||
|
{"163.com", []uint32{5}},
|
||||||
|
{"www.163.com", []uint32{5}},
|
||||||
|
{"only.full.test", []uint32{6}}, // full part only
|
||||||
|
{"nomatch.test", nil},
|
||||||
|
}
|
||||||
|
factories := map[string]DomainMatcherFactory{
|
||||||
|
"mph": &MphDomainMatcherFactory{shared: utils.NewWeakCacheMap[string, strmatcher.MphValueMatcher]()},
|
||||||
|
"compact": &CompactMphDomainMatcherFactory{shared: utils.NewWeakCacheMap[string, strmatcher.MphValueMatcher]()},
|
||||||
|
}
|
||||||
|
for name, factory := range factories {
|
||||||
|
t.Run(name, func(t *testing.T) {
|
||||||
|
matcher, err := factory.BuildMatcher(rules)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("BuildMatcher() failed: %v", err)
|
||||||
|
}
|
||||||
|
for _, c := range cases {
|
||||||
|
got := matcher.Match(c.input)
|
||||||
|
if sorted := slices.Sorted(slices.Values(got)); !slices.Equal(sorted, c.want) {
|
||||||
|
t.Fatalf("Match(%q) = %v, want %v", c.input, sorted, c.want)
|
||||||
|
}
|
||||||
|
got = got[:cap(got)]
|
||||||
|
for j := range got {
|
||||||
|
got[j] = ^uint32(0)
|
||||||
|
}
|
||||||
|
if again := slices.Sorted(slices.Values(matcher.Match(c.input))); !slices.Equal(again, c.want) {
|
||||||
|
t.Fatalf("Match(%q) after caller mutation = %v, want %v", c.input, again, c.want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
var wg sync.WaitGroup
|
||||||
|
for range 8 {
|
||||||
|
wg.Add(1)
|
||||||
|
go func() {
|
||||||
|
defer wg.Done()
|
||||||
|
for range 500 {
|
||||||
|
for _, c := range cases {
|
||||||
|
got := matcher.Match(c.input)
|
||||||
|
slices.Sort(got)
|
||||||
|
if !slices.Equal(got, c.want) {
|
||||||
|
t.Errorf("Match(%q) = %v, want %v", c.input, got, c.want)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
}
|
||||||
|
wg.Wait()
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
+213
-62
@@ -5,11 +5,14 @@ import (
|
|||||||
"bytes"
|
"bytes"
|
||||||
"io"
|
"io"
|
||||||
"runtime"
|
"runtime"
|
||||||
|
"slices"
|
||||||
"strings"
|
"strings"
|
||||||
|
"unicode/utf8"
|
||||||
|
|
||||||
"github.com/xtls/xray-core/common/errors"
|
"github.com/xtls/xray-core/common/errors"
|
||||||
"github.com/xtls/xray-core/common/platform/filesystem"
|
"github.com/xtls/xray-core/common/platform/filesystem"
|
||||||
|
|
||||||
|
"google.golang.org/protobuf/encoding/protowire"
|
||||||
"google.golang.org/protobuf/proto"
|
"google.golang.org/protobuf/proto"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -52,17 +55,56 @@ func loadIP(file, code string) ([]*CIDR, error) {
|
|||||||
return geoip.Cidr, nil
|
return geoip.Cidr, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func loadSite(file, code string) ([]*Domain, error) {
|
// loadSite calls fn, in file order, with the type and value of every domain of the geosite code
|
||||||
bs, err := loadFile(file, code)
|
// that has all the "@"-separated attrs. It decodes the entry while reading the file instead of
|
||||||
|
// unmarshalling it into a []*Domain, so value is only valid during fn.
|
||||||
|
func loadSite(file, code, attrs string, fn func(Domain_Type, []byte)) error {
|
||||||
|
runtime.GC() // peak mem
|
||||||
|
r, err := filesystem.OpenAsset(file)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return errors.New("failed to open ", file).Base(err)
|
||||||
}
|
}
|
||||||
defer runtime.GC() // peak mem
|
defer r.Close()
|
||||||
var geosite GeoSite
|
br := bufio.NewReaderSize(r, 64*1024)
|
||||||
if err := proto.Unmarshal(bs, &geosite); err != nil {
|
n, err := seek(br, []byte(code))
|
||||||
return nil, errors.New("error unmarshal Site in ", file, ":", code).Base(err)
|
if err != nil {
|
||||||
|
return errors.New("failed to load code ", code, " from ", file).Base(err)
|
||||||
}
|
}
|
||||||
return geosite.Domain, nil
|
loadErr := func(err error) error {
|
||||||
|
if err == io.EOF {
|
||||||
|
err = io.ErrUnexpectedEOF
|
||||||
|
}
|
||||||
|
return errors.New("failed to load code ", code, " from ", file).Base(err)
|
||||||
|
}
|
||||||
|
unmarshalErr := func(err error) error {
|
||||||
|
return errors.New("error unmarshal Site in ", file, ":", code).Base(err)
|
||||||
|
}
|
||||||
|
d := newSiteDecoder(attrs, fn)
|
||||||
|
for n > 0 {
|
||||||
|
w, err := br.Peek(min(n, br.Size()))
|
||||||
|
if err != nil {
|
||||||
|
return loadErr(err)
|
||||||
|
}
|
||||||
|
used, err := d.decode(w, len(w) < n)
|
||||||
|
if err != nil {
|
||||||
|
return unmarshalErr(err)
|
||||||
|
}
|
||||||
|
if used == 0 {
|
||||||
|
break // a field longer than the buffer
|
||||||
|
}
|
||||||
|
br.Discard(used)
|
||||||
|
n -= used
|
||||||
|
}
|
||||||
|
if n > 0 {
|
||||||
|
w := make([]byte, n)
|
||||||
|
if _, err := io.ReadFull(br, w); err != nil {
|
||||||
|
return loadErr(err)
|
||||||
|
}
|
||||||
|
if _, err := d.decode(w, false); err != nil {
|
||||||
|
return unmarshalErr(err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func decodeVarint(br *bufio.Reader) (uint64, error) {
|
func decodeVarint(br *bufio.Reader) (uint64, error) {
|
||||||
@@ -82,68 +124,63 @@ func decodeVarint(br *bufio.Reader) (uint64, error) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func find(r io.Reader, code []byte, readBody bool) ([]byte, error) {
|
func find(r io.Reader, code []byte, readBody bool) ([]byte, error) {
|
||||||
|
br := bufio.NewReaderSize(r, 64*1024)
|
||||||
|
bodyL, err := seek(br, code)
|
||||||
|
if err != nil || !readBody {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
out := make([]byte, bodyL)
|
||||||
|
if _, err := io.ReadFull(br, out); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
return out, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// seek advances br to the body of the entry for code and returns the body length.
|
||||||
|
func seek(br *bufio.Reader, code []byte) (int, error) {
|
||||||
codeL := len(code)
|
codeL := len(code)
|
||||||
if codeL == 0 {
|
if codeL == 0 {
|
||||||
return nil, errors.New("empty code")
|
return 0, errors.New("empty code")
|
||||||
}
|
}
|
||||||
|
|
||||||
br := bufio.NewReaderSize(r, 64*1024)
|
|
||||||
need := 2 + codeL // TODO: if code too long
|
need := 2 + codeL // TODO: if code too long
|
||||||
prefixBuf := make([]byte, need)
|
|
||||||
|
|
||||||
for {
|
for {
|
||||||
if _, err := br.ReadByte(); err != nil {
|
if _, err := br.ReadByte(); err != nil {
|
||||||
return nil, err
|
return 0, err
|
||||||
}
|
}
|
||||||
|
|
||||||
x, err := decodeVarint(br)
|
x, err := decodeVarint(br)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return 0, err
|
||||||
}
|
}
|
||||||
bodyL := int(x)
|
bodyL := int(x)
|
||||||
if bodyL <= 0 {
|
if bodyL <= 0 {
|
||||||
return nil, errors.New("invalid body length: ", bodyL)
|
return 0, errors.New("invalid body length: ", bodyL)
|
||||||
}
|
}
|
||||||
|
|
||||||
prefixL := bodyL
|
// Peek no more than the buffer holds: a code longer than the buffer cannot match a single
|
||||||
if prefixL > need {
|
// length byte anyway, so a short peek only skips it, as base find (io.ReadFull) does.
|
||||||
prefixL = need
|
prefix, err := br.Peek(min(bodyL, need, br.Size()))
|
||||||
}
|
if err != nil {
|
||||||
prefix := prefixBuf[:prefixL]
|
if err == io.EOF && len(prefix) > 0 {
|
||||||
if _, err := io.ReadFull(br, prefix); err != nil {
|
err = io.ErrUnexpectedEOF // as io.ReadFull
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
match := false
|
|
||||||
if bodyL >= need {
|
|
||||||
if int(prefix[1]) == codeL && bytes.Equal(prefix[2:need], code) {
|
|
||||||
if !readBody {
|
|
||||||
return nil, nil
|
|
||||||
}
|
|
||||||
match = true
|
|
||||||
}
|
}
|
||||||
|
return 0, err
|
||||||
}
|
}
|
||||||
|
if bodyL >= need && len(prefix) >= need && int(prefix[1]) == codeL && bytes.Equal(prefix[2:], code) {
|
||||||
remain := bodyL - prefixL
|
return bodyL, nil
|
||||||
if match {
|
|
||||||
out := make([]byte, bodyL)
|
|
||||||
copy(out, prefix)
|
|
||||||
if remain > 0 {
|
|
||||||
if _, err := io.ReadFull(br, out[prefixL:]); err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return out, nil
|
|
||||||
}
|
}
|
||||||
|
if _, err := br.Discard(bodyL); err != nil {
|
||||||
if remain > 0 {
|
return 0, err
|
||||||
if _, err := br.Discard(remain); err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// AttributeMatcher, HasAttrMatcher, AllAttrsMatcher and NewAllAttrsMatcher are the exported
|
||||||
|
// attribute helpers that have been part of this package's API since #5814. The streaming loader
|
||||||
|
// above filters attributes itself without building a *Domain, so it does not use them, but they
|
||||||
|
// are kept for external callers. Their behaviour is unchanged.
|
||||||
|
|
||||||
type AttributeMatcher interface {
|
type AttributeMatcher interface {
|
||||||
Match(*Domain) bool
|
Match(*Domain) bool
|
||||||
}
|
}
|
||||||
@@ -185,23 +222,137 @@ func NewAllAttrsMatcher(attrs string) AttributeMatcher {
|
|||||||
return m
|
return m
|
||||||
}
|
}
|
||||||
|
|
||||||
func loadSiteWithAttrs(file, code, attrs string) ([]*Domain, error) {
|
var errInvalidUTF8 = errors.New("string field contains invalid UTF-8")
|
||||||
domains, err := loadSite(file, code)
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
matcher := NewAllAttrsMatcher(attrs)
|
type siteDecoder struct {
|
||||||
if matcher == nil {
|
want []string
|
||||||
return domains, nil
|
has []bool
|
||||||
}
|
fn func(Domain_Type, []byte)
|
||||||
|
}
|
||||||
|
|
||||||
filtered := make([]*Domain, 0, len(domains))
|
func newSiteDecoder(attrs string, fn func(Domain_Type, []byte)) *siteDecoder {
|
||||||
for _, d := range domains {
|
d := &siteDecoder{fn: fn}
|
||||||
if matcher.Match(d) {
|
if attrs != "" {
|
||||||
filtered = append(filtered, d)
|
d.want = strings.Split(attrs, "@")
|
||||||
|
d.has = make([]bool, len(d.want))
|
||||||
|
}
|
||||||
|
return d
|
||||||
|
}
|
||||||
|
|
||||||
|
// decode walks the whole fields at the start of b, a part of an encoded GeoSite (see geodat.proto),
|
||||||
|
// calls fn for every domain that has all attrs and returns how many bytes it used. A field cut off
|
||||||
|
// by the end of b is an error unless more is set. It accepts and rejects what proto.Unmarshal does.
|
||||||
|
func (d *siteDecoder) decode(b []byte, more bool) (int, error) {
|
||||||
|
used := 0
|
||||||
|
for used < len(b) {
|
||||||
|
f, n, err := consumeField(b[used:])
|
||||||
|
if err == io.ErrUnexpectedEOF && more {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
if err != nil {
|
||||||
|
return used, err
|
||||||
|
}
|
||||||
|
used += n
|
||||||
|
if f.typ != protowire.BytesType {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
switch f.num {
|
||||||
|
case 1: // code
|
||||||
|
if !utf8.Valid(f.v) {
|
||||||
|
return used, errInvalidUTF8
|
||||||
|
}
|
||||||
|
case 2: // domain
|
||||||
|
t, value, err := decodeDomain(f.v, d.want, d.has)
|
||||||
|
if err != nil {
|
||||||
|
return used, err
|
||||||
|
}
|
||||||
|
if !slices.Contains(d.has, false) {
|
||||||
|
d.fn(t, value)
|
||||||
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
return used, nil
|
||||||
return filtered, nil
|
}
|
||||||
|
|
||||||
|
// decodeDomain decodes an encoded Domain and sets has[i] if one of its attributes has the key want[i].
|
||||||
|
func decodeDomain(b []byte, want []string, has []bool) (t Domain_Type, value []byte, err error) {
|
||||||
|
clear(has)
|
||||||
|
for len(b) > 0 {
|
||||||
|
f, n, err := consumeField(b)
|
||||||
|
if err != nil {
|
||||||
|
return 0, nil, err
|
||||||
|
}
|
||||||
|
b = b[n:]
|
||||||
|
switch {
|
||||||
|
case f.num == 1 && f.typ == protowire.VarintType: // type
|
||||||
|
t = Domain_Type(f.x)
|
||||||
|
case f.num == 2 && f.typ == protowire.BytesType: // value
|
||||||
|
if !utf8.Valid(f.v) {
|
||||||
|
return 0, nil, errInvalidUTF8
|
||||||
|
}
|
||||||
|
value = f.v
|
||||||
|
case f.num == 3 && f.typ == protowire.BytesType: // attribute
|
||||||
|
key, err := decodeAttributeKey(f.v)
|
||||||
|
if err != nil {
|
||||||
|
return 0, nil, err
|
||||||
|
}
|
||||||
|
for i, w := range want {
|
||||||
|
if string(key) == w {
|
||||||
|
has[i] = true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return t, value, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// decodeAttributeKey returns the key of an encoded Domain.Attribute.
|
||||||
|
func decodeAttributeKey(b []byte) ([]byte, error) {
|
||||||
|
var key []byte
|
||||||
|
for len(b) > 0 {
|
||||||
|
f, n, err := consumeField(b)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
b = b[n:]
|
||||||
|
if f.num == 1 && f.typ == protowire.BytesType {
|
||||||
|
if !utf8.Valid(f.v) {
|
||||||
|
return nil, errInvalidUTF8
|
||||||
|
}
|
||||||
|
key = f.v
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return key, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
type protoField struct {
|
||||||
|
num protowire.Number
|
||||||
|
typ protowire.Type
|
||||||
|
v []byte // payload of a length-delimited field
|
||||||
|
x uint64 // value of a varint field
|
||||||
|
}
|
||||||
|
|
||||||
|
// consumeField parses the first field of an encoded message and returns it with its length.
|
||||||
|
func consumeField(b []byte) (protoField, int, error) {
|
||||||
|
num, typ, n := protowire.ConsumeTag(b)
|
||||||
|
if n < 0 {
|
||||||
|
return protoField{}, 0, protowire.ParseError(n)
|
||||||
|
}
|
||||||
|
if num > protowire.MaxValidNumber {
|
||||||
|
return protoField{}, 0, errors.New("invalid field number ", num)
|
||||||
|
}
|
||||||
|
f := protoField{num: num, typ: typ}
|
||||||
|
var m int
|
||||||
|
switch typ {
|
||||||
|
case protowire.BytesType:
|
||||||
|
f.v, m = protowire.ConsumeBytes(b[n:])
|
||||||
|
case protowire.VarintType:
|
||||||
|
f.x, m = protowire.ConsumeVarint(b[n:])
|
||||||
|
default:
|
||||||
|
m = protowire.ConsumeFieldValue(num, typ, b[n:])
|
||||||
|
}
|
||||||
|
if m < 0 {
|
||||||
|
return protoField{}, 0, protowire.ParseError(m)
|
||||||
|
}
|
||||||
|
return f, n + m, nil
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -0,0 +1,283 @@
|
|||||||
|
package geodata
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"slices"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"google.golang.org/protobuf/encoding/protowire"
|
||||||
|
"google.golang.org/protobuf/proto"
|
||||||
|
)
|
||||||
|
|
||||||
|
type siteEntry struct {
|
||||||
|
Type Domain_Type
|
||||||
|
Value string
|
||||||
|
}
|
||||||
|
|
||||||
|
// unmarshalSite is what loadSite used to do: proto.Unmarshal, then keep the domains that have all attrs.
|
||||||
|
func unmarshalSite(b []byte, attrs string) ([]siteEntry, error) {
|
||||||
|
var site GeoSite
|
||||||
|
if err := proto.Unmarshal(b, &site); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
var entries []siteEntry
|
||||||
|
for _, d := range site.Domain {
|
||||||
|
ok := true
|
||||||
|
for _, key := range strings.Split(attrs, "@") {
|
||||||
|
ok = ok && (attrs == "" || slices.ContainsFunc(d.Attribute, func(a *Domain_Attribute) bool { return a.Key == key }))
|
||||||
|
}
|
||||||
|
if ok {
|
||||||
|
entries = append(entries, siteEntry{d.Type, d.Value})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return entries, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func checkDecodeSite(t *testing.T, name string, b []byte, attrs string) {
|
||||||
|
t.Helper()
|
||||||
|
want, wantErr := unmarshalSite(b, attrs)
|
||||||
|
var got []siteEntry
|
||||||
|
_, err := newSiteDecoder(attrs, func(typ Domain_Type, value []byte) {
|
||||||
|
got = append(got, siteEntry{typ, string(value)})
|
||||||
|
}).decode(b, false)
|
||||||
|
if (err == nil) != (wantErr == nil) {
|
||||||
|
t.Fatalf("%s@%s: error %v, proto.Unmarshal: %v", name, attrs, err, wantErr)
|
||||||
|
}
|
||||||
|
if err == nil && !slices.Equal(got, want) {
|
||||||
|
t.Fatalf("%s@%s: got %v, want %v", name, attrs, got, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestDecodeSiteMatchesUnmarshal(t *testing.T) {
|
||||||
|
bs, err := os.ReadFile(filepath.Join("..", "..", "resources", DefaultGeoSiteDat))
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
for len(bs) > 0 {
|
||||||
|
num, typ, n := protowire.ConsumeTag(bs)
|
||||||
|
if n < 0 || num != 1 || typ != protowire.BytesType {
|
||||||
|
t.Fatal("unexpected GeoSiteList field")
|
||||||
|
}
|
||||||
|
entry, m := protowire.ConsumeBytes(bs[n:])
|
||||||
|
if m < 0 {
|
||||||
|
t.Fatal(protowire.ParseError(m))
|
||||||
|
}
|
||||||
|
bs = bs[n+m:]
|
||||||
|
|
||||||
|
var site GeoSite
|
||||||
|
if err := proto.Unmarshal(entry, &site); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
queries := []string{"", "none"}
|
||||||
|
for _, d := range site.Domain {
|
||||||
|
for _, a := range d.Attribute {
|
||||||
|
if !slices.Contains(queries, a.Key) {
|
||||||
|
queries = append(queries, a.Key, a.Key+"@none")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
for _, attrs := range queries {
|
||||||
|
checkDecodeSite(t, site.Code, entry, attrs)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestDecodeSiteUnusualEncodings(t *testing.T) {
|
||||||
|
field := func(num protowire.Number, v []byte) []byte {
|
||||||
|
return protowire.AppendBytes(protowire.AppendTag(nil, num, protowire.BytesType), v)
|
||||||
|
}
|
||||||
|
typ := func(v Domain_Type) []byte {
|
||||||
|
return protowire.AppendVarint(protowire.AppendTag(nil, 1, protowire.VarintType), uint64(v))
|
||||||
|
}
|
||||||
|
value := func(s string) []byte { return field(2, []byte(s)) }
|
||||||
|
attr := func(keys ...string) []byte {
|
||||||
|
var b []byte
|
||||||
|
for _, k := range keys {
|
||||||
|
b = append(b, field(1, []byte(k))...)
|
||||||
|
}
|
||||||
|
return field(3, b)
|
||||||
|
}
|
||||||
|
domain := func(fields ...[]byte) []byte { return field(2, slices.Concat(fields...)) }
|
||||||
|
unknown := protowire.AppendFixed32(protowire.AppendTag(nil, 9, protowire.Fixed32Type), 1)
|
||||||
|
|
||||||
|
for name, b := range map[string][]byte{
|
||||||
|
"unknown field": domain(typ(Domain_Full), unknown, value("example.com")),
|
||||||
|
"repeated value": domain(value("a.com"), typ(Domain_Full), value("b.com")),
|
||||||
|
"repeated type": domain(typ(Domain_Full), value("a.com"), typ(Domain_Regex)),
|
||||||
|
"repeated key": domain(value("a.com"), attr("cn", "ads")),
|
||||||
|
"type as bytes": domain(field(1, []byte("x")), value("a.com")),
|
||||||
|
"no value": domain(typ(Domain_Domain), attr("cn")),
|
||||||
|
"truncated": domain(typ(Domain_Full), value("example.com"))[:10],
|
||||||
|
"invalid utf8": domain(value("example.\xff")),
|
||||||
|
"invalid key": domain(value("a.com"), attr("\xff")),
|
||||||
|
"bad field": protowire.AppendVarint(protowire.AppendTag(nil, protowire.MaxValidNumber+1, protowire.VarintType), 1),
|
||||||
|
"stray end group": protowire.AppendTag(nil, 5, protowire.EndGroupType),
|
||||||
|
} {
|
||||||
|
for _, attrs := range []string{"", "cn", "ads", "cn@ads"} {
|
||||||
|
checkDecodeSite(t, name, b, attrs)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestLoadSiteReadsInPieces covers what real lists never do: an entry far longer than the read
|
||||||
|
// buffer, with a field longer than the buffer in the middle, and a file cut short.
|
||||||
|
func TestLoadSiteReadsInPieces(t *testing.T) {
|
||||||
|
site := &GeoSite{Code: "BIG"}
|
||||||
|
for i := range 5000 {
|
||||||
|
d := &Domain{Type: Domain_Domain, Value: strings.Repeat("x", i%40) + ".example.com"}
|
||||||
|
if i%3 == 0 {
|
||||||
|
d.Attribute = []*Domain_Attribute{{Key: "cn"}}
|
||||||
|
}
|
||||||
|
if i == 2500 {
|
||||||
|
d = &Domain{Type: Domain_Regex, Value: strings.Repeat("a", 100_000)}
|
||||||
|
}
|
||||||
|
site.Domain = append(site.Domain, d)
|
||||||
|
}
|
||||||
|
list := &GeoSiteList{Entry: []*GeoSite{{Code: "SMALL", Domain: []*Domain{{Type: Domain_Full, Value: "a.com"}}}, site}}
|
||||||
|
bs, err := proto.Marshal(list)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
entry, err := proto.Marshal(site)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
dir := t.TempDir()
|
||||||
|
t.Setenv("xray.location.asset", dir)
|
||||||
|
write := func(b []byte) {
|
||||||
|
if err := os.WriteFile(filepath.Join(dir, "big.dat"), b, 0o644); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
for _, attrs := range []string{"", "cn"} {
|
||||||
|
want, _ := unmarshalSite(entry, attrs)
|
||||||
|
var got []siteEntry
|
||||||
|
write(bs)
|
||||||
|
err := loadSite("big.dat", "BIG", attrs, func(typ Domain_Type, value []byte) {
|
||||||
|
got = append(got, siteEntry{typ, string(value)})
|
||||||
|
})
|
||||||
|
if err != nil || !slices.Equal(got, want) {
|
||||||
|
t.Fatalf("attrs %q: %d entries, want %d, error %v", attrs, len(got), len(want), err)
|
||||||
|
}
|
||||||
|
for _, cut := range []int{30_000, len(bs) - 150_000, len(bs) - 1} {
|
||||||
|
write(bs[:cut])
|
||||||
|
if err := loadSite("big.dat", "BIG", attrs, func(Domain_Type, []byte) {}); err == nil {
|
||||||
|
t.Fatalf("file cut at %d of %d: no error", cut, len(bs))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// oneEntryGeoSiteFile wraps an encoded GeoSite as a one-entry GeoSiteList, the file loadSite reads.
|
||||||
|
func oneEntryGeoSiteFile(entry []byte) []byte {
|
||||||
|
return protowire.AppendBytes(protowire.AppendTag(nil, 1, protowire.BytesType), entry)
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestLoadSiteWindowedMatchesSingleShot checks that the windowed reader in loadSite (its Peek/Discard
|
||||||
|
// loop, the more-break when a field is cut by a window edge, the used==0 fallback for a field longer
|
||||||
|
// than the buffer, and the tail path) reaches exactly the same result as decoding the whole entry at
|
||||||
|
// once, for a category several 64 KiB windows long, valid and then mutated near a window edge and
|
||||||
|
// early in the file: same error-or-not, and the same emitted (type, value) sequence when both accept.
|
||||||
|
func TestLoadSiteWindowedMatchesSingleShot(t *testing.T) {
|
||||||
|
const window = 64 * 1024
|
||||||
|
site := &GeoSite{Code: "BIG"}
|
||||||
|
for i := range 12000 { // ~250 KiB, four windows
|
||||||
|
d := &Domain{Type: Domain_Domain, Value: fmt.Sprintf("host%d.%s.example.com", i, strings.Repeat("y", i%30))}
|
||||||
|
if i%3 == 0 {
|
||||||
|
d.Attribute = []*Domain_Attribute{{Key: "cn"}}
|
||||||
|
}
|
||||||
|
site.Domain = append(site.Domain, d)
|
||||||
|
}
|
||||||
|
// a field longer than the buffer, straddling the third window, to force the used==0 fallback
|
||||||
|
site.Domain = slices.Insert(site.Domain, 8000, &Domain{Type: Domain_Regex, Value: strings.Repeat("a", 90_000)})
|
||||||
|
entry, err := proto.Marshal(site)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
dir := t.TempDir()
|
||||||
|
t.Setenv("xray.location.asset", dir)
|
||||||
|
|
||||||
|
// mutations of the encoded entry: unchanged, a byte flipped at several offsets (early windows and
|
||||||
|
// either side of a window edge), and truncations at the same places.
|
||||||
|
type mut struct {
|
||||||
|
name string
|
||||||
|
make func([]byte) []byte
|
||||||
|
}
|
||||||
|
muts := []mut{{"valid", func(b []byte) []byte { return b }}}
|
||||||
|
for _, off := range []int{3, 40, 4000, window - 1, window, window + 1, 2*window - 2, 2 * window} {
|
||||||
|
if off < len(entry) {
|
||||||
|
off := off
|
||||||
|
muts = append(muts, mut{fmt.Sprintf("flip@%d", off), func(b []byte) []byte {
|
||||||
|
c := slices.Clone(b)
|
||||||
|
c[off] ^= 0xff
|
||||||
|
return c
|
||||||
|
}})
|
||||||
|
muts = append(muts, mut{fmt.Sprintf("cut@%d", off), func(b []byte) []byte { return slices.Clone(b[:off]) }})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, attrs := range []string{"", "cn"} {
|
||||||
|
for _, m := range muts {
|
||||||
|
e := m.make(entry)
|
||||||
|
// single-shot reference: decode the whole entry in one call
|
||||||
|
var want []siteEntry
|
||||||
|
_, wantErr := newSiteDecoder(attrs, func(typ Domain_Type, value []byte) {
|
||||||
|
want = append(want, siteEntry{typ, string(value)})
|
||||||
|
}).decode(e, false)
|
||||||
|
// windowed: loadSite reads the file 64 KiB at a time
|
||||||
|
if err := os.WriteFile(filepath.Join(dir, "w.dat"), oneEntryGeoSiteFile(e), 0o644); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
var got []siteEntry
|
||||||
|
gotErr := loadSite("w.dat", "BIG", attrs, func(typ Domain_Type, value []byte) {
|
||||||
|
got = append(got, siteEntry{typ, string(value)})
|
||||||
|
})
|
||||||
|
if (gotErr == nil) != (wantErr == nil) {
|
||||||
|
t.Fatalf("%s attrs=%q: windowed err %v, single-shot err %v", m.name, attrs, gotErr, wantErr)
|
||||||
|
}
|
||||||
|
if gotErr == nil && !slices.Equal(got, want) {
|
||||||
|
t.Fatalf("%s attrs=%q: windowed got %d entries, single-shot %d", m.name, attrs, len(got), len(want))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestLoadSiteLongCode covers a geosite entry whose code is longer than the 64 KiB read buffer. seek
|
||||||
|
// must skip it (find compares a single length byte, so it never matches such a code) and still find a
|
||||||
|
// later entry, and looking the long code up must fail cleanly, like a missing code, not panic.
|
||||||
|
func TestLoadSiteLongCode(t *testing.T) {
|
||||||
|
longCode := strings.Repeat("Z", 70000)
|
||||||
|
list := &GeoSiteList{Entry: []*GeoSite{
|
||||||
|
{Code: "FIRST", Domain: []*Domain{{Type: Domain_Full, Value: "first.com"}}},
|
||||||
|
{Code: longCode, Domain: []*Domain{{Type: Domain_Full, Value: "huge.com"}}},
|
||||||
|
{Code: "AFTER", Domain: []*Domain{{Type: Domain_Domain, Value: "after.com"}}},
|
||||||
|
}}
|
||||||
|
bs, err := proto.Marshal(list)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
dir := t.TempDir()
|
||||||
|
t.Setenv("xray.location.asset", dir)
|
||||||
|
if err := os.WriteFile(filepath.Join(dir, "lc.dat"), bs, 0o644); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
collect := func(code string) ([]siteEntry, error) {
|
||||||
|
var got []siteEntry
|
||||||
|
err := loadSite("lc.dat", code, "", func(typ Domain_Type, value []byte) {
|
||||||
|
got = append(got, siteEntry{typ, string(value)})
|
||||||
|
})
|
||||||
|
return got, err
|
||||||
|
}
|
||||||
|
if got, err := collect("FIRST"); err != nil || !slices.Equal(got, []siteEntry{{Domain_Full, "first.com"}}) {
|
||||||
|
t.Fatalf("FIRST: %v %v", got, err)
|
||||||
|
}
|
||||||
|
if got, err := collect("AFTER"); err != nil || !slices.Equal(got, []siteEntry{{Domain_Domain, "after.com"}}) {
|
||||||
|
t.Fatalf("AFTER (past the oversized entry): %v %v", got, err)
|
||||||
|
}
|
||||||
|
if _, err := collect(longCode); err == nil {
|
||||||
|
t.Fatal("oversized code: expected a not-found error, got nil")
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -1,6 +1,7 @@
|
|||||||
package strmatcher_test
|
package strmatcher_test
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"regexp"
|
||||||
"strconv"
|
"strconv"
|
||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
@@ -72,6 +73,64 @@ func BenchmarkSubstrMatcher(b *testing.B) {
|
|||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func BenchmarkRegexMatcher(b *testing.B) {
|
||||||
|
patterns := []string{ // taken from geosite
|
||||||
|
`(^|\.)91porn\.(best|com|cool|fun|group|party|plus|site|tw|work)$`,
|
||||||
|
`(^|\.)91porn[0-9]{3}\.me$`,
|
||||||
|
`(^|\.)apiproxy-device-prod-nlb-.+\.amazonaws\.com$`,
|
||||||
|
`(^|\.)dualstack\.apiproxy-.+\.amazonaws\.com$`,
|
||||||
|
`(^|\.)aqdk[0-9]{3}\.com$`,
|
||||||
|
`(^|\.)bilibili3(0[1-9]|1[0-2])\.xyz$`,
|
||||||
|
`(^|\.)byyum([3589]|2[235689]|3[34]|4[1-9]|5[1-79]|6[0134679])?\.com$`,
|
||||||
|
`(^|\.)fiftymvapi\..+$`,
|
||||||
|
`(^|\.)gossipfuli[0-9]{3,4}\.xyz$`,
|
||||||
|
`(^|\.)kpkuang\.(bond|fun|info|one|us)$`,
|
||||||
|
`(^|\.)rule34\.(asia|us|world|xxx|xyz)$`,
|
||||||
|
`(^|\.)[a-z][1-9][0-9][a-z]\.com$`,
|
||||||
|
`.+\.awsdns-[0-9][0-9]\.(co\.uk|com|net|org)$`,
|
||||||
|
`.+\.dkr\.ecr\.[^\.]+\.amazonaws\.com$`,
|
||||||
|
`^(.+\.)*zh\.okaapps\.com$`,
|
||||||
|
`^cdn\d-epicgames-\d+\.file\.myqcloud\.com$`,
|
||||||
|
`^chatgpt-async-webps-prod-\S+-\d+\.webpubsub\.azure\.com$`,
|
||||||
|
`^r+[0-9]+(---|\.)sn-(2x3|ni5|j5o)\w{5}\.googlevideo\.com$`,
|
||||||
|
`^speed\.(coe|open)\.ad\.[a-z]{2,6}\.prod\.hosts\.ooklaserver\.net$`,
|
||||||
|
`javdb\d+\.com$`,
|
||||||
|
}
|
||||||
|
domains := []string{
|
||||||
|
"www.google.com", "rr3---sn-4g5edndy.googlevideo.com", "r1---sn-2x3abcde.googlevideo.com", "i.ytimg.com",
|
||||||
|
"graph.facebook.com", "api.twitter.com", "www.baidu.com", "github.com", "objects.githubusercontent.com",
|
||||||
|
"login.microsoftonline.com", "e1234.dscb.akamaiedge.net", "d1a2b3c4d5e6f7.cloudfront.net",
|
||||||
|
"s3.us-east-1.amazonaws.com", "123456789012.dkr.ecr.us-east-1.amazonaws.com", "www.wikipedia.org",
|
||||||
|
"discord.com", "telegram.org", "store.steampowered.com", "www.91porn.com", "ns-1234.awsdns-12.org",
|
||||||
|
}
|
||||||
|
bench := func(b *testing.B, ctor func(pattern string) func(string) bool) {
|
||||||
|
var matchers []func(string) bool
|
||||||
|
for _, p := range patterns {
|
||||||
|
matchers = append(matchers, ctor(p))
|
||||||
|
}
|
||||||
|
b.ResetTimer()
|
||||||
|
for i := 0; i < b.N; i++ {
|
||||||
|
for _, d := range domains {
|
||||||
|
for _, match := range matchers {
|
||||||
|
_ = match(d)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
b.Run("regexp", func(b *testing.B) {
|
||||||
|
bench(b, func(pattern string) func(string) bool {
|
||||||
|
return regexp.MustCompile(pattern).MatchString
|
||||||
|
})
|
||||||
|
})
|
||||||
|
b.Run("prefilter", func(b *testing.B) {
|
||||||
|
bench(b, func(pattern string) func(string) bool {
|
||||||
|
m, err := Regex.New(pattern)
|
||||||
|
common.Must(err)
|
||||||
|
return m.Match
|
||||||
|
})
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
// Utility functions for benchmark
|
// Utility functions for benchmark
|
||||||
|
|
||||||
func benchmarkMatcherType(b *testing.B, t Type, ctor func() MatcherGroup) {
|
func benchmarkMatcherType(b *testing.B, t Type, ctor func() MatcherGroup) {
|
||||||
|
|||||||
@@ -52,7 +52,9 @@ func (g *MphIndexMatcher) Add(matcher Matcher) uint32 {
|
|||||||
func (g *MphIndexMatcher) Build() error {
|
func (g *MphIndexMatcher) Build() error {
|
||||||
if g.mph != nil {
|
if g.mph != nil {
|
||||||
runtime.GC() // peak mem
|
runtime.GC() // peak mem
|
||||||
g.mph.Build()
|
if err := g.mph.Build(); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
}
|
}
|
||||||
runtime.GC() // peak mem
|
runtime.GC() // peak mem
|
||||||
if g.ac != nil {
|
if g.ac != nil {
|
||||||
@@ -64,23 +66,17 @@ func (g *MphIndexMatcher) Build() error {
|
|||||||
|
|
||||||
// Match implements IndexMatcher.Match.
|
// Match implements IndexMatcher.Match.
|
||||||
func (g *MphIndexMatcher) Match(input string) []uint32 {
|
func (g *MphIndexMatcher) Match(input string) []uint32 {
|
||||||
result := make([][]uint32, 0, 5)
|
var result []uint32
|
||||||
if g.mph != nil {
|
if g.mph != nil {
|
||||||
if matches := g.mph.Match(input); len(matches) > 0 {
|
result = g.mph.Match(input) // a new slice, returned without another copy
|
||||||
result = append(result, matches)
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
if g.ac != nil {
|
if g.ac != nil {
|
||||||
if matches := g.ac.Match(input); len(matches) > 0 {
|
result = append(result, g.ac.Match(input)...)
|
||||||
result = append(result, matches)
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
if g.regex != nil {
|
if g.regex != nil {
|
||||||
if matches := g.regex.Match(input); len(matches) > 0 {
|
result = append(result, g.regex.Match(input)...)
|
||||||
result = append(result, matches)
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
return CompositeMatches(result)
|
return result
|
||||||
}
|
}
|
||||||
|
|
||||||
// MatchAny implements IndexMatcher.MatchAny.
|
// MatchAny implements IndexMatcher.MatchAny.
|
||||||
|
|||||||
@@ -78,6 +78,10 @@ func TestMphIndexMatcher(t *testing.T) {
|
|||||||
Input: "example.com",
|
Input: "example.com",
|
||||||
Output: []uint32{10, 4},
|
Output: []uint32{10, 4},
|
||||||
},
|
},
|
||||||
|
{
|
||||||
|
Input: "apis.org",
|
||||||
|
Output: []uint32{2, 6},
|
||||||
|
},
|
||||||
}
|
}
|
||||||
matcherGroup := NewMphIndexMatcher()
|
matcherGroup := NewMphIndexMatcher()
|
||||||
for _, rule := range rules {
|
for _, rule := range rules {
|
||||||
@@ -87,8 +91,13 @@ func TestMphIndexMatcher(t *testing.T) {
|
|||||||
}
|
}
|
||||||
matcherGroup.Build()
|
matcherGroup.Build()
|
||||||
for _, test := range cases {
|
for _, test := range cases {
|
||||||
if m := matcherGroup.Match(test.Input); !reflect.DeepEqual(m, test.Output) {
|
m := matcherGroup.Match(test.Input)
|
||||||
|
if !reflect.DeepEqual(m, test.Output) {
|
||||||
t.Error("unexpected output: ", m, " for test case ", test)
|
t.Error("unexpected output: ", m, " for test case ", test)
|
||||||
}
|
}
|
||||||
|
clear(m) // the caller owns the result, so this must not change the next one
|
||||||
|
if m := matcherGroup.Match(test.Input); !reflect.DeepEqual(m, test.Output) {
|
||||||
|
t.Error("unexpected output after clearing the previous one: ", m, " for test case ", test)
|
||||||
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1,198 +1,440 @@
|
|||||||
package strmatcher
|
package strmatcher
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"math/bits"
|
"bytes"
|
||||||
"runtime"
|
"cmp"
|
||||||
"sort"
|
"encoding/binary"
|
||||||
|
"errors"
|
||||||
|
"math"
|
||||||
|
"slices"
|
||||||
"strings"
|
"strings"
|
||||||
"unsafe"
|
"unsafe"
|
||||||
)
|
)
|
||||||
|
|
||||||
// PrimeRK is the prime base used in Rabin-Karp algorithm.
|
// Flags of a level1 slot, stored above the record offset.
|
||||||
const PrimeRK = 16777619
|
|
||||||
|
|
||||||
// RollingHash calculates the rolling murmurHash of given string based on a provided suffix hash.
|
|
||||||
func RollingHash(hash uint32, input string) uint32 {
|
|
||||||
for i := len(input) - 1; i >= 0; i-- {
|
|
||||||
hash = hash*PrimeRK + uint32(input[i])
|
|
||||||
}
|
|
||||||
return hash
|
|
||||||
}
|
|
||||||
|
|
||||||
// MemHash is the hash function used by go map, it utilizes available hardware instructions(behaves
|
|
||||||
// as aeshash if aes instruction is available).
|
|
||||||
// With different seed, each MemHash<seed> performs as distinct hash functions.
|
|
||||||
func MemHash(seed uint32, input string) uint32 {
|
|
||||||
return uint32(strhash(unsafe.Pointer(&input), uintptr(seed))) // nosemgrep
|
|
||||||
}
|
|
||||||
|
|
||||||
const (
|
const (
|
||||||
mphMatchTypeCount = 2 // Full and Domain
|
mphDomain = 1 << 31 // matches the pattern and its subdomains
|
||||||
|
mphFull = 1 << 30 // matches the pattern only
|
||||||
|
mphParent = 1 << 29 // matches subdomains only, from a pattern with a leading dot
|
||||||
|
mphOffMask = mphParent - 1
|
||||||
)
|
)
|
||||||
|
|
||||||
type mphRuleInfo struct {
|
// Kinds of an added pattern, indexes of mphKinds.
|
||||||
rollingHash uint32
|
const (
|
||||||
matchers [mphMatchTypeCount][]uint32
|
mphKindFull = iota
|
||||||
|
mphKindParent
|
||||||
|
mphKindDomain
|
||||||
|
)
|
||||||
|
|
||||||
|
// mphKinds are the slot flags in the order Match reports their values.
|
||||||
|
var mphKinds = [...]uint32{mphFull, mphParent, mphDomain}
|
||||||
|
|
||||||
|
// mphMultipliers are odd multipliers for the suffix hash. Build moves to the next one if two patterns collide.
|
||||||
|
var mphMultipliers = [...]uint64{0x9e3779b97f4a7c15, 0xc2b2ae3d27d4eb4f, 0x165667b19e3779f9, 0x27d4eb2f165667c5}
|
||||||
|
|
||||||
|
var (
|
||||||
|
errMphCollision = errors.New("strmatcher: suffix hash collision in MphMatcherGroup")
|
||||||
|
errMphBuilt = errors.New("strmatcher: MphMatcherGroup is already built")
|
||||||
|
)
|
||||||
|
|
||||||
|
type mphEntry struct {
|
||||||
|
off uint32 // pattern start in buf
|
||||||
|
value uint32
|
||||||
|
n uint32 // pattern length
|
||||||
|
kind uint8
|
||||||
}
|
}
|
||||||
|
|
||||||
// MphMatcherGroup is an implementation of MatcherGroup.
|
// MphMatcherGroup is an implementation of MatcherGroup for Full and Domain matchers.
|
||||||
// It implements Rabin-Karp algorithm and minimal perfect hash table for Full and Domain matcher.
|
// Each distinct pattern is stored once as a record in arena: its length (255 means a uvarint length follows),
|
||||||
|
// its bytes and, if the group holds more than one distinct value, its values. A minimal perfect hash table
|
||||||
|
// built with hash, displace and compress (http://cmph.sourceforge.net/papers/esa09.pdf) maps a pattern to its
|
||||||
|
// record. Patterns are hashed from the right, so one pass over the input hashes all its parent domains.
|
||||||
type MphMatcherGroup struct {
|
type MphMatcherGroup struct {
|
||||||
rules []string // RuleIdx -> pattern string, index 0 reserved for failed lookup
|
arena string
|
||||||
values [][]uint32 // RuleIdx -> registered matcher values for the pattern (Full Matcher takes precedence)
|
level0 []uint16 // bucket -> seed
|
||||||
level0 []uint32 // RollingHash & Mask -> seed for Memhash
|
level1 []uint32 // slot -> flags | record offset
|
||||||
level0Mask uint32 // Mask restricting RollingHash to 0 ~ len(level0)
|
fp []uint8 // slot -> low byte of its pattern's hash, rejects most misses without reading arena
|
||||||
level1 []uint32 // Memhash<seed> & Mask -> stored index for rules
|
n0, n1 uint32
|
||||||
level1Mask uint32 // Mask for restricting Memhash<seed> to 0 ~ len(level1)
|
mul uint64 // multiplier of the suffix hash
|
||||||
ruleInfos *map[string]mphRuleInfo
|
single uint32 // the only value if !multi
|
||||||
|
multi bool
|
||||||
|
|
||||||
|
buf []byte // build only, patterns in Add order
|
||||||
|
entries []mphEntry
|
||||||
}
|
}
|
||||||
|
|
||||||
func NewMphMatcherGroup() *MphMatcherGroup {
|
func NewMphMatcherGroup() *MphMatcherGroup {
|
||||||
return &MphMatcherGroup{
|
return new(MphMatcherGroup)
|
||||||
rules: []string{""},
|
|
||||||
values: [][]uint32{nil},
|
|
||||||
level0: nil,
|
|
||||||
level0Mask: 0,
|
|
||||||
level1: nil,
|
|
||||||
level1Mask: 0,
|
|
||||||
ruleInfos: &map[string]mphRuleInfo{}, // Only used for building, destroyed after build complete
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// AddFullMatcher implements MatcherGroupForFull.
|
// AddFullMatcher implements MatcherGroupForFull.
|
||||||
func (g *MphMatcherGroup) AddFullMatcher(matcher FullMatcher, value uint32) {
|
func (g *MphMatcherGroup) AddFullMatcher(matcher FullMatcher, value uint32) {
|
||||||
pattern := strings.ToLower(matcher.Pattern())
|
g.add(matcher.Pattern(), mphKindFull, value)
|
||||||
g.addPattern(0, "", pattern, matcher.Type(), value)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// AddDomainMatcher implements MatcherGroupForDomain.
|
// AddDomainMatcher implements MatcherGroupForDomain.
|
||||||
func (g *MphMatcherGroup) AddDomainMatcher(matcher DomainMatcher, value uint32) {
|
func (g *MphMatcherGroup) AddDomainMatcher(matcher DomainMatcher, value uint32) {
|
||||||
pattern := strings.ToLower(matcher.Pattern())
|
g.add(matcher.Pattern(), mphKindDomain, value)
|
||||||
hash := g.addPattern(0, "", pattern, matcher.Type(), value) // For full domain match
|
|
||||||
g.addPattern(hash, pattern, ".", matcher.Type(), value) // For partial domain match
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (g *MphMatcherGroup) addPattern(suffixHash uint32, suffixPattern string, pattern string, matcherType Type, value uint32) uint32 {
|
func (g *MphMatcherGroup) add(pattern string, kind uint8, value uint32) {
|
||||||
fullPattern := pattern + suffixPattern
|
if g.arena != "" {
|
||||||
info, found := (*g.ruleInfos)[fullPattern]
|
panic(errMphBuilt)
|
||||||
if !found {
|
}
|
||||||
info = mphRuleInfo{rollingHash: RollingHash(suffixHash, pattern)}
|
pattern = strings.ToLower(pattern)
|
||||||
g.rules = append(g.rules, fullPattern)
|
off := uint32(len(g.buf))
|
||||||
g.values = append(g.values, nil)
|
g.buf = append(g.buf, pattern...)
|
||||||
|
g.entries = append(g.entries, mphEntry{off: off, value: value, n: uint32(len(pattern)), kind: kind})
|
||||||
|
if len(pattern) > 0 && pattern[0] == '.' {
|
||||||
|
// ".x" has always matched "*.x" as well, so it also gets a parent-only record for "x"
|
||||||
|
g.entries = append(g.entries, mphEntry{off: off + 1, value: value, n: uint32(len(pattern) - 1), kind: mphKindParent})
|
||||||
}
|
}
|
||||||
info.matchers[matcherType] = append(info.matchers[matcherType], value)
|
|
||||||
(*g.ruleInfos)[fullPattern] = info
|
|
||||||
return info.rollingHash
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// Build builds a minimal perfect hash table for insert rules.
|
func (g *MphMatcherGroup) key(i uint32) []byte {
|
||||||
// Algorithm used: Hash, displace, and compress. See http://cmph.sourceforge.net/papers/esa09.pdf
|
e := &g.entries[i]
|
||||||
|
return g.buf[e.off : e.off+e.n]
|
||||||
|
}
|
||||||
|
|
||||||
|
// Build builds the hash table. It must be called once, after the last Add.
|
||||||
func (g *MphMatcherGroup) Build() error {
|
func (g *MphMatcherGroup) Build() error {
|
||||||
ruleCount := len(*g.ruleInfos)
|
if g.arena != "" {
|
||||||
g.level0 = make([]uint32, nextPow2(ruleCount/4))
|
return errMphBuilt
|
||||||
g.level0Mask = uint32(len(g.level0) - 1)
|
|
||||||
g.level1 = make([]uint32, nextPow2(ruleCount))
|
|
||||||
g.level1Mask = uint32(len(g.level1) - 1)
|
|
||||||
|
|
||||||
// Create buckets based on all rule's rolling hash
|
|
||||||
buckets := make([][]uint32, len(g.level0))
|
|
||||||
for ruleIdx := 1; ruleIdx < len(g.rules); ruleIdx++ { // Traverse rules starting from index 1 (0 reserved for failed lookup)
|
|
||||||
ruleInfo := (*g.ruleInfos)[g.rules[ruleIdx]]
|
|
||||||
bucketIdx := ruleInfo.rollingHash & g.level0Mask
|
|
||||||
buckets[bucketIdx] = append(buckets[bucketIdx], uint32(ruleIdx))
|
|
||||||
g.values[ruleIdx] = append(ruleInfo.matchers[Full], ruleInfo.matchers[Domain]...) // nolint:gocritic
|
|
||||||
}
|
}
|
||||||
g.ruleInfos = nil // Set ruleInfos nil to release memory
|
if uint64(len(g.buf)) > math.MaxUint32 {
|
||||||
runtime.GC() // peak mem
|
return errors.New("too many rules for MphMatcherGroup")
|
||||||
|
|
||||||
// Sort buckets in descending order with respect to each bucket's size
|
|
||||||
bucketIdxs := make([]int, len(buckets))
|
|
||||||
for bucketIdx := range buckets {
|
|
||||||
bucketIdxs[bucketIdx] = bucketIdx
|
|
||||||
}
|
}
|
||||||
sort.Slice(bucketIdxs, func(i, j int) bool { return len(buckets[bucketIdxs[i]]) > len(buckets[bucketIdxs[j]]) })
|
recs := g.writeRecords()
|
||||||
|
if len(g.arena) > mphOffMask {
|
||||||
|
return errors.New("too many rules for MphMatcherGroup")
|
||||||
|
}
|
||||||
|
hashes := make([]uint64, len(recs))
|
||||||
|
for _, mul := range mphMultipliers {
|
||||||
|
for i, rec := range recs {
|
||||||
|
hashes[i] = mphMix(mphHash(mul, g.recKey(rec)))
|
||||||
|
}
|
||||||
|
g.mul = mul
|
||||||
|
if err := g.place(recs, hashes); err != errMphCollision {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return errMphCollision
|
||||||
|
}
|
||||||
|
|
||||||
// Exercise Hash, Displace, and Compress algorithm to construct minimal perfect hash table
|
// writeRecords writes one record per distinct pattern to arena and returns flags | offset of each.
|
||||||
occupied := make([]bool, len(g.level1)) // Whether a second-level hash has been already used
|
func (g *MphMatcherGroup) writeRecords() []uint32 {
|
||||||
hashedBucket := make([]uint32, 0, 4) // Second-level hashes for each rule in a specific bucket
|
g.multi = false
|
||||||
for _, bucketIdx := range bucketIdxs {
|
if len(g.entries) > 0 {
|
||||||
bucket := buckets[bucketIdx]
|
g.single = g.entries[0].value
|
||||||
hashedBucket = hashedBucket[:0]
|
for _, e := range g.entries {
|
||||||
seed := uint32(0)
|
if e.value != g.single {
|
||||||
for len(hashedBucket) != len(bucket) {
|
g.multi = true
|
||||||
for _, ruleIdx := range bucket {
|
break
|
||||||
memHash := MemHash(seed, g.rules[ruleIdx]) & g.level1Mask
|
|
||||||
if occupied[memHash] { // Collision occurred with this seed
|
|
||||||
for _, hash := range hashedBucket { // Revert all values in this hashed bucket
|
|
||||||
occupied[hash] = false
|
|
||||||
g.level1[hash] = 0
|
|
||||||
}
|
|
||||||
hashedBucket = hashedBucket[:0]
|
|
||||||
seed++ // Try next seed
|
|
||||||
break
|
|
||||||
}
|
|
||||||
occupied[memHash] = true
|
|
||||||
g.level1[memHash] = ruleIdx // The final value in the hash table
|
|
||||||
hashedBucket = append(hashedBucket, memHash)
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
g.level0[bucketIdx] = seed // Displacement value for this bucket
|
}
|
||||||
|
// Equal patterns become neighbours in Add order, so their values keep their priority
|
||||||
|
order := make([]uint32, len(g.entries))
|
||||||
|
for i := range order {
|
||||||
|
order[i] = uint32(i)
|
||||||
|
}
|
||||||
|
slices.SortFunc(order, func(a, b uint32) int {
|
||||||
|
return cmp.Or(bytes.Compare(g.key(a), g.key(b)), cmp.Compare(a, b))
|
||||||
|
})
|
||||||
|
|
||||||
|
size := len(g.buf) + len(g.entries) + 2
|
||||||
|
if g.multi {
|
||||||
|
size += 3 * len(g.entries)
|
||||||
|
}
|
||||||
|
arena := make([]byte, 0, size)
|
||||||
|
recs := make([]uint32, 0, len(order))
|
||||||
|
var vals [len(mphKinds)][]uint32
|
||||||
|
for i := 0; i < len(order); {
|
||||||
|
k := g.key(order[i])
|
||||||
|
for t := range vals {
|
||||||
|
vals[t] = vals[t][:0]
|
||||||
|
}
|
||||||
|
for ; i < len(order) && bytes.Equal(g.key(order[i]), k); i++ {
|
||||||
|
e := &g.entries[order[i]]
|
||||||
|
if !slices.Contains(vals[e.kind], e.value) {
|
||||||
|
vals[e.kind] = append(vals[e.kind], e.value)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
rec := uint32(len(arena))
|
||||||
|
if len(k) < 255 {
|
||||||
|
arena = append(arena, byte(len(k)))
|
||||||
|
} else {
|
||||||
|
arena = binary.AppendUvarint(append(arena, 255), uint64(len(k)))
|
||||||
|
}
|
||||||
|
arena = append(arena, k...)
|
||||||
|
for t, v := range vals {
|
||||||
|
if len(v) == 0 {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
rec |= mphKinds[t]
|
||||||
|
if g.multi {
|
||||||
|
arena = binary.AppendUvarint(arena, uint64(len(v)))
|
||||||
|
for _, x := range v {
|
||||||
|
arena = binary.AppendUvarint(arena, uint64(x))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
recs = append(recs, rec)
|
||||||
|
}
|
||||||
|
// Lookups may point one byte past a pattern, and an empty group needs a record at offset 0 for empty slots
|
||||||
|
arena = append(arena, 0)
|
||||||
|
if len(recs) == 0 {
|
||||||
|
arena = append(arena, 0)
|
||||||
|
}
|
||||||
|
g.buf, g.entries = nil, nil
|
||||||
|
if cap(arena)-len(arena) > len(arena)/32 {
|
||||||
|
arena = slices.Clone(arena)
|
||||||
|
}
|
||||||
|
g.arena = unsafe.String(unsafe.SliceData(arena), len(arena)) // arena is not written after this
|
||||||
|
return recs
|
||||||
|
}
|
||||||
|
|
||||||
|
// place fills level0, level1 and fp: records are bucketed by hash, and each bucket, largest first, gets
|
||||||
|
// the first seed that puts all its records in free slots.
|
||||||
|
func (g *MphMatcherGroup) place(recs []uint32, hashes []uint64) error {
|
||||||
|
r := len(recs)
|
||||||
|
n0, n1 := max(1, r/3), max(1, r+r/99)
|
||||||
|
g.n0, g.n1 = uint32(n0), uint32(n1)
|
||||||
|
g.level0 = make([]uint16, n0)
|
||||||
|
g.level1 = make([]uint32, n1)
|
||||||
|
g.fp = make([]uint8, n1)
|
||||||
|
|
||||||
|
start := make([]uint32, n0+1)
|
||||||
|
for _, h := range hashes {
|
||||||
|
start[g.bucket(h)+1]++
|
||||||
|
}
|
||||||
|
for b := range n0 {
|
||||||
|
start[b+1] += start[b]
|
||||||
|
}
|
||||||
|
members := make([]uint32, r)
|
||||||
|
fill := slices.Clone(start[:n0])
|
||||||
|
for i, h := range hashes {
|
||||||
|
b := g.bucket(h)
|
||||||
|
members[fill[b]] = uint32(i)
|
||||||
|
fill[b]++
|
||||||
|
}
|
||||||
|
fill = nil
|
||||||
|
buckets := make([]uint32, n0)
|
||||||
|
for b := range buckets {
|
||||||
|
buckets[b] = uint32(b)
|
||||||
|
}
|
||||||
|
slices.SortStableFunc(buckets, func(a, b uint32) int {
|
||||||
|
return cmp.Compare(start[b+1]-start[b], start[a+1]-start[a])
|
||||||
|
})
|
||||||
|
|
||||||
|
occupied := make([]uint64, (n1+63)/64)
|
||||||
|
var slots []uint32
|
||||||
|
next:
|
||||||
|
for _, b := range buckets {
|
||||||
|
m := members[start[b]:start[b+1]]
|
||||||
|
if len(m) == 0 {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
for i := range m {
|
||||||
|
for j := range i {
|
||||||
|
if hashes[m[i]] == hashes[m[j]] {
|
||||||
|
return errMphCollision // no seed can separate them
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
search:
|
||||||
|
for seed := range math.MaxUint16 + 1 {
|
||||||
|
slots = slots[:0]
|
||||||
|
for _, ri := range m {
|
||||||
|
s := g.slot(hashes[ri], uint16(seed))
|
||||||
|
if occupied[s/64]&(1<<(s%64)) != 0 || slices.Contains(slots, s) {
|
||||||
|
continue search
|
||||||
|
}
|
||||||
|
slots = append(slots, s)
|
||||||
|
}
|
||||||
|
for k, ri := range m {
|
||||||
|
s := slots[k]
|
||||||
|
occupied[s/64] |= 1 << (s % 64)
|
||||||
|
g.level1[s] = recs[ri]
|
||||||
|
g.fp[s] = uint8(hashes[ri])
|
||||||
|
}
|
||||||
|
g.level0[b] = uint16(seed)
|
||||||
|
continue next
|
||||||
|
}
|
||||||
|
return errors.New("strmatcher: no seed found for a bucket in MphMatcherGroup")
|
||||||
}
|
}
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// Lookup searches for input in minimal perfect hash table and returns its index. 0 indicates not found.
|
// mphHash is the suffix hash of s, taken from the right: the hash of s[i:] is the state after reading s[i].
|
||||||
func (g *MphMatcherGroup) Lookup(rollingHash uint32, input string) uint32 {
|
func mphHash(mul uint64, s string) uint64 {
|
||||||
i0 := rollingHash & g.level0Mask
|
h := uint64(0)
|
||||||
seed := g.level0[i0]
|
for i := len(s) - 1; i >= 0; i-- {
|
||||||
i1 := MemHash(seed, input) & g.level1Mask
|
h = h*mul + uint64(s[i])
|
||||||
if n := g.level1[i1]; g.rules[n] == input {
|
}
|
||||||
return n
|
return h
|
||||||
|
}
|
||||||
|
|
||||||
|
// mphMix spreads the weak low bits of a suffix hash.
|
||||||
|
func mphMix(h uint64) uint64 {
|
||||||
|
h ^= h >> 32
|
||||||
|
h *= 0xd6e8feb86659fd93
|
||||||
|
return h ^ h>>32
|
||||||
|
}
|
||||||
|
|
||||||
|
func (g *MphMatcherGroup) bucket(f uint64) uint32 {
|
||||||
|
return uint32(((f >> 32) * uint64(g.n0)) >> 32)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (g *MphMatcherGroup) slot(f uint64, seed uint16) uint32 {
|
||||||
|
x := ((f ^ uint64(seed)*0x9e3779b97f4a7c15) * 0xc4ceb9fe1a85ec53) >> 32
|
||||||
|
return uint32((x * uint64(g.n1)) >> 32)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (g *MphMatcherGroup) uvarint(p uint32) (x, next uint32) {
|
||||||
|
for shift := 0; ; shift += 7 {
|
||||||
|
c := g.arena[p]
|
||||||
|
p++
|
||||||
|
x |= uint32(c&0x7f) << shift
|
||||||
|
if c < 0x80 {
|
||||||
|
return x, p
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// recSpan returns where the pattern of the record at off starts and how long it is.
|
||||||
|
func (g *MphMatcherGroup) recSpan(off uint32) (p, n uint32) {
|
||||||
|
n, p = uint32(g.arena[off]), off+1
|
||||||
|
if n == 255 {
|
||||||
|
n, p = g.uvarint(p)
|
||||||
|
}
|
||||||
|
return p, n
|
||||||
|
}
|
||||||
|
|
||||||
|
func (g *MphMatcherGroup) recKey(rec uint32) string {
|
||||||
|
p, n := g.recSpan(rec & mphOffMask)
|
||||||
|
return g.arena[p : p+n]
|
||||||
|
}
|
||||||
|
|
||||||
|
// lookup returns the level1 entry of s, or 0 if s is not a pattern. h is the suffix hash of s.
|
||||||
|
func (g *MphMatcherGroup) lookup(h uint64, s string) uint32 {
|
||||||
|
f := mphMix(h)
|
||||||
|
// bucket < n0 == len(level0) and slot < n1 == len(level1) == len(fp), skip the bounds checks
|
||||||
|
seed := *(*uint16)(unsafe.Add(unsafe.Pointer(unsafe.SliceData(g.level0)), uintptr(g.bucket(f))*2))
|
||||||
|
slot := uintptr(g.slot(f, seed))
|
||||||
|
if *(*uint8)(unsafe.Add(unsafe.Pointer(unsafe.SliceData(g.fp)), slot)) != uint8(f) {
|
||||||
|
return 0
|
||||||
|
}
|
||||||
|
e := *(*uint32)(unsafe.Add(unsafe.Pointer(unsafe.SliceData(g.level1)), slot*4))
|
||||||
|
if len(s) < 255 {
|
||||||
|
// A record whose length byte is len(s) has len(s) pattern bytes after it
|
||||||
|
p := unsafe.Add(unsafe.Pointer(unsafe.StringData(g.arena)), e&mphOffMask)
|
||||||
|
if int(*(*byte)(p)) == len(s) && unsafe.String((*byte)(unsafe.Add(p, 1)), len(s)) == s {
|
||||||
|
return e
|
||||||
|
}
|
||||||
|
return 0
|
||||||
|
}
|
||||||
|
if g.recKey(e) == s {
|
||||||
|
return e
|
||||||
}
|
}
|
||||||
return 0
|
return 0
|
||||||
}
|
}
|
||||||
|
|
||||||
// Match implements MatcherGroup.Match.
|
// appendValues appends the values of record e for the flags in want, in mphKinds order.
|
||||||
func (g *MphMatcherGroup) Match(input string) []uint32 {
|
func (g *MphMatcherGroup) appendValues(dst []uint32, e, want uint32) []uint32 {
|
||||||
matches := make([][]uint32, 0, 5)
|
if !g.multi {
|
||||||
hash := uint32(0)
|
for _, flag := range mphKinds {
|
||||||
for i := len(input) - 1; i >= 0; i-- {
|
if e&want&flag != 0 {
|
||||||
hash = hash*PrimeRK + uint32(input[i])
|
dst = append(dst, g.single)
|
||||||
if input[i] == '.' {
|
}
|
||||||
if mphIdx := g.Lookup(hash, input[i:]); mphIdx != 0 {
|
}
|
||||||
matches = append(matches, g.values[mphIdx])
|
return dst
|
||||||
|
}
|
||||||
|
if e&want == 0 {
|
||||||
|
return dst
|
||||||
|
}
|
||||||
|
p, n := g.recSpan(e & mphOffMask)
|
||||||
|
p += n
|
||||||
|
for _, flag := range mphKinds {
|
||||||
|
if e&flag == 0 {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
var count, v uint32
|
||||||
|
for count, p = g.uvarint(p); count > 0; count-- {
|
||||||
|
v, p = g.uvarint(p)
|
||||||
|
if want&flag != 0 {
|
||||||
|
dst = append(dst, v)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
if mphIdx := g.Lookup(hash, input); mphIdx != 0 {
|
return dst
|
||||||
matches = append(matches, g.values[mphIdx])
|
}
|
||||||
|
|
||||||
|
// Match implements MatcherGroup.Match. Values of an exact match come first (Full, then Domain), then those of
|
||||||
|
// the parent domains, nearest first.
|
||||||
|
func (g *MphMatcherGroup) Match(input string) []uint32 {
|
||||||
|
var stack [8]uint32
|
||||||
|
parents := stack[:0] // TLD side first
|
||||||
|
h, mul := uint64(0), g.mul
|
||||||
|
for i := len(input) - 1; i >= 0; i-- {
|
||||||
|
if input[i] == '.' {
|
||||||
|
if e := g.lookup(h, input[i+1:]); e&(mphDomain|mphParent) != 0 {
|
||||||
|
parents = append(parents, e)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
h = h*mul + uint64(input[i])
|
||||||
}
|
}
|
||||||
return CompositeMatchesReverse(matches)
|
exact := g.lookup(h, input)
|
||||||
|
if exact&(mphFull|mphDomain) == 0 && len(parents) == 0 {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
result := g.appendValues(make([]uint32, 0, len(parents)+1), exact, mphFull|mphDomain)
|
||||||
|
for k := len(parents) - 1; k >= 0; k-- {
|
||||||
|
result = g.appendValues(result, parents[k], mphParent|mphDomain)
|
||||||
|
}
|
||||||
|
return result
|
||||||
}
|
}
|
||||||
|
|
||||||
// MatchAny implements MatcherGroup.MatchAny.
|
// MatchAny implements MatcherGroup.MatchAny.
|
||||||
func (g *MphMatcherGroup) MatchAny(input string) bool {
|
func (g *MphMatcherGroup) MatchAny(input string) bool {
|
||||||
hash := uint32(0)
|
h, mul := uint64(0), g.mul
|
||||||
|
for i := len(input) - 1; i >= 0; i-- {
|
||||||
|
if input[i] == '.' && g.lookup(h, input[i+1:])&(mphDomain|mphParent) != 0 {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
h = h*mul + uint64(input[i])
|
||||||
|
}
|
||||||
|
return g.lookup(h, input)&(mphFull|mphDomain) != 0
|
||||||
|
}
|
||||||
|
|
||||||
|
// mphSuffix is the suffix hash of input[off:], a parent domain of the input.
|
||||||
|
type mphSuffix struct {
|
||||||
|
h uint64
|
||||||
|
off int
|
||||||
|
}
|
||||||
|
|
||||||
|
// mphSuffixes appends the suffix hashes of the parent domains of input to dst, TLD side first, and returns them
|
||||||
|
// with the hash of input itself: what MatchAny computes, computed once for several groups.
|
||||||
|
func mphSuffixes(dst []mphSuffix, mul uint64, input string) ([]mphSuffix, uint64) {
|
||||||
|
h := uint64(0)
|
||||||
for i := len(input) - 1; i >= 0; i-- {
|
for i := len(input) - 1; i >= 0; i-- {
|
||||||
hash = hash*PrimeRK + uint32(input[i])
|
|
||||||
if input[i] == '.' {
|
if input[i] == '.' {
|
||||||
if g.Lookup(hash, input[i:]) != 0 {
|
dst = append(dst, mphSuffix{h, i + 1})
|
||||||
return true
|
}
|
||||||
}
|
h = h*mul + uint64(input[i])
|
||||||
|
}
|
||||||
|
return dst, h
|
||||||
|
}
|
||||||
|
|
||||||
|
// matchAnyHashed is MatchAny with parents and h from mphSuffixes(_, mul, input).
|
||||||
|
func (g *MphMatcherGroup) matchAnyHashed(input string, parents []mphSuffix, h, mul uint64) bool {
|
||||||
|
if g.mul != mul {
|
||||||
|
return g.MatchAny(input) // built with a later multiplier after a collision
|
||||||
|
}
|
||||||
|
for _, p := range parents {
|
||||||
|
if g.lookup(p.h, input[p.off:])&(mphDomain|mphParent) != 0 {
|
||||||
|
return true
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
return g.Lookup(hash, input) != 0
|
return g.lookup(h, input)&(mphFull|mphDomain) != 0
|
||||||
}
|
}
|
||||||
|
|
||||||
func nextPow2(v int) int {
|
|
||||||
if v <= 1 {
|
|
||||||
return 1
|
|
||||||
}
|
|
||||||
const MaxUInt = ^uint(0)
|
|
||||||
n := (MaxUInt >> bits.LeadingZeros(uint(v))) + 1
|
|
||||||
return int(n)
|
|
||||||
}
|
|
||||||
|
|
||||||
//go:noescape
|
|
||||||
//go:linkname strhash runtime.strhash
|
|
||||||
func strhash(p unsafe.Pointer, h uintptr) uintptr
|
|
||||||
|
|||||||
@@ -0,0 +1,108 @@
|
|||||||
|
package strmatcher
|
||||||
|
|
||||||
|
import (
|
||||||
|
"slices"
|
||||||
|
"testing"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestMphMatcherGroupHashCollision(t *testing.T) {
|
||||||
|
saved := mphMultipliers
|
||||||
|
defer func() { mphMultipliers = saved }()
|
||||||
|
|
||||||
|
mphMultipliers[0] = 1 // anagrams collide
|
||||||
|
g := NewMphMatcherGroup()
|
||||||
|
g.AddFullMatcher(FullMatcher("ab.com"), 1)
|
||||||
|
g.AddDomainMatcher(DomainMatcher("ba.com"), 2)
|
||||||
|
g.AddDomainMatcher(DomainMatcher("com"), 3)
|
||||||
|
if err := g.Build(); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if g.mul != saved[1] {
|
||||||
|
t.Errorf("multiplier %#x, want the second one %#x", g.mul, saved[1])
|
||||||
|
}
|
||||||
|
for input, want := range map[string][]uint32{"ab.com": {1, 3}, "x.ba.com": {2, 3}, "x.ab.com": {3}, "ba.com": {2, 3}} {
|
||||||
|
if m := g.Match(input); !slices.Equal(m, want) {
|
||||||
|
t.Errorf("Match(%q) = %v, want %v", input, m, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Thue-Morse strings of 2048 bytes and their complements collide for every odd multiplier
|
||||||
|
mphMultipliers = saved
|
||||||
|
a, b := make([]byte, 2048), make([]byte, 2048)
|
||||||
|
for i := range a {
|
||||||
|
a[i], b[i] = "ab"[bitsOnes(i)%2], "ba"[bitsOnes(i)%2]
|
||||||
|
}
|
||||||
|
g = NewMphMatcherGroup()
|
||||||
|
g.AddFullMatcher(FullMatcher(a), 1)
|
||||||
|
g.AddFullMatcher(FullMatcher(b), 1)
|
||||||
|
if err := g.Build(); err != errMphCollision {
|
||||||
|
t.Errorf("Build() = %v, want %v", err, errMphCollision)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func bitsOnes(i int) int {
|
||||||
|
n := 0
|
||||||
|
for ; i > 0; i &= i - 1 {
|
||||||
|
n++
|
||||||
|
}
|
||||||
|
return n
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMphValueMatcherCombiner(t *testing.T) {
|
||||||
|
build := func(matchers ...Matcher) *MphValueMatcher {
|
||||||
|
m := NewMphValueMatcher()
|
||||||
|
for _, x := range matchers {
|
||||||
|
m.Add(x, 0)
|
||||||
|
}
|
||||||
|
if err := m.Build(); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
return m
|
||||||
|
}
|
||||||
|
regex, err := Regex.New(`^a\d+\.net$`)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
saved := mphMultipliers
|
||||||
|
t.Cleanup(func() { mphMultipliers = saved })
|
||||||
|
mphMultipliers[0] = 1 // anagrams collide, so this one falls back to its own hash pass
|
||||||
|
collided := build(FullMatcher("ab.com"), DomainMatcher("ba.com"))
|
||||||
|
mphMultipliers = saved
|
||||||
|
if collided.mph.mul == mphMultipliers[0] {
|
||||||
|
t.Fatal("collided matcher uses the first multiplier")
|
||||||
|
}
|
||||||
|
matchers := []*MphValueMatcher{
|
||||||
|
build(DomainMatcher("example.com"), FullMatcher("full.org"), DomainMatcher(".dot.io")),
|
||||||
|
collided,
|
||||||
|
build(regex, SubstrMatcher("keyword")),
|
||||||
|
build(),
|
||||||
|
build(DomainMatcher("com"), DomainMatcher("a.b.c.d.e.f.g.h.i.j.k.l.m.n.o.p.q.r.s")),
|
||||||
|
}
|
||||||
|
var s MphValueMatcherCombiner
|
||||||
|
for i, m := range matchers {
|
||||||
|
s.Add(m, uint32(10+i))
|
||||||
|
}
|
||||||
|
inputs := []string{
|
||||||
|
"", ".", "..", "com", "example.com", "www.example.com", "xexample.com", "example.com.", "full.org", "x.full.org",
|
||||||
|
"dot.io", "x.dot.io", ".dot.io", "ab.com", "x.ab.com", "ba.com", "x.ba.com", "a12.net", "a12.net.x", "my-keyword.org",
|
||||||
|
"a.b.c.d.e.f.g.h.i.j.k.l.m.n.o.p.q.r.s", "0.a.b.c.d.e.f.g.h.i.j.k.l.m.n.o.p.q.r.s", "b.c.d.e.f.g.h.i.j.k.l.m.n.o.p.q.r.s",
|
||||||
|
"x.y.z.1.2.3.4.5.6.7.8.9.10.11.12.13.14.15.16.17.ab.com", "x.y.z.1.2.3.4.5.6.7.8.9.10.11.12.13.14.15.16.17.org",
|
||||||
|
}
|
||||||
|
for _, input := range inputs {
|
||||||
|
var want []uint32
|
||||||
|
for i, m := range matchers {
|
||||||
|
if m.MatchAny(input) {
|
||||||
|
want = append(want, uint32(10+i))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if got := s.Match(input); !slices.Equal(got, want) {
|
||||||
|
t.Errorf("Match(%q) = %v, want %v", input, got, want)
|
||||||
|
}
|
||||||
|
if got := s.MatchAny(input); got != (len(want) > 0) {
|
||||||
|
t.Errorf("MatchAny(%q) = %v", input, got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if n := testing.AllocsPerRun(100, func() { s.MatchAny("www.a.b.c.example.org") }); n != 0 {
|
||||||
|
t.Errorf("MatchAny allocates %v times", n)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -1,7 +1,10 @@
|
|||||||
package strmatcher_test
|
package strmatcher_test
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"math/rand"
|
||||||
"reflect"
|
"reflect"
|
||||||
|
"slices"
|
||||||
|
"strings"
|
||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
"github.com/xtls/xray-core/common"
|
"github.com/xtls/xray-core/common"
|
||||||
@@ -276,3 +279,142 @@ func TestEmptyMphMatcherGroup(t *testing.T) {
|
|||||||
t.Error("Expect [], but ", r)
|
t.Error("Expect [], but ", r)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestMphMatcherGroupRandom(t *testing.T) {
|
||||||
|
inputs := []string{""} // All strings over "ab." up to 7 bytes
|
||||||
|
for i := 0; len(inputs[i]) < 7; i++ {
|
||||||
|
for _, c := range []string{"a", "b", "."} {
|
||||||
|
inputs = append(inputs, inputs[i]+c)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
for seed := int64(0); seed < 300; seed++ {
|
||||||
|
r := rand.New(rand.NewSource(seed))
|
||||||
|
g := NewMphMatcherGroup()
|
||||||
|
full, domain := map[string][]uint32{}, map[string][]uint32{} // Stored pattern -> values
|
||||||
|
for value := uint32(r.Intn(200)); value > 0; value-- {
|
||||||
|
pattern := make([]byte, r.Intn(8))
|
||||||
|
for i := range pattern {
|
||||||
|
pattern[i] = "ab."[r.Intn(3)]
|
||||||
|
}
|
||||||
|
if p := string(pattern); r.Intn(2) == 0 {
|
||||||
|
g.AddFullMatcher(FullMatcher(p), value)
|
||||||
|
full[p] = append(full[p], value)
|
||||||
|
} else {
|
||||||
|
g.AddDomainMatcher(DomainMatcher(p), value)
|
||||||
|
domain[p] = append(domain[p], value)
|
||||||
|
domain["."+p] = append(domain["."+p], value)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
common.Must(g.Build())
|
||||||
|
for _, input := range inputs {
|
||||||
|
keys := []string{input} // Whole input first, then "." suffixes from longest to shortest
|
||||||
|
for i := range len(input) {
|
||||||
|
if input[i] == '.' {
|
||||||
|
keys = append(keys, input[i:])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
var want []uint32
|
||||||
|
for _, k := range keys {
|
||||||
|
want = append(append(want, full[k]...), domain[k]...)
|
||||||
|
}
|
||||||
|
// Compared as sets: Match reports a value once per matching pattern, and orders them differently
|
||||||
|
// from want for patterns and inputs with a leading dot
|
||||||
|
m := g.Match(input)
|
||||||
|
if !slices.Equal(sortedSet(m), sortedSet(want)) {
|
||||||
|
t.Fatalf("seed %d: Match(%q) = %v, want %v", seed, input, m, want)
|
||||||
|
}
|
||||||
|
if m := g.MatchAny(input); m != (len(want) > 0) {
|
||||||
|
t.Fatalf("seed %d: MatchAny(%q) = %v", seed, input, m)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMphMatcherGroupAppend(t *testing.T) {
|
||||||
|
g := NewMphMatcherGroup()
|
||||||
|
g.AddFullMatcher(FullMatcher("a.com"), 1)
|
||||||
|
g.AddFullMatcher(FullMatcher("b.com"), 2)
|
||||||
|
g.Build()
|
||||||
|
if m := append(g.Match("a.com"), 3); !slices.Equal(m, []uint32{1, 3}) {
|
||||||
|
t.Error("expect [1 3], but ", m)
|
||||||
|
}
|
||||||
|
if m := g.Match("b.com"); !slices.Equal(m, []uint32{2}) {
|
||||||
|
t.Error("expect [2], but ", m)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func sortedSet(v []uint32) []uint32 {
|
||||||
|
v = slices.Clone(v)
|
||||||
|
slices.Sort(v)
|
||||||
|
return slices.Compact(v)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMphMatcherGroupLongPattern(t *testing.T) {
|
||||||
|
long := strings.Repeat("a", 300) + ".com"
|
||||||
|
for _, values := range [][4]uint32{{1, 2, 3, 4}, {7, 7, 7, 7}} {
|
||||||
|
g := NewMphMatcherGroup()
|
||||||
|
g.AddDomainMatcher(DomainMatcher(long), values[0])
|
||||||
|
g.AddFullMatcher(FullMatcher("x."+long), values[1])
|
||||||
|
g.AddFullMatcher(FullMatcher(long[:255]), values[2]) // the shortest pattern stored with a long length
|
||||||
|
g.AddFullMatcher(FullMatcher(long[:254]), values[3])
|
||||||
|
common.Must(g.Build())
|
||||||
|
cases := []struct {
|
||||||
|
input string
|
||||||
|
want []uint32
|
||||||
|
}{
|
||||||
|
{long, []uint32{values[0]}},
|
||||||
|
{"www." + long, []uint32{values[0]}},
|
||||||
|
{"x." + long, []uint32{values[1], values[0]}},
|
||||||
|
{long[1:], nil},
|
||||||
|
{"a" + long, nil},
|
||||||
|
{long[:255], []uint32{values[2]}},
|
||||||
|
{long[:254], []uint32{values[3]}},
|
||||||
|
{long[:256], nil},
|
||||||
|
{long[:253], nil},
|
||||||
|
}
|
||||||
|
for _, c := range cases {
|
||||||
|
if m := g.Match(c.input); !slices.Equal(m, c.want) {
|
||||||
|
t.Errorf("Match(%d bytes) = %v, want %v", len(c.input), m, c.want)
|
||||||
|
}
|
||||||
|
if m := g.MatchAny(c.input); m != (c.want != nil) {
|
||||||
|
t.Errorf("MatchAny(%d bytes) = %v", len(c.input), m)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// A pattern longer than 65535 bytes builds and matches: a record's length is a uvarint,
|
||||||
|
// so the only cap was the build-time length field, now widened to uint32.
|
||||||
|
huge := strings.Repeat("a", 70000)
|
||||||
|
g := NewMphMatcherGroup()
|
||||||
|
g.AddFullMatcher(FullMatcher(strings.Repeat("a", 65535)), 1)
|
||||||
|
g.AddDomainMatcher(DomainMatcher(huge+".com"), 2)
|
||||||
|
g.AddFullMatcher(FullMatcher("a.com"), 3)
|
||||||
|
common.Must(g.Build())
|
||||||
|
if !g.MatchAny(strings.Repeat("a", 65535)) || g.MatchAny(strings.Repeat("a", 65534)) {
|
||||||
|
t.Error("wrong answer for a 65535-byte pattern")
|
||||||
|
}
|
||||||
|
if m := g.Match(huge + ".com"); !slices.Equal(m, []uint32{2}) {
|
||||||
|
t.Errorf("Match(%d-byte input) = %v, want [2]", len(huge)+4, m)
|
||||||
|
}
|
||||||
|
if m := g.Match("x." + huge + ".com"); !slices.Equal(m, []uint32{2}) {
|
||||||
|
t.Errorf("Match(subdomain of a %d-byte pattern) = %v, want [2]", len(huge)+4, m)
|
||||||
|
}
|
||||||
|
if g.MatchAny(huge) { // the 70000-byte label on its own is not a rule
|
||||||
|
t.Error("unexpected match for the bare 70000-byte label")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMphMatcherGroupBuildOnce(t *testing.T) {
|
||||||
|
g := NewMphMatcherGroup()
|
||||||
|
g.AddFullMatcher(FullMatcher("a.com"), 1)
|
||||||
|
common.Must(g.Build())
|
||||||
|
if err := g.Build(); err == nil || !g.MatchAny("a.com") {
|
||||||
|
t.Errorf("second Build() = %v, MatchAny(a.com) = %v", err, g.MatchAny("a.com"))
|
||||||
|
}
|
||||||
|
defer func() {
|
||||||
|
if recover() == nil {
|
||||||
|
t.Error("Add after Build did not panic")
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
g.AddDomainMatcher(DomainMatcher("b.com"), 2)
|
||||||
|
}
|
||||||
|
|||||||
@@ -2,9 +2,12 @@ package strmatcher
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"errors"
|
"errors"
|
||||||
|
"math/bits"
|
||||||
"regexp"
|
"regexp"
|
||||||
|
"regexp/syntax"
|
||||||
"slices"
|
"slices"
|
||||||
"strings"
|
"strings"
|
||||||
|
"unicode"
|
||||||
"unicode/utf8"
|
"unicode/utf8"
|
||||||
|
|
||||||
"golang.org/x/net/idna"
|
"golang.org/x/net/idna"
|
||||||
@@ -73,7 +76,274 @@ func (m SubstrMatcher) Match(s string) bool {
|
|||||||
|
|
||||||
// RegexMatcher is an implementation of Matcher.
|
// RegexMatcher is an implementation of Matcher.
|
||||||
type RegexMatcher struct {
|
type RegexMatcher struct {
|
||||||
pattern *regexp.Regexp
|
pattern *regexp.Regexp
|
||||||
|
literals []string // every match contains all of them, longest first
|
||||||
|
tail []byteSet // tail[i] holds the bytes a matching input can have i bytes before its end
|
||||||
|
rest *byteSet // the bytes it can have further before, nil if any
|
||||||
|
}
|
||||||
|
|
||||||
|
func newRegexMatcher(pattern string) (Matcher, error) {
|
||||||
|
regex, err := regexp.Compile(pattern)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
m := &RegexMatcher{pattern: regex}
|
||||||
|
if re, err := syntax.Parse(pattern, syntax.Perl); err == nil { // same flags as regexp.Compile
|
||||||
|
m.literals = requiredLiterals(re, nil)
|
||||||
|
slices.SortStableFunc(m.literals, func(a, b string) int { return len(b) - len(a) })
|
||||||
|
m.tail, m.rest = tailGuard(re)
|
||||||
|
}
|
||||||
|
return m, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// byteSet is a set of bytes. The bytes >= 0x80 share one bit with 0x7f.
|
||||||
|
type byteSet [4]uint32
|
||||||
|
|
||||||
|
func (s *byteSet) add(c byte) { c = min(c, 0x7f); s[c>>5] |= 1 << (c & 31) }
|
||||||
|
func (s *byteSet) has(c byte) bool { c = min(c, 0x7f); return s[c>>5]&(1<<(c&31)) != 0 }
|
||||||
|
func (s *byteSet) or(t *byteSet) {
|
||||||
|
for i := range s {
|
||||||
|
s[i] |= t[i]
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
var allBytes = byteSet{^uint32(0), ^uint32(0), ^uint32(0), ^uint32(0)}
|
||||||
|
|
||||||
|
// tailLen is how many positions before the end of the input tailGuard tells apart.
|
||||||
|
const tailLen = 8
|
||||||
|
|
||||||
|
// tailBudget caps how many repetition steps tailGuard walks. Only nested repeats can make the
|
||||||
|
// walk explode, so only they are charged: a flat pattern, however long, is walked once and keeps
|
||||||
|
// its guard.
|
||||||
|
const tailBudget = 100000
|
||||||
|
|
||||||
|
// tailWalk is a set of positions in the input, counted in bytes before its end.
|
||||||
|
type tailWalk struct {
|
||||||
|
at uint32 // bit i: exactly i bytes before the end, for i < tailLen
|
||||||
|
far bool // tailLen or more bytes before the end
|
||||||
|
free bool // not tied to the end of the input yet
|
||||||
|
}
|
||||||
|
|
||||||
|
func (w tailWalk) union(v tailWalk) tailWalk {
|
||||||
|
return tailWalk{w.at | v.at, w.far || v.far, w.free || v.free}
|
||||||
|
}
|
||||||
|
|
||||||
|
type tailBuilder struct {
|
||||||
|
tail [tailLen]byteSet
|
||||||
|
rest byteSet
|
||||||
|
void bool
|
||||||
|
work int
|
||||||
|
}
|
||||||
|
|
||||||
|
// tailGuard walks re backwards from the end of the input and collects the bytes an input
|
||||||
|
// matching re can have at each position before its end. It returns nil, nil when a branch
|
||||||
|
// of re does not end with $ or when nested repeats push the walk past tailBudget.
|
||||||
|
func tailGuard(re *syntax.Regexp) ([]byteSet, *byteSet) {
|
||||||
|
var b tailBuilder
|
||||||
|
w := b.walk(re, tailWalk{free: true})
|
||||||
|
b.stop(w)
|
||||||
|
if b.void {
|
||||||
|
return nil, nil
|
||||||
|
}
|
||||||
|
if w.at != 0 { // a match can start here, so any bytes can come before
|
||||||
|
for i := bits.TrailingZeros32(w.at); i < tailLen; i++ {
|
||||||
|
b.tail[i] = allBytes
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if w.at != 0 || w.far {
|
||||||
|
b.rest = allBytes
|
||||||
|
}
|
||||||
|
n := tailLen
|
||||||
|
for n > 0 && b.tail[n-1] == b.rest {
|
||||||
|
n--
|
||||||
|
}
|
||||||
|
var tail []byteSet
|
||||||
|
if n > 0 {
|
||||||
|
tail = slices.Clone(b.tail[:n])
|
||||||
|
}
|
||||||
|
if b.rest != allBytes {
|
||||||
|
rest := b.rest
|
||||||
|
return tail, &rest
|
||||||
|
}
|
||||||
|
return tail, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// stop ends the paths of w. One that never met $ lets its match be followed by anything.
|
||||||
|
func (b *tailBuilder) stop(w tailWalk) {
|
||||||
|
if w.free {
|
||||||
|
b.void = true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (b *tailBuilder) walk(re *syntax.Regexp, w tailWalk) tailWalk {
|
||||||
|
if w == (tailWalk{}) || b.void {
|
||||||
|
return w
|
||||||
|
}
|
||||||
|
switch re.Op {
|
||||||
|
case syntax.OpNoMatch:
|
||||||
|
return tailWalk{}
|
||||||
|
case syntax.OpLiteral:
|
||||||
|
for i := len(re.Rune) - 1; i >= 0; i-- {
|
||||||
|
var set byteSet
|
||||||
|
set.add(byte(min(re.Rune[i], utf8.RuneSelf)))
|
||||||
|
if re.Flags&syntax.FoldCase != 0 {
|
||||||
|
for f := unicode.SimpleFold(re.Rune[i]); f != re.Rune[i]; f = unicode.SimpleFold(f) {
|
||||||
|
set.add(byte(min(f, utf8.RuneSelf)))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
w = b.step(w, &set)
|
||||||
|
}
|
||||||
|
return w
|
||||||
|
case syntax.OpCharClass:
|
||||||
|
var set byteSet
|
||||||
|
for i := 0; i+1 < len(re.Rune); i += 2 {
|
||||||
|
for r := min(re.Rune[i], utf8.RuneSelf); r <= min(re.Rune[i+1], utf8.RuneSelf); r++ {
|
||||||
|
set.add(byte(r))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return b.step(w, &set)
|
||||||
|
case syntax.OpAnyChar, syntax.OpAnyCharNotNL: // a domain has no \n to reject
|
||||||
|
return b.step(w, &allBytes)
|
||||||
|
case syntax.OpBeginText: // nothing comes before
|
||||||
|
b.stop(w)
|
||||||
|
return tailWalk{}
|
||||||
|
case syntax.OpEndText:
|
||||||
|
out := tailWalk{at: w.at & 1}
|
||||||
|
if w.free {
|
||||||
|
out.at = 1
|
||||||
|
}
|
||||||
|
return out
|
||||||
|
case syntax.OpCapture:
|
||||||
|
return b.walk(re.Sub[0], w)
|
||||||
|
case syntax.OpConcat:
|
||||||
|
for i := len(re.Sub) - 1; i >= 0; i-- {
|
||||||
|
w = b.walk(re.Sub[i], w)
|
||||||
|
}
|
||||||
|
return w
|
||||||
|
case syntax.OpAlternate:
|
||||||
|
var out tailWalk
|
||||||
|
for _, sub := range re.Sub {
|
||||||
|
out = out.union(b.walk(sub, w))
|
||||||
|
}
|
||||||
|
return out
|
||||||
|
case syntax.OpQuest:
|
||||||
|
return b.repeat(re.Sub[0], w, 1)
|
||||||
|
case syntax.OpStar:
|
||||||
|
return b.repeat(re.Sub[0], w, -1)
|
||||||
|
case syntax.OpPlus:
|
||||||
|
return b.repeat(re.Sub[0], b.walk(re.Sub[0], w), -1)
|
||||||
|
case syntax.OpRepeat:
|
||||||
|
for i := 0; i < re.Min; i++ {
|
||||||
|
if b.charge() {
|
||||||
|
return w
|
||||||
|
}
|
||||||
|
w = b.walk(re.Sub[0], w)
|
||||||
|
}
|
||||||
|
if re.Max < 0 {
|
||||||
|
return b.repeat(re.Sub[0], w, -1)
|
||||||
|
}
|
||||||
|
return b.repeat(re.Sub[0], w, re.Max-re.Min)
|
||||||
|
}
|
||||||
|
return w // empty match, line and word boundaries: no constraint
|
||||||
|
}
|
||||||
|
|
||||||
|
// charge counts one repetition step and reports whether the walk has run out of budget. Only
|
||||||
|
// repeats re-walk their body, so charging them alone bounds the blow-up of nested repeats while
|
||||||
|
// leaving a single linear pass, of any length, free.
|
||||||
|
func (b *tailBuilder) charge() bool {
|
||||||
|
b.work++
|
||||||
|
if b.work > tailBudget {
|
||||||
|
b.void = true
|
||||||
|
}
|
||||||
|
return b.void
|
||||||
|
}
|
||||||
|
|
||||||
|
// repeat walks back over up to n more repetitions of re, any number if n < 0.
|
||||||
|
func (b *tailBuilder) repeat(re *syntax.Regexp, w tailWalk, n int) tailWalk {
|
||||||
|
for ; n != 0; n-- {
|
||||||
|
if b.charge() {
|
||||||
|
return w
|
||||||
|
}
|
||||||
|
next := w.union(b.walk(re, w))
|
||||||
|
if next == w {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
w = next
|
||||||
|
}
|
||||||
|
return w
|
||||||
|
}
|
||||||
|
|
||||||
|
// step walks back over one character whose last byte is in set. A character that can be
|
||||||
|
// non-ASCII can take up to 4 bytes, all >= 0x80; regexp matches an invalid byte as U+FFFD.
|
||||||
|
func (b *tailBuilder) step(w tailWalk, set *byteSet) tailWalk {
|
||||||
|
out := tailWalk{far: w.far, free: w.free}
|
||||||
|
if w.far {
|
||||||
|
b.rest.or(set)
|
||||||
|
}
|
||||||
|
width := 1
|
||||||
|
if set.has(0x80) {
|
||||||
|
width = utf8.UTFMax
|
||||||
|
}
|
||||||
|
for i := 0; i < tailLen; i++ {
|
||||||
|
if w.at&(1<<i) == 0 {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
b.tail[i].or(set)
|
||||||
|
for n := 1; n <= width; n++ {
|
||||||
|
if j := i + n; j < tailLen {
|
||||||
|
out.at |= 1 << j
|
||||||
|
if n < width {
|
||||||
|
b.tail[j].add(0x80)
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
out.far = true
|
||||||
|
if n < width {
|
||||||
|
b.rest.add(0x80)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return out
|
||||||
|
}
|
||||||
|
|
||||||
|
// mayMatch reports whether s passes the tail guard.
|
||||||
|
func (m *RegexMatcher) mayMatch(s string) bool {
|
||||||
|
n := len(s)
|
||||||
|
if m.rest == nil {
|
||||||
|
n = min(n, len(m.tail))
|
||||||
|
}
|
||||||
|
for i := 0; i < n; i++ {
|
||||||
|
set := m.rest
|
||||||
|
if i < len(m.tail) {
|
||||||
|
set = &m.tail[i]
|
||||||
|
}
|
||||||
|
if !set.has(s[len(s)-1-i]) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
// requiredLiterals appends to dst the case-sensitive strings that every match of re contains.
|
||||||
|
func requiredLiterals(re *syntax.Regexp, dst []string) []string {
|
||||||
|
switch re.Op {
|
||||||
|
case syntax.OpLiteral:
|
||||||
|
// regexp matches U+FFFD against invalid UTF-8 bytes, strings.Contains does not
|
||||||
|
if re.Flags&syntax.FoldCase == 0 && !slices.Contains(re.Rune, utf8.RuneError) {
|
||||||
|
dst = append(dst, string(re.Rune))
|
||||||
|
}
|
||||||
|
case syntax.OpCapture, syntax.OpPlus:
|
||||||
|
dst = requiredLiterals(re.Sub[0], dst)
|
||||||
|
case syntax.OpRepeat:
|
||||||
|
if re.Min > 0 {
|
||||||
|
dst = requiredLiterals(re.Sub[0], dst)
|
||||||
|
}
|
||||||
|
case syntax.OpConcat:
|
||||||
|
for _, sub := range re.Sub {
|
||||||
|
dst = requiredLiterals(sub, dst)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return dst
|
||||||
}
|
}
|
||||||
|
|
||||||
func (*RegexMatcher) Type() Type {
|
func (*RegexMatcher) Type() Type {
|
||||||
@@ -89,6 +359,14 @@ func (m *RegexMatcher) String() string {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (m *RegexMatcher) Match(s string) bool {
|
func (m *RegexMatcher) Match(s string) bool {
|
||||||
|
if !m.mayMatch(s) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
for _, l := range m.literals {
|
||||||
|
if !strings.Contains(s, l) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
}
|
||||||
return m.pattern.MatchString(s)
|
return m.pattern.MatchString(s)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -102,11 +380,7 @@ func (t Type) New(pattern string) (Matcher, error) {
|
|||||||
case Domain:
|
case Domain:
|
||||||
return DomainMatcher(pattern), nil
|
return DomainMatcher(pattern), nil
|
||||||
case Regex: // 1. regex matching is case-sensitive
|
case Regex: // 1. regex matching is case-sensitive
|
||||||
regex, err := regexp.Compile(pattern)
|
return newRegexMatcher(pattern)
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
return &RegexMatcher{pattern: regex}, nil
|
|
||||||
default:
|
default:
|
||||||
return nil, errors.New("unknown matcher type")
|
return nil, errors.New("unknown matcher type")
|
||||||
}
|
}
|
||||||
@@ -135,11 +409,7 @@ func (t Type) NewDomainPattern(pattern string) (Matcher, error) {
|
|||||||
}
|
}
|
||||||
return DomainMatcher(pattern), nil
|
return DomainMatcher(pattern), nil
|
||||||
case Regex: // Regex's charset not in LDH subset
|
case Regex: // Regex's charset not in LDH subset
|
||||||
regex, err := regexp.Compile(pattern)
|
return newRegexMatcher(pattern)
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
return &RegexMatcher{pattern: regex}, nil
|
|
||||||
default:
|
default:
|
||||||
return nil, errors.New("unknown matcher type")
|
return nil, errors.New("unknown matcher type")
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -0,0 +1,233 @@
|
|||||||
|
package strmatcher
|
||||||
|
|
||||||
|
import (
|
||||||
|
"hash/fnv"
|
||||||
|
"math/rand/v2"
|
||||||
|
"regexp"
|
||||||
|
"regexp/syntax"
|
||||||
|
"slices"
|
||||||
|
"strconv"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
"unicode"
|
||||||
|
"unicode/utf8"
|
||||||
|
)
|
||||||
|
|
||||||
|
var regexLiteralCases = []struct {
|
||||||
|
pattern string
|
||||||
|
literals []string
|
||||||
|
}{
|
||||||
|
{`(^|\.)91porn\.(best|com)$`, []string{"91porn."}},
|
||||||
|
{`.+\.awsdns-cn-[0-9][0-9]\.(biz|com|net|top)$`, []string{".awsdns-cn-", "."}},
|
||||||
|
{`^r+[0-9]+(---|\.)sn-(2x3|ni5|j5o)\w{5}\.googlevideo\.com$`, []string{".googlevideo.com", "sn-", "r"}},
|
||||||
|
{`(?i)abc`, nil},
|
||||||
|
{`ab(?i:CD)ef`, []string{"ab", "ef"}},
|
||||||
|
{`(abc)?x`, []string{"x"}},
|
||||||
|
{`(abc)*x`, []string{"x"}},
|
||||||
|
{`x{0,3}yy`, []string{"yy"}},
|
||||||
|
{`(ab)+c{2}`, []string{"ab", "c"}},
|
||||||
|
{`abc|abd`, []string{"ab"}},
|
||||||
|
{`\Qa.b\E`, []string{"a.b"}},
|
||||||
|
{`a\x{FFFD}b`, nil},
|
||||||
|
{`^[^.]+$`, nil},
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRegexRequiredLiterals(t *testing.T) {
|
||||||
|
for _, test := range regexLiteralCases {
|
||||||
|
m, err := newRegexMatcher(test.pattern)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if got := m.(*RegexMatcher).literals; !slices.Equal(got, test.literals) {
|
||||||
|
t.Errorf("%s: got %q, want %q", test.pattern, got, test.literals)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
var regexTailCases = []struct {
|
||||||
|
pattern string
|
||||||
|
guard bool
|
||||||
|
match []string // inputs the pattern matches
|
||||||
|
reject []string // inputs the tail guard alone rejects
|
||||||
|
}{
|
||||||
|
{`^[a-z]([a-z0-9-]{0,61}[a-z0-9])?$`, true, []string{"a", "localhost", "x-1"}, []string{"www.example.com", "localhost.", "LOCALHOST", "a b"}},
|
||||||
|
{`(^|\.)[a-z][1-9][0-9][a-z]\.com$`, true, []string{"a12b.com", "x.q10z.com"}, []string{"google.com", "a12b.co", "a12b.com.", "ab12.com"}},
|
||||||
|
{`^hses[1-7]?\.akamaized\.net$`, true, []string{"hses.akamaized.net", "hses3.akamaized.net"}, []string{"xhses.akamaized.net", "www.hses.akamaized.net"}},
|
||||||
|
{`(?i)k\.net$`, true, []string{"k.net", "K.NET", "\u212a.net"}, []string{"x.net", "k.nex"}},
|
||||||
|
{`[^.]+\.cn$`, true, []string{"a.cn", "\xff.cn", "\u4e2d.cn"}, []string{"a.cnn", "a.c"}},
|
||||||
|
{`\x{FFFD}$`, true, []string{"\xff", "a\xc3", "\uFFFD"}, []string{"a", "\xff."}},
|
||||||
|
{`^.\.cn$`, true, []string{"a.cn", "\u4E2D.cn", "\xff.cn"}, []string{"ab.cn"}},
|
||||||
|
{`^$`, true, []string{""}, []string{"a"}},
|
||||||
|
{`(^|\.)youyuapi\..+$`, false, []string{"youyuapi.com"}, nil},
|
||||||
|
{`abc`, false, []string{"abc", "xabcx"}, nil},
|
||||||
|
{`^ab`, false, []string{"ab", "abc"}, nil},
|
||||||
|
{`a$|b`, false, []string{"a", "bx"}, nil},
|
||||||
|
{`(?m)a$`, false, []string{"a", "a\nb"}, nil},
|
||||||
|
{strings.Repeat(`(?:abcdefgh(?:a`, 20) + strings.Repeat(`)*)*`, 20) + `\.com$`, false, []string{".com", "abcdefgha.com"}, nil}, // over tailBudget
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRegexTailGuard(t *testing.T) {
|
||||||
|
for _, test := range regexTailCases {
|
||||||
|
m, err := newRegexMatcher(test.pattern)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
rm := m.(*RegexMatcher)
|
||||||
|
if guard := rm.tail != nil || rm.rest != nil; guard != test.guard {
|
||||||
|
t.Errorf("%s: guard %v, want %v", test.pattern, guard, test.guard)
|
||||||
|
}
|
||||||
|
for _, s := range test.match {
|
||||||
|
if !rm.pattern.MatchString(s) || !rm.Match(s) {
|
||||||
|
t.Errorf("%s: %q does not match", test.pattern, s)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
for _, s := range test.reject {
|
||||||
|
if rm.pattern.MatchString(s) || rm.mayMatch(s) {
|
||||||
|
t.Errorf("%s: %q passes the guard", test.pattern, s)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestRegexTailGuardFlatAlternation checks that a long but non-recursive pattern keeps its
|
||||||
|
// guard. Only nested repeats are charged against tailBudget, so a flat alternation of many
|
||||||
|
// names, however large, is walked once and guarded; its guard is checked against regexp.
|
||||||
|
func TestRegexTailGuardFlatAlternation(t *testing.T) {
|
||||||
|
var sb strings.Builder
|
||||||
|
sb.WriteString("(?:")
|
||||||
|
for i := 0; i < 20000; i++ {
|
||||||
|
if i > 0 {
|
||||||
|
sb.WriteByte('|')
|
||||||
|
}
|
||||||
|
sb.WriteString("name")
|
||||||
|
sb.WriteString(strconv.Itoa(i))
|
||||||
|
}
|
||||||
|
sb.WriteString(`)\.example\.com$`)
|
||||||
|
m, err := newRegexMatcher(sb.String())
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
rm := m.(*RegexMatcher)
|
||||||
|
if rm.tail == nil && rm.rest == nil {
|
||||||
|
t.Fatal("flat alternation of 20000 names lost its guard")
|
||||||
|
}
|
||||||
|
for _, s := range []string{"name0.example.com", "name19999.example.com", "x.name12345.example.com"} {
|
||||||
|
if !rm.pattern.MatchString(s) || !rm.Match(s) {
|
||||||
|
t.Errorf("%q should match", s)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
for _, s := range []string{"name0.example.org", "name0.example.com.", "name0.example.con", "google.com"} {
|
||||||
|
if rm.pattern.MatchString(s) {
|
||||||
|
t.Fatalf("test bug: %q matches the pattern", s)
|
||||||
|
}
|
||||||
|
if rm.mayMatch(s) {
|
||||||
|
t.Errorf("%q should be rejected by the guard", s)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// sampleMatch appends a string that re matches, assertions aside, unless it runs out of
|
||||||
|
// budget, which it spends one per call so that nested repeats stay cheap.
|
||||||
|
func sampleMatch(sb *strings.Builder, re *syntax.Regexp, rnd *rand.Rand, budget *int) {
|
||||||
|
if *budget <= 0 {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
*budget--
|
||||||
|
switch re.Op {
|
||||||
|
case syntax.OpLiteral:
|
||||||
|
for _, r := range re.Rune {
|
||||||
|
if re.Flags&syntax.FoldCase != 0 {
|
||||||
|
for n := rnd.IntN(4); n > 0; n-- {
|
||||||
|
r = unicode.SimpleFold(r)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
sampleRune(sb, r, rnd)
|
||||||
|
}
|
||||||
|
case syntax.OpCharClass:
|
||||||
|
if len(re.Rune) > 0 {
|
||||||
|
i := rnd.IntN(len(re.Rune)/2) * 2
|
||||||
|
sampleRune(sb, re.Rune[i]+rnd.Int32N(min(re.Rune[i+1]-re.Rune[i]+1, 300)), rnd)
|
||||||
|
}
|
||||||
|
case syntax.OpAnyChar, syntax.OpAnyCharNotNL:
|
||||||
|
sampleRune(sb, []rune{'a', '.', '\n', 0xe9, 0x212a, utf8.RuneError}[rnd.IntN(6)], rnd)
|
||||||
|
case syntax.OpCapture:
|
||||||
|
sampleMatch(sb, re.Sub[0], rnd, budget)
|
||||||
|
case syntax.OpConcat:
|
||||||
|
for _, sub := range re.Sub {
|
||||||
|
sampleMatch(sb, sub, rnd, budget)
|
||||||
|
}
|
||||||
|
case syntax.OpAlternate:
|
||||||
|
sampleMatch(sb, re.Sub[rnd.IntN(len(re.Sub))], rnd, budget)
|
||||||
|
case syntax.OpQuest, syntax.OpStar, syntax.OpPlus, syntax.OpRepeat:
|
||||||
|
lo, hi := 0, 3
|
||||||
|
switch re.Op {
|
||||||
|
case syntax.OpQuest:
|
||||||
|
hi = 1
|
||||||
|
case syntax.OpPlus:
|
||||||
|
lo = 1
|
||||||
|
case syntax.OpRepeat:
|
||||||
|
lo, hi = re.Min, re.Min+3
|
||||||
|
if re.Max >= 0 {
|
||||||
|
hi = min(hi, re.Max)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
for n := lo + rnd.IntN(hi-lo+1); n > 0; n-- {
|
||||||
|
sampleMatch(sb, re.Sub[0], rnd, budget)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func sampleRune(sb *strings.Builder, r rune, rnd *rand.Rand) {
|
||||||
|
if r == utf8.RuneError && rnd.IntN(2) == 0 {
|
||||||
|
sb.WriteByte(0x80 | byte(rnd.IntN(0x80))) // regexp matches an invalid byte as U+FFFD
|
||||||
|
return
|
||||||
|
}
|
||||||
|
sb.WriteRune(r)
|
||||||
|
}
|
||||||
|
|
||||||
|
func FuzzRegexMatcher(f *testing.F) {
|
||||||
|
inputs := []string{
|
||||||
|
"", "x", "yy", "abd", "ccc", "ABC", "abCDef", "abcdef", "abababcc", "a.b", "a\xffb", "a\uFFFDb",
|
||||||
|
"www.91porn.com", "ns1.awsdns-cn-01.top", "r1---sn-2x3abcde.googlevideo.com",
|
||||||
|
}
|
||||||
|
for _, test := range regexLiteralCases {
|
||||||
|
for _, s := range inputs {
|
||||||
|
f.Add(test.pattern, s)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
for _, test := range regexTailCases {
|
||||||
|
for _, s := range append(test.match, test.reject...) {
|
||||||
|
f.Add(test.pattern, s)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
f.Fuzz(func(t *testing.T, pattern, s string) {
|
||||||
|
re, err := regexp.Compile(pattern)
|
||||||
|
if err != nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
m, _ := newRegexMatcher(pattern)
|
||||||
|
check := func(s string) {
|
||||||
|
if got, want := m.Match(s), re.MatchString(s); got != want {
|
||||||
|
t.Errorf("pattern %q, input %q: got %v, want %v", pattern, s, got, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
check(s)
|
||||||
|
// random inputs seldom match, so also try strings built from the pattern
|
||||||
|
parsed, _ := syntax.Parse(pattern, syntax.Perl)
|
||||||
|
h := fnv.New64a()
|
||||||
|
h.Write([]byte(s))
|
||||||
|
rnd := rand.New(rand.NewPCG(h.Sum64(), 1))
|
||||||
|
for range 8 {
|
||||||
|
var sb strings.Builder
|
||||||
|
budget := 256
|
||||||
|
sampleMatch(&sb, parsed, rnd, &budget)
|
||||||
|
sample := sb.String()
|
||||||
|
check(sample)
|
||||||
|
check(s + sample)
|
||||||
|
if len(sample) > 0 && len(s) > 0 {
|
||||||
|
i := rnd.IntN(len(sample))
|
||||||
|
check(sample[:i] + s[:1] + sample[i+1:])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
@@ -46,7 +46,9 @@ func (g *MphValueMatcher) Add(matcher Matcher, value uint32) {
|
|||||||
func (g *MphValueMatcher) Build() error {
|
func (g *MphValueMatcher) Build() error {
|
||||||
if g.mph != nil {
|
if g.mph != nil {
|
||||||
runtime.GC() // peak mem
|
runtime.GC() // peak mem
|
||||||
g.mph.Build()
|
if err := g.mph.Build(); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
}
|
}
|
||||||
runtime.GC() // peak mem
|
runtime.GC() // peak mem
|
||||||
if g.ac != nil {
|
if g.ac != nil {
|
||||||
@@ -58,23 +60,17 @@ func (g *MphValueMatcher) Build() error {
|
|||||||
|
|
||||||
// Match implements ValueMatcher.Match.
|
// Match implements ValueMatcher.Match.
|
||||||
func (g *MphValueMatcher) Match(input string) []uint32 {
|
func (g *MphValueMatcher) Match(input string) []uint32 {
|
||||||
result := make([][]uint32, 0, 5)
|
var result []uint32
|
||||||
if g.mph != nil {
|
if g.mph != nil {
|
||||||
if matches := g.mph.Match(input); len(matches) > 0 {
|
result = g.mph.Match(input) // a new slice, returned without another copy
|
||||||
result = append(result, matches)
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
if g.ac != nil {
|
if g.ac != nil {
|
||||||
if matches := g.ac.Match(input); len(matches) > 0 {
|
result = append(result, g.ac.Match(input)...)
|
||||||
result = append(result, matches)
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
if g.regex != nil {
|
if g.regex != nil {
|
||||||
if matches := g.regex.Match(input); len(matches) > 0 {
|
result = append(result, g.regex.Match(input)...)
|
||||||
result = append(result, matches)
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
return CompositeMatches(result)
|
return result
|
||||||
}
|
}
|
||||||
|
|
||||||
// MatchAny implements ValueMatcher.MatchAny.
|
// MatchAny implements ValueMatcher.MatchAny.
|
||||||
@@ -87,3 +83,62 @@ func (g *MphValueMatcher) MatchAny(input string) bool {
|
|||||||
}
|
}
|
||||||
return g.regex != nil && g.regex.MatchAny(input)
|
return g.regex != nil && g.regex.MatchAny(input)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (g *MphValueMatcher) matchAnyHashed(input string, parents []mphSuffix, h, mul uint64) bool {
|
||||||
|
if g.mph != nil && g.mph.matchAnyHashed(input, parents, h, mul) {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
if g.ac != nil && g.ac.MatchAny(input) {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
return g.regex != nil && g.regex.MatchAny(input)
|
||||||
|
}
|
||||||
|
|
||||||
|
// MphValueMatcherCombiner combines several built MphValueMatchers, each bound to one value, and matches an input
|
||||||
|
// against them as their MatchAny would, hashing the input once for all of them.
|
||||||
|
type MphValueMatcherCombiner struct {
|
||||||
|
matchers []*MphValueMatcher
|
||||||
|
values []uint32
|
||||||
|
}
|
||||||
|
|
||||||
|
// Add adds a built matcher that stands for value.
|
||||||
|
func (s *MphValueMatcherCombiner) Add(m *MphValueMatcher, value uint32) {
|
||||||
|
s.matchers = append(s.matchers, m)
|
||||||
|
s.values = append(s.values, value)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Match returns the values of the matchers that match input, in Add order.
|
||||||
|
func (s *MphValueMatcherCombiner) Match(input string) []uint32 {
|
||||||
|
if len(s.matchers) == 0 {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
var stack [16]mphSuffix
|
||||||
|
mul := mphMultipliers[0]
|
||||||
|
parents, h := mphSuffixes(stack[:0], mul, input)
|
||||||
|
var result []uint32
|
||||||
|
for i, m := range s.matchers {
|
||||||
|
if m.matchAnyHashed(input, parents, h, mul) {
|
||||||
|
result = append(result, s.values[i])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return result
|
||||||
|
}
|
||||||
|
|
||||||
|
// MatchAny returns true as soon as one matcher matches input.
|
||||||
|
func (s *MphValueMatcherCombiner) MatchAny(input string) bool {
|
||||||
|
switch len(s.matchers) {
|
||||||
|
case 0:
|
||||||
|
return false
|
||||||
|
case 1:
|
||||||
|
return s.matchers[0].MatchAny(input) // nothing to share, and it stops at the first matching suffix
|
||||||
|
}
|
||||||
|
var stack [16]mphSuffix
|
||||||
|
mul := mphMultipliers[0]
|
||||||
|
parents, h := mphSuffixes(stack[:0], mul, input)
|
||||||
|
for _, m := range s.matchers {
|
||||||
|
if m.matchAnyHashed(input, parents, h, mul) {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|||||||
+21
-25
@@ -1,7 +1,7 @@
|
|||||||
package log // import "github.com/xtls/xray-core/common/log"
|
package log // import "github.com/xtls/xray-core/common/log"
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"sync"
|
"sync/atomic"
|
||||||
|
|
||||||
"github.com/xtls/xray-core/common/serial"
|
"github.com/xtls/xray-core/common/serial"
|
||||||
)
|
)
|
||||||
@@ -29,36 +29,32 @@ func (m *GeneralMessage) String() string {
|
|||||||
|
|
||||||
// Record writes a message into log stream.
|
// Record writes a message into log stream.
|
||||||
func Record(msg Message) {
|
func Record(msg Message) {
|
||||||
logHandler.Handle(msg)
|
if h := logHandler.Load(); h != nil {
|
||||||
|
(*h).Handle(msg)
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
var logHandler syncHandler
|
type SeverityLogger interface {
|
||||||
|
Handler
|
||||||
|
Severity() Severity
|
||||||
|
}
|
||||||
|
|
||||||
|
func GetSeverity() Severity {
|
||||||
|
if h := logHandler.Load(); h != nil {
|
||||||
|
if sh, ok := (*h).(SeverityLogger); ok {
|
||||||
|
return sh.Severity()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
// log everything by default
|
||||||
|
return Severity_Debug
|
||||||
|
}
|
||||||
|
|
||||||
|
var logHandler atomic.Pointer[Handler]
|
||||||
|
|
||||||
// RegisterHandler registers a new handler as current log handler. Previous registered handler will be discarded.
|
// RegisterHandler registers a new handler as current log handler. Previous registered handler will be discarded.
|
||||||
func RegisterHandler(handler Handler) {
|
func RegisterHandler(handler Handler) {
|
||||||
if handler == nil {
|
if handler == nil {
|
||||||
panic("Log handler is nil")
|
panic("Log handler is nil")
|
||||||
}
|
}
|
||||||
logHandler.Set(handler)
|
logHandler.Store(&handler)
|
||||||
}
|
|
||||||
|
|
||||||
type syncHandler struct {
|
|
||||||
sync.RWMutex
|
|
||||||
Handler
|
|
||||||
}
|
|
||||||
|
|
||||||
func (h *syncHandler) Handle(msg Message) {
|
|
||||||
h.RLock()
|
|
||||||
defer h.RUnlock()
|
|
||||||
|
|
||||||
if h.Handler != nil {
|
|
||||||
h.Handler.Handle(msg)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func (h *syncHandler) Set(handler Handler) {
|
|
||||||
h.Lock()
|
|
||||||
defer h.Unlock()
|
|
||||||
|
|
||||||
h.Handler = handler
|
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -68,6 +68,10 @@ func (l *serverityLogger) Handle(msg Message) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (l *serverityLogger) Severity() Severity {
|
||||||
|
return l.logLevel
|
||||||
|
}
|
||||||
|
|
||||||
func (l *generalLogger) run() {
|
func (l *generalLogger) run() {
|
||||||
defer l.access.Signal()
|
defer l.access.Signal()
|
||||||
|
|
||||||
|
|||||||
@@ -38,7 +38,7 @@ func (m *ClientManager) Dispatch(ctx context.Context, link *transport.Link) erro
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
return errors.New("unable to find an available mux client").AtWarning()
|
return errors.New("unable to find an available mux client")
|
||||||
}
|
}
|
||||||
|
|
||||||
type WorkerPicker interface {
|
type WorkerPicker interface {
|
||||||
|
|||||||
+1
-1
@@ -117,7 +117,7 @@ func (f *FrameMetadata) Unmarshal(reader io.Reader, readSourceAndLocal bool) err
|
|||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
if metaLen > 512 {
|
if metaLen > 512 {
|
||||||
return errors.New("invalid metalen ", metaLen).AtError()
|
return errors.New("invalid metalen ", metaLen)
|
||||||
}
|
}
|
||||||
|
|
||||||
b := buf.New()
|
b := buf.New()
|
||||||
|
|||||||
@@ -351,7 +351,7 @@ func (w *ServerWorker) handleFrame(ctx context.Context, reader *buf.BufferedRead
|
|||||||
err = w.handleStatusKeep(&meta, reader)
|
err = w.handleStatusKeep(&meta, reader)
|
||||||
default:
|
default:
|
||||||
status := meta.SessionStatus
|
status := meta.SessionStatus
|
||||||
return errors.New("unknown status: ", status).AtError()
|
return errors.New("unknown status: ", status)
|
||||||
}
|
}
|
||||||
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
|||||||
@@ -0,0 +1,20 @@
|
|||||||
|
package net
|
||||||
|
|
||||||
|
// PacketConnWrapper wraps a PacketConn into a Conn with a fixed destination address.
|
||||||
|
type PacketConnWrapper struct {
|
||||||
|
PacketConn
|
||||||
|
Dest Addr
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *PacketConnWrapper) Read(p []byte) (int, error) {
|
||||||
|
n, _, err := c.PacketConn.ReadFrom(p)
|
||||||
|
return n, err
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *PacketConnWrapper) Write(p []byte) (int, error) {
|
||||||
|
return c.PacketConn.WriteTo(p, c.Dest)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *PacketConnWrapper) RemoteAddr() Addr {
|
||||||
|
return c.Dest
|
||||||
|
}
|
||||||
@@ -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) {
|
||||||
|
|||||||
+1
-1
@@ -20,7 +20,7 @@ import (
|
|||||||
var (
|
var (
|
||||||
Version_x byte = 26
|
Version_x byte = 26
|
||||||
Version_y byte = 9
|
Version_y byte = 9
|
||||||
Version_z byte = 9
|
Version_z byte = 30
|
||||||
)
|
)
|
||||||
|
|
||||||
var (
|
var (
|
||||||
|
|||||||
@@ -13,7 +13,7 @@ type FakeDNSEngine interface {
|
|||||||
|
|
||||||
var (
|
var (
|
||||||
FakeIPv4Pool = "198.18.0.0/15"
|
FakeIPv4Pool = "198.18.0.0/15"
|
||||||
FakeIPv6Pool = "fc00::/18"
|
FakeIPv6Pool = "2001:2::/48"
|
||||||
)
|
)
|
||||||
|
|
||||||
type FakeDNSEngineRev0 interface {
|
type FakeDNSEngineRev0 interface {
|
||||||
|
|||||||
@@ -97,6 +97,9 @@ func New() *Client {
|
|||||||
r := &net.Resolver{
|
r := &net.Resolver{
|
||||||
PreferGo: true,
|
PreferGo: true,
|
||||||
Dial: func(ctx context.Context, network, address string) (net.Conn, error) {
|
Dial: func(ctx context.Context, network, address string) (net.Conn, error) {
|
||||||
|
if internet.IsSkippedDNSServer(address) {
|
||||||
|
return nil, errors.New("skipped DNS server ", address)
|
||||||
|
}
|
||||||
return d.DialContext(ctx, network, address)
|
return d.DialContext(ctx, network, address)
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -0,0 +1,23 @@
|
|||||||
|
package localdns
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"net/netip"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/xtls/xray-core/transport/internet"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestSkippedDNSServers(t *testing.T) {
|
||||||
|
internet.SkipDNSServers([]netip.Addr{netip.MustParseAddr("203.0.113.53")})
|
||||||
|
t.Cleanup(func() { internet.SkipDNSServers(nil) })
|
||||||
|
c := New()
|
||||||
|
if _, err := c.r.Dial(context.Background(), "udp", "203.0.113.53:53"); err == nil {
|
||||||
|
t.Error("a skipped DNS server was dialed")
|
||||||
|
}
|
||||||
|
conn, err := c.r.Dial(context.Background(), "udp", "127.0.0.1:53")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
conn.Close()
|
||||||
|
}
|
||||||
@@ -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())
|
||||||
|
|||||||
@@ -1,6 +1,7 @@
|
|||||||
package conf
|
package conf
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"context"
|
||||||
"crypto/x509"
|
"crypto/x509"
|
||||||
"encoding/base64"
|
"encoding/base64"
|
||||||
"encoding/hex"
|
"encoding/hex"
|
||||||
@@ -14,7 +15,7 @@ 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/common/serial"
|
||||||
"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"
|
||||||
@@ -82,7 +83,7 @@ var (
|
|||||||
"noise": func() interface{} { return new(NoiseMask) },
|
"noise": func() interface{} { return new(NoiseMask) },
|
||||||
"salamander": func() interface{} { return new(Salamander) },
|
"salamander": func() interface{} { return new(Salamander) },
|
||||||
"sudoku": func() interface{} { return new(Sudoku) },
|
"sudoku": func() interface{} { return new(Sudoku) },
|
||||||
"xdns": func() interface{} { return new(Xdns) },
|
"xdns": func() interface{} { return new(XDNS) },
|
||||||
"xicmp": func() interface{} { return new(Xicmp) },
|
"xicmp": func() interface{} { return new(Xicmp) },
|
||||||
"realm": func() interface{} { return new(Realm) },
|
"realm": func() interface{} { return new(Realm) },
|
||||||
"udphop": func() interface{} { return new(UDPHop) },
|
"udphop": func() interface{} { return new(UDPHop) },
|
||||||
@@ -309,14 +310,27 @@ type NoiseMask struct {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (c *NoiseMask) Build() (proto.Message, error) {
|
func (c *NoiseMask) Build() (proto.Message, error) {
|
||||||
|
noiseSlice := make([]*noise.Item, 0, len(c.Noise))
|
||||||
for _, item := range c.Noise {
|
for _, item := range c.Noise {
|
||||||
if len(item.Packet) > 0 && item.Rand.To > 0 {
|
if len(item.Packet) > 0 && item.Rand.To > 0 {
|
||||||
return nil, errors.New("len(item.Packet) > 0 && item.Rand.To > 0")
|
return nil, errors.New("len(item.Packet) > 0 && item.Rand.To > 0")
|
||||||
}
|
}
|
||||||
}
|
if strings.ToLower(item.Type) == "exp" {
|
||||||
|
var exp string
|
||||||
noiseSlice := make([]*noise.Item, 0, len(c.Noise))
|
if err := json.Unmarshal(item.Packet, &exp); err != nil {
|
||||||
for _, item := range c.Noise {
|
return nil, errors.New(`"packet" of noise "type": "exp" must be a string`).Base(err)
|
||||||
|
}
|
||||||
|
segments, err := parseNoiseExp(exp)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
noiseSlice = append(noiseSlice, &noise.Item{
|
||||||
|
Segments: segments,
|
||||||
|
DelayMin: int64(item.Delay.From),
|
||||||
|
DelayMax: int64(item.Delay.To),
|
||||||
|
})
|
||||||
|
continue
|
||||||
|
}
|
||||||
if item.RandRange == nil {
|
if item.RandRange == nil {
|
||||||
item.RandRange = &Int32Range{From: 0, To: 255}
|
item.RandRange = &Int32Range{From: 0, To: 255}
|
||||||
}
|
}
|
||||||
@@ -345,6 +359,88 @@ func (c *NoiseMask) Build() (proto.Message, error) {
|
|||||||
}, nil
|
}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
var noiseExpPattern = regexp.MustCompile(`<\s*([a-z]+)(?:\s+([^>]*?))?\s*>`)
|
||||||
|
|
||||||
|
func parseNoiseExp(exp string) ([]*noise.Segment, error) {
|
||||||
|
var segments []*noise.Segment
|
||||||
|
matches := noiseExpPattern.FindAllStringSubmatchIndex(exp, -1)
|
||||||
|
last := 0
|
||||||
|
for _, m := range matches {
|
||||||
|
if strings.TrimSpace(exp[last:m[0]]) != "" {
|
||||||
|
return nil, errors.New("invalid noise exp near ", exp[last:m[0]])
|
||||||
|
}
|
||||||
|
last = m[1]
|
||||||
|
key := exp[m[2]:m[3]]
|
||||||
|
arg := ""
|
||||||
|
if m[4] >= 0 {
|
||||||
|
arg = exp[m[4]:m[5]]
|
||||||
|
}
|
||||||
|
segment, err := buildNoiseSegment(key, arg)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
segments = append(segments, segment)
|
||||||
|
}
|
||||||
|
if strings.TrimSpace(exp[last:]) != "" {
|
||||||
|
return nil, errors.New("invalid noise exp near ", exp[last:])
|
||||||
|
}
|
||||||
|
if len(segments) == 0 {
|
||||||
|
return nil, errors.New("empty noise exp: ", exp)
|
||||||
|
}
|
||||||
|
return segments, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func buildNoiseSegment(key, arg string) (*noise.Segment, error) {
|
||||||
|
sizeSegment := func(kind noise.Segment_Kind) (*noise.Segment, error) {
|
||||||
|
if arg == "" {
|
||||||
|
return nil, errors.New("<", key, "> in noise exp needs a size")
|
||||||
|
}
|
||||||
|
lo, hi, err := ParseRangeString(arg)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
if lo < 0 || hi < lo || hi > 65535 {
|
||||||
|
return nil, errors.New("invalid size in noise exp: ", arg)
|
||||||
|
}
|
||||||
|
return &noise.Segment{Kind: kind, MinSize: int64(lo), MaxSize: int64(hi)}, nil
|
||||||
|
}
|
||||||
|
switch key {
|
||||||
|
case "b":
|
||||||
|
hexStr := strings.TrimPrefix(strings.TrimPrefix(strings.Join(strings.Fields(arg), ""), "0x"), "0X")
|
||||||
|
if len(hexStr) == 0 {
|
||||||
|
return nil, errors.New("empty bytes in noise exp")
|
||||||
|
}
|
||||||
|
raw, err := hex.DecodeString(hexStr)
|
||||||
|
if err != nil {
|
||||||
|
return nil, errors.New("invalid hex in noise exp: ", arg).Base(err)
|
||||||
|
}
|
||||||
|
return &noise.Segment{Kind: noise.Segment_BYTES, Bytes: raw}, nil
|
||||||
|
case "r":
|
||||||
|
return sizeSegment(noise.Segment_RANDOM)
|
||||||
|
case "rc":
|
||||||
|
return sizeSegment(noise.Segment_RANDOM_ASCII)
|
||||||
|
case "rd":
|
||||||
|
return sizeSegment(noise.Segment_RANDOM_DIGIT)
|
||||||
|
case "t":
|
||||||
|
if arg != "" {
|
||||||
|
return nil, errors.New("<t> in noise exp takes no argument")
|
||||||
|
}
|
||||||
|
return &noise.Segment{Kind: noise.Segment_TIMESTAMP}, nil
|
||||||
|
case "c":
|
||||||
|
if arg != "" {
|
||||||
|
return nil, errors.New("<c> in noise exp takes no argument")
|
||||||
|
}
|
||||||
|
return &noise.Segment{Kind: noise.Segment_COUNTER}, nil
|
||||||
|
case "n":
|
||||||
|
if arg != "" {
|
||||||
|
return nil, errors.New("<n> in noise exp takes no argument")
|
||||||
|
}
|
||||||
|
return &noise.Segment{Kind: noise.Segment_NONCE}, nil
|
||||||
|
default:
|
||||||
|
return nil, errors.New("unknown <", key, "> in noise exp")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
type UDPItem struct {
|
type UDPItem struct {
|
||||||
Rand int32 `json:"rand"`
|
Rand int32 `json:"rand"`
|
||||||
RandRange *Int32Range `json:"randRange"`
|
RandRange *Int32Range `json:"randRange"`
|
||||||
@@ -695,32 +791,88 @@ func (c *Sudoku) Build() (proto.Message, error) {
|
|||||||
}, nil
|
}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
type Xdns struct {
|
type XDNSDomain struct {
|
||||||
Domain json.RawMessage `json:"domain"`
|
Name string `json:"name"`
|
||||||
|
LenLimit int32 `json:"lenLimit"`
|
||||||
Domains []string `json:"domains"`
|
LabelLimit int32 `json:"labelLimit"`
|
||||||
Resolvers []string `json:"resolvers"`
|
Types []int32 `json:"types"`
|
||||||
|
Edns0 int32 `json:"edns0"`
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *Xdns) Build() (proto.Message, error) {
|
type XDNSResolverTCP struct {
|
||||||
if c.Domain != nil {
|
Addr string `json:"addr"`
|
||||||
return nil, errors.PrintRemovedFeatureError("domain", "domains(server) & resolvers(client)")
|
}
|
||||||
}
|
|
||||||
|
|
||||||
if len(c.Domains) == 0 && len(c.Resolvers) == 0 {
|
func (c *XDNSResolverTCP) Build() (proto.Message, error) {
|
||||||
return nil, errors.New("empty domains & empty resolvers")
|
return &xdns.TCPResolverProto{Addr: c.Addr}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
for _, r := range c.Resolvers {
|
type XDNSResolverUDP struct {
|
||||||
if !strings.Contains(r, "+udp://") {
|
Addr string `json:"addr"`
|
||||||
return nil, errors.New("invalid resolver ", r)
|
}
|
||||||
|
|
||||||
|
func (c *XDNSResolverUDP) Build() (proto.Message, error) {
|
||||||
|
return &xdns.UDPResolverProto{Addr: c.Addr}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
var xdnsLoader = NewJSONConfigLoader(ConfigCreatorCache{
|
||||||
|
"tcp": func() interface{} { return new(XDNSResolverTCP) },
|
||||||
|
"udp": func() interface{} { return new(XDNSResolverUDP) },
|
||||||
|
}, "type", "settings")
|
||||||
|
|
||||||
|
type XDNSResolver struct {
|
||||||
|
Type string `json:"type"`
|
||||||
|
Settings json.RawMessage `json:"settings"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type XDNS struct {
|
||||||
|
Domains []XDNSDomain `json:"domains"`
|
||||||
|
Resolvers []XDNSResolver `json:"resolvers"`
|
||||||
|
ExtraPoll int32 `json:"extraPoll"`
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *XDNS) Build() (proto.Message, error) {
|
||||||
|
var domains []*xdns.DomainProto
|
||||||
|
var resolvers []*serial.TypedMessage
|
||||||
|
for i := range c.Domains {
|
||||||
|
if c.Domains[i].LenLimit == 0 {
|
||||||
|
c.Domains[i].LenLimit = 255
|
||||||
}
|
}
|
||||||
|
if c.Domains[i].LabelLimit == 0 {
|
||||||
|
c.Domains[i].LabelLimit = 63
|
||||||
|
}
|
||||||
|
types := make([]uint16, 0, len(c.Domains[i].Types))
|
||||||
|
for j := range c.Domains[i].Types {
|
||||||
|
types = append(types, uint16(c.Domains[i].Types[j]))
|
||||||
|
}
|
||||||
|
domain, err := xdns.NewDomain(c.Domains[i].Name, int(c.Domains[i].LenLimit), int(c.Domains[i].LabelLimit), types, uint16(c.Domains[i].Edns0))
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
errors.LogInfo(context.Background(), domain.Show())
|
||||||
|
domains = append(domains, &xdns.DomainProto{
|
||||||
|
Name: c.Domains[i].Name,
|
||||||
|
LenLimit: c.Domains[i].LenLimit,
|
||||||
|
LabelLimit: c.Domains[i].LabelLimit,
|
||||||
|
Types: c.Domains[i].Types,
|
||||||
|
Edns0: c.Domains[i].Edns0,
|
||||||
|
})
|
||||||
}
|
}
|
||||||
|
for i := range c.Resolvers {
|
||||||
return &xdns.Config{
|
config, err := xdnsLoader.LoadWithID(c.Resolvers[i].Settings, c.Resolvers[i].Type)
|
||||||
Domains: c.Domains,
|
if err != nil {
|
||||||
Resolvers: c.Resolvers,
|
return nil, err
|
||||||
}, nil
|
}
|
||||||
|
pm, err := config.(interface{ Build() (proto.Message, error) }).Build()
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
resolvers = append(resolvers, serial.ToTypedMessage(pm))
|
||||||
|
}
|
||||||
|
if c.ExtraPoll < 0 || c.ExtraPoll > 3 {
|
||||||
|
return nil, errors.New("c.ExtraPoll < 0 || c.ExtraPoll > 3")
|
||||||
|
}
|
||||||
|
return &xdns.Config{Domains: domains, Resolvers: resolvers, ExtraPoll: c.ExtraPoll}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
type XMC struct {
|
type XMC struct {
|
||||||
@@ -909,22 +1061,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 +1095,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
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,136 @@
|
|||||||
|
package conf
|
||||||
|
|
||||||
|
import (
|
||||||
|
"encoding/json"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/xtls/xray-core/transport/internet/finalmask/noise"
|
||||||
|
)
|
||||||
|
|
||||||
|
func expPacket(exp string) json.RawMessage {
|
||||||
|
b, _ := json.Marshal(exp)
|
||||||
|
return b
|
||||||
|
}
|
||||||
|
|
||||||
|
func buildNoiseExp(exp string) (*noise.Config, error) {
|
||||||
|
msg, err := (&NoiseMask{Noise: []NoiseItem{{Type: "exp", Packet: expPacket(exp)}}}).Build()
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
return msg.(*noise.Config), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestNoiseExp(t *testing.T) {
|
||||||
|
cfg, err := buildNoiseExp("<b 0d0a0d0a><t><r 24><rc 20-40><rd 8><c><n>")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
segments := cfg.Items[0].Segments
|
||||||
|
if len(segments) != 7 {
|
||||||
|
t.Fatalf("got %d segments, want 7", len(segments))
|
||||||
|
}
|
||||||
|
want := []struct {
|
||||||
|
kind noise.Segment_Kind
|
||||||
|
bytes []byte
|
||||||
|
min, max int64
|
||||||
|
}{
|
||||||
|
{noise.Segment_BYTES, []byte{0x0d, 0x0a, 0x0d, 0x0a}, 0, 0},
|
||||||
|
{noise.Segment_TIMESTAMP, nil, 0, 0},
|
||||||
|
{noise.Segment_RANDOM, nil, 24, 24},
|
||||||
|
{noise.Segment_RANDOM_ASCII, nil, 20, 40},
|
||||||
|
{noise.Segment_RANDOM_DIGIT, nil, 8, 8},
|
||||||
|
{noise.Segment_COUNTER, nil, 0, 0},
|
||||||
|
{noise.Segment_NONCE, nil, 0, 0},
|
||||||
|
}
|
||||||
|
for i, w := range want {
|
||||||
|
s := segments[i]
|
||||||
|
if s.Kind != w.kind || s.MinSize != w.min || s.MaxSize != w.max || string(s.Bytes) != string(w.bytes) {
|
||||||
|
t.Errorf("segment %d = %+v, want %+v", i, s, w)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestNoiseExpStripsHexPrefix(t *testing.T) {
|
||||||
|
cfg, err := buildNoiseExp("<b 0x16030100>")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if got := cfg.Items[0].Segments[0].Bytes; string(got) != string([]byte{0x16, 0x03, 0x01, 0x00}) {
|
||||||
|
t.Errorf("got %x", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestNoiseExpWhitespace(t *testing.T) {
|
||||||
|
if _, err := buildNoiseExp(" <b 00> <t> "); err != nil {
|
||||||
|
t.Errorf("surrounding whitespace should be allowed: %v", err)
|
||||||
|
}
|
||||||
|
cfg, err := buildNoiseExp("<b 0d 0a 0d 0a>")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if got := cfg.Items[0].Segments[0].Bytes; string(got) != "\r\n\r\n" {
|
||||||
|
t.Errorf("got %x", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestNoiseExpRejects(t *testing.T) {
|
||||||
|
for _, exp := range []string{
|
||||||
|
"<x 1>",
|
||||||
|
"<b>",
|
||||||
|
"<b zz>",
|
||||||
|
"<b 0d0>",
|
||||||
|
"<r>",
|
||||||
|
"<r -1>",
|
||||||
|
"<r 40-20>",
|
||||||
|
"<r 70000>",
|
||||||
|
"<t 5>",
|
||||||
|
"<n 5>",
|
||||||
|
"garbage<t>",
|
||||||
|
"<t> tail",
|
||||||
|
"<t><b>",
|
||||||
|
} {
|
||||||
|
if _, err := buildNoiseExp(exp); err == nil {
|
||||||
|
t.Errorf("expected an error for %q", exp)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestNoiseExpConflicts(t *testing.T) {
|
||||||
|
if _, err := (&NoiseMask{Noise: []NoiseItem{{Type: "exp", Packet: expPacket("<t>"), Rand: Int32Range{From: 10, To: 20}}}}).Build(); err == nil {
|
||||||
|
t.Error("exp with rand should be rejected")
|
||||||
|
}
|
||||||
|
for _, packet := range []string{``, `[1, 2]`, `5`} {
|
||||||
|
if _, err := (&NoiseMask{Noise: []NoiseItem{{Type: "exp", Packet: json.RawMessage(packet)}}}).Build(); err == nil {
|
||||||
|
t.Errorf("expected an error for packet %q", packet)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestNoiseExpFromJSON(t *testing.T) {
|
||||||
|
var mask NoiseMask
|
||||||
|
if err := json.Unmarshal([]byte(`{"noise": [
|
||||||
|
{"type": "exp", "packet": "<b 504f5354><rd 10-20>", "delay": "1-3"},
|
||||||
|
{"type": "EXP", "packet": "<t>"},
|
||||||
|
{"type": "str", "packet": "<t>"},
|
||||||
|
{"rand": "10-20"}
|
||||||
|
]}`), &mask); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
msg, err := mask.Build()
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
items := msg.(*noise.Config).Items
|
||||||
|
if len(items[0].Segments) != 2 || items[0].DelayMin != 1 || items[0].DelayMax != 3 {
|
||||||
|
t.Errorf("item 0 = %+v", items[0])
|
||||||
|
}
|
||||||
|
if len(items[1].Segments) != 1 || items[1].Segments[0].Kind != noise.Segment_TIMESTAMP {
|
||||||
|
t.Errorf("item 1 = %+v", items[1])
|
||||||
|
}
|
||||||
|
if len(items[2].Segments) != 0 || string(items[2].Packet) != "<t>" {
|
||||||
|
t.Errorf("item 2 = %+v", items[2])
|
||||||
|
}
|
||||||
|
if len(items[3].Segments) != 0 || items[3].RandMin != 10 || items[3].RandMax != 20 {
|
||||||
|
t.Errorf("item 3 = %+v", items[3])
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -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")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
@@ -5,8 +5,12 @@ import (
|
|||||||
"fmt"
|
"fmt"
|
||||||
"math/big"
|
"math/big"
|
||||||
"net"
|
"net"
|
||||||
|
"runtime"
|
||||||
|
"slices"
|
||||||
"strconv"
|
"strconv"
|
||||||
|
"strings"
|
||||||
|
|
||||||
|
"github.com/xtls/xray-core/common/errors"
|
||||||
"github.com/xtls/xray-core/proxy/tun"
|
"github.com/xtls/xray-core/proxy/tun"
|
||||||
"google.golang.org/protobuf/proto"
|
"google.golang.org/protobuf/proto"
|
||||||
)
|
)
|
||||||
@@ -20,6 +24,8 @@ 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"`
|
||||||
|
AutoSystemDnsToGateway bool `json:"autoSystemDnsToGateway"`
|
||||||
|
AutoSystemWfpBlockLeak []string `json:"autoSystemWfpBlockLeak"`
|
||||||
}
|
}
|
||||||
|
|
||||||
func (v *TunConfig) Build() (proto.Message, error) {
|
func (v *TunConfig) Build() (proto.Message, error) {
|
||||||
@@ -31,6 +37,32 @@ 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,
|
||||||
|
AutoSystemDnsToGateway: v.AutoSystemDnsToGateway,
|
||||||
|
}
|
||||||
|
for _, leak := range v.AutoSystemWfpBlockLeak {
|
||||||
|
switch leak := strings.ToLower(leak); leak {
|
||||||
|
case "dns", "misconfigtun":
|
||||||
|
config.AutoSystemWfpBlockLeak = append(config.AutoSystemWfpBlockLeak, leak)
|
||||||
|
default:
|
||||||
|
return nil, errors.New("unknown autoSystemWfpBlockLeak value: ", leak)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
// Each option needs other settings on the system it takes effect on: the
|
||||||
|
// filters go along with the routes of autoSystemRoutingTable, "dns" lets
|
||||||
|
// DNS through the TUN only, and autoSystemDnsToGateway points the system
|
||||||
|
// DNS at the gateway.
|
||||||
|
switch runtime.GOOS {
|
||||||
|
case "windows":
|
||||||
|
if len(config.AutoSystemWfpBlockLeak) > 0 && len(v.AutoSystemRoutingTable) == 0 {
|
||||||
|
return nil, errors.New("autoSystemWfpBlockLeak needs autoSystemRoutingTable to be set")
|
||||||
|
}
|
||||||
|
if slices.Contains(config.AutoSystemWfpBlockLeak, "dns") && len(v.DNS) == 0 {
|
||||||
|
return nil, errors.New(`autoSystemWfpBlockLeak "dns" needs dns to be set`)
|
||||||
|
}
|
||||||
|
case "linux":
|
||||||
|
if v.AutoSystemDnsToGateway && len(v.Gateway) == 0 {
|
||||||
|
return nil, errors.New("autoSystemDnsToGateway needs gateway to be set")
|
||||||
|
}
|
||||||
}
|
}
|
||||||
if v.AutoOutboundsInterface != nil {
|
if v.AutoOutboundsInterface != nil {
|
||||||
config.AutoOutboundsInterface = *v.AutoOutboundsInterface
|
config.AutoOutboundsInterface = *v.AutoOutboundsInterface
|
||||||
|
|||||||
@@ -0,0 +1,71 @@
|
|||||||
|
package conf_test
|
||||||
|
|
||||||
|
import (
|
||||||
|
"encoding/json"
|
||||||
|
"runtime"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
. "github.com/xtls/xray-core/infra/conf"
|
||||||
|
"github.com/xtls/xray-core/proxy/tun"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestTunConfigAutoSystem(t *testing.T) {
|
||||||
|
creator := func() Buildable {
|
||||||
|
return new(TunConfig)
|
||||||
|
}
|
||||||
|
|
||||||
|
runMultiTestCase(t, []TestCase{
|
||||||
|
{
|
||||||
|
Input: `{"name": "xray0"}`,
|
||||||
|
Parser: loadJSON(creator),
|
||||||
|
Output: &tun.Config{Name: "xray0", Desc: "Wintun", MTU: 1500},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
Input: `{"name": "xray0", "gateway": ["10.0.0.1/24"], "autoSystemDnsToGateway": true}`,
|
||||||
|
Parser: loadJSON(creator),
|
||||||
|
Output: &tun.Config{Name: "xray0", Desc: "Wintun", MTU: 1500, Gateway: []string{"10.0.0.1/24"}, AutoSystemDnsToGateway: true},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
Input: `{"name": "xray0", "dns": ["1.1.1.1"], "autoSystemRoutingTable": ["0.0.0.0/0"], "autoSystemWfpBlockLeak": ["dns", "misconfigtun"]}`,
|
||||||
|
Parser: loadJSON(creator),
|
||||||
|
Output: &tun.Config{Name: "xray0", Desc: "Wintun", MTU: 1500, DNS: []string{"1.1.1.1"}, AutoSystemRoutingTable: []string{"0.0.0.0/0"}, AutoOutboundsInterface: "auto", AutoSystemWfpBlockLeak: []string{"dns", "misconfigtun"}},
|
||||||
|
},
|
||||||
|
{
|
||||||
|
Input: `{"name": "xray0", "dns": ["1.1.1.1"], "autoSystemRoutingTable": ["0.0.0.0/0"], "autoSystemWfpBlockLeak": ["DNS"]}`,
|
||||||
|
Parser: loadJSON(creator),
|
||||||
|
Output: &tun.Config{Name: "xray0", Desc: "Wintun", MTU: 1500, DNS: []string{"1.1.1.1"}, AutoSystemRoutingTable: []string{"0.0.0.0/0"}, AutoOutboundsInterface: "auto", AutoSystemWfpBlockLeak: []string{"dns"}},
|
||||||
|
},
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestTunConfigAutoSystemNeeds checks that an option is rejected without the
|
||||||
|
// setting it needs, only on the system it takes effect on.
|
||||||
|
func TestTunConfigAutoSystemNeeds(t *testing.T) {
|
||||||
|
for _, c := range []struct {
|
||||||
|
input string
|
||||||
|
goos string // where it is rejected
|
||||||
|
}{
|
||||||
|
{`{"name": "xray0", "autoSystemWfpBlockLeak": ["misconfigtun"]}`, "windows"},
|
||||||
|
{`{"name": "xray0", "autoSystemRoutingTable": ["0.0.0.0/0"], "autoSystemWfpBlockLeak": ["misconfigtun"]}`, ""},
|
||||||
|
{`{"name": "xray0", "autoSystemRoutingTable": ["0.0.0.0/0"], "autoSystemWfpBlockLeak": ["dns"]}`, "windows"},
|
||||||
|
{`{"name": "xray0", "autoSystemDnsToGateway": true}`, "linux"},
|
||||||
|
} {
|
||||||
|
config := new(TunConfig)
|
||||||
|
if err := json.Unmarshal([]byte(c.input), config); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if _, err := config.Build(); (err != nil) != (runtime.GOOS == c.goos) {
|
||||||
|
t.Errorf("%s on %s: error = %v", c.input, runtime.GOOS, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestTunConfigAutoSystemWfpBlockLeakUnknown(t *testing.T) {
|
||||||
|
config := new(TunConfig)
|
||||||
|
if err := json.Unmarshal([]byte(`{"name": "xray0", "autoSystemWfpBlockLeak": ["dns", "ip"]}`), config); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if _, err := config.Build(); err == nil {
|
||||||
|
t.Error("an unknown autoSystemWfpBlockLeak value was accepted")
|
||||||
|
}
|
||||||
|
}
|
||||||
+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,8 @@ 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/hysteria"
|
||||||
|
"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 +90,10 @@ 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
|
||||||
|
case *hysteria.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"
|
||||||
|
|||||||
@@ -467,7 +467,7 @@ func NewPacketReader(conn net.Conn, h *Handler, defaultRule *FinalRule, UDPOverr
|
|||||||
if statConn != nil {
|
if statConn != nil {
|
||||||
counter = statConn.ReadCounter
|
counter = statConn.ReadCounter
|
||||||
}
|
}
|
||||||
if c, ok := iConn.(*internet.PacketConnWrapper); ok {
|
if c, ok := iConn.(*net.PacketConnWrapper); ok {
|
||||||
isOverridden := false
|
isOverridden := false
|
||||||
if UDPOverride.Address != nil || UDPOverride.Port != 0 {
|
if UDPOverride.Address != nil || UDPOverride.Port != 0 {
|
||||||
isOverridden = true
|
isOverridden = true
|
||||||
@@ -487,7 +487,7 @@ func NewPacketReader(conn net.Conn, h *Handler, defaultRule *FinalRule, UDPOverr
|
|||||||
}
|
}
|
||||||
|
|
||||||
type PacketReader struct {
|
type PacketReader struct {
|
||||||
*internet.PacketConnWrapper
|
*net.PacketConnWrapper
|
||||||
stats.Counter
|
stats.Counter
|
||||||
Handler *Handler
|
Handler *Handler
|
||||||
DefaultRule *FinalRule
|
DefaultRule *FinalRule
|
||||||
@@ -542,7 +542,7 @@ func NewPacketWriter(conn net.Conn, h *Handler, defaultRule *FinalRule, UDPOverr
|
|||||||
if statConn != nil {
|
if statConn != nil {
|
||||||
counter = statConn.WriteCounter
|
counter = statConn.WriteCounter
|
||||||
}
|
}
|
||||||
if c, ok := iConn.(*internet.PacketConnWrapper); ok {
|
if c, ok := iConn.(*net.PacketConnWrapper); ok {
|
||||||
// If DialDest is a domain, it will be resolved in dialer
|
// If DialDest is a domain, it will be resolved in dialer
|
||||||
// check this behavior and add it to map
|
// check this behavior and add it to map
|
||||||
resolvedUDPAddr := utils.NewTypedSyncMap[string, net.Address]()
|
resolvedUDPAddr := utils.NewTypedSyncMap[string, net.Address]()
|
||||||
@@ -563,7 +563,7 @@ func NewPacketWriter(conn net.Conn, h *Handler, defaultRule *FinalRule, UDPOverr
|
|||||||
}
|
}
|
||||||
|
|
||||||
type PacketWriter struct {
|
type PacketWriter struct {
|
||||||
*internet.PacketConnWrapper
|
*net.PacketConnWrapper
|
||||||
stats.Counter
|
stats.Counter
|
||||||
*Handler
|
*Handler
|
||||||
DefaultRule *FinalRule
|
DefaultRule *FinalRule
|
||||||
|
|||||||
@@ -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.(*net.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),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
+131
-129
@@ -4,23 +4,16 @@ import (
|
|||||||
"context"
|
"context"
|
||||||
"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"
|
||||||
"github.com/xtls/xray-core/common/net"
|
"github.com/xtls/xray-core/common/net"
|
||||||
"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/core"
|
||||||
"github.com/xtls/xray-core/common/singbridge"
|
"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 +25,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 +42,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 +81,105 @@ 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 {
|
||||||
if err != nil {
|
defer conn.Close()
|
||||||
buf.ReleaseMulti(mb)
|
|
||||||
return singbridge.ReturnError(err)
|
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)
|
||||||
|
}
|
||||||
|
|
||||||
|
// 1. Single read call for Salt + Fixed-length header chunk per SIP022 §3.1.4
|
||||||
|
headerLen := i.method.KeySaltLength + RequestHeaderFixedChunkLength + AEADTagSize
|
||||||
|
headerBuf := make([]byte, headerLen)
|
||||||
|
n, err := conn.Read(headerBuf)
|
||||||
|
if err != nil || n < headerLen {
|
||||||
|
ResetTCPConn(conn)
|
||||||
|
return errors.New("failed to read complete handshake header")
|
||||||
|
}
|
||||||
|
|
||||||
|
var salt [32]byte
|
||||||
|
copy(salt[:i.method.KeySaltLength], headerBuf[:i.method.KeySaltLength])
|
||||||
|
saltSlice := salt[:i.method.KeySaltLength]
|
||||||
|
fixedChunk := headerBuf[i.method.KeySaltLength:]
|
||||||
|
|
||||||
|
reader, reqHeader, err := InitServerStream(conn, i.method, i.psk, saltSlice, salt, fixedChunk, i.saltFilter)
|
||||||
|
if err != nil {
|
||||||
|
ResetTCPConn(conn)
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
dest := reqHeader.Destination
|
||||||
|
|
||||||
|
writer := NewServerStreamWriter(conn, i.method, i.psk, saltSlice)
|
||||||
|
|
||||||
|
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 {
|
||||||
|
mb := buf.MergeBytes(nil, reqHeader.EarlyData)
|
||||||
|
if err := link.Writer.WriteMultiBuffer(mb); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return TransportTCP(ctx, i.policyManager.ForLevel(uint32(i.user.Level)), reader, writer, link)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (i *Inbound) processUDP(ctx context.Context, conn stat.Connection, dispatcher routing.Dispatcher) error {
|
||||||
|
reader := buf.NewPacketReader(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())
|
||||||
|
b.Release()
|
||||||
|
if err != nil || decoded.HeaderType != HeaderTypeClient {
|
||||||
|
continue
|
||||||
}
|
}
|
||||||
for _, buffer := range mb {
|
|
||||||
packet := B.As(buffer.Bytes()).ToOwned()
|
sessionItem := i.udpCodec.GetSession(decoded.SessionID)
|
||||||
buffer.Release()
|
if sessionItem.User == nil {
|
||||||
err = i.service.NewPacket(ctx, pc, packet, metadata)
|
sessionItem.Lock()
|
||||||
if err != nil {
|
if sessionItem.User == nil {
|
||||||
packet.Release()
|
sessionItem.User = i.user
|
||||||
buf.ReleaseMulti(mb)
|
|
||||||
return err
|
|
||||||
}
|
}
|
||||||
|
sessionItem.Unlock()
|
||||||
}
|
}
|
||||||
|
link, err := sessionItem.EnsureLink(ctx, conn, decoded.Destination, dispatcher, i.policyManager, func(dest net.Destination, payload []byte) ([]byte, error) {
|
||||||
|
return i.udpCodec.EncodeServerPacket(decoded.SessionID, dest, payload)
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
payloadBuf := buf.New()
|
||||||
|
payloadBuf.Write(decoded.Payload)
|
||||||
|
payloadBuf.UDP = &decoded.Destination
|
||||||
|
_ = 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,30 +2,26 @@ package shadowsocks_2022
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"encoding/base64"
|
"crypto/cipher"
|
||||||
|
"encoding/binary"
|
||||||
"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"
|
||||||
"github.com/xtls/xray-core/common/net"
|
"github.com/xtls/xray-core/common/net"
|
||||||
"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/utils"
|
||||||
"github.com/xtls/xray-core/common/singbridge"
|
|
||||||
"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 +34,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 +54,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 +190,167 @@ 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. Single read call for Salt + EIH + Fixed-length header chunk per SIP022 §3.1.4
|
||||||
|
headerLen := i.method.KeySaltLength + AESBlockSize + RequestHeaderFixedChunkLength + AEADTagSize
|
||||||
|
headerBuf := make([]byte, headerLen)
|
||||||
|
n, err := conn.Read(headerBuf)
|
||||||
|
if err != nil || n < headerLen {
|
||||||
|
ResetTCPConn(conn)
|
||||||
|
return errors.New("failed to read complete handshake header")
|
||||||
|
}
|
||||||
|
|
||||||
if network == net.Network_TCP {
|
var salt [32]byte
|
||||||
return singbridge.ReturnError(i.service.NewConnection(ctx, connection, metadata))
|
copy(salt[:i.method.KeySaltLength], headerBuf[:i.method.KeySaltLength])
|
||||||
} else {
|
saltSlice := salt[:i.method.KeySaltLength]
|
||||||
reader := buf.NewReader(connection)
|
eih := headerBuf[i.method.KeySaltLength : i.method.KeySaltLength+AESBlockSize]
|
||||||
pc := &natPacketConn{connection}
|
fixedChunk := headerBuf[i.method.KeySaltLength+AESBlockSize:]
|
||||||
for {
|
|
||||||
mb, err := reader.ReadMultiBuffer()
|
decryptedHash, err := DecryptEIH(i.method, i.masterPSK, saltSlice, eih)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
buf.ReleaseMulti(mb)
|
ResetTCPConn(conn)
|
||||||
return singbridge.ReturnError(err)
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
// Lookup user
|
||||||
|
user, ok := i.usersByHash.Load(decryptedHash)
|
||||||
|
if !ok {
|
||||||
|
ResetTCPConn(conn)
|
||||||
|
return ErrInvalidRequest
|
||||||
|
}
|
||||||
|
userPSK := user.Account.(*MemoryAccount).Key
|
||||||
|
|
||||||
|
reader, reqHeader, err := InitServerStream(conn, i.method, userPSK, saltSlice, salt, fixedChunk, i.saltFilter)
|
||||||
|
if err != nil {
|
||||||
|
ResetTCPConn(conn)
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
dest := reqHeader.Destination
|
||||||
|
|
||||||
|
writer := NewServerStreamWriter(conn, i.method, userPSK, saltSlice)
|
||||||
|
|
||||||
|
// 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 {
|
||||||
|
mb := buf.MergeBytes(nil, reqHeader.EarlyData)
|
||||||
|
if err := link.Writer.WriteMultiBuffer(mb); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return TransportTCP(ctx, i.policyManager.ForLevel(user.Level), reader, writer, link)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (i *MultiUserInbound) processUDP(ctx context.Context, conn stat.Connection, dispatcher routing.Dispatcher) error {
|
||||||
|
reader := buf.NewPacketReader(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)
|
|
||||||
if err != nil {
|
sessionID := binary.BigEndian.Uint64(rawHeader[:8])
|
||||||
packet.Release()
|
packetID := binary.BigEndian.Uint64(rawHeader[8:16])
|
||||||
buf.ReleaseMulti(mb)
|
|
||||||
return err
|
sessionItem := i.udpSessions.GetOrCreate(sessionID)
|
||||||
|
|
||||||
|
if !sessionItem.CheckPacketID(packetID) {
|
||||||
|
b.Release()
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
var userPSK []byte
|
||||||
|
var currentUser *protocol.MemoryUser
|
||||||
|
sessionItem.Lock()
|
||||||
|
currentUser = sessionItem.User
|
||||||
|
userPSK = sessionItem.UserPSK
|
||||||
|
sessionItem.Unlock()
|
||||||
|
|
||||||
|
if currentUser == nil {
|
||||||
|
// Decrypt EIH
|
||||||
|
decryptedHash := DecryptUDPEIH(i.udpMasterCipher, rawHeader[:], packetBytes[16:32])
|
||||||
|
|
||||||
|
user, ok := i.usersByHash.Load(decryptedHash)
|
||||||
|
if !ok {
|
||||||
|
b.Release()
|
||||||
|
continue
|
||||||
}
|
}
|
||||||
|
currentUser = user
|
||||||
|
userPSK = user.Account.(*MemoryAccount).Key
|
||||||
}
|
}
|
||||||
|
|
||||||
|
decoded, err := sessionItem.DecryptAESPayload(i.method, userPSK, sessionID, packetID, rawHeader[:], packetBytes[32:])
|
||||||
|
b.Release()
|
||||||
|
if err != nil {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
sessionItem.Lock()
|
||||||
|
if sessionItem.User == nil {
|
||||||
|
sessionItem.User = currentUser
|
||||||
|
sessionItem.UserPSK = userPSK
|
||||||
|
}
|
||||||
|
sessionItem.Unlock()
|
||||||
|
|
||||||
|
link, err := sessionItem.EnsureLink(ctx, conn, decoded.Destination, dispatcher, i.policyManager, func(replyDest net.Destination, payload []byte) ([]byte, error) {
|
||||||
|
return i.encodeServerUDPPacket(sessionID, userPSK, replyDest, payload)
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
pBuf := buf.New()
|
||||||
|
pBuf.Write(decoded.Payload)
|
||||||
|
pBuf.UDP = &decoded.Destination
|
||||||
|
_ = 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)
|
return i.udpSessions.EncodeServerPacket(i.method, userPSK, clientSessionID, dest, payload)
|
||||||
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 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, 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,11 @@ package shadowsocks_2022
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
|
"crypto/cipher"
|
||||||
|
"encoding/binary"
|
||||||
"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"
|
||||||
@@ -21,9 +14,9 @@ import (
|
|||||||
"github.com/xtls/xray-core/common/net"
|
"github.com/xtls/xray-core/common/net"
|
||||||
"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/singbridge"
|
|
||||||
"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 +27,21 @@ func init() {
|
|||||||
}))
|
}))
|
||||||
}
|
}
|
||||||
|
|
||||||
|
type relayDest struct {
|
||||||
|
destination net.Destination
|
||||||
|
email string
|
||||||
|
level uint32
|
||||||
|
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
|
||||||
|
udpSessions *UDPSessionManager
|
||||||
|
policyManager policy.Manager
|
||||||
}
|
}
|
||||||
|
|
||||||
func NewRelayServer(ctx context.Context, config *RelayServerConfig) (*RelayInbound, error) {
|
func NewRelayServer(ctx context.Context, config *RelayServerConfig) (*RelayInbound, error) {
|
||||||
@@ -48,39 +52,62 @@ 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),
|
||||||
|
udpSessions: NewUDPSessionManager(500 * time.Second),
|
||||||
|
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),
|
||||||
|
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 +119,147 @@ 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 initial handshake in a single read call per SIP022 §3.1.3 & §3.1.4
|
||||||
|
needed := i.method.KeySaltLength + AESBlockSize
|
||||||
|
requestHeader := buf.New()
|
||||||
|
n, err := requestHeader.ReadFrom(conn)
|
||||||
|
if err != nil {
|
||||||
|
requestHeader.Release()
|
||||||
|
ResetTCPConn(conn)
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if int(n) < needed {
|
||||||
|
requestHeader.Release()
|
||||||
|
ResetTCPConn(conn)
|
||||||
|
return ErrInvalidRequest
|
||||||
|
}
|
||||||
|
|
||||||
if network == net.Network_TCP {
|
headerSlice := requestHeader.Bytes()
|
||||||
return singbridge.ReturnError(i.service.NewConnection(ctx, connection, metadata))
|
salt := headerSlice[:i.method.KeySaltLength]
|
||||||
} else {
|
eih := headerSlice[i.method.KeySaltLength:needed]
|
||||||
reader := buf.NewReader(connection)
|
|
||||||
pc := &natPacketConn{connection}
|
decryptedHash, err := DecryptEIH(i.method, i.relayPSK, salt, eih)
|
||||||
for {
|
if err != nil {
|
||||||
mb, err := reader.ReadMultiBuffer()
|
requestHeader.Release()
|
||||||
if err != nil {
|
ResetTCPConn(conn)
|
||||||
buf.ReleaseMulti(mb)
|
return err
|
||||||
return singbridge.ReturnError(err)
|
}
|
||||||
|
|
||||||
|
targetDest, ok := i.destinations[decryptedHash]
|
||||||
|
if !ok {
|
||||||
|
requestHeader.Release()
|
||||||
|
ResetTCPConn(conn)
|
||||||
|
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 {
|
||||||
|
requestHeader.Release()
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
// Unwrap outer EIH: send client salt and remaining handshake bytes to next hop
|
||||||
|
// in a single write call, satisfying downstream server's single-read handshake expectation (SIP022 §3.1.3).
|
||||||
|
var saltCopy [32]byte
|
||||||
|
copy(saltCopy[:i.method.KeySaltLength], salt)
|
||||||
|
copy(requestHeader.Bytes()[AESBlockSize:AESBlockSize+i.method.KeySaltLength], saltCopy[:i.method.KeySaltLength])
|
||||||
|
requestHeader.Advance(AESBlockSize)
|
||||||
|
|
||||||
|
if err := link.Writer.WriteMultiBuffer(buf.MultiBuffer{requestHeader}); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
return TransportTCP(ctx, i.policyManager.ForLevel(targetDest.level), buf.NewReader(conn), buf.NewWriter(conn), link)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (i *RelayInbound) processUDP(ctx context.Context, conn stat.Connection, dispatcher routing.Dispatcher) error {
|
||||||
|
reader := buf.NewPacketReader(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)
|
|
||||||
if err != nil {
|
eiHeader := DecryptUDPEIH(i.relayBlock, packetHeader[:], data[AESBlockSize:2*AESBlockSize])
|
||||||
packet.Release()
|
|
||||||
buf.ReleaseMulti(mb)
|
targetDest, ok := i.destinations[eiHeader]
|
||||||
return err
|
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
|
||||||
|
|
||||||
|
sessionItem := i.udpSessions.GetOrCreate(sessionID)
|
||||||
|
if sessionItem.User == nil {
|
||||||
|
sessionItem.Lock()
|
||||||
|
if sessionItem.User == nil {
|
||||||
|
sessionItem.User = &protocol.MemoryUser{
|
||||||
|
Email: targetDest.email,
|
||||||
|
Level: targetDest.level,
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
sessionItem.Unlock()
|
||||||
}
|
}
|
||||||
|
link, err := sessionItem.EnsureLink(ctx, conn, dest, dispatcher, i.policyManager, nil)
|
||||||
|
if err != nil {
|
||||||
|
b.Release()
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
_ = link.Writer.WriteMultiBuffer(buf.MultiBuffer{b})
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
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,74 @@
|
|||||||
|
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
|
||||||
|
}
|
||||||
|
|
||||||
|
func DecryptEIH(method *CipherMethod, key, salt, eih []byte) ([AESBlockSize]byte, error) {
|
||||||
|
identitySubkey := DeriveIdentitySubKey(key, salt, method.KeySaltLength)
|
||||||
|
block, err := method.NewBlock(identitySubkey)
|
||||||
|
if err != nil {
|
||||||
|
return [AESBlockSize]byte{}, err
|
||||||
|
}
|
||||||
|
var decryptedHash [AESBlockSize]byte
|
||||||
|
block.Decrypt(decryptedHash[:], eih)
|
||||||
|
return decryptedHash, nil
|
||||||
|
}
|
||||||
@@ -2,21 +2,19 @@ package shadowsocks_2022
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"time"
|
"crypto/rand"
|
||||||
|
"io"
|
||||||
|
|
||||||
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 +26,51 @@ 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)
|
||||||
|
}
|
||||||
|
|
||||||
|
if method.IsChaCha && len(pskList) > 1 {
|
||||||
|
return nil, errors.New("multi-key is not supported for chacha20-poly1305")
|
||||||
|
}
|
||||||
|
|
||||||
|
finalPSK := pskList[len(pskList)-1]
|
||||||
|
udpCodec, err := NewUDPPacketCodec(method, pskList)
|
||||||
|
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 +85,140 @@ 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)
|
||||||
|
|
||||||
|
var initialPayload []byte
|
||||||
|
var firstBuf *buf.Buffer
|
||||||
|
var remainingMB buf.MultiBuffer
|
||||||
|
if timeoutReader, ok := link.Reader.(buf.TimeoutReader); ok {
|
||||||
|
if mb, err := timeoutReader.ReadMultiBufferTimeout(0); err == nil && !mb.IsEmpty() {
|
||||||
|
remainingMB, firstBuf = buf.SplitFirst(mb)
|
||||||
|
initialPayload = firstBuf.Bytes()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
bodyWriter, err := WriteTCPRequest(conn, o.method, o.pskList, destination, clientSaltSlice, initialPayload)
|
||||||
|
if firstBuf != nil {
|
||||||
|
firstBuf.Release()
|
||||||
|
}
|
||||||
|
if err != nil {
|
||||||
|
buf.ReleaseMulti(remainingMB)
|
||||||
|
return errors.New("failed to write request").Base(err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if !remainingMB.IsEmpty() {
|
||||||
|
if err := bodyWriter.WriteMultiBuffer(remainingMB); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
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 {
|
||||||
|
session, err := o.udpCodec.NewClientSession()
|
||||||
|
if err != nil {
|
||||||
|
return errors.New("failed to create client udp session").Base(err)
|
||||||
|
}
|
||||||
|
|
||||||
|
requestDone := func() error {
|
||||||
|
defer timer.SetTimeout(sessionPolicy.Timeouts.DownlinkOnly)
|
||||||
|
|
||||||
|
writer := &UDPWriter{
|
||||||
|
Writer: conn,
|
||||||
|
Destination: destination,
|
||||||
|
Session: session,
|
||||||
|
}
|
||||||
|
|
||||||
|
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,
|
||||||
|
Session: session,
|
||||||
|
}
|
||||||
|
|
||||||
|
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,766 @@
|
|||||||
|
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
|
||||||
|
pskList [][]byte
|
||||||
|
psk []byte
|
||||||
|
blockCipher cipher.Block
|
||||||
|
blockCiphers []cipher.Block
|
||||||
|
chachaCipher cipher.AEAD
|
||||||
|
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, pskList [][]byte) (*UDPCodec, error) {
|
||||||
|
if method.IsChaCha && len(pskList) > 1 {
|
||||||
|
return nil, errors.New("multi-key is not supported for chacha20-poly1305")
|
||||||
|
}
|
||||||
|
finalPSK := pskList[len(pskList)-1]
|
||||||
|
c, err := newUDPCodec(method, finalPSK)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
c.pskList = pskList
|
||||||
|
if len(pskList) > 1 {
|
||||||
|
c.blockCiphers = make([]cipher.Block, len(pskList))
|
||||||
|
for i, psk := range pskList {
|
||||||
|
c.blockCiphers[i], err = method.NewBlock(psk)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
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) Sessions() *UDPSessionManager {
|
||||||
|
return c.sessions
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *UDPCodec) GetSession(sessionID uint64) *ServerUDPSession {
|
||||||
|
if c.sessions == nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
return c.sessions.GetOrCreate(sessionID)
|
||||||
|
}
|
||||||
|
|
||||||
|
type DecodedUDPPacket struct {
|
||||||
|
SessionID uint64
|
||||||
|
PacketID uint64
|
||||||
|
HeaderType byte
|
||||||
|
Timestamp uint64
|
||||||
|
ClientSessionID uint64
|
||||||
|
Destination net.Destination
|
||||||
|
Payload []byte
|
||||||
|
}
|
||||||
|
|
||||||
|
func DecryptUDPEIH(block cipher.Block, rawHeader, eih []byte) [AESBlockSize]byte {
|
||||||
|
var decryptedHash [AESBlockSize]byte
|
||||||
|
block.Decrypt(decryptedHash[:], eih)
|
||||||
|
for k := 0; k < AESBlockSize; k++ {
|
||||||
|
decryptedHash[k] ^= rawHeader[k]
|
||||||
|
}
|
||||||
|
return decryptedHash
|
||||||
|
}
|
||||||
|
|
||||||
|
func ParseAddressPort(data []byte) (net.Destination, int, error) {
|
||||||
|
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]
|
||||||
|
if headerType != HeaderTypeClient && headerType != HeaderTypeServer {
|
||||||
|
return DecodedUDPPacket{}, ErrBadHeaderType
|
||||||
|
}
|
||||||
|
epoch := binary.BigEndian.Uint64(bodyPlain[1:9])
|
||||||
|
diff := int(math.Abs(float64(time.Now().Unix() - int64(epoch))))
|
||||||
|
if diff > 30 {
|
||||||
|
return DecodedUDPPacket{}, ErrBadTimestamp
|
||||||
|
}
|
||||||
|
|
||||||
|
offset := 9
|
||||||
|
var clientSessionID uint64
|
||||||
|
if headerType == HeaderTypeServer {
|
||||||
|
if len(bodyPlain) < offset+8+2 {
|
||||||
|
return DecodedUDPPacket{}, ErrPacketTooShort
|
||||||
|
}
|
||||||
|
clientSessionID = binary.BigEndian.Uint64(bodyPlain[offset : offset+8])
|
||||||
|
offset += 8
|
||||||
|
}
|
||||||
|
|
||||||
|
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,
|
||||||
|
ClientSessionID: clientSessionID,
|
||||||
|
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(nil, nonce, ciphertext, nil)
|
||||||
|
if err != nil {
|
||||||
|
return DecodedUDPPacket{}, errors.New("failed to decrypt chacha udp packet").Base(err)
|
||||||
|
}
|
||||||
|
if len(plain) < 16+1+8+2 {
|
||||||
|
return DecodedUDPPacket{}, ErrPacketTooShort
|
||||||
|
}
|
||||||
|
|
||||||
|
sessionID := binary.BigEndian.Uint64(plain[:8])
|
||||||
|
packetID := binary.BigEndian.Uint64(plain[8:16])
|
||||||
|
|
||||||
|
sessionItem := c.sessions.GetOrCreate(sessionID)
|
||||||
|
if !sessionItem.CheckPacketID(packetID) {
|
||||||
|
return DecodedUDPPacket{}, ErrPacketIdNotUnique
|
||||||
|
}
|
||||||
|
|
||||||
|
decoded, err := parsePlainUDPPacket(sessionID, packetID, plain[16:])
|
||||||
|
if err != nil {
|
||||||
|
return DecodedUDPPacket{}, err
|
||||||
|
}
|
||||||
|
|
||||||
|
if decoded.HeaderType != HeaderTypeClient {
|
||||||
|
return DecodedUDPPacket{}, ErrBadHeaderType
|
||||||
|
}
|
||||||
|
|
||||||
|
sessionItem.AddPacketID(packetID)
|
||||||
|
return decoded, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// 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])
|
||||||
|
|
||||||
|
sessionItem := c.sessions.GetOrCreate(sessionID)
|
||||||
|
if !sessionItem.CheckPacketID(packetID) {
|
||||||
|
return DecodedUDPPacket{}, ErrPacketIdNotUnique
|
||||||
|
}
|
||||||
|
|
||||||
|
return sessionItem.DecryptAESPayload(c.method, c.psk, sessionID, packetID, rawHeader[:], data[16:])
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *ServerUDPSession) DecryptAESPayload(method *CipherMethod, psk []byte, sessionID, packetID uint64, rawHeader, bodyCipher []byte) (DecodedUDPPacket, error) {
|
||||||
|
bodyAead := s.clientBodyCipher
|
||||||
|
isNewCipher := false
|
||||||
|
if bodyAead == nil {
|
||||||
|
bodyKey := DeriveSessionSubKey(psk, rawHeader[:8], method.KeySaltLength)
|
||||||
|
var err error
|
||||||
|
bodyAead, err = method.NewAEAD(bodyKey)
|
||||||
|
if err != nil {
|
||||||
|
return DecodedUDPPacket{}, err
|
||||||
|
}
|
||||||
|
isNewCipher = true
|
||||||
|
}
|
||||||
|
|
||||||
|
bodyNonce := rawHeader[4:16]
|
||||||
|
bodyPlain, err := bodyAead.Open(nil, bodyNonce, bodyCipher, nil)
|
||||||
|
if err != nil {
|
||||||
|
return DecodedUDPPacket{}, errors.New("failed to decrypt aes udp body").Base(err)
|
||||||
|
}
|
||||||
|
|
||||||
|
decoded, err := parsePlainUDPPacket(sessionID, packetID, bodyPlain)
|
||||||
|
if err != nil {
|
||||||
|
return DecodedUDPPacket{}, err
|
||||||
|
}
|
||||||
|
|
||||||
|
if decoded.HeaderType != HeaderTypeClient {
|
||||||
|
return DecodedUDPPacket{}, ErrBadHeaderType
|
||||||
|
}
|
||||||
|
|
||||||
|
s.AddPacketID(packetID)
|
||||||
|
|
||||||
|
if isNewCipher {
|
||||||
|
s.clientBodyCipher = bodyAead
|
||||||
|
}
|
||||||
|
|
||||||
|
return decoded, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *ServerUDPSession) EnsureServerState(method *CipherMethod, 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 {
|
||||||
|
var err error
|
||||||
|
s.serverChaCha, err = method.NewUDPCipher(psk)
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
var err error
|
||||||
|
s.serverHeaderBlock, err = method.NewBlock(psk)
|
||||||
|
if err != nil {
|
||||||
|
s.ServerSessionID = 0
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
bodyKey := DeriveSessionSubKey(psk, sidBuf[:], method.KeySaltLength)
|
||||||
|
s.serverBodyCipher, err = method.NewAEAD(bodyKey)
|
||||||
|
if err != nil {
|
||||||
|
s.ServerSessionID = 0
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *ServerUDPSession) EncodeServerPacket(method *CipherMethod, clientSessionID uint64, dest net.Destination, payload []byte) ([]byte, error) {
|
||||||
|
serverSessionID := s.ServerSessionID
|
||||||
|
serverPacketID := s.ServerPacketID.Add(1) - 1
|
||||||
|
|
||||||
|
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.serverHeaderBlock.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.serverBodyCipher.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) {
|
||||||
|
return c.sessions.EncodeServerPacket(c.method, c.psk, clientSessionID, dest, payload)
|
||||||
|
}
|
||||||
|
|
||||||
|
type serverSessionState struct {
|
||||||
|
sessionID uint64
|
||||||
|
window *SlidingWindow
|
||||||
|
cipher cipher.AEAD
|
||||||
|
lastSeen atomic.Int64
|
||||||
|
}
|
||||||
|
|
||||||
|
func (st *serverSessionState) check(packetID uint64) bool {
|
||||||
|
if st.window == nil {
|
||||||
|
st.window = new(SlidingWindow)
|
||||||
|
}
|
||||||
|
return st.window.Check(packetID)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (st *serverSessionState) add(packetID uint64) {
|
||||||
|
if st.window == nil {
|
||||||
|
st.window = new(SlidingWindow)
|
||||||
|
}
|
||||||
|
st.window.Add(packetID)
|
||||||
|
}
|
||||||
|
|
||||||
|
type ClientUDPSession struct {
|
||||||
|
codec *UDPCodec
|
||||||
|
clientSessionID uint64
|
||||||
|
nextPacketID atomic.Uint64
|
||||||
|
clientBodyCipher cipher.AEAD
|
||||||
|
current atomic.Pointer[serverSessionState]
|
||||||
|
old atomic.Pointer[serverSessionState]
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *UDPCodec) NewClientSession() (*ClientUDPSession, error) {
|
||||||
|
var sessID [8]byte
|
||||||
|
if _, err := io.ReadFull(rand.Reader, sessID[:]); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
clientSessionID := binary.BigEndian.Uint64(sessID[:])
|
||||||
|
|
||||||
|
var clientBodyCipher cipher.AEAD
|
||||||
|
var err error
|
||||||
|
if !c.method.IsChaCha {
|
||||||
|
finalPSK := c.psk
|
||||||
|
clientBodyKey := DeriveSessionSubKey(finalPSK, sessID[:], c.method.KeySaltLength)
|
||||||
|
clientBodyCipher, err = c.method.NewAEAD(clientBodyKey)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return &ClientUDPSession{
|
||||||
|
codec: c,
|
||||||
|
clientSessionID: clientSessionID,
|
||||||
|
clientBodyCipher: clientBodyCipher,
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *ClientUDPSession) getServerSession(sessionID uint64, now int64) (*serverSessionState, error) {
|
||||||
|
cur := s.current.Load()
|
||||||
|
if cur != nil && cur.sessionID == sessionID {
|
||||||
|
return cur, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
old := s.old.Load()
|
||||||
|
if old != nil && old.sessionID == sessionID {
|
||||||
|
if now-old.lastSeen.Load() > 60 {
|
||||||
|
s.old.CompareAndSwap(old, nil)
|
||||||
|
return nil, errors.New("old server session expired")
|
||||||
|
}
|
||||||
|
return old, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// New server session:
|
||||||
|
// Spec §3.2.4: reject newer server sessions when the last packet received from the old session is less than 1 minute old.
|
||||||
|
if old != nil && now-old.lastSeen.Load() < 60 {
|
||||||
|
return nil, errors.New("newer server session rejected: old session is less than 1 minute old")
|
||||||
|
}
|
||||||
|
|
||||||
|
var bodyAead cipher.AEAD
|
||||||
|
if !s.codec.method.IsChaCha {
|
||||||
|
var sessBytes [8]byte
|
||||||
|
binary.BigEndian.PutUint64(sessBytes[:], sessionID)
|
||||||
|
bodyKey := DeriveSessionSubKey(s.codec.psk, sessBytes[:], s.codec.method.KeySaltLength)
|
||||||
|
var err error
|
||||||
|
bodyAead, err = s.codec.method.NewAEAD(bodyKey)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
newState := &serverSessionState{
|
||||||
|
sessionID: sessionID,
|
||||||
|
cipher: bodyAead,
|
||||||
|
}
|
||||||
|
newState.lastSeen.Store(now)
|
||||||
|
|
||||||
|
if cur == nil {
|
||||||
|
s.current.CompareAndSwap(nil, newState)
|
||||||
|
return s.current.Load(), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
s.old.Store(cur)
|
||||||
|
s.current.Store(newState)
|
||||||
|
return newState, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *ClientUDPSession) ClientSessionID() uint64 {
|
||||||
|
return s.clientSessionID
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *ClientUDPSession) EncodePacket(dest net.Destination, payload []byte) (*buf.Buffer, error) {
|
||||||
|
packetID := s.nextPacketID.Add(1) - 1
|
||||||
|
sessID := s.clientSessionID
|
||||||
|
|
||||||
|
var paddingLen int
|
||||||
|
if dest.Port == 53 && len(payload) < MaxPaddingLength {
|
||||||
|
paddingLen = mrand.IntN(MaxPaddingLength) + 1
|
||||||
|
}
|
||||||
|
|
||||||
|
addrPortLen := AddrPortLength(dest)
|
||||||
|
|
||||||
|
if s.codec.method.IsChaCha {
|
||||||
|
totalLen := PacketNonceSize + 27 + paddingLen + addrPortLen + len(payload) + AEADTagSize
|
||||||
|
if totalLen > buf.Size {
|
||||||
|
return nil, ErrPacketTooLarge
|
||||||
|
}
|
||||||
|
|
||||||
|
outBuf := buf.New()
|
||||||
|
|
||||||
|
var nonce [PacketNonceSize]byte
|
||||||
|
if _, err := io.ReadFull(rand.Reader, nonce[:]); err != nil {
|
||||||
|
outBuf.Release()
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
outBuf.Write(nonce[:])
|
||||||
|
|
||||||
|
var hdr [16 + 1 + 8 + 2]byte
|
||||||
|
binary.BigEndian.PutUint64(hdr[0:8], sessID)
|
||||||
|
binary.BigEndian.PutUint64(hdr[8:16], packetID)
|
||||||
|
hdr[16] = HeaderTypeClient
|
||||||
|
binary.BigEndian.PutUint64(hdr[17:25], uint64(time.Now().Unix()))
|
||||||
|
binary.BigEndian.PutUint16(hdr[25:27], uint16(paddingLen))
|
||||||
|
outBuf.Write(hdr[:])
|
||||||
|
if paddingLen > 0 {
|
||||||
|
outBuf.Write(zeroPadding[:paddingLen])
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := WriteAddressPort(outBuf, dest); err != nil {
|
||||||
|
outBuf.Release()
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
outBuf.Write(payload)
|
||||||
|
|
||||||
|
plainBytes := outBuf.Bytes()[PacketNonceSize:]
|
||||||
|
outBuf.Extend(int32(s.codec.chachaCipher.Overhead()))
|
||||||
|
s.codec.chachaCipher.Seal(plainBytes[:0], nonce[:], plainBytes, nil)
|
||||||
|
return outBuf, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// AES mode
|
||||||
|
var sessBytes [8]byte
|
||||||
|
binary.BigEndian.PutUint64(sessBytes[:], sessID)
|
||||||
|
|
||||||
|
var rawHeader [16]byte
|
||||||
|
copy(rawHeader[:8], sessBytes[:])
|
||||||
|
binary.BigEndian.PutUint64(rawHeader[8:16], packetID)
|
||||||
|
|
||||||
|
eihCount := 0
|
||||||
|
if len(s.codec.pskList) > 1 {
|
||||||
|
eihCount = len(s.codec.pskList) - 1
|
||||||
|
}
|
||||||
|
|
||||||
|
totalLen := 16 + eihCount*16 + 11 + paddingLen + addrPortLen + len(payload) + AEADTagSize
|
||||||
|
if totalLen > buf.Size {
|
||||||
|
return nil, ErrPacketTooLarge
|
||||||
|
}
|
||||||
|
|
||||||
|
outBuf := buf.New()
|
||||||
|
|
||||||
|
if len(s.codec.pskList) > 1 {
|
||||||
|
var encryptedHeader [16]byte
|
||||||
|
s.codec.blockCiphers[0].Encrypt(encryptedHeader[:], rawHeader[:])
|
||||||
|
outBuf.Write(encryptedHeader[:])
|
||||||
|
|
||||||
|
for i := 0; i < len(s.codec.pskList)-1; i++ {
|
||||||
|
nextPSK := s.codec.pskList[i+1]
|
||||||
|
pskHash := DeriveUserPSKHash(nextPSK)
|
||||||
|
var eihPlain [16]byte
|
||||||
|
for k := 0; k < 16; k++ {
|
||||||
|
eihPlain[k] = pskHash[k] ^ rawHeader[k]
|
||||||
|
}
|
||||||
|
var encryptedEIH [16]byte
|
||||||
|
s.codec.blockCiphers[i].Encrypt(encryptedEIH[:], eihPlain[:])
|
||||||
|
outBuf.Write(encryptedEIH[:])
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
var encryptedHeader [16]byte
|
||||||
|
s.codec.blockCipher.Encrypt(encryptedHeader[:], rawHeader[:])
|
||||||
|
outBuf.Write(encryptedHeader[:])
|
||||||
|
}
|
||||||
|
|
||||||
|
bodyAead := s.clientBodyCipher
|
||||||
|
|
||||||
|
var hdr [1 + 8 + 2]byte
|
||||||
|
hdr[0] = HeaderTypeClient
|
||||||
|
binary.BigEndian.PutUint64(hdr[1:9], uint64(time.Now().Unix()))
|
||||||
|
binary.BigEndian.PutUint16(hdr[9:11], uint16(paddingLen))
|
||||||
|
outBuf.Write(hdr[:])
|
||||||
|
if paddingLen > 0 {
|
||||||
|
outBuf.Write(zeroPadding[:paddingLen])
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := WriteAddressPort(outBuf, dest); err != nil {
|
||||||
|
outBuf.Release()
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
outBuf.Write(payload)
|
||||||
|
|
||||||
|
headerOffset := 16 + eihCount*16
|
||||||
|
plainBytes := outBuf.Bytes()[headerOffset:]
|
||||||
|
bodyNonce := rawHeader[4:16]
|
||||||
|
outBuf.Extend(int32(bodyAead.Overhead()))
|
||||||
|
bodyAead.Seal(plainBytes[:0], bodyNonce, plainBytes, nil)
|
||||||
|
return outBuf, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *ClientUDPSession) DecodePacket(data []byte) (DecodedUDPPacket, error) {
|
||||||
|
if len(data) < PacketMinimalHeaderSize {
|
||||||
|
return DecodedUDPPacket{}, ErrPacketTooShort
|
||||||
|
}
|
||||||
|
|
||||||
|
if s.codec.method.IsChaCha {
|
||||||
|
if len(data) < PacketNonceSize+AEADTagSize {
|
||||||
|
return DecodedUDPPacket{}, ErrPacketTooShort
|
||||||
|
}
|
||||||
|
nonce := data[:PacketNonceSize]
|
||||||
|
ciphertext := data[PacketNonceSize:]
|
||||||
|
plain, err := s.codec.chachaCipher.Open(nil, nonce, ciphertext, nil)
|
||||||
|
if err != nil {
|
||||||
|
return DecodedUDPPacket{}, errors.New("failed to decrypt chacha udp packet").Base(err)
|
||||||
|
}
|
||||||
|
if len(plain) < 16+1+8+2 {
|
||||||
|
return DecodedUDPPacket{}, ErrPacketTooShort
|
||||||
|
}
|
||||||
|
|
||||||
|
sessionID := binary.BigEndian.Uint64(plain[:8])
|
||||||
|
packetID := binary.BigEndian.Uint64(plain[8:16])
|
||||||
|
|
||||||
|
now := time.Now().Unix()
|
||||||
|
st, err := s.getServerSession(sessionID, now)
|
||||||
|
if err != nil {
|
||||||
|
return DecodedUDPPacket{}, err
|
||||||
|
}
|
||||||
|
if !st.check(packetID) {
|
||||||
|
return DecodedUDPPacket{}, ErrPacketIdNotUnique
|
||||||
|
}
|
||||||
|
|
||||||
|
decoded, err := parsePlainUDPPacket(sessionID, packetID, plain[16:])
|
||||||
|
if err != nil {
|
||||||
|
return DecodedUDPPacket{}, err
|
||||||
|
}
|
||||||
|
|
||||||
|
if decoded.HeaderType != HeaderTypeServer {
|
||||||
|
return DecodedUDPPacket{}, ErrBadHeaderType
|
||||||
|
}
|
||||||
|
if decoded.ClientSessionID != s.clientSessionID {
|
||||||
|
return DecodedUDPPacket{}, errors.New("client session ID mismatch")
|
||||||
|
}
|
||||||
|
|
||||||
|
st.add(packetID)
|
||||||
|
st.lastSeen.Store(now)
|
||||||
|
|
||||||
|
return decoded, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// AES mode
|
||||||
|
var rawHeader [16]byte
|
||||||
|
s.codec.blockCipher.Decrypt(rawHeader[:], data[:16])
|
||||||
|
sessionID := binary.BigEndian.Uint64(rawHeader[:8])
|
||||||
|
packetID := binary.BigEndian.Uint64(rawHeader[8:16])
|
||||||
|
|
||||||
|
now := time.Now().Unix()
|
||||||
|
st, err := s.getServerSession(sessionID, now)
|
||||||
|
if err != nil {
|
||||||
|
return DecodedUDPPacket{}, err
|
||||||
|
}
|
||||||
|
if !st.check(packetID) {
|
||||||
|
return DecodedUDPPacket{}, ErrPacketIdNotUnique
|
||||||
|
}
|
||||||
|
bodyAead := st.cipher
|
||||||
|
|
||||||
|
bodyNonce := rawHeader[4:16]
|
||||||
|
bodyCipher := data[16:]
|
||||||
|
bodyPlain, err := bodyAead.Open(nil, bodyNonce, bodyCipher, nil)
|
||||||
|
if err != nil {
|
||||||
|
return DecodedUDPPacket{}, errors.New("failed to decrypt aes udp body").Base(err)
|
||||||
|
}
|
||||||
|
|
||||||
|
decoded, err := parsePlainUDPPacket(sessionID, packetID, bodyPlain)
|
||||||
|
if err != nil {
|
||||||
|
return DecodedUDPPacket{}, err
|
||||||
|
}
|
||||||
|
|
||||||
|
if decoded.HeaderType != HeaderTypeServer {
|
||||||
|
return DecodedUDPPacket{}, ErrBadHeaderType
|
||||||
|
}
|
||||||
|
if decoded.ClientSessionID != s.clientSessionID {
|
||||||
|
return DecodedUDPPacket{}, errors.New("client session ID mismatch")
|
||||||
|
}
|
||||||
|
|
||||||
|
st.add(packetID)
|
||||||
|
st.lastSeen.Store(now)
|
||||||
|
|
||||||
|
return decoded, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
type UDPWriter struct {
|
||||||
|
Writer io.Writer
|
||||||
|
Destination net.Destination
|
||||||
|
Session *ClientUDPSession
|
||||||
|
}
|
||||||
|
|
||||||
|
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.Session.EncodePacket(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
|
||||||
|
Session *ClientUDPSession
|
||||||
|
}
|
||||||
|
|
||||||
|
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.Session.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,376 @@
|
|||||||
|
package shadowsocks_2022_test
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"crypto/rand"
|
||||||
|
"encoding/base64"
|
||||||
|
"encoding/binary"
|
||||||
|
"errors"
|
||||||
|
"io"
|
||||||
|
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
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRelayTCPHandshakeForwarding(t *testing.T) {
|
||||||
|
methods := []string{MethodAES128GCM, MethodAES256GCM}
|
||||||
|
for _, methodName := range methods {
|
||||||
|
t.Run(methodName, func(t *testing.T) {
|
||||||
|
method, err := GetCipherMethod(methodName)
|
||||||
|
common.Must(err)
|
||||||
|
|
||||||
|
relayKey := make([]byte, method.KeySaltLength)
|
||||||
|
destKey := make([]byte, method.KeySaltLength)
|
||||||
|
_, _ = io.ReadFull(rand.Reader, relayKey)
|
||||||
|
_, _ = io.ReadFull(rand.Reader, destKey)
|
||||||
|
|
||||||
|
targetPort := uint32(54321)
|
||||||
|
relayConfig := &RelayServerConfig{
|
||||||
|
Method: methodName,
|
||||||
|
Key: base64.StdEncoding.EncodeToString(relayKey),
|
||||||
|
Destinations: []*RelayDestination{
|
||||||
|
{
|
||||||
|
Key: base64.StdEncoding.EncodeToString(destKey),
|
||||||
|
Address: net.NewIPOrDomain(net.LocalHostIP),
|
||||||
|
Port: targetPort,
|
||||||
|
Email: "test@xray.com",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
testCtx := newTestContext()
|
||||||
|
inbound, err := NewRelayServer(testCtx, relayConfig)
|
||||||
|
common.Must(err)
|
||||||
|
|
||||||
|
targetDest := net.TCPDestination(net.LocalHostIP, net.Port(targetPort))
|
||||||
|
|
||||||
|
downstreamR, downstreamW := gonet.Pipe()
|
||||||
|
defer downstreamR.Close()
|
||||||
|
defer downstreamW.Close()
|
||||||
|
|
||||||
|
disp := &dummyDispatcher{
|
||||||
|
onDispatch: func(ctx context.Context, dest net.Destination) (*transport.Link, error) {
|
||||||
|
inLink := &transport.Link{
|
||||||
|
Reader: buf.NewReader(downstreamR),
|
||||||
|
Writer: &customWriter{
|
||||||
|
write: func(mb buf.MultiBuffer) error {
|
||||||
|
defer buf.ReleaseMulti(mb)
|
||||||
|
for _, b := range mb {
|
||||||
|
if _, err := downstreamW.Write(b.Bytes()); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
return inLink, nil
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
clientConn, relayConn := gonet.Pipe()
|
||||||
|
defer clientConn.Close()
|
||||||
|
defer relayConn.Close()
|
||||||
|
|
||||||
|
go func() {
|
||||||
|
_ = inbound.Process(testCtx, net.Network_TCP, &dummyStatConn{Conn: relayConn}, disp)
|
||||||
|
}()
|
||||||
|
|
||||||
|
clientSalt := make([]byte, method.KeySaltLength)
|
||||||
|
_, _ = io.ReadFull(rand.Reader, clientSalt)
|
||||||
|
pskList := [][]byte{relayKey, destKey}
|
||||||
|
|
||||||
|
go func() {
|
||||||
|
_, err := WriteTCPRequest(clientConn, method, pskList, targetDest, clientSalt, []byte("relay payload"))
|
||||||
|
if err != nil {
|
||||||
|
t.Errorf("WriteTCPRequest failed: %v", err)
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
|
||||||
|
// Downstream server must be able to read Salt + Fixed chunk in a single Read call!
|
||||||
|
headerLen := method.KeySaltLength + RequestHeaderFixedChunkLength + AEADTagSize
|
||||||
|
headerBuf := make([]byte, headerLen)
|
||||||
|
n, err := downstreamR.Read(headerBuf)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("downstream failed to read handshake: %v", err)
|
||||||
|
}
|
||||||
|
if n < headerLen {
|
||||||
|
t.Fatalf("downstream expected single read >= %d bytes, got %d", headerLen, n)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Verify downstream can decode the fixed chunk and subsequent payload
|
||||||
|
sessionKey := DeriveSessionSubKey(destKey, headerBuf[:method.KeySaltLength], method.KeySaltLength)
|
||||||
|
aead, err := method.NewAEAD(sessionKey)
|
||||||
|
common.Must(err)
|
||||||
|
|
||||||
|
reader := NewStreamReader(downstreamR, aead)
|
||||||
|
reqHeader, err := ReadClientRequestHeaderWithFixed(reader, headerBuf[method.KeySaltLength:])
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("downstream failed to parse client request header: %v", err)
|
||||||
|
}
|
||||||
|
if string(reqHeader.EarlyData) != "relay payload" {
|
||||||
|
t.Fatalf("payload mismatch: expected 'relay payload', got '%s'", string(reqHeader.EarlyData))
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user